pulvis-v2 / modeling_pulvis.py
TobiasLogic's picture
pulvis-v2
cbe17d2 verified
Raw History Blame Contribute Delete
5.31 kB
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}