1-layer-addition / attention.py
melephant's picture
Publish addition-transformer run s85nnxtf
58223a8 verified
Raw
History Blame Contribute Delete
2.25 kB
from __future__ import annotations
from dataclasses import dataclass
import torch
import torch.nn.functional as F
from torch import nn
@dataclass(frozen=True)
class AttentionOutput:
values: torch.Tensor
pattern: torch.Tensor | None = None
class CausalSelfAttention(nn.Module):
def __init__(self, d_model: int, n_heads: int, bias: bool = False) -> None:
super().__init__()
if d_model % n_heads != 0:
raise ValueError("d_model must be divisible by n_heads.")
self.d_model = d_model
self.n_heads = n_heads
self.d_head = d_model // n_heads
self.W_Q = nn.Linear(d_model, d_model, bias=bias)
self.W_K = nn.Linear(d_model, d_model, bias=bias)
self.W_V = nn.Linear(d_model, d_model, bias=bias)
self.W_O = nn.Linear(d_model, d_model, bias=bias)
self.register_buffer("_causal_mask", torch.empty(0, 0, dtype=torch.bool), persistent=False)
def forward(self, x: torch.Tensor, return_pattern: bool = False) -> AttentionOutput:
batch, seq_len, _ = x.shape
q = self._split_heads(self.W_Q(x))
k = self._split_heads(self.W_K(x))
v = self._split_heads(self.W_V(x))
scores = torch.matmul(q, k.transpose(-1, -2)) / self.d_head
if self._causal_mask.shape[0] < seq_len or self._causal_mask.device != x.device:
self._causal_mask = torch.triu(
torch.ones(seq_len, seq_len, dtype=torch.bool, device=x.device),
diagonal=1,
)
causal_mask = self._causal_mask[:seq_len, :seq_len]
scores = scores.masked_fill(causal_mask, float("-inf"))
pattern = F.softmax(scores, dim=-1)
attended = torch.matmul(pattern, v) / 3.0
values = self.W_O(self._merge_heads(attended, batch, seq_len))
return AttentionOutput(values=values, pattern=pattern if return_pattern else None)
def _split_heads(self, x: torch.Tensor) -> torch.Tensor:
batch, seq_len, _ = x.shape
return x.view(batch, seq_len, self.n_heads, self.d_head).transpose(1, 2)
def _merge_heads(self, x: torch.Tensor, batch: int, seq_len: int) -> torch.Tensor:
return x.transpose(1, 2).contiguous().view(batch, seq_len, self.d_model)