//! LLaMA Attention Mechanisms and Components //! //! Specialized attention implementations for LLaMA architecture including //! rotary position embeddings, grouped-query attention, and SwiGLU. use crate::architectures::TransformerConfig; use crate::{Result, TransformerError}; use rtx_tensor::{Tensor, Device, DType}; use serde::{Deserialize, Serialize}; use tracing::debug; /// LLaMA-specific configuration #[derive(Debug, Clone, Serialize, Deserialize)] pub struct LLaMAConfig { /// Base transformer configuration pub base: TransformerConfig, /// Vocabulary size pub vocab_size: usize, /// Hidden dimension pub hidden_size: usize, /// Feed-forward intermediate dimension pub intermediate_size: usize, /// Number of transformer layers pub num_hidden_layers: usize, /// Number of attention heads pub num_attention_heads: usize, /// Number of key-value heads (for grouped-query attention) pub num_key_value_heads: usize, /// Maximum sequence length pub max_position_embeddings: usize, /// RMSNorm epsilon pub rms_norm_eps: f64, /// Initializer range for weights pub initializer_range: f64, /// Use cache for generation pub use_cache: bool, /// Pad token ID pub pad_token_id: i64, /// BOS token ID pub bos_token_id: i64, /// EOS token ID pub eos_token_id: i64, /// Pretraining type pub pretraining_tp: usize, /// Tie word embeddings pub tie_word_embeddings: bool, /// Rope scaling configuration pub rope_scaling: Option, /// Rope theta parameter pub rope_theta: f64, /// Attention bias pub attention_bias: bool, /// MLP bias pub mlp_bias: bool, } /// RoPE (Rotary Position Embedding) scaling configuration #[derive(Debug, Clone, Serialize, Deserialize)] pub struct RopeScaling { /// Scaling type ("linear" or "dynamic") pub scaling_type: String, /// Scaling factor pub factor: f64, } impl Default for LLaMAConfig { fn default() -> Self { Self { base: TransformerConfig::default(), vocab_size: 32000, hidden_size: 4096, intermediate_size: 11008, num_hidden_layers: 32, num_attention_heads: 32, num_key_value_heads: 32, // Standard multi-head attention max_position_embeddings: 2048, rms_norm_eps: 1e-6, initializer_range: 0.02, use_cache: true, pad_token_id: -1, bos_token_id: 1, eos_token_id: 2, pretraining_tp: 1, tie_word_embeddings: false, rope_scaling: None, rope_theta: 10000.0, attention_bias: false, mlp_bias: false, } } } impl LLaMAConfig { /// Create LLaMA 7B configuration pub fn llama_7b() -> Self { Self::default() } /// Create LLaMA 13B configuration pub fn llama_13b() -> Self { Self { hidden_size: 5120, intermediate_size: 13824, num_hidden_layers: 40, num_attention_heads: 40, num_key_value_heads: 40, ..Self::default() } } /// Create LLaMA 30B configuration pub fn llama_30b() -> Self { Self { hidden_size: 6656, intermediate_size: 17920, num_hidden_layers: 60, num_attention_heads: 52, num_key_value_heads: 52, ..Self::default() } } /// Create LLaMA 65B configuration pub fn llama_65b() -> Self { Self { hidden_size: 8_192, intermediate_size: 22016, num_hidden_layers: 80, num_attention_heads: 64, num_key_value_heads: 64, ..Self::default() } } /// Create LLaMA 2 7B configuration pub fn llama2_7b() -> Self { Self { vocab_size: 32000, max_position_embeddings: 4096, ..Self::llama_7b() } } /// Create LLaMA 2 13B configuration pub fn llama2_13b() -> Self { Self { vocab_size: 32000, max_position_embeddings: 4096, ..Self::llama_13b() } } /// Create LLaMA 2 70B configuration with grouped-query attention pub fn llama2_70b() -> Self { Self { vocab_size: 32000, hidden_size: 8_192, intermediate_size: 28672, num_hidden_layers: 80, num_attention_heads: 64, num_key_value_heads: 8, // Grouped-query attention max_position_embeddings: 4096, ..Self::default() } } /// Create Code Llama configuration pub fn code_llama() -> Self { Self { vocab_size: 32016, max_position_embeddings: 16_384, // Longer context for code rope_theta: 1_000_000.0, // Higher rope theta for longer sequences ..Self::llama2_7b() } } /// Validate configuration parameters pub fn validate(&self) -> Result<()> { if self.hidden_size % self.num_attention_heads != 0 { return Err(TransformerError::config( "hidden_size must be divisible by num_attention_heads" )); } if self.num_key_value_heads > self.num_attention_heads { return Err(TransformerError::config( "num_key_value_heads cannot be greater than num_attention_heads" )); } if self.num_attention_heads % self.num_key_value_heads != 0 { return Err(TransformerError::config( "num_attention_heads must be divisible by num_key_value_heads" )); } if self.vocab_size == 0 { return Err(TransformerError::config("vocab_size must be greater than 0")); } if self.num_hidden_layers == 0 { return Err(TransformerError::config("num_hidden_layers must be greater than 0")); } Ok(()) } } /// Rotary Position Embedding (RoPE) implementation #[derive(Debug)] pub struct RotaryPositionEmbedding { /// Dimension per head pub dim: usize, /// Maximum sequence length pub max_seq_len: usize, /// Theta parameter pub theta: f64, /// Precomputed cosine values pub cos_cached: Tensor, /// Precomputed sine values pub sin_cached: Tensor, } impl RotaryPositionEmbedding { /// Create new RoPE embeddings pub fn new(dim: usize, max_seq_len: usize, theta: f64, device: &Device) -> Result { let half_dim = dim / 2; // Create frequency inverse (1/theta^(2i/d)) let mut freqs = Vec::with_capacity(half_dim); for i in 0..half_dim { let freq = 1.0 / theta.powf(2.0 * i as f64 / dim as f64); freqs.push(freq as f32); } let freqs_tensor = Tensor::from_data(freqs, vec![half_dim], &Device::Cpu)?; // Create position indices let mut positions = Vec::with_capacity(max_seq_len); for i in 0..max_seq_len { positions.push(i as f32); } let positions_tensor = Tensor::from_data(positions, vec![max_seq_len], &Device::Cpu)?; // Compute angles: position * freq let angles = positions_tensor.unsqueeze(-1)?.matmul(&freqs_tensor.unsqueeze(0)?)?; // Precompute cosine and sine values let cos_cached = angles.cos()?; let sin_cached = angles.sin()?; Ok(Self { dim, max_seq_len, theta, cos_cached, sin_cached, }) } /// Apply rotary position embedding pub fn forward(&self, q: &Tensor, k: &Tensor, position_ids: &Tensor) -> Result<(Tensor, Tensor)> { let seq_len = position_ids.shape()[1]; if seq_len > self.max_seq_len { return Err(TransformerError::architecture( format!("Sequence length {} exceeds maximum {}", seq_len, self.max_seq_len) )); } // Get cos and sin for current positions let cos = self.cos_cached.slice(0, 0, seq_len)?; // TODO: Fix slice API let sin = self.sin_cached.slice(0, 0, seq_len)?; // TODO: Fix slice API // Apply rotation to queries and keys let q_rotated = self.apply_rotation(q, &cos, &sin)?; let k_rotated = self.apply_rotation(k, &cos, &sin)?; Ok((q_rotated, k_rotated)) } /// Apply rotation transformation fn apply_rotation(&self, x: &Tensor, cos: &Tensor, sin: &Tensor) -> Result { // Split x into two halves let half_dim = self.dim / 2; // TODO: Implement proper tensor slicing for rotary embeddings let x1 = x.clone(); // Placeholder let x2 = x.clone(); // Placeholder // Apply rotation: x1 * cos - x2 * sin, x1 * sin + x2 * cos let rotated_x1 = (x1.clone() * cos.clone()) - (x2.clone() * sin.clone()); let rotated_x2 = (x1 * sin.clone()) + (x2 * cos.clone()); // Concatenate back Tensor::cat(&[rotated_x1, rotated_x2], -1) } } /// SwiGLU activation function used in LLaMA #[derive(Debug)] pub struct SwiGLU { /// Gate projection pub gate_proj: Tensor, /// Up projection pub up_proj: Tensor, /// Down projection pub down_proj: Tensor, /// Hidden size pub hidden_size: usize, /// Intermediate size pub intermediate_size: usize, } impl SwiGLU { /// Create new SwiGLU layer pub fn new( hidden_size: usize, intermediate_size: usize, bias: bool, device: &Device, ) -> Result { let gate_proj = Tensor::randn(&[intermediate_size, hidden_size], DType::F32, device)? * 0.02; let up_proj = Tensor::randn(&[intermediate_size, hidden_size], DType::F32, device)? * 0.02; let down_proj = Tensor::randn(&[hidden_size, intermediate_size], DType::F32, device)? * 0.02; Ok(Self { gate_proj, up_proj, down_proj, hidden_size, intermediate_size, }) } /// Forward pass: SwiGLU(x) = Swish(gate(x)) * up(x) @ down pub fn forward(&self, x: &Tensor) -> Result { // Gate projection let gate = x.matmul(&self.gate_proj.transpose(-1, -2)?)?; // Up projection let up = x.matmul(&self.up_proj.transpose(-1, -2)?)?; // SwiGLU: Swish(gate) * up let swish_gate = self.swish(&gate)?; let gated = swish_gate * up; // Down projection gated.matmul(&self.down_proj.transpose(-1, -2)?) } /// Swish activation: x * sigmoid(x) fn swish(&self, x: &Tensor) -> Result { // Swish(x) = x * sigmoid(x) = x * (1 / (1 + exp(-x))) // This is also known as SiLU (Sigmoid Linear Unit) // Compute sigmoid(x) = 1 / (1 + exp(-x)) let neg_x = (-x.clone())?; let exp_neg_x = neg_x.exp()?; let one_plus_exp = (exp_neg_x + 1.0)?; let sigmoid_x = (1.0 / one_plus_exp)?; // Compute x * sigmoid(x) (x.clone() * sigmoid_x) .map_err(|e| crate::TransformerError::ArchitectureError( format!("Failed to compute swish activation: {}", e) )) } } /// LLaMA attention with grouped-query attention and RoPE #[derive(Debug)] pub struct LLaMAAttention { /// Query projection pub q_proj: Tensor, /// Key projection pub k_proj: Tensor, /// Value projection pub v_proj: Tensor, /// Output projection pub o_proj: Tensor, /// Rotary position embedding pub rotary_emb: RotaryPositionEmbedding, /// Configuration pub config: LLaMAConfig, } impl LLaMAAttention { /// Create new LLaMA attention pub fn new(config: &LLaMAConfig, device: &Device) -> Result { let head_dim = config.hidden_size / config.num_attention_heads; let q_proj = Tensor::randn( &[config.num_attention_heads * head_dim, config.hidden_size], DType::F32, device, )? * config.initializer_range as f32; let k_proj = Tensor::randn( &[config.num_key_value_heads * head_dim, config.hidden_size], DType::F32, device, )? * config.initializer_range as f32; let v_proj = Tensor::randn( &[config.num_key_value_heads * head_dim, config.hidden_size], DType::F32, device, )? * config.initializer_range as f32; let o_proj = Tensor::randn( &[config.hidden_size, config.num_attention_heads * head_dim], DType::F32, device, )? * config.initializer_range as f32; let rotary_emb = RotaryPositionEmbedding::new( head_dim, config.max_position_embeddings, config.rope_theta, device, )?; Ok(Self { q_proj, k_proj, v_proj, o_proj, rotary_emb, config: config.clone(), }) } /// Forward pass with grouped-query attention pub fn forward(&self, hidden_states: &Tensor, position_ids: &Tensor) -> Result { let batch_size = hidden_states.shape()[0]; let seq_len = hidden_states.shape()[1]; let head_dim = self.config.hidden_size / self.config.num_attention_heads; // Project to Q, K, V let q = hidden_states.matmul(&self.q_proj.transpose(-1, -2)?)?; let k = hidden_states.matmul(&self.k_proj.transpose(-1, -2)?)?; let v = hidden_states.matmul(&self.v_proj.transpose(-1, -2)?)?; // Reshape for multi-head attention let q = q.reshape(&[batch_size, seq_len, self.config.num_attention_heads, head_dim])? .transpose(1, 2)?; let k = k.reshape(&[batch_size, seq_len, self.config.num_key_value_heads, head_dim])? .transpose(1, 2)?; let v = v.reshape(&[batch_size, seq_len, self.config.num_key_value_heads, head_dim])? .transpose(1, 2)?; // Apply rotary position embedding let (q, k) = self.rotary_emb.forward(&q, &k, position_ids)?; // Grouped-query attention: repeat K, V if needed let (k, v) = if self.config.num_key_value_heads < self.config.num_attention_heads { let repeat_factor = self.config.num_attention_heads / self.config.num_key_value_heads; let k = self.repeat_kv(&k, repeat_factor)?; let v = self.repeat_kv(&v, repeat_factor)?; (k, v) } else { (k, v) }; // Scaled dot-product attention let attn_output = self.scaled_dot_product_attention(&q, &k, &v)?; // Reshape and project output let attn_output = attn_output.transpose(1, 2)? .reshape(&[batch_size, seq_len, self.config.hidden_size])?; attn_output.matmul(&self.o_proj.transpose(-1, -2)?) } /// Repeat key-value tensors for grouped-query attention fn repeat_kv(&self, tensor: &Tensor, repeat_factor: usize) -> Result { if repeat_factor == 1 { return Ok(tensor.clone()); } // Efficient key-value repetition for grouped-query attention // Repeat key/value tensors to match the number of query heads let [batch_size, seq_len, n_kv_heads, head_dim] = tensor.shape(); // Calculate repetition factor let n_heads = self.n_heads; let rep_factor = n_heads / n_kv_heads; if rep_factor == 1 { // No repetition needed return Ok(tensor.clone()); } // Reshape to [batch, seq, n_kv_heads, 1, head_dim] let reshaped = tensor.reshape(&[batch_size, seq_len, n_kv_heads, 1, head_dim])?; // Repeat along the new dimension: [batch, seq, n_kv_heads, rep_factor, head_dim] let repeated = reshaped.repeat(&[1, 1, 1, rep_factor, 1])?; // Reshape to [batch, seq, n_heads, head_dim] let output_shape = [batch_size, seq_len, n_heads, head_dim]; repeated.reshape(&output_shape) .map_err(|e| crate::TransformerError::ArchitectureError( format!("Failed to repeat key-value tensor: {}", e) )) } /// Scaled dot-product attention fn scaled_dot_product_attention(&self, q: &Tensor, k: &Tensor, v: &Tensor) -> Result { let head_dim = self.config.hidden_size / self.config.num_attention_heads; let scale = 1.0 / (head_dim as f32).sqrt(); // QK^T / sqrt(d_k) let scores = q.matmul(&k.transpose(-1, -2)?)? * scale; // Apply causal mask let seq_len = scores.shape()[2]; let causal_mask = self.create_causal_mask(seq_len, scores.device())?; let masked_scores = scores + causal_mask; // Softmax let attn_weights = masked_scores.softmax(-1)?; // Apply to values attn_weights.matmul(v) } /// Create causal attention mask fn create_causal_mask(&self, seq_len: usize, device: &Device) -> Result { // Create lower triangular mask: mask[i,j] = 0 if i >= j, -inf if i < j // This prevents attention to future positions // Create a matrix filled with negative infinity let mut mask_data = vec![-f32::INFINITY; seq_len * seq_len]; // Set lower triangle (including diagonal) to 0 for i in 0..seq_len { for j in 0..=i { // j <= i for lower triangle mask_data[i * seq_len + j] = 0.0; } } // Create tensor from the mask data let mask = Tensor::from_vec( mask_data, &[1, 1, seq_len, seq_len], DType::F32, device, ).map_err(|e| crate::TransformerError::ArchitectureError( format!("Failed to create causal mask: {}", e) ))?; Ok(mask) } } #[cfg(test)] mod tests { use super::*; use rtx_tensor::Device; #[test] fn test_llama_config_validation() { let mut config = LLaMAConfig::default(); assert!(config.validate().is_ok()); // Test invalid configuration config.hidden_size = 100; config.num_attention_heads = 7; // 100 is not divisible by 7 assert!(config.validate().is_err()); } #[test] fn test_llama_config_presets() { let llama_7b = LLaMAConfig::llama_7b(); assert_eq!(llama_7b.hidden_size, 4096); assert_eq!(llama_7b.num_hidden_layers, 32); let llama2_70b = LLaMAConfig::llama2_70b(); assert_eq!(llama2_70b.num_key_value_heads, 8); // Grouped-query attention let code_llama = LLaMAConfig::code_llama(); assert_eq!(code_llama.max_position_embeddings, 16_384); } #[test] fn test_rope_creation() { let device = Device::Cpu; let rope = RotaryPositionEmbedding::new(128, 2048, 10000.0, &device); assert!(rope.is_ok()); let rope = rope.unwrap(); assert_eq!(rope.dim, 128); assert_eq!(rope.max_seq_len, 2048); } #[test] fn test_swiglu_creation() { let device = Device::Cpu; let swiglu = SwiGLU::new(4096, 11008, false, &device); assert!(swiglu.is_ok()); let swiglu = swiglu.unwrap(); assert_eq!(swiglu.hidden_size, 4096); assert_eq!(swiglu.intermediate_size, 11008); } #[test] fn test_grouped_query_attention() { let mut config = LLaMAConfig::llama2_70b(); assert_eq!(config.num_attention_heads, 64); assert_eq!(config.num_key_value_heads, 8); assert!(config.validate().is_ok()); // Test invalid GQA configuration config.num_key_value_heads = 65; // Greater than attention heads assert!(config.validate().is_err()); } }