def model_forward(params, tokens, return_attention=False):
mask = make_padding_mask(tokens)
b, t = tokens.shape
d = CFG["d_model"]
n_heads = CFG["n_heads"]
assert d % n_heads == 0
x = params["token_embed"][tokens]
x = x + params["pos_embed"][None, :t, :]
residual = x
attn_input = layer_norm(
x, params["ln1_g"], params["ln1_b"]
)
attended, weights = multi_head_attention(
attn_input, mask,
params["wq"], params["wk"], params["wv"], params["wo"],
n_heads,
)
x = residual + attended
residual = x
ff_input = layer_norm(
x, params["ln2_g"], params["ln2_b"]
)
ff = gelu(ff_input @ params["ff1_w"] + params["ff1_b"])
ff = ff @ params["ff2_w"] + params["ff2_b"]
x = residual + ff
x = layer_norm(x, params["ln_f_g"], params["ln_f_b"])
float_mask = mask[..., None].astype(x.dtype)
valid_count = jnp.maximum(
jnp.sum(float_mask, axis=1), 1.0
)
pooled = jnp.sum(x * float_mask, axis=1) / valid_count
logits = pooled @ params["head_w"] + params["head_b"]
if return_attention:
return logits, weights
return logits