LogSentry-LM / architecture.py
sankath's picture
Upload architecture.py with huggingface_hub
cd1ec5e verified
Raw
History Blame Contribute Delete
3.3 kB
# file: architecture.py
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class MultiHeadAttention(nn.Module):
"""Scaled dot-product self-attention with a fused QKV projection."""
def __init__(self, d_model, n_heads):
super().__init__()
assert d_model % n_heads == 0, "d_model must be divisible by n_heads"
self.d_model = d_model
self.n_heads = n_heads
self.d_k = d_model // n_heads
# One matmul for Q, K and V is cheaper than three separate projections
self.qkv_projection = nn.Linear(d_model, 3 * d_model, bias=False)
self.out_projection = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, is_causal=True):
batch_size, seq_len, d_model = x.shape
qkv = self.qkv_projection(x)
q, k, v = qkv.chunk(3, dim=-1)
# (B, T, d_model) -> (B, heads, T, d_k) so each head runs independently
q = q.view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2)
k = k.view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2)
v = v.view(batch_size, seq_len, self.n_heads, self.d_k).transpose(1, 2)
# Scale by sqrt(d_k) — keeps the scores from blowing up and the softmax stable
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
if is_causal:
# Block attention to future tokens — required for next-token prediction,
# otherwise the model would just read the answer.
mask = torch.tril(torch.ones(seq_len, seq_len, device=x.device)).view(1, 1, seq_len, seq_len)
scores = scores.masked_fill(mask == 0, float('-inf'))
attention_weights = F.softmax(scores, dim=-1)
context = torch.matmul(attention_weights, v)
# Merge heads back: (B, heads, T, d_k) -> (B, T, d_model)
context = context.transpose(1, 2).contiguous().view(batch_size, seq_len, d_model)
return self.out_projection(context)
class FeedForwardNetwork(nn.Module):
"""Per-token MLP: expand to d_ff, non-linearity, project back."""
def __init__(self, d_model, d_ff, dropout=0.1):
super().__init__()
self.linear1 = nn.Linear(d_model, d_ff)
self.dropout = nn.Dropout(dropout)
self.linear2 = nn.Linear(d_ff, d_model)
def forward(self, x):
return self.linear2(self.dropout(F.gelu(self.linear1(x))))
class TransformerBlock(nn.Module):
"""One decoder block: attention + feed-forward, pre-norm with residuals."""
def __init__(self, d_model, n_heads, dropout=0.1):
super().__init__()
self.attention = MultiHeadAttention(d_model, n_heads)
self.feed_forward = FeedForwardNetwork(d_model, d_ff=4 * d_model, dropout=dropout) # 4x is the usual ratio
self.ln1 = nn.LayerNorm(d_model)
self.ln2 = nn.LayerNorm(d_model)
self.dropout = nn.Dropout(dropout)
def forward(self, x, is_causal=True):
# Pre-norm (LayerNorm before each sub-layer) trains more stably than post-norm.
# The `x + ...` residuals let gradients flow and each layer refine rather than replace.
x = x + self.dropout(self.attention(self.ln1(x), is_causal=is_causal))
x = x + self.dropout(self.feed_forward(self.ln2(x)))
return x