import torch import torch.nn as nn import torch.nn.functional as F from transformers import GenerationMixin, PreTrainedModel from transformers.modeling_outputs import CausalLMOutput from .configuration_pulvis import PulvisConfig def _rope(x, cos, sin): x1, x2 = x.chunk(2, dim=-1) return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1) class PulvisAttention(nn.Module): def __init__(self, c): super().__init__() self.h, self.kv, self.hd = c.num_attention_heads, c.num_key_value_heads, c.head_dim self.qkv = nn.Linear(c.hidden_size, (self.h + 2 * self.kv) * self.hd, bias=False) self.o = nn.Linear(self.h * self.hd, c.hidden_size, bias=False) self.eps = c.rms_norm_eps self.q_w = nn.Parameter(torch.ones(self.hd)) self.k_w = nn.Parameter(torch.ones(self.hd)) def forward(self, x, cos, sin): B, T, _ = x.shape q, k, v = self.qkv(x).view(B, T, self.h + 2 * self.kv, self.hd).split([self.h, self.kv, self.kv], dim=2) q = F.rms_norm(q, (self.hd,), self.q_w, self.eps) k = F.rms_norm(k, (self.hd,), self.k_w, self.eps) q = _rope(q, cos, sin).transpose(1, 2) k = _rope(k, cos, sin).transpose(1, 2) v = v.transpose(1, 2) y = F.scaled_dot_product_attention(q, k, v, is_causal=True, enable_gqa=self.kv != self.h) r = self.h // self.kv y5 = y.view(B, self.kv, r, T, self.hd) v5 = v.unsqueeze(2) coef = (y5 * v5).sum(-1, keepdim=True) / (v5 * v5).sum(-1, keepdim=True).clamp_min(1e-6) y = (y5 - coef * v5).view(B, self.h, T, self.hd) return self.o(y.transpose(1, 2).reshape(B, T, self.h * self.hd)) class PulvisMLP(nn.Module): def __init__(self, c): super().__init__() self.up = nn.Linear(c.hidden_size, 2 * c.intermediate_size, bias=False) self.down = nn.Linear(c.intermediate_size, c.hidden_size, bias=False) def forward(self, x): g, u = self.up(x).chunk(2, dim=-1) return self.down(F.silu(g) * u) class PulvisBlock(nn.Module): def __init__(self, c): super().__init__() self.n1 = nn.Parameter(torch.ones(c.hidden_size)) self.n2 = nn.Parameter(torch.ones(c.hidden_size)) self.attn = PulvisAttention(c) self.mlp = PulvisMLP(c) self.eps = c.rms_norm_eps def forward(self, x, cos, sin): x = x + self.attn(F.rms_norm(x, (x.size(-1),), self.n1, self.eps), cos, sin) return x + self.mlp(F.rms_norm(x, (x.size(-1),), self.n2, self.eps)) class PulvisPreTrainedModel(PreTrainedModel): config_class = PulvisConfig base_model_prefix = "model" _no_split_modules = ["PulvisBlock"] _supports_sdpa = True class PulvisForCausalLM(PulvisPreTrainedModel, GenerationMixin): def __init__(self, config): super().__init__(config) c = config self.embed = nn.Embedding(c.vocab_size, c.hidden_size) self.prelude = nn.ModuleList([PulvisBlock(c) for _ in range(c.prelude_layers)]) self.core = nn.ModuleList([PulvisBlock(c) for _ in range(c.core_layers)]) self.coda = nn.ModuleList([PulvisBlock(c) for _ in range(c.coda_layers)]) self.loop_emb = nn.Parameter(torch.zeros(c.core_loops, c.hidden_size)) if c.core_loops > 1 else None self.norm_out = nn.Parameter(torch.ones(c.hidden_size)) self.post_init() def _rope_tables(self, T, device, dtype): c = self.config inv = 1.0 / (c.rope_theta ** (torch.arange(0, c.head_dim, 2, device=device).float() / c.head_dim)) fr = torch.outer(torch.arange(T, device=device).float(), inv)[None, :, None, :] return fr.cos().to(torch.bfloat16).to(dtype), fr.sin().to(torch.bfloat16).to(dtype) def _init_weights(self, module): pass def get_input_embeddings(self): return self.embed def set_input_embeddings(self, value): self.embed = value def get_output_embeddings(self): return None def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): c = self.config T = input_ids.size(1) if T > c.max_position_embeddings: input_ids = input_ids[:, -c.max_position_embeddings:] T = input_ids.size(1) cos, sin = self._rope_tables(T, input_ids.device, self.embed.weight.dtype) x = self.embed(input_ids) for b in self.prelude: x = b(x, cos, sin) for i in range(c.core_loops): if self.loop_emb is not None: x = x + self.loop_emb[i] for b in self.core: x = b(x, cos, sin) for b in self.coda: x = b(x, cos, sin) h = F.rms_norm(x, (x.size(-1),), self.norm_out, c.rms_norm_eps) logits = F.linear(h, self.embed.weight) if c.logit_cap: logits = c.logit_cap * torch.tanh(logits / c.logit_cap) loss = None if labels is not None: loss = F.cross_entropy(logits[:, :-1].float().reshape(-1, logits.size(-1)), labels[:, 1:].reshape(-1), ignore_index=-100) return CausalLMOutput(loss=loss, logits=logits) def prepare_inputs_for_generation(self, input_ids, attention_mask=None, **kwargs): return {"input_ids": input_ids}