CORe-Predetermined-v1 / model /modeling_core.py
COReTechnologies's picture
Upload 21 files
8fbe2d2 verified
Raw History Blame Contribute Delete
6.82 kB
"""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}