import math import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.modeling_outputs import CausalLMOutputWithPast from .configuration_picolm import PicoLMConfig class RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-5): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def forward(self, x): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.weight def precompute_rope_cis(dim: int, max_seq_len: int, theta: float = 10000.0): freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) t = torch.arange(max_seq_len, dtype=torch.float32) freqs = torch.outer(t, freqs) return torch.cos(freqs), torch.sin(freqs) def apply_rope(x, cos, sin): B, H, S, D = x.shape cos = cos[:S, :].to(x.device).unsqueeze(0).unsqueeze(0) sin = sin[:S, :].to(x.device).unsqueeze(0).unsqueeze(0) x1, x2 = x[..., : D // 2], x[..., D // 2 :] return torch.cat([x1 * cos - x2 * sin, x1 * sin + x2 * cos], dim=-1) class Attention(nn.Module): def __init__(self, args): super().__init__() self.args = args self.head_dim = args.hidden_size // args.num_attention_heads self.n_heads = args.num_attention_heads self.n_kv_heads = args.num_key_value_heads self.num_reps = self.n_heads // self.n_kv_heads self.q_proj = nn.Linear(args.hidden_size, self.n_heads * self.head_dim, bias=False) self.k_proj = nn.Linear(args.hidden_size, self.n_kv_heads * self.head_dim, bias=False) self.v_proj = nn.Linear(args.hidden_size, self.n_kv_heads * self.head_dim, bias=False) self.out_proj = nn.Linear(args.n_heads * self.head_dim, args.hidden_size, bias=False) self.q_norm = RMSNorm(self.head_dim, eps=args.rms_norm_eps) self.k_norm = RMSNorm(self.head_dim, eps=args.rms_norm_eps) def forward(self, x, cos, sin): B, S, C = x.shape q = apply_rope(self.q_norm(self.q_proj(x).view(B, S, self.n_heads, self.head_dim).transpose(1, 2)), cos, sin) k = apply_rope(self.k_norm(self.k_proj(x).view(B, S, self.n_kv_heads, self.head_dim).transpose(1, 2)), cos, sin) v = self.v_proj(x).view(B, S, self.n_kv_heads, self.head_dim).transpose(1, 2) if self.num_reps > 1: k = k[:, :, None, :, :].expand(B, self.n_kv_heads, self.num_reps, S, self.head_dim).reshape(B, self.n_heads, S, self.head_dim) v = v[:, :, None, :, :].expand(B, self.n_kv_heads, self.num_reps, S, self.head_dim).reshape(B, self.n_heads, S, self.head_dim) attn_out = F.scaled_dot_product_attention(q, k, v, is_causal=True) return self.out_proj(attn_out.transpose(1, 2).contiguous().view(B, S, C)) class SwiGLUMLP(nn.Module): def __init__(self, args): super().__init__() self.gate_proj = nn.Linear(args.hidden_size, args.intermediate_size, bias=False) self.up_proj = nn.Linear(args.hidden_size, args.intermediate_size, bias=False) self.down_proj = nn.Linear(args.intermediate_size, args.hidden_size, bias=False) def forward(self, x): return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) class TransformerBlock(nn.Module): def __init__(self, args): super().__init__() self.attn_norm = RMSNorm(args.hidden_size, eps=args.rms_norm_eps) self.attn = Attention(args) self.mlp_norm = RMSNorm(args.hidden_size, eps=args.rms_norm_eps) self.mlp = SwiGLUMLP(args) def forward(self, x, cos, sin): return x + self.mlp(self.mlp_norm(x + self.attn(self.attn_norm(x), cos, sin))) class PicoLMForCausalLM(PreTrainedModel): config_class = PicoLMConfig def __init__(self, config): super().__init__(config) self.config = config self.tok_embeddings = nn.Embedding(config.vocab_size, config.hidden_size) self.blocks = nn.ModuleList([TransformerBlock(config) for _ in range(config.num_hidden_layers)]) self.final_norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.lm_head.weight = self.tok_embeddings.weight cos, sin = precompute_rope_cis(config.hidden_size // config.num_attention_heads, config.max_position_embeddings, config.rope_theta) self.register_buffer("rope_cos", cos, persistent=False) self.register_buffer("rope_sin", sin, persistent=False) def forward(self, input_ids, labels=None, **kwargs): B, S = input_ids.shape x = self.tok_embeddings(input_ids) cos, sin = self.rope_cos[:S], self.rope_sin[:S] for block in self.blocks: x = block(x, cos, sin) logits = self.lm_head(self.final_norm(x)) loss = None if labels is not None: shift_logits = logits[..., :-1, :].contiguous() shift_labels = labels[..., 1:].contiguous() loss = F.cross_entropy(shift_logits.view(-1, self.config.vocab_size), shift_labels.view(-1), ignore_index=-100) return CausalLMOutputWithPast(loss=loss, logits=logits)