| """Capability-first Vortex 175M decoder. |
| |
| This file intentionally uses only standard PyTorch CUDA primitives. The |
| model is deep-and-thin, uses GQA and SwiGLU, and ties the input/output |
| embedding. Those choices are much easier to train and export than a custom |
| SSM kernel while retaining the main sub-billion-parameter wins reported by |
| MobileLLM-style studies. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import math |
| from dataclasses import asdict, dataclass |
|
|
| import torch |
| import torch.nn.functional as F |
| from torch import nn |
| from torch.utils.checkpoint import checkpoint |
|
|
|
|
| @dataclass |
| class VortexConfig: |
| vocab_size: int = 8_192 |
| max_seq_len: int = 4_096 |
| n_layer: int = 12 |
| n_embd: int = 1024 |
| n_head: int = 16 |
| n_kv_head: int = 4 |
| head_dim: int = 64 |
| intermediate_size: int = 3_664 |
| rope_theta: float = 100_000.0 |
| norm_eps: float = 1e-5 |
| logits_chunk_tokens: int = 16_384 |
| gradient_checkpointing: bool = False |
| use_transformer_engine: bool = False |
| attn_input_format: str = "bshd" |
|
|
| def to_dict(self) -> dict: |
| return asdict(self) |
|
|
|
|
| class RMSNorm(nn.Module): |
| def __init__(self, dim: int, eps: float) -> None: |
| super().__init__() |
| self.weight = nn.Parameter(torch.ones(dim)) |
| self.eps = eps |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return F.rms_norm(x, (x.shape[-1],), self.weight, self.eps) |
|
|
|
|
| def _linear(config: VortexConfig, in_features: int, out_features: int) -> nn.Module: |
| """Create a bias-free projection, optionally backed by Transformer Engine.""" |
| if not config.use_transformer_engine: |
| return nn.Linear(in_features, out_features, bias=False) |
| try: |
| import transformer_engine.pytorch as te |
| except ImportError as exc: |
| raise RuntimeError( |
| "use_transformer_engine=True requires transformer-engine[pytorch]" |
| ) from exc |
| return te.Linear( |
| in_features, |
| out_features, |
| bias=False, |
| params_dtype=torch.bfloat16, |
| device="cuda", |
| ) |
|
|
|
|
| class RotaryEmbedding(nn.Module): |
| def __init__(self, dim: int, max_seq_len: int, theta: float) -> None: |
| super().__init__() |
| inv_freq = 1.0 / (theta ** (torch.arange(0, dim, 2).float() / dim)) |
| positions = torch.arange(max_seq_len, dtype=torch.float32) |
| frequencies = torch.outer(positions, inv_freq) |
| |
| |
| |
| angles = torch.cat((frequencies, frequencies), dim=-1) |
| self.register_buffer("cos_cached", angles.cos()[None, None], persistent=False) |
| self.register_buffer("sin_cached", angles.sin()[None, None], persistent=False) |
|
|
| @staticmethod |
| def rotate_half(x: torch.Tensor) -> torch.Tensor: |
| half = x.shape[-1] // 2 |
| return torch.cat((-x[..., half:], x[..., :half]), dim=-1) |
|
|
| def forward(self, q: torch.Tensor, k: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: |
| seq_len = q.shape[-2] |
| if seq_len > self.cos_cached.shape[-2]: |
| raise ValueError(f"sequence length {seq_len} exceeds configured maximum") |
| cos = self.cos_cached[:, :, :seq_len].to(dtype=q.dtype) |
| sin = self.sin_cached[:, :, :seq_len].to(dtype=q.dtype) |
| return ( |
| q * cos + self.rotate_half(q) * sin, |
| k * cos + self.rotate_half(k) * sin, |
| ) |
|
|
|
|
| class GQAAttention(nn.Module): |
| def __init__(self, config: VortexConfig) -> None: |
| super().__init__() |
| if config.n_head % config.n_kv_head: |
| raise ValueError("n_head must be divisible by n_kv_head") |
| if config.n_head * config.head_dim != config.n_embd: |
| raise ValueError("n_head * head_dim must equal n_embd") |
| self.n_head = config.n_head |
| self.n_kv_head = config.n_kv_head |
| self.head_dim = config.head_dim |
| kv_dim = config.n_kv_head * config.head_dim |
| self.q_proj = _linear(config, config.n_embd, config.n_embd) |
| self.k_proj = _linear(config, config.n_embd, kv_dim) |
| self.v_proj = _linear(config, config.n_embd, kv_dim) |
| self.o_proj = _linear(config, config.n_embd, config.n_embd) |
| |
| |
| |
| self.q_norm = RMSNorm(config.head_dim, config.norm_eps) |
| self.k_norm = RMSNorm(config.head_dim, config.norm_eps) |
| self.rope = RotaryEmbedding(config.head_dim, config.max_seq_len, config.rope_theta) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| batch, seq_len, _ = x.shape |
| q = self.q_proj(x).view(batch, seq_len, self.n_head, self.head_dim).transpose(1, 2) |
| k = self.k_proj(x).view(batch, seq_len, self.n_kv_head, self.head_dim).transpose(1, 2) |
| v = self.v_proj(x).view(batch, seq_len, self.n_kv_head, self.head_dim).transpose(1, 2) |
| q, k = self.rope(self.q_norm(q), self.k_norm(k)) |
| |
| |
| |
| y = F.scaled_dot_product_attention( |
| q, k, v, is_causal=True, enable_gqa=True |
| ) |
| y = y.transpose(1, 2).contiguous().view(batch, seq_len, -1) |
| return self.o_proj(y) |
|
|
|
|
| class SwiGLU(nn.Module): |
| def __init__(self, config: VortexConfig) -> None: |
| super().__init__() |
| self.gate_proj = _linear(config, config.n_embd, config.intermediate_size) |
| self.up_proj = _linear(config, config.n_embd, config.intermediate_size) |
| self.down_proj = _linear(config, config.intermediate_size, config.n_embd) |
|
|
| def forward(self, x: torch.Tensor) -> torch.Tensor: |
| return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) |
|
|
|
|
| class VortexBlock(nn.Module): |
| def __init__(self, config: VortexConfig, layer_number: int | None = None) -> None: |
| super().__init__() |
| self.use_transformer_engine = config.use_transformer_engine |
| self.attn_input_format = config.attn_input_format |
| if self.use_transformer_engine: |
| try: |
| import transformer_engine.pytorch as te |
| except ImportError as exc: |
| raise RuntimeError( |
| "use_transformer_engine=True requires transformer-engine[pytorch]" |
| ) from exc |
| |
| |
| |
| |
| self.te_layer = te.TransformerLayer( |
| hidden_size=config.n_embd, |
| ffn_hidden_size=config.intermediate_size, |
| num_attention_heads=config.n_head, |
| num_gqa_groups=config.n_kv_head, |
| layernorm_epsilon=config.norm_eps, |
| hidden_dropout=0.0, |
| attention_dropout=0.0, |
| kv_channels=config.head_dim, |
| layer_number=layer_number, |
| bias=False, |
| activation="swiglu", |
| normalization="RMSNorm", |
| qk_norm_type="RMSNorm", |
| qk_norm_before_rope=True, |
| fuse_qkv_params=True, |
| self_attn_mask_type="causal", |
| attn_input_format=config.attn_input_format, |
| params_dtype=torch.bfloat16, |
| device="cuda", |
| ) |
| return |
| self.norm1 = RMSNorm(config.n_embd, config.norm_eps) |
| self.attn = GQAAttention(config) |
| self.norm2 = RMSNorm(config.n_embd, config.norm_eps) |
| self.ffn = SwiGLU(config) |
|
|
| def forward( |
| self, |
| x: torch.Tensor, |
| rotary_pos_emb: torch.Tensor | None = None, |
| is_first_microbatch: bool | None = None, |
| inference_params=None, |
| attention_mask: torch.Tensor | None = None, |
| inference_decode_bshd: bool = False, |
| ) -> torch.Tensor: |
| if self.use_transformer_engine: |
| if inference_params is not None: |
| |
| |
| |
| |
| |
| effective_format = "bshd" if inference_decode_bshd else self.attn_input_format |
| self.te_layer.self_attention.qkv_format = effective_format |
| te_kwargs = { |
| "attention_mask": attention_mask, |
| "self_attn_mask_type": ( |
| "padding_causal" if inference_params is not None else None |
| ), |
| "rotary_pos_emb": rotary_pos_emb, |
| "is_first_microbatch": is_first_microbatch, |
| "inference_params": inference_params, |
| } |
| if inference_params is not None and not inference_decode_bshd and self.attn_input_format == "thd": |
| batch_size = len(inference_params.sequences) |
| cu_seqlens = inference_params.cu_seqlens_q[: batch_size + 1] |
| sequence_lengths = cu_seqlens[1:] - cu_seqlens[:-1] |
| te_kwargs.update( |
| { |
| "cu_seqlens_q": cu_seqlens, |
| "cu_seqlens_q_padded": cu_seqlens, |
| "max_seqlen_q": int(sequence_lengths.max().item()), |
| "max_seqlen_kv": int(sequence_lengths.max().item()), |
| } |
| ) |
| return self.te_layer( |
| x, |
| **te_kwargs, |
| ) |
| x = x + self.attn(self.norm1(x)) |
| x = x + self.ffn(self.norm2(x)) |
| return x |
|
|
|
|
| class VortexForCausalLM(nn.Module): |
| def __init__(self, config: VortexConfig | None = None) -> None: |
| super().__init__() |
| self.config = config or VortexConfig() |
| self.embed_tokens = nn.Embedding(self.config.vocab_size, self.config.n_embd) |
| self.layers = nn.ModuleList( |
| VortexBlock(self.config, layer_number=index + 1) |
| for index in range(self.config.n_layer) |
| ) |
| self.norm = RMSNorm(self.config.n_embd, self.config.norm_eps) |
| if self.config.use_transformer_engine: |
| import transformer_engine.pytorch as te |
|
|
| self.rotary = te.RotaryPositionEmbedding( |
| self.config.head_dim, |
| rotary_base=self.config.rope_theta, |
| interleaved=False, |
| ) |
| self._initialize_weights() |
|
|
| def _initialize_weights(self) -> None: |
| |
| |
| output_std = 0.02 / math.sqrt(2.0 * self.config.n_layer) |
| nn.init.normal_(self.embed_tokens.weight, mean=0.0, std=0.02) |
| for block in self.layers: |
| if self.config.use_transformer_engine: |
| for name, parameter in block.named_parameters(): |
| if parameter.ndim == 1: |
| nn.init.ones_(parameter) |
| else: |
| is_output = ( |
| name.endswith("self_attention.proj.weight") |
| or name.endswith("layernorm_mlp.fc2_weight") |
| ) |
| nn.init.normal_( |
| parameter, |
| mean=0.0, |
| std=output_std if is_output else 0.02, |
| ) |
| else: |
| for child in block.modules(): |
| if isinstance(child, nn.Linear): |
| is_output = child is block.attn.o_proj or child is block.ffn.down_proj |
| nn.init.normal_( |
| child.weight, |
| mean=0.0, |
| std=output_std if is_output else 0.02, |
| ) |
|
|
| def parameter_count(self) -> int: |
| return sum(parameter.numel() for parameter in self.parameters()) |
|
|
| def parameter_breakdown(self) -> dict[str, int]: |
| c = self.config |
| embedding = c.vocab_size * c.n_embd |
| q = c.n_layer * c.n_embd * c.n_embd |
| k = c.n_layer * c.n_embd * (c.n_kv_head * c.head_dim) |
| v = k |
| o = q |
| |
| |
| |
| qk_norm = c.n_layer * 2 * c.head_dim |
| ffn = c.n_layer * 3 * c.n_embd * c.intermediate_size |
| block_norm = c.n_layer * 2 * c.n_embd |
| final_norm = c.n_embd |
| return { |
| "input_embedding_and_tied_output": embedding, |
| "attention_q_projection": q, |
| "attention_k_projection": k, |
| "attention_v_projection": v, |
| "attention_o_projection": o, |
| "attention_qk_norm": qk_norm, |
| "ffn_swiglu": ffn, |
| "block_rmsnorm": block_norm, |
| "final_rmsnorm": final_norm, |
| "total": self.parameter_count(), |
| } |
|
|
| def _chunked_tied_loss(self, hidden: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: |
| chunk = self.config.logits_chunk_tokens |
| total = hidden.new_zeros((), dtype=torch.float32) |
| for start in range(0, hidden.shape[0], chunk): |
| end = min(hidden.shape[0], start + chunk) |
| logits = F.linear(hidden[start:end], self.embed_tokens.weight) |
| total = total + F.cross_entropy(logits, targets[start:end]).float() * (end - start) |
| return total / hidden.shape[0] |
|
|
| def forward( |
| self, |
| input_ids: torch.Tensor, |
| labels: torch.Tensor | None = None, |
| is_first_microbatch: bool | None = None, |
| inference_params=None, |
| inference_attention_mask=None, |
| inference_decode_bshd: bool = False, |
| ) -> tuple[torch.Tensor | None, torch.Tensor | None]: |
| x = self.embed_tokens(input_ids) |
| rotary_pos_emb = None |
| attention_mask = None |
| if self.config.use_transformer_engine: |
| if inference_params is not None: |
| |
| |
| |
| |
| rotary_pos_emb = self.rotary(self.config.max_seq_len) |
| if self.config.attn_input_format == "thd" and not inference_decode_bshd: |
| attention_mask = None |
| elif inference_attention_mask is None: |
| query_padding_mask = torch.zeros( |
| (x.shape[0], 1, 1, x.shape[1]), |
| dtype=torch.bool, |
| device=x.device, |
| ) |
| key_padding_mask = torch.ones( |
| (x.shape[0], 1, 1, inference_params.max_sequence_length), |
| dtype=torch.bool, |
| device=x.device, |
| ) |
| for batch_index, sequence_length in enumerate( |
| inference_params.sequences.values() |
| ): |
| key_padding_mask[batch_index, :, :, :sequence_length] = False |
| |
| |
| |
| attention_mask = (query_padding_mask, key_padding_mask) |
| else: |
| attention_mask = inference_attention_mask |
| else: |
| rotary_pos_emb = self.rotary(x.shape[1]) |
| for block in self.layers: |
| if self.training and self.config.gradient_checkpointing: |
| x = checkpoint( |
| block, |
| x, |
| rotary_pos_emb, |
| is_first_microbatch, |
| use_reentrant=False, |
| ) |
| else: |
| x = block( |
| x, |
| rotary_pos_emb, |
| is_first_microbatch, |
| inference_params, |
| attention_mask, |
| inference_decode_bshd, |
| ) |
| x = self.norm(x) |
| if labels is None: |
| return F.linear(x, self.embed_tokens.weight), None |
| hidden = x[:, :-1].reshape(-1, x.shape[-1]) |
| targets = labels[:, 1:].reshape(-1) |
| return None, self._chunked_tied_loss(hidden, targets) |
|
|
|
|
| if __name__ == "__main__": |
| config = VortexConfig() |
| model = VortexForCausalLM(config) |
| print(config.to_dict()) |
| print(model.parameter_breakdown()) |
|
|