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)