class TransformerBlock(nnx.Module):
"""Configurable transformer block: attention + FFN on a residual
stream."""
def __init__(self, num_hiddens, num_heads, dropout=0, norm='rms',
act='swiglu', pre_norm=True, bias=False, attn_factory=None,
ffn_factory=None, rngs=None):
rngs = nnx.Rngs(params=0, dropout=1) if rngs is None else rngs
assert norm in ('rms', 'layer'), f'unknown norm: {norm!r}'
assert num_hiddens % num_heads == 0
self.pre_norm = pre_norm
make_norm = nnx.RMSNorm if norm == 'rms' else nnx.LayerNorm
self.norm1 = make_norm(num_hiddens, rngs=rngs)
self.norm2 = make_norm(num_hiddens, rngs=rngs)
self.attention = (d2l.MultiHeadAttention(num_hiddens, num_heads,
dropout, bias=bias,
rngs=rngs)
if attn_factory is None else attn_factory(rngs))
self.ffn = (FeedForward(num_hiddens, act, bias=bias, rngs=rngs)
if ffn_factory is None else ffn_factory(rngs))
self.dropout = nnx.Dropout(dropout, rngs=rngs)
def __call__(self, X, valid_lens=None):
if self.pre_norm:
Y = self.norm1(X)
X = X + self.dropout(self.attention(Y, Y, Y, valid_lens)[0])
return X + self.dropout(self.ffn(self.norm2(X)))
X = self.norm1(X + self.dropout(
self.attention(X, X, X, valid_lens)[0]))
return self.norm2(X + self.dropout(self.ffn(X)))