"""CORe architecture: a compact decoder-only transformer. COReForCausalLM is a from-scratch causal LM with weight-tied embeddings, pre-norm transformer blocks, GELU MLPs, and either learned absolute positions or RoPE. It subclasses PreTrainedModel, so it works with the standard transformers API (generate, save_pretrained, from_pretrained). """ import math import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.generation import GenerationMixin from transformers.modeling_outputs import CausalLMOutputWithPast try: from .configuration_core import COReConfig except ImportError: # direct script import (conversion tools) from configuration_core import COReConfig def build_rope_cache(head_dim, max_seq, device, base=10000.0): inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2, device=device).float() / head_dim)) t = torch.arange(max_seq, device=device).float() freqs = torch.outer(t, inv_freq) return torch.cos(freqs), torch.sin(freqs) def apply_rope(x, cos, sin): B, H, T, D = x.shape x1 = x[..., : D // 2] x2 = x[..., D // 2:] c = cos[:T].unsqueeze(0).unsqueeze(0) s = sin[:T].unsqueeze(0).unsqueeze(0) out1 = x1 * c - x2 * s out2 = x1 * s + x2 * c return torch.cat([out1, out2], dim=-1).to(x.dtype) class COReAttention(nn.Module): def __init__(self, config): super().__init__() assert config.n_embd % config.n_head == 0 self.c_attn = nn.Linear(config.n_embd, 3 * config.n_embd) self.c_proj = nn.Linear(config.n_embd, config.n_embd) self.attn_dropout = nn.Dropout(config.dropout) self.resid_dropout = nn.Dropout(config.dropout) self.n_head = config.n_head self.head_dim = config.n_embd // config.n_head self.rope = config.rope self.register_buffer( "causal_mask", torch.tril(torch.ones(config.block_size, config.block_size)) .view(1, 1, config.block_size, config.block_size), persistent=False, ) def forward(self, x, rope_cache=None): B, T, C = x.size() q, k, v = self.c_attn(x).split(C, dim=2) q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2) k = k.view(B, T, self.n_head, self.head_dim).transpose(1, 2) v = v.view(B, T, self.n_head, self.head_dim).transpose(1, 2) if self.rope and rope_cache is not None: cos, sin = rope_cache q = apply_rope(q, cos, sin) k = apply_rope(k, cos, sin) try: y = F.scaled_dot_product_attention( q, k, v, attn_mask=None, dropout_p=self.attn_dropout.p if self.training else 0.0, is_causal=True, ) except Exception: att = (q @ k.transpose(-2, -1)) / math.sqrt(self.head_dim) att = att.masked_fill(self.causal_mask[:, :, :T, :T] == 0, float("-inf")) att = F.softmax(att, dim=-1) att = self.attn_dropout(att) y = att @ v y = y.transpose(1, 2).contiguous().view(B, T, C) return self.resid_dropout(self.c_proj(y)) class COReMLP(nn.Module): def __init__(self, config): super().__init__() self.c_fc = nn.Linear(config.n_embd, 4 * config.n_embd) self.c_proj = nn.Linear(4 * config.n_embd, config.n_embd) self.dropout = nn.Dropout(config.dropout) def forward(self, x): return self.dropout(self.c_proj(F.gelu(self.c_fc(x)))) class COReBlock(nn.Module): def __init__(self, config): super().__init__() self.ln_1 = nn.LayerNorm(config.n_embd) self.attn = COReAttention(config) self.ln_2 = nn.LayerNorm(config.n_embd) self.mlp = COReMLP(config) def forward(self, x, rope_cache=None): x = x + self.attn(self.ln_1(x), rope_cache) x = x + self.mlp(self.ln_2(x)) return x class CORePreTrainedModel(PreTrainedModel): config_class = COReConfig base_model_prefix = "core" supports_gradient_checkpointing = False def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) class COReForCausalLM(CORePreTrainedModel, GenerationMixin): _tied_weights_keys = {"": ["head.weight"]} def __init__(self, config): super().__init__(config) self.tok_emb = nn.Embedding(config.vocab_size, config.n_embd) self.pos_emb = None if config.rope else nn.Embedding(config.block_size, config.n_embd) self.drop = nn.Dropout(config.dropout) self.blocks = nn.ModuleList(COReBlock(config) for _ in range(config.n_layer)) self.ln_f = nn.LayerNorm(config.n_embd) self.head = nn.Linear(config.n_embd, config.vocab_size, bias=False) self.rope_cache = None if config.rope: head_dim = config.n_embd // config.n_head self.rope_cache = build_rope_cache( head_dim, config.block_size, torch.device("cpu"), base=config.rope_base) # Don't tie in __init__: load_state_dict needs each key to have its # own tensor. tie_weights() is called by post_init instead. self.post_init() def tie_weights(self, **kwargs): self.head.weight = self.tok_emb.weight def get_input_embeddings(self): return self.tok_emb def set_input_embeddings(self, value): self.tok_emb = value self.head.weight = value def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs): B, T = input_ids.size() assert T <= self.config.block_size, ( f"sequence length {T} exceeds block size {self.config.block_size}") if self.config.rope: x = self.drop(self.tok_emb(input_ids)) else: pos = torch.arange(0, T, device=input_ids.device) x = self.drop(self.tok_emb(input_ids) + self.pos_emb(pos)) for block in self.blocks: x = block(x, self.rope_cache) x = self.ln_f(x) logits = self.head(x) loss = None if labels is not None: loss = F.cross_entropy( logits.view(-1, logits.size(-1)), labels.view(-1), ignore_index=-1) return CausalLMOutputWithPast(logits=logits, loss=loss) def prepare_inputs_for_generation(self, input_ids, **kwargs): if input_ids.size(1) > self.config.block_size: input_ids = input_ids[:, -self.config.block_size:] return {"input_ids": input_ids}