Download splitbit_llm/model/model.py from hermescures1/splitbit-llm: direct link, hf CLI and curl.
- Browser
- Download file 16.9 kB
-
https://huggingface.co/hermescures1/splitbit-llm/resolve/main/splitbit_llm/model/model.py
- Command line
-
hf download hf://hermescures1/splitbit-llm/splitbit_llm/model/model.py
-
curl -L -o model.py https://huggingface.co/hermescures1/splitbit-llm/resolve/main/splitbit_llm/model/model.py
16.9 kB
| """SplitBit LLM — full autoregressive transformer model. | |
| Pure NumPy implementation with: | |
| - Token embedding + positional encoding (RoPE) | |
| - Stack of transformer layers | |
| - LM head for next-token prediction | |
| - KV cache for fast autoregressive generation | |
| - Streaming generation (token-by-token) | |
| - SplitBit weight quantization support | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import math | |
| import os | |
| import time | |
| from typing import Any, Iterator | |
| import numpy as np | |
| from .tokenizer import BPETokenizer, BOS_ID, EOS_ID, PAD_ID | |
| from .layers import TransformerLayer, Embedding, Linear, layer_norm | |
| from .quantization import SplitBitQuantizer | |
| logger = logging.getLogger(__name__) | |
| class SplitBitLLM: | |
| """Full autoregressive transformer LLM. | |
| forward(tokens) → logits | |
| generate(prompt, max_tokens, temperature) → text | |
| generate_stream(prompt) → iterator yielding tokens | |
| """ | |
| def __init__(self, config: Any = None, tokenizer: BPETokenizer | None = None) -> None: | |
| if config is None: | |
| from ..config import Settings, get_model_config, detect_hardware | |
| config = get_model_config(detect_hardware()) | |
| self.config = config | |
| self.tokenizer = tokenizer | |
| # Model architecture | |
| self.embedding = Embedding(config.vocab_size, config.d_model) | |
| self.layers = [ | |
| TransformerLayer(config.d_model, config.n_heads, config.d_ff, config.max_seq_len) | |
| for _ in range(config.n_layers) | |
| ] | |
| # Final layer norm | |
| self.ln_f_gamma = np.ones(config.d_model, dtype=np.float32) | |
| self.ln_f_beta = np.zeros(config.d_model, dtype=np.float32) | |
| # LM head (tied with embedding) | |
| self.lm_head = Linear(config.d_model, config.vocab_size, bias=False) | |
| # Generation state | |
| self._kv_cache_active = False | |
| self._inference_count = 0 | |
| self._total_tokens_generated = 0 | |
| self._total_inference_time_s = 0.0 | |
| def param_count(self) -> int: | |
| """Total parameter count.""" | |
| total = self.embedding.weight.size + self.lm_head.weight.size | |
| total += self.ln_f_gamma.size + self.ln_f_beta.size | |
| for layer in self.layers: | |
| total += sum(p.size for p in layer.get_params().values()) | |
| return total | |
| def forward(self, token_ids: np.ndarray, use_cache: bool = False, past_len: int = 0) -> np.ndarray: | |
| """ | |
| token_ids: [batch, seq_len] | |
| Returns logits: [batch, seq_len, vocab_size] | |
| """ | |
| batch, seq_len = token_ids.shape | |
| # Truncate to max_seq_len | |
| if seq_len > self.config.max_seq_len: | |
| token_ids = token_ids[:, -self.config.max_seq_len:] | |
| seq_len = self.config.max_seq_len | |
| if past_len > 0: | |
| past_len = max(0, past_len - (seq_len - self.config.max_seq_len)) | |
| # Embedding | |
| x = self.embedding.forward(token_ids) # [batch, seq_len, d_model] | |
| # Transformer layers | |
| for i, layer in enumerate(self.layers): | |
| x = layer.forward(x, layer_idx=i, use_cache=use_cache, past_len=past_len) | |
| # Final layer norm | |
| x = layer_norm(x, self.ln_f_gamma, self.ln_f_beta) | |
| # LM head | |
| logits = self.lm_head.forward(x) # [batch, seq_len, vocab_size] | |
| return logits | |
| def reset_cache(self) -> None: | |
| """Reset KV cache for all layers.""" | |
| for layer in self.layers: | |
| layer.attn.reset_cache() | |
| self._kv_cache_active = False | |
| def generate( | |
| self, | |
| prompt: str, | |
| max_tokens: int = 128, | |
| temperature: float = 0.7, | |
| top_k: int = 40, | |
| use_cache: bool = True, | |
| ) -> str: | |
| """Generate text from a prompt. | |
| Args: | |
| prompt: input text | |
| max_tokens: max tokens to generate | |
| temperature: sampling temperature (0 = greedy) | |
| top_k: top-k sampling (0 = disabled) | |
| use_cache: use KV cache for faster generation | |
| Returns: | |
| Generated text (prompt + completion) | |
| """ | |
| tokens = self._encode_prompt(prompt) | |
| if not tokens: | |
| return prompt | |
| generated = list(tokens) | |
| self.reset_cache() | |
| # Initial forward pass | |
| token_arr = np.array([generated], dtype=np.int64) | |
| logits = self.forward(token_arr, use_cache=use_cache, past_len=0) | |
| past_len = len(generated) | |
| for _ in range(max_tokens): | |
| # Get logits for last token | |
| next_logits = logits[0, -1, :] # [vocab_size] | |
| # Sample next token | |
| next_token = self._sample(next_logits, temperature, top_k) | |
| if next_token == EOS_ID: | |
| break | |
| generated.append(next_token) | |
| self._total_tokens_generated += 1 | |
| # Forward just the new token with cache | |
| if use_cache: | |
| new_arr = np.array([[next_token]], dtype=np.int64) | |
| logits = self.forward(new_arr, use_cache=True, past_len=past_len) | |
| past_len += 1 | |
| else: | |
| token_arr = np.array([generated[-self.config.max_seq_len:]], dtype=np.int64) | |
| logits = self.forward(token_arr, use_cache=False, past_len=0) | |
| self.reset_cache() | |
| return self._decode(generated) | |
| def generate_stream( | |
| self, | |
| prompt: str, | |
| max_tokens: int = 128, | |
| temperature: float = 0.7, | |
| top_k: int = 40, | |
| use_cache: bool = True, | |
| ) -> Iterator[str]: | |
| """Streaming generation — yields text chunks as they're generated. | |
| Yields decoded text chunks (may be partial words). | |
| """ | |
| tokens = self._encode_prompt(prompt) | |
| if not tokens: | |
| return | |
| # Yield prompt first | |
| yield self.tokenizer.decode(tokens) if self.tokenizer else "" | |
| generated = list(tokens) | |
| self.reset_cache() | |
| # Initial forward pass | |
| token_arr = np.array([generated], dtype=np.int64) | |
| logits = self.forward(token_arr, use_cache=use_cache, past_len=0) | |
| past_len = len(generated) | |
| for _ in range(max_tokens): | |
| next_logits = logits[0, -1, :] | |
| next_token = self._sample(next_logits, temperature, top_k) | |
| if next_token == EOS_ID: | |
| break | |
| generated.append(next_token) | |
| self._total_tokens_generated += 1 | |
| # Decode just this token | |
| if self.tokenizer: | |
| chunk = self.tokenizer.decode([next_token]) | |
| if chunk: | |
| yield chunk | |
| if use_cache: | |
| new_arr = np.array([[next_token]], dtype=np.int64) | |
| logits = self.forward(new_arr, use_cache=True, past_len=past_len) | |
| past_len += 1 | |
| else: | |
| token_arr = np.array([generated[-self.config.max_seq_len:]], dtype=np.int64) | |
| logits = self.forward(token_arr, use_cache=False, past_len=0) | |
| self.reset_cache() | |
| def generate_stream_sentences( | |
| self, | |
| prompt: str, | |
| max_tokens: int = 128, | |
| temperature: float = 0.7, | |
| top_k: int = 40, | |
| ) -> Iterator[str]: | |
| """Streaming generation that yields complete sentences. | |
| Used for voice/TTS — first sentence comes out ASAP. | |
| """ | |
| buffer = "" | |
| for chunk in self.generate_stream(prompt, max_tokens, temperature, top_k): | |
| buffer += chunk | |
| # Check for sentence boundaries | |
| while buffer: | |
| # Find sentence end | |
| end_idx = -1 | |
| for delim in [". ", "! ", "? ", ".\n", "!\n", "?\n"]: | |
| idx = buffer.find(delim) | |
| if idx >= 0 and (end_idx < 0 or idx < end_idx): | |
| end_idx = idx + len(delim) | |
| if end_idx > 0: | |
| yield buffer[:end_idx] | |
| buffer = buffer[end_idx:] | |
| else: | |
| break | |
| if buffer: | |
| yield buffer | |
| def _encode_prompt(self, prompt: str) -> list[int]: | |
| """Encode prompt to token IDs.""" | |
| if self.tokenizer: | |
| return self.tokenizer.encode(prompt, add_bos=True) | |
| # Fallback: simple char-level encoding | |
| return [BOS_ID] + [min(ord(c), self.config.vocab_size - 1) for c in prompt[:self.config.max_seq_len - 1]] | |
| def _decode(self, tokens: list[int]) -> str: | |
| """Decode tokens to text.""" | |
| if self.tokenizer: | |
| return self.tokenizer.decode(tokens) | |
| return "".join(chr(t) for t in tokens if t < 128 and t not in (PAD_ID, BOS_ID, EOS_ID)) | |
| def _sample(self, logits: np.ndarray, temperature: float, top_k: int) -> int: | |
| """Sample next token from logits.""" | |
| if temperature <= 0: | |
| return int(np.argmax(logits)) | |
| # Apply temperature | |
| logits = logits / max(temperature, 1e-8) | |
| # Top-k filtering | |
| if top_k > 0 and top_k < len(logits): | |
| top_indices = np.argpartition(logits, -top_k)[-top_k:] | |
| mask = np.full_like(logits, -1e9) | |
| mask[top_indices] = logits[top_indices] | |
| logits = mask | |
| # Softmax and sample | |
| probs = np.exp(logits - np.max(logits)) | |
| probs = probs / np.sum(probs) | |
| return int(np.random.choice(len(probs), p=probs)) | |
| def save(self, path: str, quantizer: SplitBitQuantizer | None = None) -> None: | |
| """Save model to disk. If quantizer provided, weights are quantized.""" | |
| os.makedirs(os.path.dirname(path) or ".", exist_ok=True) | |
| data = { | |
| "config": { | |
| "n_layers": self.config.n_layers, | |
| "n_heads": self.config.n_heads, | |
| "d_model": self.config.d_model, | |
| "d_ff": self.config.d_ff, | |
| "vocab_size": self.config.vocab_size, | |
| "max_seq_len": self.config.max_seq_len, | |
| }, | |
| "embedding": self.embedding.weight, | |
| "lm_head": self.lm_head.weight, | |
| "ln_f_gamma": self.ln_f_gamma, | |
| "ln_f_beta": self.ln_f_beta, | |
| "layers": [], | |
| "quantized": quantizer is not None, | |
| } | |
| for layer in self.layers: | |
| params = layer.get_params() | |
| if quantizer: | |
| layer_data = {} | |
| for k, v in params.items(): | |
| if "ln" in k: | |
| layer_data[k] = v # Don't quantize layer norm | |
| else: | |
| packed = quantizer.quantize(v) | |
| layer_data[k] = packed | |
| else: | |
| layer_data = params | |
| data["layers"].append(layer_data) | |
| if quantizer: | |
| data["embedding"] = quantizer.quantize(self.embedding.weight) | |
| data["lm_head"] = quantizer.quantize(self.lm_head.weight) | |
| np.savez(path, **self._flatten_save_dict(data)) | |
| logger.info("Model saved to %s (quantized=%s)", path, quantizer is not None) | |
| def _flatten_save_dict(self, data: dict, prefix: str = "") -> dict: | |
| """Flatten nested dict for np.savez.""" | |
| flat = {} | |
| for k, v in data.items(): | |
| key = f"{prefix}_{k}" if prefix else k | |
| if isinstance(v, dict) and "data" not in v: | |
| flat.update(self._flatten_save_dict(v, key)) | |
| elif isinstance(v, list): | |
| for i, item in enumerate(v): | |
| flat.update(self._flatten_save_dict(item, f"{key}_{i}")) | |
| elif isinstance(v, np.ndarray): | |
| flat[key] = v | |
| elif isinstance(v, dict): | |
| # Packed quantized data | |
| for pk, pv in v.items(): | |
| if isinstance(pv, np.ndarray): | |
| flat[f"{key}_{pk}"] = pv | |
| elif pv is not None: | |
| flat[f"{key}_{pk}"] = np.array(pv) | |
| elif v is not None: | |
| flat[key] = np.array(v) | |
| return flat | |
| def load(cls, path: str, tokenizer: BPETokenizer | None = None, | |
| quantizer: SplitBitQuantizer | None = None) -> "SplitBitLLM": | |
| """Load model from disk.""" | |
| from ..config import ModelConfig | |
| npz = np.load(path, allow_pickle=True) | |
| config = ModelConfig( | |
| n_layers=int(npz["config_n_layers"]), | |
| n_heads=int(npz["config_n_heads"]), | |
| d_model=int(npz["config_d_model"]), | |
| d_ff=int(npz["config_d_ff"]), | |
| vocab_size=int(npz["config_vocab_size"]), | |
| max_seq_len=int(npz["config_max_seq_len"]), | |
| ) | |
| model = cls(config=config, tokenizer=tokenizer) | |
| # Load embedding and LM head | |
| if quantizer and "embedding_data" in npz: | |
| model.embedding.weight = quantizer.dequantize({ | |
| "data": npz["embedding_data"], | |
| "scale": npz["embedding_scale"] if "embedding_scale" in npz else None, | |
| "shape": npz["embedding_shape"], | |
| "format": str(npz["embedding_format"]) if "embedding_format" in npz else "q4_k_m", | |
| "bits": int(npz["embedding_bits"]) if "embedding_bits" in npz else 4, | |
| "n_blocks": int(npz["embedding_n_blocks"]) if "embedding_n_blocks" in npz else 0, | |
| "block_size": int(npz["embedding_block_size"]) if "embedding_block_size" in npz else 32, | |
| "pad_len": int(npz["embedding_pad_len"]) if "embedding_pad_len" in npz else 0, | |
| }) | |
| model.lm_head.weight = quantizer.dequantize({ | |
| "data": npz["lm_head_data"], | |
| "scale": npz["lm_head_scale"] if "lm_head_scale" in npz else None, | |
| "shape": npz["lm_head_shape"], | |
| "format": str(npz["lm_head_format"]) if "lm_head_format" in npz else "q4_k_m", | |
| "bits": int(npz["lm_head_bits"]) if "lm_head_bits" in npz else 4, | |
| "n_blocks": int(npz["lm_head_n_blocks"]) if "lm_head_n_blocks" in npz else 0, | |
| "block_size": int(npz["lm_head_block_size"]) if "lm_head_block_size" in npz else 32, | |
| "pad_len": int(npz["lm_head_pad_len"]) if "lm_head_pad_len" in npz else 0, | |
| }) | |
| else: | |
| model.embedding.weight = npz["embedding"] | |
| model.lm_head.weight = npz["lm_head"] | |
| model.ln_f_gamma = npz["ln_f_gamma"] | |
| model.ln_f_beta = npz["ln_f_beta"] | |
| # Load layers | |
| for i, layer in enumerate(model.layers): | |
| params = {} | |
| for key in ["wq", "wk", "wv", "wo", "w1", "w2"]: | |
| full_key = f"layers_{i}_{key}" | |
| if quantizer and f"{full_key}_data" in npz: | |
| params[key] = quantizer.dequantize({ | |
| "data": npz[f"{full_key}_data"], | |
| "scale": npz[f"{full_key}_scale"] if f"{full_key}_scale" in npz else None, | |
| "shape": npz[f"{full_key}_shape"], | |
| "format": str(npz[f"{full_key}_format"]) if f"{full_key}_format" in npz else "q4_k_m", | |
| "bits": int(npz[f"{full_key}_bits"]) if f"{full_key}_bits" in npz else 4, | |
| "n_blocks": int(npz[f"{full_key}_n_blocks"]) if f"{full_key}_n_blocks" in npz else 0, | |
| "block_size": int(npz[f"{full_key}_block_size"]) if f"{full_key}_block_size" in npz else 32, | |
| "pad_len": int(npz[f"{full_key}_pad_len"]) if f"{full_key}_pad_len" in npz else 0, | |
| }) | |
| elif full_key in npz: | |
| params[key] = npz[full_key] | |
| for key in ["ln1_gamma", "ln1_beta", "ln2_gamma", "ln2_beta"]: | |
| full_key = f"layers_{i}_{key}" | |
| if full_key in npz: | |
| params[key] = npz[full_key] | |
| layer.set_params(params) | |
| logger.info("Model loaded from %s (%d params)", path, model.param_count) | |
| return model | |
| def get_stats(self) -> dict[str, Any]: | |
| """Get model statistics.""" | |
| avg_time = self._total_inference_time_s / max(self._inference_count, 1) | |
| return { | |
| "param_count": self.param_count, | |
| "config": { | |
| "n_layers": self.config.n_layers, | |
| "n_heads": self.config.n_heads, | |
| "d_model": self.config.d_model, | |
| "d_ff": self.config.d_ff, | |
| "vocab_size": self.config.vocab_size, | |
| "max_seq_len": self.config.max_seq_len, | |
| }, | |
| "inference_count": self._inference_count, | |
| "total_tokens_generated": self._total_tokens_generated, | |
| "avg_inference_time_s": round(avg_time, 4), | |
| "tokens_per_second": round(self._total_tokens_generated / max(self._total_inference_time_s, 0.001), 2), | |
| } | |