"""CORe model architecture for HuggingFace transformers.""" import math from typing import Optional, Tuple import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel, PretrainedConfig, GenerationMixin from transformers.modeling_outputs import CausalLMOutput class COReConfig(PretrainedConfig): model_type = "core" def __init__( self, n_layer=12, n_head=16, n_embd=1024, block_size=512, vocab_size=16384, rope=False, dropout=0.0, **kwargs, ): super().__init__(**kwargs) self.n_layer = n_layer self.n_head = n_head self.n_embd = n_embd self.block_size = block_size self.vocab_size = vocab_size self.rope = rope self.dropout = dropout class CausalSelfAttention(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.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): 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) y = F.scaled_dot_product_attention( q, k, v, dropout_p=self.attn_dropout.p if self.training else 0.0, is_causal=True, ) y = y.transpose(1, 2).contiguous().view(B, T, C) return self.resid_dropout(self.c_proj(y)) class MLP(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 Block(nn.Module): def __init__(self, config): super().__init__() self.ln_1 = nn.LayerNorm(config.n_embd) self.attn = CausalSelfAttention(config) self.ln_2 = nn.LayerNorm(config.n_embd) self.mlp = MLP(config) def forward(self, x): x = x + self.attn(self.ln_1(x)) x = x + self.mlp(self.ln_2(x)) return x class COReModel(PreTrainedModel): config_class = COReConfig base_model_prefix = "core" def __init__(self, config): super().__init__(config) self.config = 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(Block(config) for _ in range(config.n_layer)) self.ln_f = nn.LayerNorm(config.n_embd) self.post_init() def forward(self, input_ids, attention_mask=None, **kwargs): B, T = input_ids.size() if self.pos_emb is not None: pos = torch.arange(0, T, device=input_ids.device) x = self.drop(self.tok_emb(input_ids) + self.pos_emb(pos)) else: x = self.drop(self.tok_emb(input_ids)) for block in self.blocks: x = block(x) return self.ln_f(x) class COReForCausalLM(PreTrainedModel, GenerationMixin): config_class = COReConfig base_model_prefix = "core" main_input_name = "input_ids" _supports_cache_class = False _supports_static_cache = False def __init__(self, config): super().__init__(config) self.config = 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(Block(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.tok_emb.weight = self.head.weight self.post_init() def forward(self, input_ids, attention_mask=None, labels=None, **kwargs): B, T = input_ids.size() if self.pos_emb is not None: pos = torch.arange(0, T, device=input_ids.device) x = self.drop(self.tok_emb(input_ids) + self.pos_emb(pos)) else: x = self.drop(self.tok_emb(input_ids)) for block in self.blocks: x = block(x) x = self.ln_f(x) logits = self.head(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, shift_logits.size(-1)), shift_labels.view(-1), ignore_index=-100, ) return CausalLMOutput(loss=loss, logits=logits) 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} def _reorder_cache(self, past, beam_idx): return past