TaoNet-mini-A2 / taonet_model.py
Lobakkang's picture
Upload folder using huggingface_hub
14531a0 verified
Raw
History Blame Contribute Delete
6.41 kB
"""Standalone TaoNet model used by the HF wrapper."""
from types import SimpleNamespace
import torch
import torch.nn as nn
import torch.nn.functional as F
try:
from .embeddings import FactorizedEmbedding
from .mla_components import AttentionBlock
except ImportError:
from embeddings import FactorizedEmbedding
from mla_components import AttentionBlock
class SimpleLLM(nn.Module):
"""Pure attention TaoNet language model."""
def __init__(self, config):
super().__init__()
self.config = config
self.vocab_size = config.vocab_size
self.d_model = config.hidden_dim
self.n_layers = config.num_layers
self.n_heads = config.num_heads
self.dropout = config.dropout
self.d_latent_kv = config.d_latent_kv if config.d_latent_kv is not None else int(self.d_model * 0.75)
self.d_rope = config.d_rope if config.d_rope is not None else (self.d_model // self.n_heads)
self.d_ff = config.hidden_dim_ff if config.hidden_dim_ff is not None else (self.d_model * 4)
self.gqa_groups = getattr(config, "gqa_groups", 1)
self.use_factorized_embedding = getattr(config, "use_factorized_embedding", False)
self.d_embed_rank = getattr(config, "d_embed_rank", 96)
self.rope_scale = getattr(config, "rope_scale", 40.0)
self.yarn_enabled = getattr(config, "yarn_enabled", False)
self.yarn_original_max_seq_length = getattr(config, "yarn_original_max_seq_length", None)
self.yarn_alpha = getattr(config, "yarn_alpha", 1.0)
self.max_seq_length = config.max_seq_length
if self.d_model % self.n_heads != 0:
raise ValueError("hidden_dim must be divisible by num_heads")
if self.d_latent_kv % self.n_heads != 0:
raise ValueError("d_latent_kv must be divisible by num_heads")
if self.use_factorized_embedding:
self.token_embedding = FactorizedEmbedding(
self.vocab_size,
self.d_model,
self.d_embed_rank,
)
else:
self.token_embedding = nn.Embedding(self.vocab_size, self.d_model)
self.embedding_dropout = nn.Dropout(self.dropout)
self.blocks = nn.ModuleList(
[
AttentionBlock(
d_model=self.d_model,
d_latent_kv=self.d_latent_kv,
n_heads=self.n_heads,
d_rope=self.d_rope,
d_ff=int(self.d_ff),
dropout=self.dropout,
gqa_groups=self.gqa_groups,
rope_scale=self.rope_scale,
max_seq_length=self.max_seq_length,
yarn_enabled=self.yarn_enabled,
yarn_original_max_seq_length=self.yarn_original_max_seq_length,
yarn_alpha=self.yarn_alpha,
)
for _ in range(self.n_layers)
]
)
self.final_norm = nn.LayerNorm(self.d_model)
self.output_head = nn.Linear(self.d_model, self.vocab_size, bias=False)
self.register_buffer("causal_mask_cache", None, persistent=False)
self.apply(self._init_weights)
def _init_weights(self, module):
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=0.02)
def _get_causal_mask(self, seq_len, device):
if self.causal_mask_cache is None or self.causal_mask_cache.size(-1) < seq_len:
mask = torch.tril(torch.ones(seq_len, seq_len, device=device, dtype=torch.bool))
self.register_buffer("causal_mask_cache", mask, persistent=False)
return self.causal_mask_cache[:seq_len, :seq_len]
def forward(
self,
input_ids=None,
attention_mask=None,
labels=None,
inputs_embeds=None,
pixel_values=None,
):
del pixel_values
if inputs_embeds is None:
if input_ids is None:
raise ValueError("Either input_ids or inputs_embeds must be provided")
_, seq_len = input_ids.shape
device = input_ids.device
x = self.token_embedding(input_ids)
else:
_, seq_len, _ = inputs_embeds.shape
device = inputs_embeds.device
x = inputs_embeds
causal_mask = self._get_causal_mask(seq_len, device)
if attention_mask is not None:
padding_mask = attention_mask.unsqueeze(1).unsqueeze(1).bool()
combined_mask = causal_mask.unsqueeze(0).unsqueeze(0) & padding_mask
combined_mask = combined_mask.float()
else:
combined_mask = causal_mask.unsqueeze(0).unsqueeze(0).float()
x = self.embedding_dropout(x)
for block in self.blocks:
x = block(x, attention_mask=combined_mask)
x = self.final_norm(x)
logits = self.output_head(x)
loss = None
if labels is not None:
logits_flat = logits.view(-1, logits.size(-1))
labels_flat = labels.view(-1)
if not torch.any(labels_flat != -100):
raise ValueError("All labels are masked out (-100), so loss cannot be computed")
loss = F.cross_entropy(logits_flat, labels_flat, reduction="mean", ignore_index=-100)
return {"logits": logits, "loss": loss}
def build_runtime_config(config):
"""Create a light-weight runtime config object for the standalone model."""
return SimpleNamespace(
vocab_size=config.vocab_size,
hidden_dim=config.hidden_dim,
num_layers=config.num_layers,
num_heads=config.num_heads,
dropout=config.dropout,
max_seq_length=config.max_seq_length,
d_latent_kv=config.d_latent_kv,
d_rope=config.d_rope,
gqa_groups=config.gqa_groups,
hidden_dim_ff=config.hidden_dim_ff,
use_factorized_embedding=config.use_factorized_embedding,
d_embed_rank=config.d_embed_rank,
rope_scale=config.rope_scale,
yarn_enabled=config.yarn_enabled,
yarn_original_max_seq_length=config.yarn_original_max_seq_length,
yarn_alpha=config.yarn_alpha,
)