Download source/src/vimeml/training/model_v2.py from Voltline/vimeml-tiny-ja-v2.1: direct link, hf CLI and curl.
- Browser
- Download file 5.55 kB
-
https://huggingface.co/Voltline/vimeml-tiny-ja-v2.1/resolve/main/source/src/vimeml/training/model_v2.py
- Command line
-
hf download hf://Voltline/vimeml-tiny-ja-v2.1/source/src/vimeml/training/model_v2.py
-
curl -L -o model_v2.py https://huggingface.co/Voltline/vimeml-tiny-ja-v2.1/resolve/main/source/src/vimeml/training/model_v2.py
5.55 kB
| """V2 decoder: primitive RMSNorm, bias-free attention and SwiGLU.""" | |
| import math | |
| from dataclasses import asdict, dataclass | |
| import torch | |
| from torch import nn | |
| from torch.nn import functional as F | |
| from vimeml.training.model import GPTConfig | |
| class GPTV2Config(GPTConfig): | |
| d_model: int = 320 | |
| n_heads: int = 5 | |
| n_layers: int = 6 | |
| d_ff: int = 832 | |
| norm: str = "rmsnorm" | |
| norm_eps: float = 1e-5 | |
| activation: str = "swiglu" | |
| bias: bool = False | |
| tie_embeddings: bool = True | |
| position_embedding: str = "learned" | |
| def __post_init__(self): | |
| super().__post_init__() | |
| if ( | |
| self.norm != "rmsnorm" | |
| or self.activation != "swiglu" | |
| or self.bias is not False | |
| or self.tie_embeddings is not True | |
| or self.position_embedding != "learned" | |
| ): | |
| raise ValueError( | |
| "V2 requires RMSNorm, SwiGLU, bias-free Linear, tied weights and learned positions." | |
| ) | |
| if not math.isfinite(self.norm_eps) or self.norm_eps <= 0: | |
| raise ValueError("norm_eps must be finite and positive.") | |
| class RMSNorm(nn.Module): | |
| def __init__(self, width, eps=1e-5): | |
| super().__init__() | |
| self.weight = nn.Parameter(torch.ones(width)) | |
| self.eps = eps | |
| def forward(self, hidden): | |
| # Accumulate the mean square in FP32 during BF16 training. | |
| normalized = hidden.float() | |
| variance = normalized.pow(2).mean(-1, keepdim=True) | |
| normalized = normalized * torch.rsqrt(variance + self.eps) | |
| return normalized.to(hidden.dtype) * self.weight | |
| class CausalAttentionV2(nn.Module): | |
| def __init__(self, config): | |
| super().__init__() | |
| self.heads = config.n_heads | |
| self.dropout = config.dropout | |
| self.qkv = nn.Linear(config.d_model, 3 * config.d_model, bias=False) | |
| self.projection = nn.Linear(config.d_model, config.d_model, bias=False) | |
| def forward(self, hidden): | |
| batch, length, width = hidden.shape | |
| qkv = self.qkv(hidden).view(batch, length, 3, self.heads, width // self.heads) | |
| query, key, value = qkv.permute(2, 0, 3, 1, 4).unbind(0) | |
| attended = F.scaled_dot_product_attention( | |
| query, key, value, is_causal=True, dropout_p=self.dropout if self.training else 0.0 | |
| ) | |
| return self.projection(attended.transpose(1, 2).reshape(batch, length, width)) | |
| class SwiGLU(nn.Module): | |
| def __init__(self, config): | |
| super().__init__() | |
| self.gate = nn.Linear(config.d_model, config.d_ff, bias=False) | |
| self.up = nn.Linear(config.d_model, config.d_ff, bias=False) | |
| self.down = nn.Linear(config.d_ff, config.d_model, bias=False) | |
| def forward(self, hidden): | |
| return self.down(F.silu(self.gate(hidden)) * self.up(hidden)) | |
| class DecoderBlockV2(nn.Module): | |
| def __init__(self, config): | |
| super().__init__() | |
| self.attention_norm = RMSNorm(config.d_model, config.norm_eps) | |
| self.attention = CausalAttentionV2(config) | |
| self.ffn_norm = RMSNorm(config.d_model, config.norm_eps) | |
| self.ffn = SwiGLU(config) | |
| self.dropout = nn.Dropout(config.dropout) | |
| def forward(self, hidden): | |
| hidden = hidden + self.dropout(self.attention(self.attention_norm(hidden))) | |
| return hidden + self.dropout(self.ffn(self.ffn_norm(hidden))) | |
| class TinyGPTV2(nn.Module): | |
| def __init__(self, config=GPTV2Config()): | |
| super().__init__() | |
| self.config = config | |
| self.token_embedding = nn.Embedding(config.vocab_size, config.d_model) | |
| self.position_embedding = nn.Embedding(config.context_length, config.d_model) | |
| self.dropout = nn.Dropout(config.dropout) | |
| self.blocks = nn.ModuleList(DecoderBlockV2(config) for _ in range(config.n_layers)) | |
| self.final_norm = RMSNorm(config.d_model, config.norm_eps) | |
| self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False) | |
| self.apply(self._initialize) | |
| self.lm_head.weight = self.token_embedding.weight | |
| for block in self.blocks: | |
| for projection in (block.attention.projection, block.ffn.down): | |
| nn.init.normal_(projection.weight, std=0.02 / math.sqrt(2 * config.n_layers)) | |
| def _initialize(module): | |
| if isinstance(module, (nn.Linear, nn.Embedding)): | |
| nn.init.normal_(module.weight, std=0.02) | |
| def forward(self, input_ids, labels=None): | |
| if input_ids.ndim != 2 or not 1 <= input_ids.shape[1] <= self.config.context_length: | |
| raise ValueError("Expected [batch, time] input within model context.") | |
| positions = torch.arange(input_ids.shape[1], device=input_ids.device) | |
| hidden = self.dropout(self.token_embedding(input_ids) + self.position_embedding(positions)) | |
| for block in self.blocks: | |
| hidden = block(hidden) | |
| hidden = self.final_norm(hidden) | |
| if labels is None: | |
| return self.lm_head(hidden) | |
| if labels.shape != input_ids.shape: | |
| raise ValueError("Labels must match inputs; labels are already shifted by DataLoader.") | |
| valid = labels != -100 | |
| logits = self.lm_head(hidden[valid]) | |
| return { | |
| "loss_sum": F.cross_entropy(logits, labels[valid], reduction="sum"), | |
| "token_count": valid.sum(), | |
| } | |
| def parameter_count(self): | |
| return sum(parameter.numel() for parameter in self.parameters()) | |
| def configuration(self): | |
| return asdict(self.config) | |