Repository navigation
Replies: 1 comment
|
Ah, I figured out the problem. I was passing the same instance to each TransformerBlock right below. Will edit the original post to show. For anyone with the same problem, this is the correct pattern if each layer needs an instance (kind of obvious in hindsight): class TransformerEncoder(nnx.Module):
def __init__(
self,
num_layers: int,
d_model: int,
n_heads: int,
d_ff: int,
*,
dropout: float = 0.0,
rel_pos_bias_cls: Type[AbstractRelativeBias] = RelativePositionBias,
rngs
):
super().__init__()
@nnx.split_rngs(splits=num_layers)
@nnx.vmap(in_axes=(0,), out_axes=0)
def create_block(rngs: nnx.Rngs):
rel_pos_bias = rel_pos_bias_cls(n_heads, rngs=rngs) # <-- give each block it's own bias
return TransformerBlock(
d_model,
n_heads,
d_ff,
dropout=dropout,
rel_pos_bias_module=rel_pos_bias,
rngs=rngs
)
self.blocks = create_block(rngs)
self.final_ln = nnx.LayerNorm(d_model, rngs=rngs)If you need the parameters to be shared, this ain't for you. Not sure how that is done, but I'll have to figure it out soon. |
0 replies
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
I feel a little stupid, but I can't figure it out and I can't find anything online either (which most likely has to do with me not knowing how to properly describe what I want to do to a search engine).
Basically I just want to pass a nnx.Module as a parameter when initializing my Model.
What I have (only bottom 5 lines should matter):
What I'm getting:
Maybe important for context: My train step is jitted using jax.jit (with nnx.merge and nnx.split to make it work) and tracing (and running) worked fine before I added this functionality.
If this is the completely wrong approach, just let me know, I'd actually be happy to learn the cleanest way to do this.
All reactions