Download splitbit_llm/model/layers.py from hermescures1/splitbit-llm: direct link, hf CLI and curl.
- Browser
- Download file 10.5 kB
-
https://huggingface.co/hermescures1/splitbit-llm/resolve/main/splitbit_llm/model/layers.py
- Command line
-
hf download hf://hermescures1/splitbit-llm/splitbit_llm/model/layers.py
-
curl -L -o layers.py https://huggingface.co/hermescures1/splitbit-llm/resolve/main/splitbit_llm/model/layers.py
10.5 kB
| """Neural network layers for SplitBit LLM — pure NumPy implementation. | |
| Layers: | |
| - Embedding (token ID → dense vector) | |
| - Multi-head self-attention with RoPE | |
| - Feed-forward network (MLP with GELU) | |
| - Layer normalization (pre-norm) | |
| - KV cache for fast autoregressive generation | |
| All weights stored as NumPy arrays, quantized via SplitBitQuantizer. | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import math | |
| from typing import Any | |
| import numpy as np | |
| from .quantization import SplitBitQuantizer | |
| logger = logging.getLogger(__name__) | |
| def gelu(x: np.ndarray) -> np.ndarray: | |
| """GELU activation — Gaussian Error Linear Unit.""" | |
| return 0.5 * x * (1.0 + np.tanh(math.sqrt(2.0 / math.pi) * (x + 0.044715 * x ** 3))) | |
| def softmax(x: np.ndarray, axis: int = -1) -> np.ndarray: | |
| """Numerically stable softmax.""" | |
| x_max = np.max(x, axis=axis, keepdims=True) | |
| exp_x = np.exp(x - x_max) | |
| return exp_x / np.sum(exp_x, axis=axis, keepdims=True) | |
| def layer_norm(x: np.ndarray, gamma: np.ndarray, beta: np.ndarray, eps: float = 1e-5) -> np.ndarray: | |
| """Layer normalization.""" | |
| mean = np.mean(x, axis=-1, keepdims=True) | |
| var = np.var(x, axis=-1, keepdims=True) | |
| return gamma * (x - mean) / np.sqrt(var + eps) + beta | |
| def rope(pos: np.ndarray, d_head: int, base: float = 10000.0) -> tuple[np.ndarray, np.ndarray]: | |
| """Rotary Position Embedding (RoPE). | |
| Returns cos and sin tensors for rotating Q and K. | |
| """ | |
| inv_freq = 1.0 / (base ** (np.arange(0, d_head, 2) / d_head)) | |
| # pos: [seq_len], inv_freq: [d_head/2] | |
| freqs = np.outer(pos, inv_freq) # [seq_len, d_head/2] | |
| cos = np.cos(freqs) | |
| sin = np.sin(freqs) | |
| # Repeat to match d_head | |
| cos = np.repeat(cos, 2, axis=-1) # [seq_len, d_head] | |
| sin = np.repeat(sin, 2, axis=-1) | |
| return cos, sin | |
| def apply_rope(x: np.ndarray, cos: np.ndarray, sin: np.ndarray) -> np.ndarray: | |
| """Apply rotary embedding to tensor x. | |
| x: [batch, n_heads, seq_len, d_head] | |
| cos/sin: [seq_len, d_head] | |
| """ | |
| x1 = x[..., 0::2] # even indices | |
| x2 = x[..., 1::2] # odd indices | |
| # Rotate | |
| cos = cos[None, None, :, :] # [1, 1, seq_len, d_head] | |
| sin = sin[None, None, :, :] | |
| rotated = np.empty_like(x) | |
| rotated[..., 0::2] = x1 * cos[..., 0::2] - x2 * sin[..., 0::2] | |
| rotated[..., 1::2] = x1 * sin[..., 1::2] + x2 * cos[..., 1::2] | |
| return rotated | |
| class Embedding: | |
| """Token embedding layer.""" | |
| def __init__(self, vocab_size: int, d_model: int) -> None: | |
| # Xavier/Glorot initialization | |
| std = math.sqrt(2.0 / (vocab_size + d_model)) | |
| self.weight = np.random.randn(vocab_size, d_model).astype(np.float32) * std | |
| self.d_model = d_model | |
| self.vocab_size = vocab_size | |
| def forward(self, token_ids: np.ndarray) -> np.ndarray: | |
| """token_ids: [batch, seq_len] → [batch, seq_len, d_model]""" | |
| return self.weight[token_ids] | |
| def backward(self, grad: np.ndarray, token_ids: np.ndarray) -> np.ndarray: | |
| """Gradient w.r.t. embedding weights.""" | |
| grad_weight = np.zeros_like(self.weight) | |
| np.add.at(grad_weight, token_ids, grad) | |
| return grad_weight | |
| class Linear: | |
| """Linear layer: y = x @ W^T + b, with SplitBit quantization support.""" | |
| def __init__(self, in_features: int, out_features: int, bias: bool = True) -> None: | |
| std = math.sqrt(2.0 / (in_features + out_features)) | |
| self.weight = np.random.randn(out_features, in_features).astype(np.float32) * std | |
| self.bias = np.zeros(out_features, dtype=np.float32) if bias else None | |
| self.in_features = in_features | |
| self.out_features = out_features | |
| self.use_bias = bias | |
| self._quantized = None | |
| def quantize(self, quantizer: SplitBitQuantizer) -> None: | |
| """Quantize weights for storage/inference.""" | |
| self._quantized = { | |
| "weight": quantizer.quantize(self.weight), | |
| "bias": self.bias.copy() if self.bias is not None else None, | |
| } | |
| def dequantize(self) -> None: | |
| """Restore full-precision weights.""" | |
| self._quantized = None | |
| def forward(self, x: np.ndarray) -> np.ndarray: | |
| """x: [..., in_features] → [..., out_features]""" | |
| w = self.weight | |
| out = x @ w.T | |
| if self.bias is not None: | |
| out = out + self.bias | |
| return out | |
| def forward_with_cache(self, x: np.ndarray, kv_cache: dict | None = None, layer_idx: int = 0, | |
| is_kv: bool = False) -> np.ndarray: | |
| """Forward pass that optionally uses/appends to KV cache.""" | |
| return self.forward(x) | |
| class MultiHeadAttention: | |
| """Multi-head self-attention with RoPE and KV cache.""" | |
| def __init__(self, d_model: int, n_heads: int, max_seq_len: int = 512) -> None: | |
| self.d_model = d_model | |
| self.n_heads = n_heads | |
| self.d_head = d_model // n_heads | |
| self.max_seq_len = max_seq_len | |
| self.wq = Linear(d_model, d_model, bias=False) | |
| self.wk = Linear(d_model, d_model, bias=False) | |
| self.wv = Linear(d_model, d_model, bias=False) | |
| self.wo = Linear(d_model, d_model, bias=False) | |
| # Precompute RoPE | |
| pos = np.arange(max_seq_len, dtype=np.float32) | |
| self._cos, self._sin = rope(pos, self.d_head) | |
| # KV cache: {layer_idx: (k, v)} | |
| self._kv_cache: dict[int, tuple[np.ndarray, np.ndarray]] = {} | |
| def forward( | |
| self, | |
| x: np.ndarray, | |
| layer_idx: int = 0, | |
| use_cache: bool = False, | |
| past_len: int = 0, | |
| ) -> np.ndarray: | |
| """ | |
| x: [batch, seq_len, d_model] | |
| Returns: [batch, seq_len, d_model] | |
| """ | |
| batch, seq_len, _ = x.shape | |
| # Project to Q, K, V | |
| q = self.wq.forward(x) # [batch, seq_len, d_model] | |
| k = self.wk.forward(x) | |
| v = self.wv.forward(x) | |
| # Reshape to [batch, n_heads, seq_len, d_head] | |
| q = q.reshape(batch, seq_len, self.n_heads, self.d_head).transpose(0, 2, 1, 3) | |
| k = k.reshape(batch, seq_len, self.n_heads, self.d_head).transpose(0, 2, 1, 3) | |
| v = v.reshape(batch, seq_len, self.n_heads, self.d_head).transpose(0, 2, 1, 3) | |
| # Apply RoPE to Q and K | |
| pos_start = past_len | |
| pos_end = past_len + seq_len | |
| if pos_end > self.max_seq_len: | |
| # Extend RoPE tables dynamically | |
| pos = np.arange(pos_end, dtype=np.float32) | |
| cos_ext, sin_ext = rope(pos, self.d_head) | |
| cos = cos_ext[pos_start:pos_end] | |
| sin = sin_ext[pos_start:pos_end] | |
| self._cos = cos_ext | |
| self._sin = sin_ext | |
| else: | |
| cos = self._cos[pos_start:pos_end] | |
| sin = self._sin[pos_start:pos_end] | |
| q = apply_rope(q, cos, sin) | |
| k = apply_rope(k, cos, sin) | |
| # KV cache | |
| if use_cache: | |
| if layer_idx in self._kv_cache: | |
| past_k, past_v = self._kv_cache[layer_idx] | |
| k = np.concatenate([past_k, k], axis=2) | |
| v = np.concatenate([past_v, v], axis=2) | |
| self._kv_cache[layer_idx] = (k, v) | |
| # Scaled dot-product attention | |
| # q: [batch, n_heads, seq_len, d_head] | |
| # k: [batch, n_heads, total_len, d_head] | |
| scores = q @ k.transpose(0, 1, 3, 2) / math.sqrt(self.d_head) | |
| # Causal mask | |
| total_len = k.shape[2] | |
| causal = np.triu(np.ones((seq_len, total_len), dtype=bool), k=total_len - seq_len) | |
| scores = np.where(causal[None, None, :, :], -1e9, scores) | |
| attn = softmax(scores, axis=-1) | |
| # Apply attention to V | |
| out = attn @ v # [batch, n_heads, seq_len, d_head] | |
| out = out.transpose(0, 2, 1, 3).reshape(batch, seq_len, self.d_model) | |
| return self.wo.forward(out) | |
| def reset_cache(self) -> None: | |
| self._kv_cache.clear() | |
| class FeedForward: | |
| """Feed-forward network: 2-layer MLP with GELU.""" | |
| def __init__(self, d_model: int, d_ff: int) -> None: | |
| self.w1 = Linear(d_model, d_ff, bias=False) | |
| self.w2 = Linear(d_ff, d_model, bias=False) | |
| def forward(self, x: np.ndarray) -> np.ndarray: | |
| """x: [..., d_model] → [..., d_model]""" | |
| return self.w2.forward(gelu(self.w1.forward(x))) | |
| class TransformerLayer: | |
| """Single transformer layer: pre-norm attention + pre-norm FFN.""" | |
| def __init__(self, d_model: int, n_heads: int, d_ff: int, max_seq_len: int = 512) -> None: | |
| self.attn = MultiHeadAttention(d_model, n_heads, max_seq_len) | |
| self.ffn = FeedForward(d_model, d_ff) | |
| # Layer norm parameters | |
| self.ln1_gamma = np.ones(d_model, dtype=np.float32) | |
| self.ln1_beta = np.zeros(d_model, dtype=np.float32) | |
| self.ln2_gamma = np.ones(d_model, dtype=np.float32) | |
| self.ln2_beta = np.zeros(d_model, dtype=np.float32) | |
| def forward( | |
| self, | |
| x: np.ndarray, | |
| layer_idx: int = 0, | |
| use_cache: bool = False, | |
| past_len: int = 0, | |
| ) -> np.ndarray: | |
| """Pre-norm transformer layer.""" | |
| # Attention with residual | |
| normed = layer_norm(x, self.ln1_gamma, self.ln1_beta) | |
| attn_out = self.attn.forward(normed, layer_idx=layer_idx, use_cache=use_cache, past_len=past_len) | |
| x = x + attn_out | |
| # FFN with residual | |
| normed = layer_norm(x, self.ln2_gamma, self.ln2_beta) | |
| ffn_out = self.ffn.forward(normed) | |
| x = x + ffn_out | |
| return x | |
| def get_params(self) -> dict[str, Any]: | |
| """Get all parameters as a dict (for saving/quantization).""" | |
| return { | |
| "wq": self.attn.wq.weight, | |
| "wk": self.attn.wk.weight, | |
| "wv": self.attn.wv.weight, | |
| "wo": self.attn.wo.weight, | |
| "w1": self.ffn.w1.weight, | |
| "w2": self.ffn.w2.weight, | |
| "ln1_gamma": self.ln1_gamma, | |
| "ln1_beta": self.ln1_beta, | |
| "ln2_gamma": self.ln2_gamma, | |
| "ln2_beta": self.ln2_beta, | |
| } | |
| def set_params(self, params: dict[str, Any]) -> None: | |
| """Set parameters from a dict.""" | |
| self.attn.wq.weight = params["wq"] | |
| self.attn.wk.weight = params["wk"] | |
| self.attn.wv.weight = params["wv"] | |
| self.attn.wo.weight = params["wo"] | |
| self.ffn.w1.weight = params["w1"] | |
| self.ffn.w2.weight = params["w2"] | |
| self.ln1_gamma = params["ln1_gamma"] | |
| self.ln1_beta = params["ln1_beta"] | |
| self.ln2_gamma = params["ln2_gamma"] | |
| self.ln2_beta = params["ln2_beta"] | |