PicoLM-80M-Instruct / modeling_picolm.py
aethertp's picture
Upload folder using huggingface_hub
8ca632a verified
Raw History Blame Contribute Delete
5.22 kB
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)