class TransformerBlock(nn.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):
super().__init__()
assert norm in ('rms', 'layer'), f'unknown norm: {norm!r}'
assert num_hiddens % num_heads == 0
self.pre_norm = pre_norm
make_norm = nn.RMSNorm if norm == 'rms' else nn.LayerNorm
self.norm1, self.norm2 = make_norm(num_hiddens), make_norm(num_hiddens)
self.attention = (d2l.MultiHeadAttention(num_hiddens, num_heads,
dropout, bias=bias)
if attn_factory is None else attn_factory())
self.ffn = (FeedForward(num_hiddens, act, bias=bias)
if ffn_factory is None else ffn_factory())
self.dropout = nn.Dropout(dropout)
def forward(self, X, valid_lens=None):
if self.pre_norm:
Y = self.norm1(X)
X = X + self.dropout(self.attention(Y, Y, Y, valid_lens))
return X + self.dropout(self.ffn(self.norm2(X)))
X = self.norm1(X + self.dropout(self.attention(X, X, X, valid_lens)))
return self.norm2(X + self.dropout(self.ffn(X)))