from models.attention import CausalSelfAttention import torch.nn as nn import torch class TransformerBlock(nn.Module): def __init__(self, embed_dim): super().__init__() self.ln1 = nn.LayerNorm(embed_dim) self.attention = CausalSelfAttention( embed_dim ) self.ln2 = nn.LayerNorm(embed_dim) # FFN self.ffn = nn.Sequential( nn.Linear( in_features=embed_dim, out_features=4*embed_dim ), nn.GELU(), nn.Linear( in_features=4*embed_dim, out_features=embed_dim ) ) def forward(self, x): x = x + self.attention( self.ln1(x) ) x = x + self.ffn( self.ln2(x) ) return x