//! WASM Model Loading and Inference //! //! Model loading and forward pass for WebAssembly inference. use crate::tensor::WasmKvCache; use serde::{Deserialize, Serialize}; use wasm_bindgen::prelude::*; /// Model configuration #[wasm_bindgen] #[derive(Debug, Clone, Serialize, Deserialize)] pub struct ModelConfig { /// Vocabulary size pub vocab_size: usize, /// Hidden dimension pub hidden_size: usize, /// Number of layers pub num_layers: usize, /// Number of attention heads pub num_heads: usize, /// Head dimension pub head_dim: usize, /// Intermediate size (FFN) pub intermediate_size: usize, /// Maximum sequence length pub max_seq_len: usize, /// RMS norm epsilon pub rms_norm_eps: f32, /// Whether model is quantized pub quantized: bool, } impl Default for ModelConfig { fn default() -> Self { Self { vocab_size: 32000, hidden_size: 4096, num_layers: 32, num_heads: 32, head_dim: 128, intermediate_size: 11008, max_seq_len: 2048, rms_norm_eps: 1e-6, quantized: false, } } } #[wasm_bindgen] impl ModelConfig { /// Create a small model config for testing #[wasm_bindgen] pub fn tiny() -> Self { Self { vocab_size: 32000, hidden_size: 512, num_layers: 4, num_heads: 8, head_dim: 64, intermediate_size: 1376, max_seq_len: 512, rms_norm_eps: 1e-6, quantized: false, } } /// Create a 1B-like model config #[wasm_bindgen] pub fn small() -> Self { Self { vocab_size: 32000, hidden_size: 2048, num_layers: 16, num_heads: 16, head_dim: 128, intermediate_size: 5504, max_seq_len: 2048, rms_norm_eps: 1e-6, quantized: false, } } } /// Layer weights #[derive(Debug, Clone)] struct LayerWeights { /// Attention QKV projection attn_qkv: Vec, /// Attention output projection attn_out: Vec, /// FFN gate projection ffn_gate: Vec, /// FFN up projection ffn_up: Vec, /// FFN down projection ffn_down: Vec, /// Pre-attention norm attn_norm: Vec, /// Pre-FFN norm ffn_norm: Vec, } /// WASM model for inference #[wasm_bindgen] pub struct WasmModel { /// Model configuration config: ModelConfig, /// Embedding weights embeddings: Vec, /// Layer weights layers: Vec, /// Final norm final_norm: Vec, /// Output projection (lm_head) lm_head: Vec, /// Whether model is loaded loaded: bool, } #[wasm_bindgen] impl WasmModel { /// Create a new uninitialized model #[wasm_bindgen(constructor)] pub fn new(config: ModelConfig) -> Self { Self { config, embeddings: Vec::new(), layers: Vec::new(), final_norm: Vec::new(), lm_head: Vec::new(), loaded: false, } } /// Load model from bytes pub fn from_bytes(bytes: &[u8], quantized: bool) -> Result { // Parse model format (simplified - would need actual format parsing) // For now, create a dummy model for structure validation if bytes.len() < 16 { return Err(JsError::new("Invalid model data: too short")); } // Read config from header (simplified) let config = ModelConfig { quantized, ..ModelConfig::tiny() }; let mut model = Self::new(config.clone()); // Initialize with random weights for now // In production, parse actual weights from bytes model.init_random_weights(); model.loaded = true; Ok(model) } /// Initialize with random weights (for testing) fn init_random_weights(&mut self) { let c = &self.config; // Embeddings self.embeddings = (0..c.vocab_size * c.hidden_size) .map(|_| (rand::random::() - 0.5) * 0.02) .collect(); // Layers self.layers = (0..c.num_layers) .map(|_| LayerWeights { attn_qkv: (0..c.hidden_size * c.hidden_size * 3) .map(|_| (rand::random::() - 0.5) * 0.02) .collect(), attn_out: (0..c.hidden_size * c.hidden_size) .map(|_| (rand::random::() - 0.5) * 0.02) .collect(), ffn_gate: (0..c.hidden_size * c.intermediate_size) .map(|_| (rand::random::() - 0.5) * 0.02) .collect(), ffn_up: (0..c.hidden_size * c.intermediate_size) .map(|_| (rand::random::() - 0.5) * 0.02) .collect(), ffn_down: (0..c.intermediate_size * c.hidden_size) .map(|_| (rand::random::() - 0.5) * 0.02) .collect(), attn_norm: vec![1.0; c.hidden_size], ffn_norm: vec![1.0; c.hidden_size], }) .collect(); // Final norm self.final_norm = vec![1.0; c.hidden_size]; // LM head self.lm_head = (0..c.hidden_size * c.vocab_size) .map(|_| (rand::random::() - 0.5) * 0.02) .collect(); } /// Check if model is loaded #[wasm_bindgen] pub fn is_loaded(&self) -> bool { self.loaded } /// Get model config #[wasm_bindgen] pub fn config(&self) -> ModelConfig { self.config.clone() } /// Memory usage in bytes #[wasm_bindgen] pub fn memory_usage(&self) -> usize { let mut total = 0; total += self.embeddings.len() * 4; for layer in &self.layers { total += layer.attn_qkv.len() * 4; total += layer.attn_out.len() * 4; total += layer.ffn_gate.len() * 4; total += layer.ffn_up.len() * 4; total += layer.ffn_down.len() * 4; total += layer.attn_norm.len() * 4; total += layer.ffn_norm.len() * 4; } total += self.final_norm.len() * 4; total += self.lm_head.len() * 4; total } /// Get parameter count #[wasm_bindgen] pub fn param_count(&self) -> usize { let c = &self.config; let embeddings = c.vocab_size * c.hidden_size; let per_layer = c.hidden_size * c.hidden_size * 3 // QKV + c.hidden_size * c.hidden_size // O + c.hidden_size * c.intermediate_size * 3 // gate, up, down + c.hidden_size * 2; // norms let lm_head = c.hidden_size * c.vocab_size; embeddings + per_layer * c.num_layers + c.hidden_size + lm_head } } // Internal methods (not exposed to WASM) impl WasmModel { /// Forward pass pub fn forward( &self, input_ids: &[u32], _kv_cache: Option<&mut WasmKvCache>, ) -> Result, JsError> { if !self.loaded { return Err(JsError::new("Model not loaded")); } let seq_len = input_ids.len(); let c = &self.config; // Embed tokens let mut hidden = vec![0.0; seq_len * c.hidden_size]; for (i, &token_id) in input_ids.iter().enumerate() { if token_id as usize >= c.vocab_size { return Err(JsError::new(&format!("Token ID {} out of range", token_id))); } let start = token_id as usize * c.hidden_size; let end = start + c.hidden_size; hidden[i * c.hidden_size..(i + 1) * c.hidden_size] .copy_from_slice(&self.embeddings[start..end]); } // Process layers (simplified) for layer in &self.layers { // RMS Norm + Attention (simplified) hidden = self.rms_norm(&hidden, &layer.attn_norm); hidden = self.attention(&hidden); // RMS Norm + FFN (simplified) hidden = self.rms_norm(&hidden, &layer.ffn_norm); hidden = self.ffn(&hidden, layer); } // Final norm hidden = self.rms_norm(&hidden, &self.final_norm); // Compute logits for last token only let last_hidden = &hidden[(seq_len - 1) * c.hidden_size..seq_len * c.hidden_size]; let logits = self.compute_logits(last_hidden); Ok(logits) } fn rms_norm(&self, x: &[f32], weight: &[f32]) -> Vec { let hidden_size = self.config.hidden_size; let num_tokens = x.len() / hidden_size; let mut result = vec![0.0; x.len()]; for t in 0..num_tokens { let start = t * hidden_size; let end = start + hidden_size; let slice = &x[start..end]; // Compute RMS let rms: f32 = (slice.iter().map(|v| v * v).sum::() / hidden_size as f32 + self.config.rms_norm_eps) .sqrt(); // Normalize and apply weight for (i, &v) in slice.iter().enumerate() { result[start + i] = (v / rms) * weight[i]; } } result } fn attention(&self, x: &[f32]) -> Vec { // Simplified attention - in production would use actual attention mechanism // For now, just return input scaled (placeholder) x.iter().map(|v| v * 0.9).collect() } fn ffn(&self, x: &[f32], _layer: &LayerWeights) -> Vec { // Simplified FFN: SiLU(x @ gate) * (x @ up) @ down // For now, just apply activation (placeholder) x.iter() .map(|v| { let silu = v / (1.0 + (-v).exp()); silu * 0.9 }) .collect() } fn compute_logits(&self, hidden: &[f32]) -> Vec { let c = &self.config; let mut logits = vec![0.0; c.vocab_size]; // Matrix multiply: hidden @ lm_head^T for v in 0..c.vocab_size { let mut sum = 0.0; for h in 0..c.hidden_size { sum += hidden[h] * self.lm_head[v * c.hidden_size + h]; } logits[v] = sum; } logits } } #[cfg(test)] mod tests { use super::*; #[test] fn test_model_config() { let config = ModelConfig::tiny(); assert_eq!(config.num_layers, 4); assert_eq!(config.hidden_size, 512); } #[test] fn test_model_memory() { let config = ModelConfig::tiny(); let mut model = WasmModel::new(config); model.init_random_weights(); model.loaded = true; assert!(model.memory_usage() > 0); assert!(model.param_count() > 0); } }