Flash-V1 / flash-v1 /modeling_core.py
COReTechnologies's picture
Upload 11 files
d52dd01 verified
Raw History Blame Contribute Delete
6.01 kB
"""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