import torch.nn as nn import torch class CausalSelfAttention(nn.Module): def __init__(self, embed_dim): super().__init__() self.embed_dim = embed_dim # Query | Key | Value self.query = nn.Linear( embed_dim, embed_dim ) self.key = nn.Linear( embed_dim, embed_dim ) self.value = nn.Linear( embed_dim, embed_dim ) self.out = nn.Linear(embed_dim, embed_dim) def forward(self, x): batch_size, seq_len, embed_dim = x.shape Q = self.query(x) K = self.key(x) V = self.value(x) scores = Q @ K.transpose(-2, -1) scores = scores / (embed_dim ** 0.5) # Causal Mask mask = torch.triu( torch.ones(seq_len, seq_len, device=x.device), diagonal=1 ).bool() scores = scores.masked_fill(mask, float("-inf")) attention_weight = torch.softmax(scores, dim=-1) output = attention_weight @ V output = self.out(output) return output