"""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 @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 # 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 @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) # 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), }