import torch import torch.nn as nn class TransformerBlock(nn.Module): def __init__(self, d_model, n_heads, d_ff, dropout=0.1): super().__init__() self.norm1 = nn.LayerNorm(d_model) self.attention = MultiHeadAttention(d_model, n_heads) self.norm2 = nn.LayerNorm(d_model) self.feed_forward = nn.Sequential( nn.Linear(d_model, d_ff), nn.GELU(), nn.Linear(d_ff, d_model) ) self.dropout = nn.Dropout(dropout) def forward(self, x): h = x + self.dropout(self.attention(self.norm1(x))) x = h + self.dropout(self.feed_forward(self.norm2(h))) return x