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

    # Token identity and position enter the residual stream.
    x = params["token_embed"][tokens]
    x = x + params["pos_embed"][None, :t, :]

    # Pre-LN attention sublayer.
    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

    # Pre-LN feed-forward sublayer.
    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

    # A Pre-LN stack normally has a final norm before its head.
    x = layer_norm(x, params["ln_f_g"], params["ln_f_b"])

    # Masked mean pooling for one prediction per sequence.
    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