from __future__ import annotations from dataclasses import dataclass import torch from torch import nn from .attention import CausalSelfAttention from .bilinear_mlp import BilinearMLP from .model_config import AdditionModelConfig @dataclass(frozen=True) class TransformerBlockOutput: residual_pre_attention: torch.Tensor attention_out: torch.Tensor attention_pattern: torch.Tensor | None residual_after_attention: torch.Tensor mlp_out: torch.Tensor residual_after_mlp: torch.Tensor class TransformerBlock(nn.Module): def __init__(self, config: AdditionModelConfig) -> None: super().__init__() self.residual_alpha = config.residual_alpha self.attention = CausalSelfAttention(config.d_model, config.n_heads, bias=False) self.mlp = BilinearMLP(config.d_model, config.d_mlp, bias=False) def forward(self, x: torch.Tensor, return_pattern: bool = False) -> TransformerBlockOutput: attention_output = self.attention(x, return_pattern=return_pattern) residual_after_attention = torch.lerp(x, attention_output.values, self.residual_alpha) mlp_out = self.mlp(residual_after_attention) residual_after_mlp = torch.lerp(residual_after_attention, mlp_out, self.residual_alpha) return TransformerBlockOutput( residual_pre_attention=x, attention_out=attention_output.values, attention_pattern=attention_output.pattern, residual_after_attention=residual_after_attention, mlp_out=mlp_out, residual_after_mlp=residual_after_mlp, )