| """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 |
|
|
| |
| 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) |
| ] |
| |
| self.ln_f_gamma = np.ones(config.d_model, dtype=np.float32) |
| self.ln_f_beta = np.zeros(config.d_model, dtype=np.float32) |
| |
| self.lm_head = Linear(config.d_model, config.vocab_size, bias=False) |
|
|
| |
| self._kv_cache_active = False |
| self._inference_count = 0 |
| self._total_tokens_generated = 0 |
| self._total_inference_time_s = 0.0 |
|
|
| @property |
| 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 |
|
|
| |
| 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)) |
|
|
| |
| x = self.embedding.forward(token_ids) |
|
|
| |
| for i, layer in enumerate(self.layers): |
| x = layer.forward(x, layer_idx=i, use_cache=use_cache, past_len=past_len) |
|
|
| |
| x = layer_norm(x, self.ln_f_gamma, self.ln_f_beta) |
|
|
| |
| logits = self.lm_head.forward(x) |
|
|
| 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() |
|
|
| |
| 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 |
|
|
| |
| 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 self.tokenizer.decode(tokens) if self.tokenizer else "" |
|
|
| generated = list(tokens) |
| self.reset_cache() |
|
|
| |
| 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 |
|
|
| |
| 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 |
| |
| while buffer: |
| |
| 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) |
| |
| 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)) |
|
|
| |
| logits = logits / max(temperature, 1e-8) |
|
|
| |
| 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 |
|
|
| |
| 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 |
| 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): |
| |
| 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 |
|
|
| @classmethod |
| 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) |
|
|
| |
| 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"] |
|
|
| |
| 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), |
| } |
|
|