custom-optimizer-model-test / modeling_bananamind2nano.py
Banaxi-Tech's picture
Upload Aspect-Cautious Muon checkpoint at 24,999,591,936 tokens
51115cb verified
Raw
History Blame Contribute Delete
10.9 kB
"""BananaMind 2 Nano implementation for Hugging Face Transformers."""
import math
from typing import Optional
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import PreTrainedModel
from transformers.cache_utils import Cache, DynamicCache
from transformers.generation.utils import GenerationMixin
from transformers.modeling_outputs import CausalLMOutputWithPast
from .configuration_bananamind2nano import BananaMind2NanoConfig
class RMSNorm(nn.Module):
def __init__(self, dim, eps=1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
x_float = x.float()
rms = torch.rsqrt(x_float.pow(2).mean(-1, keepdim=True) + self.eps)
return (x_float * rms * self.weight.float()).type_as(x)
def build_rope_inv_freq(head_dim, theta=100000.0):
return 1.0 / (theta ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim))
def precompute_freqs_cis(head_dim, seq_len, theta=100000.0):
freqs = build_rope_inv_freq(head_dim, theta)
positions = torch.arange(seq_len, dtype=torch.float32)
freqs = torch.outer(positions, freqs)
return torch.polar(torch.ones_like(freqs), freqs)
def apply_rotary_emb(q, k, freqs_cis):
q_complex = torch.view_as_complex(q.float().reshape(*q.shape[:-1], -1, 2))
k_complex = torch.view_as_complex(k.float().reshape(*k.shape[:-1], -1, 2))
freqs_cis = freqs_cis.unsqueeze(0).unsqueeze(0)
q_out = torch.view_as_real(q_complex * freqs_cis).flatten(-2)
k_out = torch.view_as_real(k_complex * freqs_cis).flatten(-2)
return q_out.type_as(q), k_out.type_as(k)
class BananaMind2NanoAttention(nn.Module):
def __init__(self, config, layer_idx):
super().__init__()
self.layer_idx = layer_idx
self.n_head = config.num_attention_heads
self.n_kv_heads = config.num_key_value_heads
self.head_dim = config.head_dim
self.n_rep = self.n_head // self.n_kv_heads
self.q_proj = nn.Linear(config.hidden_size, self.n_head * self.head_dim, bias=False)
self.k_proj = nn.Linear(config.hidden_size, self.n_kv_heads * self.head_dim, bias=False)
self.v_proj = nn.Linear(config.hidden_size, self.n_kv_heads * self.head_dim, bias=False)
self.o_proj = nn.Linear(self.n_head * self.head_dim, config.hidden_size, bias=False)
self.q_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = RMSNorm(self.head_dim, eps=config.rms_norm_eps)
def forward(
self,
x,
freqs_cis,
attention_mask=None,
past_key_values=None,
use_cache=False,
):
batch_size, seq_len, _ = x.shape
q = self.q_proj(x).view(
batch_size,
seq_len,
self.n_head,
self.head_dim,
).transpose(1, 2)
k = self.k_proj(x).view(
batch_size,
seq_len,
self.n_kv_heads,
self.head_dim,
).transpose(1, 2)
v = self.v_proj(x).view(
batch_size,
seq_len,
self.n_kv_heads,
self.head_dim,
).transpose(1, 2)
q = self.q_norm(q)
k = self.k_norm(k)
q, k = apply_rotary_emb(q, k, freqs_cis)
past_length = 0
if use_cache and past_key_values is not None:
past_length = past_key_values.get_seq_length(self.layer_idx)
k, v = past_key_values.update(k, v, self.layer_idx)
kv_len = k.size(-2)
k = k.unsqueeze(2).expand(
batch_size,
self.n_kv_heads,
self.n_rep,
kv_len,
self.head_dim,
).reshape(batch_size, self.n_head, kv_len, self.head_dim)
v = v.unsqueeze(2).expand(
batch_size,
self.n_kv_heads,
self.n_rep,
kv_len,
self.head_dim,
).reshape(batch_size, self.n_head, kv_len, self.head_dim)
attn_mask = None
is_causal = past_length == 0 and attention_mask is None
if not is_causal:
query_positions = past_length + torch.arange(seq_len, device=x.device)
key_positions = torch.arange(kv_len, device=x.device)
causal = key_positions.unsqueeze(0) <= query_positions.unsqueeze(1)
attn_mask = causal[None, None, :, :]
if attention_mask is not None:
key_padding = attention_mask.to(torch.bool)
if key_padding.size(-1) < kv_len:
cached_padding = torch.ones(
key_padding.size(0),
kv_len - key_padding.size(-1),
dtype=torch.bool,
device=key_padding.device,
)
key_padding = torch.cat((cached_padding, key_padding), dim=-1)
else:
key_padding = key_padding[:, -kv_len:]
attn_mask = attn_mask & key_padding[:, None, None, :]
is_causal = False
y = F.scaled_dot_product_attention(
q,
k,
v,
attn_mask=attn_mask,
is_causal=is_causal,
)
y = y.transpose(1, 2).contiguous().view(
batch_size,
seq_len,
self.n_head * self.head_dim,
)
return self.o_proj(y)
class BananaMind2NanoSwiGLUMLP(nn.Module):
def __init__(self, config):
super().__init__()
self.w_gate = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
self.w_up = nn.Linear(config.hidden_size, config.intermediate_size, bias=False)
self.w_down = nn.Linear(config.intermediate_size, config.hidden_size, bias=False)
def forward(self, x):
return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))
class BananaMind2NanoBlock(nn.Module):
def __init__(self, config, layer_idx):
super().__init__()
self.ln_1 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.attn = BananaMind2NanoAttention(config, layer_idx)
self.ln_2 = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.mlp = BananaMind2NanoSwiGLUMLP(config)
def forward(
self,
x,
freqs_cis,
attention_mask=None,
past_key_values=None,
use_cache=False,
):
x = x + self.attn(
self.ln_1(x),
freqs_cis,
attention_mask=attention_mask,
past_key_values=past_key_values,
use_cache=use_cache,
)
return x + self.mlp(self.ln_2(x))
class BananaMind2NanoPreTrainedModel(PreTrainedModel):
config_class = BananaMind2NanoConfig
base_model_prefix = "transformer"
supports_gradient_checkpointing = False
def _init_weights(self, module):
std = 0.02
if hasattr(module, "NANOGPT_SCALE_INIT"):
std *= 2 * self.config.num_hidden_layers ** -0.5
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=std)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
class BananaMind2NanoForCausalLM(BananaMind2NanoPreTrainedModel, GenerationMixin):
_tied_weights_keys = {"lm_head.weight": "transformer.wte.weight"}
def __init__(self, config):
super().__init__(config)
self.config = config
self.transformer = nn.ModuleDict(
{
"wte": nn.Embedding(config.vocab_size, config.hidden_size),
"h": nn.ModuleList(
[
BananaMind2NanoBlock(config, layer_idx)
for layer_idx in range(config.num_hidden_layers)
]
),
"ln_f": RMSNorm(config.hidden_size, eps=config.rms_norm_eps),
}
)
self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
if config.tie_word_embeddings:
self.lm_head.weight = self.transformer["wte"].weight
self._embd_scale = math.sqrt(config.hidden_size)
self._freqs_cis_cache = None
self.post_init()
def get_input_embeddings(self):
return self.transformer["wte"]
def set_input_embeddings(self, value):
self.transformer["wte"] = value
def get_output_embeddings(self):
return self.lm_head
def set_output_embeddings(self, new_embeddings):
self.lm_head = new_embeddings
def _get_freqs_cis(self, seq_len, device):
cache = self._freqs_cis_cache
if cache is None or cache.device != device or cache.size(0) < seq_len:
cache = precompute_freqs_cis(
self.config.head_dim,
seq_len,
self.config.rope_theta,
).to(device)
self._freqs_cis_cache = cache
return cache[:seq_len]
def forward(
self,
input_ids,
attention_mask=None,
labels=None,
past_key_values: Optional[Cache] = None,
use_cache=None,
**kwargs,
):
_, seq_len = input_ids.shape
if use_cache is None:
use_cache = self.config.use_cache and labels is None
if use_cache and past_key_values is None:
past_key_values = DynamicCache(config=self.config)
past_length = past_key_values.get_seq_length() if use_cache else 0
total_length = past_length + seq_len
if total_length > self.config.max_position_embeddings:
raise ValueError(
f"Sequence length {total_length} exceeds the configured maximum "
f"of {self.config.max_position_embeddings}"
)
x = self.transformer["wte"](input_ids) * self._embd_scale
freqs_cis = self._get_freqs_cis(total_length, input_ids.device)[past_length:]
for block in self.transformer["h"]:
x = block(
x,
freqs_cis,
attention_mask=attention_mask,
past_key_values=past_key_values,
use_cache=use_cache,
)
x = self.transformer["ln_f"](x)
logits = self.lm_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.float().reshape(-1, shift_logits.size(-1)),
shift_labels.reshape(-1),
)
if self.config.z_loss_coeff:
loss = loss + self.config.z_loss_coeff * logits.float().pow(2).mean()
return CausalLMOutputWithPast(
loss=loss,
logits=logits,
past_key_values=past_key_values if use_cache else None,
)