1-layer-addition / transformer_block.py
melephant's picture
Publish addition-transformer run s85nnxtf
58223a8 verified
Raw
History Blame Contribute Delete
1.6 kB
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,
)