Skip to main content

entrenar/transformer/
model.rs

1//! Complete transformer model module
2//!
3//! This module provides the full transformer model for language modeling.
4
5use crate::autograd::matmul_nt;
6use crate::error::{Error, Result};
7use crate::Tensor;
8use provable_contracts_macros::{ensures, requires};
9use std::collections::HashMap;
10use std::path::Path;
11
12use super::block::TransformerBlock;
13use super::config::TransformerConfig;
14use super::embedding::Embedding;
15use super::norm::RMSNorm;
16use super::weights::{load_safetensors_weights, validate_weights, Architecture};
17
18/// Complete transformer model
19pub struct Transformer {
20    /// Configuration
21    pub config: TransformerConfig,
22    /// Token embedding layer
23    pub embed_tokens: Embedding,
24    /// Transformer layers
25    pub layers: Vec<TransformerBlock>,
26    /// Final layer normalization
27    pub norm: RMSNorm,
28    /// Language model head (tied to embeddings or separate)
29    pub lm_head: Option<Tensor>,
30}
31
32impl Transformer {
33    /// Create new transformer with initialized weights
34    pub fn new(config: &TransformerConfig) -> Self {
35        let layers: Vec<TransformerBlock> =
36            (0..config.num_hidden_layers).map(|i| TransformerBlock::new(config, i)).collect();
37
38        Self {
39            config: config.clone(),
40            embed_tokens: Embedding::new(config.vocab_size, config.hidden_size),
41            layers,
42            norm: RMSNorm::new(config.hidden_size, config.rms_norm_eps),
43            lm_head: None, // Use tied embeddings by default
44        }
45    }
46
47    /// Create transformer from parameter map
48    ///
49    /// Expected parameter names (following HuggingFace LLaMA convention):
50    /// - `model.embed_tokens.weight`
51    /// - `model.layers.{i}.*`
52    /// - `model.norm.weight`
53    /// - `lm_head.weight` (optional, uses tied embeddings if not present)
54    pub fn from_params(
55        config: &TransformerConfig,
56        params: &HashMap<String, Tensor>,
57    ) -> Option<Self> {
58        let embed_tokens = Embedding::from_params(
59            params,
60            "model.embed_tokens.weight",
61            config.vocab_size,
62            config.hidden_size,
63        )?;
64
65        let layers: Option<Vec<TransformerBlock>> = (0..config.num_hidden_layers)
66            .map(|i| TransformerBlock::from_params(config, params, i))
67            .collect();
68        let layers = layers?;
69
70        let norm =
71            RMSNorm::from_params(params, "model.norm", config.rms_norm_eps, config.hidden_size)?;
72
73        // PMAT-329: Validate lm_head shape if present
74        let lm_head = if let Some(tensor) = params.get("lm_head.weight") {
75            let expected = config.hidden_size * config.vocab_size;
76            if tensor.len() != expected {
77                eprintln!(
78                    "[PMAT-329] lm_head.weight: shape mismatch — got {} elements, expected {expected} ({hidden}x{vocab})",
79                    tensor.len(),
80                    hidden = config.hidden_size,
81                    vocab = config.vocab_size,
82                );
83                return None;
84            }
85            Some(tensor.clone())
86        } else {
87            None
88        };
89
90        Some(Self { config: config.clone(), embed_tokens, layers, norm, lm_head })
91    }
92
93    /// Load transformer from SafeTensors file(s)
94    ///
95    /// Reads SafeTensors weights from `model_path`, converts BF16/F16 to F32,
96    /// validates shapes against `config`, checks for NaN/Inf, and constructs
97    /// the complete `Transformer`.
98    ///
99    /// # Arguments
100    ///
101    /// * `model_path` - Path to model directory or single SafeTensors file
102    /// * `config` - Transformer configuration specifying model dimensions
103    ///
104    /// # Errors
105    ///
106    /// Returns `Error::ConfigError` if:
107    /// - No SafeTensors files found
108    /// - Required weight tensors are missing
109    /// - Weight shapes do not match config dimensions
110    /// - Weights contain NaN or Inf values
111    /// - Layer count does not match config
112    pub fn from_safetensors(
113        model_path: impl AsRef<Path>,
114        config: &TransformerConfig,
115    ) -> Result<Self> {
116        let model_path = model_path.as_ref();
117
118        // Load and convert all weights from SafeTensors files
119        let weights = load_safetensors_weights(model_path, Architecture::Auto)?;
120
121        // Structural validation: all required keys present
122        validate_weights(&weights, config.num_hidden_layers)?;
123
124        // Shape validation against config dimensions
125        Self::validate_weight_shapes(&weights, config)?;
126
127        // NaN/Inf validation
128        Self::validate_weight_values(&weights)?;
129
130        // Build transformer from validated weights
131        Self::from_params(config, &weights).ok_or_else(|| {
132            Error::ConfigError(
133                "Failed to construct Transformer from loaded weights \
134                 (internal from_params returned None after validation passed)"
135                    .into(),
136            )
137        })
138    }
139
140    /// Load transformer from APR file (.apr format)
141    ///
142    /// Reads tensor data from an APR binary file, dequantizing from any stored
143    /// dtype (F16, Q4K, etc.) to F32. Uses the same validation pipeline as
144    /// `from_safetensors`: structural, shape, and NaN/Inf checks.
145    ///
146    /// # Arguments
147    /// * `apr_path` - Path to the .apr model file
148    /// * `config` - Transformer configuration (typically read from APR metadata)
149    ///
150    /// # Errors
151    /// Returns `Error::ConfigError` if tensors are missing, shapes mismatch, or
152    /// weights contain NaN/Inf values.
153    pub fn from_apr(apr_path: impl AsRef<Path>, config: &TransformerConfig) -> Result<Self> {
154        use aprender::serialization::apr::AprReader;
155        use rayon::prelude::*;
156
157        // PMAT-FINETUNE-CONSTRUCT: phase-timed construction (was a 15+min unprofiled wall
158        // before any training step). Set APR_FROM_APR_TIMING=1 to see per-phase ms.
159        let timing = std::env::var("APR_FROM_APR_TIMING").is_ok();
160        let t0 = std::time::Instant::now();
161
162        let apr_path = apr_path.as_ref();
163        let reader = AprReader::open(apr_path).map_err(|e| {
164            Error::ConfigError(format!("Failed to open APR file '{}': {e}", apr_path.display()))
165        })?;
166        if timing {
167            eprintln!(
168                "[from_apr-timing] open+header: {:?} ({} tensors)",
169                t0.elapsed(),
170                reader.tensors.len()
171            );
172        }
173
174        // Build weight map from APR tensors — map GGUF names to HF convention (PMAT-489)
175        let is_gguf_names = reader.tensors.iter().any(|t| t.name == "token_embd.weight");
176        if is_gguf_names {
177            eprintln!(
178                "[PMAT-489] Detected GGUF tensor names in APR file, mapping to HF convention"
179            );
180        }
181        // PMAT-FINETUNE-CONSTRUCT: dequant every tensor to F32 IN PARALLEL. Previously this
182        // was a serial loop over hundreds of tensors (each a Q4K/Q6K/F16 dequant of up to
183        // ~vocab*hidden f32s) — the dominant term in the construction wall. Rayon-parallel
184        // dequant scales with cores; the AprReader.read_tensor_as_f32 path is read-only
185        // (offsets into the mmap'd/owned byte buffer) so it is data-parallel-safe.
186        let t_dequant = std::time::Instant::now();
187        let dequanted: Vec<(String, Vec<f32>)> = reader
188            .tensors
189            .par_iter()
190            .map(|desc| {
191                let data = reader.read_tensor_as_f32(&desc.name).map_err(|e| {
192                    Error::ConfigError(format!("Failed to read tensor '{}': {e}", desc.name))
193                })?;
194                let mapped_name = if is_gguf_names {
195                    super::weights::mapping::map_weight_name(
196                        &desc.name,
197                        super::weights::Architecture::Gguf,
198                    )
199                } else {
200                    desc.name.clone()
201                };
202                Ok::<_, Error>((mapped_name, data))
203            })
204            .collect::<Result<Vec<_>>>()?;
205        let mut weights = HashMap::with_capacity(dequanted.len());
206        for (name, data) in dequanted {
207            weights.insert(name, Tensor::from_vec(data, false));
208        }
209        if timing {
210            eprintln!("[from_apr-timing] parallel dequant->F32: {:?}", t_dequant.elapsed());
211        }
212
213        // #2441: honor the tied-word-embedding placeholder. An `.apr` written from
214        // a `tie_word_embeddings=true` checkpoint stores `lm_head.weight` as a
215        // shape-only, ZERO-BYTE descriptor — the real matrix lives in
216        // `model.embed_tokens.weight`. Dequantizing that descriptor yields an empty
217        // f32 vec, which `validate_weight_shapes` then rejects with
218        // "Shape mismatch for 'lm_head.weight': expected N elements, got 0".
219        // `from_params` already implements the tie (`lm_head: None` falls back to
220        // `embed_tokens.weight` in both `forward` and `lm_head_weight_slice`), so
221        // resolving the tie here means dropping the placeholder.
222        Self::resolve_tied_lm_head(&mut weights);
223
224        // Same validation pipeline as from_safetensors
225        let t_val = std::time::Instant::now();
226        validate_weights(&weights, config.num_hidden_layers)?;
227        Self::validate_weight_shapes(&weights, config)?;
228        Self::validate_weight_values(&weights)?;
229        if timing {
230            eprintln!("[from_apr-timing] validate (struct+shape+nan/inf): {:?}", t_val.elapsed());
231        }
232
233        let t_params = std::time::Instant::now();
234        let model = Self::from_params(config, &weights).ok_or_else(|| {
235            Error::ConfigError(
236                "Failed to construct Transformer from APR weights \
237                 (from_params returned None after validation passed)"
238                    .into(),
239            )
240        });
241        if timing {
242            eprintln!(
243                "[from_apr-timing] from_params(build tensors): {:?} | TOTAL {:?}",
244                t_params.elapsed(),
245                t0.elapsed()
246            );
247        }
248        model
249    }
250
251    /// Resolve a tied-word-embedding `lm_head.weight` placeholder (#2441, #2309).
252    ///
253    /// APR files written from a `tie_word_embeddings=true` checkpoint record
254    /// `lm_head.weight` with its full shape but ZERO bytes of data, because the
255    /// matrix is the same one stored under `model.embed_tokens.weight`. A loader
256    /// that takes the descriptor at face value materializes an empty tensor and
257    /// then fails shape validation.
258    ///
259    /// Dropping the placeholder is the tie: `from_params` maps a missing
260    /// `lm_head.weight` to `lm_head: None`, and every consumer of the output
261    /// projection (`forward`, `forward_hidden`, `lm_head_weight`,
262    /// `lm_head_weight_slice`) already falls back to `embed_tokens.weight`.
263    ///
264    /// Returns `true` when a placeholder was resolved. An empty `lm_head.weight`
265    /// with no usable `model.embed_tokens.weight` is left in place so the shape
266    /// validator still reports the genuinely broken file.
267    fn resolve_tied_lm_head(weights: &mut HashMap<String, Tensor>) -> bool {
268        if weights.get("lm_head.weight").is_none_or(|t| t.len() != 0) {
269            return false;
270        }
271        if weights.get("model.embed_tokens.weight").is_none_or(|t| t.len() == 0) {
272            return false;
273        }
274        weights.remove("lm_head.weight");
275        eprintln!(
276            "[#2441] lm_head.weight is a 0-byte tied-embedding placeholder — \
277             tying the output projection to model.embed_tokens.weight"
278        );
279        true
280    }
281
282    /// Validate that all weight tensor shapes match the config dimensions
283    fn validate_weight_shapes(
284        weights: &HashMap<String, Tensor>,
285        config: &TransformerConfig,
286    ) -> Result<()> {
287        let hidden = config.hidden_size;
288        let q_dim = config.q_dim();
289        let kv_hidden = config.num_kv_heads * config.head_dim();
290        let intermediate = config.intermediate_size;
291        let vocab = config.vocab_size;
292
293        // Helper closure for shape checking
294        let check = |name: &str, expected: usize| -> Result<()> {
295            if let Some(tensor) = weights.get(name) {
296                if tensor.len() != expected {
297                    return Err(Error::ConfigError(format!(
298                        "Shape mismatch for '{name}': expected {expected} elements, got {}",
299                        tensor.len()
300                    )));
301                }
302            }
303            // Missing keys are caught by validate_weights
304            Ok(())
305        };
306
307        // Global weights
308        check("model.embed_tokens.weight", vocab * hidden)?;
309        check("model.norm.weight", hidden)?;
310
311        // Optional lm_head
312        if weights.contains_key("lm_head.weight") {
313            check("lm_head.weight", vocab * hidden)?;
314        }
315
316        // Per-layer weights
317        for i in 0..config.num_hidden_layers {
318            let p = format!("model.layers.{i}");
319
320            // Layer norms
321            check(&format!("{p}.input_layernorm.weight"), hidden)?;
322            check(&format!("{p}.post_attention_layernorm.weight"), hidden)?;
323
324            // Attention projections: Q/O use q_dim, K/V use kv_hidden
325            check(&format!("{p}.self_attn.q_proj.weight"), q_dim * hidden)?;
326            check(&format!("{p}.self_attn.k_proj.weight"), kv_hidden * hidden)?;
327            check(&format!("{p}.self_attn.v_proj.weight"), kv_hidden * hidden)?;
328            check(&format!("{p}.self_attn.o_proj.weight"), hidden * q_dim)?;
329
330            // Optional attention biases (Qwen2 etc.)
331            check(&format!("{p}.self_attn.q_proj.bias"), q_dim)?;
332            check(&format!("{p}.self_attn.k_proj.bias"), kv_hidden)?;
333            check(&format!("{p}.self_attn.v_proj.bias"), kv_hidden)?;
334
335            // MLP projections
336            check(&format!("{p}.mlp.gate_proj.weight"), hidden * intermediate)?;
337            check(&format!("{p}.mlp.up_proj.weight"), hidden * intermediate)?;
338            check(&format!("{p}.mlp.down_proj.weight"), intermediate * hidden)?;
339        }
340
341        Ok(())
342    }
343
344    /// Validate that no weight tensors contain NaN or Inf values
345    fn validate_weight_values(weights: &HashMap<String, Tensor>) -> Result<()> {
346        for (name, tensor) in weights {
347            let data = tensor.data();
348            for (i, &val) in data.iter().enumerate() {
349                if val.is_nan() {
350                    return Err(Error::ConfigError(format!(
351                        "NaN detected in weight '{name}' at index {i}"
352                    )));
353                }
354                if val.is_infinite() {
355                    return Err(Error::ConfigError(format!(
356                        "Inf detected in weight '{name}' at index {i}"
357                    )));
358                }
359            }
360        }
361        Ok(())
362    }
363
364    /// Forward pass for language modeling
365    ///
366    /// # Arguments
367    /// * `token_ids` - Input token IDs
368    ///
369    /// # Returns
370    /// Logits tensor (seq_len * vocab_size, flattened)
371    #[requires(!token_ids.is_empty())]
372    #[ensures(ret.len() == token_ids.len() * self.config.vocab_size)]
373    pub fn forward(&self, token_ids: &[u32]) -> Tensor {
374        contract_pre_embedding_lookup!(token_ids);
375        let seq_len = token_ids.len();
376        let hidden_size = self.config.hidden_size;
377
378        // Embed tokens
379        let mut hidden = self.embed_tokens.forward(token_ids);
380
381        // Pass through transformer layers
382        for layer in &self.layers {
383            hidden = layer.forward(&hidden, seq_len);
384        }
385
386        // Final normalization
387        let normalized = self.norm.forward_batched(&hidden, seq_len, hidden_size);
388
389        // Language model head
390        let lm_weight = self.lm_head.as_ref().unwrap_or(&self.embed_tokens.weight);
391
392        // lm_head / tied embed_tokens is [vocab_size, hidden_size] in HF (ENT-269)
393        let result =
394            matmul_nt(&normalized, lm_weight, seq_len, hidden_size, self.config.vocab_size);
395        contract_post_embedding_lookup!(result.data().as_slice().unwrap_or(&[]));
396        result
397    }
398
399    /// Forward pass returning hidden states (before lm_head)
400    ///
401    /// # Arguments
402    /// * `token_ids` - Input token IDs
403    ///
404    /// # Returns
405    /// Hidden states tensor (seq_len * hidden_size, flattened)
406    #[requires(!token_ids.is_empty())]
407    #[ensures(ret.len() == token_ids.len() * self.config.hidden_size)]
408    pub fn forward_hidden(&self, token_ids: &[u32]) -> Tensor {
409        contract_pre_embedding_lookup!(token_ids);
410        let seq_len = token_ids.len();
411        let hidden_size = self.config.hidden_size;
412
413        // Embed tokens
414        let mut hidden = self.embed_tokens.forward(token_ids);
415
416        // Pass through transformer layers
417        for layer in &self.layers {
418            hidden = layer.forward(&hidden, seq_len);
419        }
420
421        // Final normalization
422        let result = self.norm.forward_batched(&hidden, seq_len, hidden_size);
423        contract_post_embedding_lookup!(result.data().as_slice().unwrap_or(&[]));
424        result
425    }
426
427    /// Forward pass returning hidden states with LoRA adjusts (KAIZEN-011)
428    ///
429    /// Like `forward_hidden` but applies LoRA adapters to Q/V projections in
430    /// each transformer layer's attention. Enables non-CUDA LoRA training by
431    /// putting LoRA parameters into the autograd graph.
432    ///
433    /// # Arguments
434    /// * `token_ids` - Input token IDs
435    /// * `lora_layers` - LoRA layers in [Q_0, V_0, Q_1, V_1, ...] order
436    ///
437    /// # Returns
438    /// Hidden states tensor (seq_len * hidden_size, flattened)
439    pub fn forward_hidden_with_lora(
440        &self,
441        token_ids: &[u32],
442        lora_layers: &[crate::lora::LoRALayer],
443    ) -> Tensor {
444        contract_pre_embedding_lookup!(token_ids);
445        let seq_len = token_ids.len();
446        let hidden_size = self.config.hidden_size;
447
448        let mut hidden = self.embed_tokens.forward(token_ids);
449
450        for (layer_idx, layer) in self.layers.iter().enumerate() {
451            let norm1 = layer.input_norm.forward_batched(&hidden, seq_len, hidden_size);
452
453            // KAIZEN-011: Apply LoRA to attention Q/V projections
454            let q_idx = layer_idx * 2;
455            let v_idx = layer_idx * 2 + 1;
456            let attn_out = if v_idx < lora_layers.len() {
457                layer.self_attn.forward_with_lora(
458                    &norm1,
459                    seq_len,
460                    lora_layers[q_idx].lora_a(),
461                    lora_layers[q_idx].lora_b(),
462                    lora_layers[v_idx].lora_a(),
463                    lora_layers[v_idx].lora_b(),
464                    lora_layers[q_idx].rank(),
465                    lora_layers[q_idx].scale(),
466                )
467            } else {
468                layer.self_attn.forward(&norm1, seq_len)
469            };
470
471            let residual = crate::autograd::add(&hidden, &attn_out);
472            let norm2 = layer.post_attn_norm.forward_batched(&residual, seq_len, hidden_size);
473            let ffn_out = layer.ffn.forward(&norm2, seq_len);
474            hidden = crate::autograd::add(&residual, &ffn_out);
475        }
476
477        let result = self.norm.forward_batched(&hidden, seq_len, hidden_size);
478        contract_post_embedding_lookup!(result.data().as_slice().unwrap_or(&[]));
479        result
480    }
481
482    /// Forward pass with LoRA adapters (ENT-LoRA-001)
483    ///
484    /// Like `forward` but applies LoRA adapters to Q/V projections.
485    /// Returns full logits (seq_len * vocab_size).
486    ///
487    /// # Arguments
488    /// * `token_ids` - Input token IDs
489    /// * `lora_layers` - LoRA layers in [Q_0, V_0, Q_1, V_1, ...] order
490    pub fn forward_with_lora(
491        &self,
492        token_ids: &[u32],
493        lora_layers: &[crate::lora::LoRALayer],
494    ) -> Tensor {
495        contract_pre_embedding_lookup!(token_ids);
496        let seq_len = token_ids.len();
497        let hidden_size = self.config.hidden_size;
498
499        let hidden = self.forward_hidden_with_lora(token_ids, lora_layers);
500        let lm_weight = self.lm_head.as_ref().unwrap_or(&self.embed_tokens.weight);
501        let result = matmul_nt(&hidden, lm_weight, seq_len, hidden_size, self.config.vocab_size);
502        contract_post_embedding_lookup!(result.data().as_slice().unwrap_or(&[]));
503        result
504    }
505
506    /// Get the last token's logits (for generation)
507    pub fn forward_last(&self, token_ids: &[u32]) -> Tensor {
508        contract_pre_embedding_lookup!(token_ids);
509        let logits = self.forward(token_ids);
510        let seq_len = token_ids.len();
511        let vocab_size = self.config.vocab_size;
512
513        // Extract last position
514        let start = (seq_len - 1) * vocab_size;
515        let end = start + vocab_size;
516        let last_logits: Vec<f32> =
517            logits.data().as_slice().expect("logits must be contiguous")[start..end].to_vec();
518
519        let result = Tensor::from_vec(last_logits, logits.requires_grad());
520        contract_post_embedding_lookup!(result.data().as_slice().unwrap_or(&[]));
521        result
522    }
523
524    /// Get all parameters as a vector
525    pub fn parameters(&self) -> Vec<&Tensor> {
526        let mut params = vec![&self.embed_tokens.weight, &self.norm.weight];
527        for layer in &self.layers {
528            params.extend(layer.parameters());
529        }
530        if let Some(lm_head) = &self.lm_head {
531            params.push(lm_head);
532        }
533        params
534    }
535
536    /// Get all parameters as mutable references for optimizer
537    pub fn parameters_mut(&mut self) -> Vec<&mut Tensor> {
538        let mut params: Vec<&mut Tensor> = Vec::new();
539        params.push(&mut self.embed_tokens.weight);
540        params.push(&mut self.norm.weight);
541        for layer in &mut self.layers {
542            params.extend(layer.parameters_mut());
543        }
544        if let Some(lm_head) = &mut self.lm_head {
545            params.push(lm_head);
546        }
547        params
548    }
549
550    /// Get configuration
551    pub fn config(&self) -> &TransformerConfig {
552        &self.config
553    }
554
555    /// Embed a single token, returning hidden_size floats.
556    pub fn embed_token(&self, token_id: u32) -> Vec<f32> {
557        let w = self.embed_tokens.weight.data();
558        let data = w.as_slice().expect("contiguous embedding");
559        let h = self.config.hidden_size;
560        let offset = (token_id as usize) * h;
561        data[offset..offset + h].to_vec()
562    }
563
564    /// Get the output norm weight as a slice.
565    pub fn output_norm_weight_slice(&self) -> &[f32] {
566        self.norm.weight.data().as_slice().expect("contiguous norm weight")
567    }
568
569    /// Get the lm_head weight as a slice (vocab_size × hidden_size, row-major).
570    pub fn lm_head_weight_slice(&self) -> &[f32] {
571        let w = self.lm_head.as_ref().unwrap_or(&self.embed_tokens.weight);
572        w.data().as_slice().expect("contiguous lm_head")
573    }
574
575    /// Get the language model head weight tensor.
576    ///
577    /// Returns the dedicated `lm_head` weight if present, otherwise falls back
578    /// to tied embedding weights.
579    pub fn lm_head_weight(&self) -> &Tensor {
580        self.lm_head.as_ref().unwrap_or(&self.embed_tokens.weight)
581    }
582
583    /// Get named parameters for checkpoint serialization.
584    ///
585    /// Returns (name, tensor) pairs matching HuggingFace weight conventions.
586    /// This handles variable parameter counts (e.g., models with/without attention biases)
587    /// correctly, unlike the hardcoded 9-params-per-layer assumption.
588    pub fn named_parameters(&self) -> Vec<(String, &Tensor)> {
589        let mut params = vec![
590            ("model.embed_tokens.weight".to_string(), &self.embed_tokens.weight),
591            ("model.norm.weight".to_string(), &self.norm.weight),
592        ];
593        for layer in &self.layers {
594            params.extend(layer.named_parameters());
595        }
596        if let Some(ref lm_head) = self.lm_head {
597            params.push(("lm_head.weight".to_string(), lm_head));
598        }
599        params
600    }
601
602    /// ENT-282: Set a named parameter by name (for delta checkpoint overlay).
603    ///
604    /// Returns true if the parameter was found and set.
605    pub fn set_named_parameter(&mut self, name: &str, value: Tensor) -> bool {
606        if name == "model.embed_tokens.weight" {
607            self.embed_tokens.weight = value;
608            return true;
609        }
610        if name == "model.norm.weight" {
611            self.norm.weight = value;
612            return true;
613        }
614        if name == "lm_head.weight" {
615            self.lm_head = Some(value);
616            return true;
617        }
618        // Per-layer parameters: model.layers.{idx}.{suffix}
619        if let Some(rest) = name.strip_prefix("model.layers.") {
620            if let Some(dot_pos) = rest.find('.') {
621                if let Ok(idx) = rest[..dot_pos].parse::<usize>() {
622                    if idx < self.layers.len() {
623                        let suffix = &rest[dot_pos + 1..];
624                        return self.layers[idx].set_named_parameter(suffix, value);
625                    }
626                }
627            }
628        }
629        false
630    }
631}
632
633#[cfg(test)]
634mod tests {
635    use super::*;
636
637    #[test]
638    fn test_transformer_tiny_forward() {
639        let config = TransformerConfig::tiny();
640        let transformer = Transformer::new(&config);
641        let tokens = vec![1, 2, 3];
642        let logits = transformer.forward(&tokens);
643        assert_eq!(logits.len(), 3 * config.vocab_size);
644    }
645
646    /// FALSIFY-APR-PRETRAIN-ARCH-004 (smoke level) — GQA-7:1 forward pass
647    /// runs without panic and produces finite output of correct shape.
648    ///
649    /// Per `apr-pretrain-arch-polymorphic-v1` (PR #1473), the §49 fine-tune
650    /// path uses Qwen2.5-Coder-0.5B's GQA-7:1 ratio (kv_heads=2,
651    /// query_heads=14). The Llama370M codepath only exercised GQA-4:1.
652    /// This test pins that the existing attention kernel handles the new
653    /// ratio without per-ratio specialization.
654    ///
655    /// Tiny shape (hidden=112=14*8, head_dim=8) keeps the test under 1ms.
656    /// Full numerical-parity vs GQA-1:1 reference (cosine ≥ 0.9999) is a
657    /// FUNCTIONAL-level discharge, not algorithm-level. PARTIAL_ALGORITHM_LEVEL
658    /// only requires that the kernel COMPILES and PRODUCES finite output for
659    /// the new ratio — both proven here.
660    ///
661    /// Spec: SPEC-SHIP-TWO-001 §50.4 step 5e.
662    #[test]
663    fn falsify_apr_pretrain_arch_004_gqa_7_1_forward_pass_smoke() {
664        // Tiny GQA-7:1 config: kv_heads=2, num_attention_heads=14, head_dim=8.
665        // hidden = num_attention_heads * head_dim = 14 * 8 = 112.
666        let config = TransformerConfig {
667            hidden_size: 112,
668            num_attention_heads: 14,
669            num_kv_heads: 2,
670            intermediate_size: 64,
671            num_hidden_layers: 1,
672            vocab_size: 256,
673            max_position_embeddings: 512,
674            rms_norm_eps: 1e-6,
675            rope_theta: 1_000_000.0, // Qwen2 ROPE convention
676            use_bias: true,          // Qwen2 quirk
677            head_dim_override: None, // 112 / 14 = 8, no override needed
678            architecture: crate::transformer::config::ModelArchitecture::Decoder,
679            hf_architecture: None,
680            hf_model_type: None,
681            tie_word_embeddings: true, // Qwen2.5-0.5B convention
682        };
683
684        // Drift-prevention: verify GQA-7:1 ratio holds at construction time.
685        // If a future refactor flips num_attention_heads or num_kv_heads such
686        // that the ratio is no longer 7, this test catches it before the
687        // forward pass even runs.
688        assert_eq!(
689            config.num_attention_heads / config.num_kv_heads,
690            7,
691            "GQA-7:1 ratio must be 14/2=7 (Qwen2.5-0.5B canonical)"
692        );
693
694        let transformer = Transformer::new(&config);
695        let tokens = vec![1u32, 2, 3, 4]; // seq_len=4 short prefix
696        let logits = transformer.forward(&tokens);
697
698        // Shape invariant: forward returns seq_len * vocab_size logits.
699        assert_eq!(
700            logits.len(),
701            4 * config.vocab_size,
702            "GQA-7:1 forward must return seq_len * vocab_size logits"
703        );
704
705        // Numerical invariant: all logits finite (no NaN, no Inf).
706        // The §24 retrospective showed silent NaN propagation through GQA can
707        // produce loss=NaN that the divergence guard catches LATE (after
708        // multiple steps). FALSIFY-004's smoke level catches it at the
709        // first forward pass, before any optimizer state corrupts.
710        assert!(
711            logits.data().iter().all(|&v| v.is_finite()),
712            "GQA-7:1 forward must produce all-finite logits — silent NaN \
713             would corrupt the §49 fine-tune trajectory before FALSIFY-006 \
714             (init_loss < 6.0) could measure it"
715        );
716    }
717
718    #[test]
719    fn test_transformer_tiny_forward_last() {
720        let config = TransformerConfig::tiny();
721        let transformer = Transformer::new(&config);
722        let tokens = vec![1, 2, 3];
723        let logits = transformer.forward_last(&tokens);
724        assert_eq!(logits.len(), config.vocab_size);
725    }
726
727    #[test]
728    fn test_transformer_parameters() {
729        let config = TransformerConfig::tiny();
730        let transformer = Transformer::new(&config);
731        let params = transformer.parameters();
732        // embed_tokens + norm + (layers * (input_norm + post_attn_norm + 4 attn weights + 3 ffn weights))
733        // = 2 + 2 * (2 + 4 + 3) = 2 + 2 * 9 = 20
734        assert_eq!(params.len(), 20);
735    }
736
737    #[test]
738    fn test_transformer_config_accessor() {
739        let config = TransformerConfig::tiny();
740        let transformer = Transformer::new(&config);
741        assert_eq!(transformer.config().hidden_size, config.hidden_size);
742        assert_eq!(transformer.config().vocab_size, config.vocab_size);
743    }
744
745    #[test]
746    fn test_transformer_single_token() {
747        let config = TransformerConfig::tiny();
748        let transformer = Transformer::new(&config);
749        let tokens = vec![42];
750        let logits = transformer.forward(&tokens);
751        assert_eq!(logits.len(), config.vocab_size);
752    }
753
754    #[test]
755    fn test_output_finite_values() {
756        let config = TransformerConfig::tiny();
757        let transformer = Transformer::new(&config);
758        let tokens = vec![1, 2, 3, 4, 5];
759        let logits = transformer.forward(&tokens);
760        // All outputs should be finite (no NaN or Inf)
761        assert!(logits.data().iter().all(|&v| v.is_finite()));
762    }
763
764    #[test]
765    fn test_transformer_empty_lm_head_uses_tied_weights() {
766        let config = TransformerConfig::tiny();
767        let transformer = Transformer::new(&config);
768        // Default transformer should have no separate lm_head
769        assert!(transformer.lm_head.is_none());
770        // But should still produce valid logits
771        let tokens = vec![1, 2];
772        let logits = transformer.forward(&tokens);
773        assert_eq!(logits.len(), 2 * config.vocab_size);
774    }
775
776    /// FALSIFY-FINETUNE-CONSTRUCT-001 — `Transformer::from_apr` parallel F32
777    /// dequant loads EVERY tensor correctly (no dropped/garbled weight from the
778    /// rayon `par_iter` collect that replaced the serial dequant loop).
779    ///
780    /// PMAT-FINETUNE-CONSTRUCT: the construction path was changed from a serial
781    /// `for desc in &reader.tensors { read_tensor_as_f32 }` loop to a
782    /// `par_iter().map(read_tensor_as_f32).collect()`. `read_tensor_as_f32` is a
783    /// pure read into an owned/mmap'd byte buffer (no shared mutable state), so
784    /// the parallel collect MUST yield the identical (name -> f32 data) map.
785    ///
786    /// Falsifier: build a tiny but complete Qwen2-shaped APR with KNOWN F32 weight
787    /// values, load it through `from_apr`, and assert the round-tripped weights are
788    /// bit-exact and the model produces finite logits. If the parallel collect ever
789    /// drops a tensor, mis-maps a name, or corrupts data, validate_weights /
790    /// validate_weight_values / the value assertions below go RED.
791    #[test]
792    fn falsify_finetune_construct_001_from_apr_parallel_dequant_loads_all_weights() {
793        use aprender::serialization::apr::AprWriter;
794
795        let config = TransformerConfig::tiny(); // hidden=64, heads=2, kv=2, inter=256, layers=2, vocab=1000
796        let hidden = config.hidden_size;
797        let q_dim = config.q_dim();
798        let kv_hidden = config.num_kv_heads * config.head_dim();
799        let inter = config.intermediate_size;
800        let vocab = config.vocab_size;
801
802        // Distinct, exactly-representable f32 fill per tensor so a swap/corruption
803        // (not just a drop) is also caught. 0.5 is bit-exact under any dequant.
804        let mut w = AprWriter::new();
805        w.add_tensor_f32(
806            "model.embed_tokens.weight",
807            vec![vocab, hidden],
808            &vec![0.5; vocab * hidden],
809        );
810        w.add_tensor_f32("model.norm.weight", vec![hidden], &vec![1.0; hidden]);
811        w.add_tensor_f32("lm_head.weight", vec![vocab, hidden], &vec![0.25; vocab * hidden]);
812        for i in 0..config.num_hidden_layers {
813            let p = format!("model.layers.{i}");
814            w.add_tensor_f32(
815                &format!("{p}.input_layernorm.weight"),
816                vec![hidden],
817                &vec![1.0; hidden],
818            );
819            w.add_tensor_f32(
820                &format!("{p}.post_attention_layernorm.weight"),
821                vec![hidden],
822                &vec![1.0; hidden],
823            );
824            w.add_tensor_f32(
825                &format!("{p}.self_attn.q_proj.weight"),
826                vec![q_dim, hidden],
827                &vec![0.01; q_dim * hidden],
828            );
829            w.add_tensor_f32(
830                &format!("{p}.self_attn.k_proj.weight"),
831                vec![kv_hidden, hidden],
832                &vec![0.02; kv_hidden * hidden],
833            );
834            w.add_tensor_f32(
835                &format!("{p}.self_attn.v_proj.weight"),
836                vec![kv_hidden, hidden],
837                &vec![0.03; kv_hidden * hidden],
838            );
839            w.add_tensor_f32(
840                &format!("{p}.self_attn.o_proj.weight"),
841                vec![hidden, q_dim],
842                &vec![0.04; hidden * q_dim],
843            );
844            w.add_tensor_f32(
845                &format!("{p}.mlp.gate_proj.weight"),
846                vec![inter, hidden],
847                &vec![0.05; inter * hidden],
848            );
849            w.add_tensor_f32(
850                &format!("{p}.mlp.up_proj.weight"),
851                vec![inter, hidden],
852                &vec![0.06; inter * hidden],
853            );
854            w.add_tensor_f32(
855                &format!("{p}.mlp.down_proj.weight"),
856                vec![hidden, inter],
857                &vec![0.07; hidden * inter],
858            );
859        }
860
861        let dir = std::env::temp_dir().join(format!("apr_from_apr_falsify_{}", std::process::id()));
862        std::fs::create_dir_all(&dir).unwrap();
863        let apr_path = dir.join("tiny.apr");
864        std::fs::write(&apr_path, w.to_bytes().unwrap()).unwrap();
865
866        // Exercises the parallel dequant + validation construction path.
867        let model = Transformer::from_apr(&apr_path, &config)
868            .expect("from_apr must load every tensor via the parallel dequant collect");
869
870        // Structural: forward produces finite logits of the right shape — only
871        // possible if ALL per-layer weights were loaded (a dropped tensor would
872        // have failed validate_weights before we get here).
873        let logits = model.forward(&[1u32, 2, 3]);
874        assert_eq!(logits.len(), 3 * vocab, "all weights present => correct logit shape");
875        assert!(
876            logits.data().iter().all(|v| v.is_finite()),
877            "parallel-dequanted weights must produce finite logits"
878        );
879
880        let _ = std::fs::remove_dir_all(&dir);
881    }
882
883    /// Write a tiny but complete Qwen2-shaped APR for the tied-embedding
884    /// falsifiers (#2441).
885    ///
886    /// `lm_head` selects how the output projection is recorded:
887    /// - `Some(fill)` — a genuine `lm_head.weight` holding `vocab*hidden` values
888    /// - `None` — the tied case: full `[vocab, hidden]` shape, ZERO bytes of data,
889    ///   exactly what a converter emits for `tie_word_embeddings=true`
890    fn write_tied_fixture_apr(
891        config: &TransformerConfig,
892        embed: &[f32],
893        lm_head: Option<f32>,
894        path: &std::path::Path,
895    ) {
896        use aprender::serialization::apr::AprWriter;
897
898        let hidden = config.hidden_size;
899        let q_dim = config.q_dim();
900        let kv_hidden = config.num_kv_heads * config.head_dim();
901        let inter = config.intermediate_size;
902        let vocab = config.vocab_size;
903
904        let mut w = AprWriter::new();
905        w.add_tensor_f32("model.embed_tokens.weight", vec![vocab, hidden], embed);
906        w.add_tensor_f32("model.norm.weight", vec![hidden], &vec![1.0; hidden]);
907        match lm_head {
908            Some(fill) => {
909                w.add_tensor_f32("lm_head.weight", vec![vocab, hidden], &vec![fill; vocab * hidden])
910            }
911            // The defect fixture: shape-only descriptor, no data.
912            None => w.add_tensor_f32("lm_head.weight", vec![vocab, hidden], &[]),
913        }
914        for i in 0..config.num_hidden_layers {
915            let p = format!("model.layers.{i}");
916            w.add_tensor_f32(
917                &format!("{p}.input_layernorm.weight"),
918                vec![hidden],
919                &vec![1.0; hidden],
920            );
921            w.add_tensor_f32(
922                &format!("{p}.post_attention_layernorm.weight"),
923                vec![hidden],
924                &vec![1.0; hidden],
925            );
926            w.add_tensor_f32(
927                &format!("{p}.self_attn.q_proj.weight"),
928                vec![q_dim, hidden],
929                &vec![0.01; q_dim * hidden],
930            );
931            w.add_tensor_f32(
932                &format!("{p}.self_attn.k_proj.weight"),
933                vec![kv_hidden, hidden],
934                &vec![0.02; kv_hidden * hidden],
935            );
936            w.add_tensor_f32(
937                &format!("{p}.self_attn.v_proj.weight"),
938                vec![kv_hidden, hidden],
939                &vec![0.03; kv_hidden * hidden],
940            );
941            w.add_tensor_f32(
942                &format!("{p}.self_attn.o_proj.weight"),
943                vec![hidden, q_dim],
944                &vec![0.04; hidden * q_dim],
945            );
946            w.add_tensor_f32(
947                &format!("{p}.mlp.gate_proj.weight"),
948                vec![inter, hidden],
949                &vec![0.05; inter * hidden],
950            );
951            w.add_tensor_f32(
952                &format!("{p}.mlp.up_proj.weight"),
953                vec![inter, hidden],
954                &vec![0.06; inter * hidden],
955            );
956            w.add_tensor_f32(
957                &format!("{p}.mlp.down_proj.weight"),
958                vec![hidden, inter],
959                &vec![0.07; hidden * inter],
960            );
961        }
962        std::fs::write(path, w.to_bytes().expect("apr bytes")).expect("write fixture");
963    }
964
965    /// Embedding fill that varies BETWEEN ROWS, so a model that ties its output
966    /// projection to the embeddings produces per-vocab-entry logits. A pattern
967    /// that only varies within a row repeats identically for every row and would
968    /// make the logits degenerate no matter which matrix the head used.
969    fn varied_embed(vocab: usize, hidden: usize) -> Vec<f32> {
970        (0..vocab * hidden)
971            .map(|i| {
972                let row = i / hidden;
973                let col = i % hidden;
974                ((row % 7) as f32 - 3.0) * 0.05 + ((col % 3) as f32 - 1.0) * 0.01
975            })
976            .collect()
977    }
978
979    /// #2441 RED→GREEN falsifier: `apr finetune` on a tied-embedding `.apr`.
980    ///
981    /// `apr finetune qwen2.5-coder-0.5b-instruct.apr --task classify` died with
982    ///
983    ///   error: Shape mismatch for 'lm_head.weight': expected 136134656 elements, got 0
984    ///
985    /// because the `.apr` records `lm_head.weight` as a shape-only, ZERO-BYTE
986    /// placeholder (its weights ARE `model.embed_tokens.weight`, stored once) and
987    /// `Transformer::from_apr` dequantized that descriptor into an empty tensor
988    /// before running it through shape validation. The sibling-directory load that
989    /// #2436 removed had been masking this — the wrong model loaded fine.
990    ///
991    /// RED before the fix: `from_apr` returns that ConfigError.
992    /// GREEN after: the tie is resolved and the output projection is the embedding.
993    #[test]
994    fn falsify_2441_from_apr_resolves_tied_lm_head_placeholder() {
995        let config = TransformerConfig::tiny();
996        let vocab = config.vocab_size;
997        let hidden = config.hidden_size;
998        let embed = varied_embed(vocab, hidden);
999
1000        let dir = std::env::temp_dir().join(format!("apr_tied_lm_head_{}", std::process::id()));
1001        std::fs::create_dir_all(&dir).expect("tmp dir");
1002        let apr_path = dir.join("tied.apr");
1003        write_tied_fixture_apr(&config, &embed, None, &apr_path);
1004
1005        let model = Transformer::from_apr(&apr_path, &config)
1006            .expect("#2441: a tied-embedding .apr must load, not fail shape validation");
1007
1008        assert!(
1009            model.lm_head.is_none(),
1010            "#2441: a 0-byte lm_head placeholder means TIED — no separate head"
1011        );
1012        assert_eq!(
1013            model.lm_head_weight_slice(),
1014            embed.as_slice(),
1015            "#2441: the output projection must be the embedding matrix itself"
1016        );
1017
1018        // Behaviour: the logits the finetune loop consumes are real numbers coming
1019        // off the embedding matrix, not an empty/zero head.
1020        let logits = model.forward(&[1u32, 2, 3]);
1021        assert_eq!(logits.len(), 3 * vocab);
1022        assert!(
1023            logits.data().iter().all(|v| v.is_finite()),
1024            "tied lm_head must produce finite logits"
1025        );
1026        let first_row: Vec<f32> = logits.data().iter().take(vocab).copied().collect();
1027        let first = first_row[0];
1028        assert!(
1029            first_row.iter().any(|v| (*v - first).abs() > 1e-6),
1030            "tied lm_head must produce non-degenerate logits — an all-equal row would \
1031             mean the head is empty/constant rather than the embedding matrix"
1032        );
1033
1034        let _ = std::fs::remove_dir_all(&dir);
1035    }
1036
1037    /// #2441 negative control: a REAL `lm_head.weight` must survive untouched.
1038    ///
1039    /// Guards against an over-broad fix that drops or overwrites every lm_head:
1040    /// an untied model's output projection must stay its own matrix.
1041    #[test]
1042    fn falsify_2441_untied_lm_head_is_preserved() {
1043        let config = TransformerConfig::tiny();
1044        let vocab = config.vocab_size;
1045        let hidden = config.hidden_size;
1046        let embed = varied_embed(vocab, hidden);
1047
1048        let dir = std::env::temp_dir().join(format!("apr_untied_lm_head_{}", std::process::id()));
1049        std::fs::create_dir_all(&dir).expect("tmp dir");
1050        let apr_path = dir.join("untied.apr");
1051        write_tied_fixture_apr(&config, &embed, Some(0.25), &apr_path);
1052
1053        let model = Transformer::from_apr(&apr_path, &config).expect("untied .apr must load");
1054
1055        assert!(
1056            model.lm_head.is_some(),
1057            "#2441: a materialized lm_head must NOT be dropped as if it were tied"
1058        );
1059        assert!(
1060            model.lm_head_weight_slice().iter().all(|v| (*v - 0.25).abs() < 1e-9),
1061            "#2441: the untied output projection must keep its own weights"
1062        );
1063
1064        let _ = std::fs::remove_dir_all(&dir);
1065    }
1066
1067    /// #2441: a 0-byte `lm_head.weight` with no usable embedding matrix is a
1068    /// genuinely broken file and must still be rejected — the tie resolution must
1069    /// not turn "unloadable" into "silently loaded with a wrong head".
1070    #[test]
1071    fn falsify_2441_empty_lm_head_without_tie_target_still_rejected() {
1072        let config = TransformerConfig::tiny();
1073        let mut weights: HashMap<String, Tensor> = HashMap::new();
1074        weights.insert("lm_head.weight".to_string(), Tensor::from_vec(vec![], false));
1075
1076        assert!(
1077            !Transformer::resolve_tied_lm_head(&mut weights),
1078            "no usable embedding matrix => nothing to tie to"
1079        );
1080        assert!(
1081            weights.contains_key("lm_head.weight"),
1082            "the placeholder must be left in place so shape validation still reports it"
1083        );
1084        let err = Transformer::validate_weight_shapes(&weights, &config)
1085            .expect_err("an empty lm_head with no tie target must fail shape validation");
1086        assert!(
1087            format!("{err}").contains("lm_head.weight"),
1088            "the error must still name lm_head.weight, got: {err}"
1089        );
1090    }
1091
1092    #[test]
1093    fn test_from_params_returns_none_on_missing() {
1094        let config = TransformerConfig::tiny();
1095        let params: HashMap<String, Tensor> = HashMap::new();
1096        let result = Transformer::from_params(&config, &params);
1097        assert!(result.is_none());
1098    }
1099
1100    #[test]
1101    fn test_transformer_from_params_with_lm_head() {
1102        let config = TransformerConfig::tiny();
1103        let hidden_size = config.hidden_size;
1104        let vocab_size = config.vocab_size;
1105        let kv_hidden_size = config.num_kv_heads * config.head_dim();
1106        let intermediate_size = config.intermediate_size;
1107
1108        let mut params = HashMap::new();
1109
1110        // Embedding
1111        params.insert(
1112            "model.embed_tokens.weight".to_string(),
1113            Tensor::from_vec(vec![0.1; vocab_size * hidden_size], true),
1114        );
1115
1116        // All layers
1117        for layer_idx in 0..config.num_hidden_layers {
1118            let prefix = format!("model.layers.{layer_idx}");
1119            params.insert(
1120                format!("{prefix}.input_layernorm.weight"),
1121                Tensor::from_vec(vec![1.0; hidden_size], true),
1122            );
1123            params.insert(
1124                format!("{prefix}.self_attn.q_proj.weight"),
1125                Tensor::from_vec(vec![0.1; hidden_size * hidden_size], true),
1126            );
1127            params.insert(
1128                format!("{prefix}.self_attn.k_proj.weight"),
1129                Tensor::from_vec(vec![0.1; hidden_size * kv_hidden_size], true),
1130            );
1131            params.insert(
1132                format!("{prefix}.self_attn.v_proj.weight"),
1133                Tensor::from_vec(vec![0.1; hidden_size * kv_hidden_size], true),
1134            );
1135            params.insert(
1136                format!("{prefix}.self_attn.o_proj.weight"),
1137                Tensor::from_vec(vec![0.1; hidden_size * hidden_size], true),
1138            );
1139            params.insert(
1140                format!("{prefix}.post_attention_layernorm.weight"),
1141                Tensor::from_vec(vec![1.0; hidden_size], true),
1142            );
1143            params.insert(
1144                format!("{prefix}.mlp.gate_proj.weight"),
1145                Tensor::from_vec(vec![0.1; hidden_size * intermediate_size], true),
1146            );
1147            params.insert(
1148                format!("{prefix}.mlp.up_proj.weight"),
1149                Tensor::from_vec(vec![0.1; hidden_size * intermediate_size], true),
1150            );
1151            params.insert(
1152                format!("{prefix}.mlp.down_proj.weight"),
1153                Tensor::from_vec(vec![0.1; intermediate_size * hidden_size], true),
1154            );
1155        }
1156
1157        // Final norm
1158        params.insert(
1159            "model.norm.weight".to_string(),
1160            Tensor::from_vec(vec![1.0; hidden_size], true),
1161        );
1162
1163        // LM head (separate, not tied)
1164        params.insert(
1165            "lm_head.weight".to_string(),
1166            Tensor::from_vec(vec![0.1; hidden_size * vocab_size], true),
1167        );
1168
1169        let transformer = Transformer::from_params(&config, &params);
1170        assert!(transformer.is_some());
1171        let transformer = transformer.expect("operation should succeed");
1172        assert!(transformer.lm_head.is_some());
1173        assert_eq!(transformer.layers.len(), config.num_hidden_layers);
1174    }
1175
1176    #[test]
1177    fn test_transformer_from_params_without_lm_head() {
1178        let config = TransformerConfig::tiny();
1179        let hidden_size = config.hidden_size;
1180        let vocab_size = config.vocab_size;
1181        let kv_hidden_size = config.num_kv_heads * config.head_dim();
1182        let intermediate_size = config.intermediate_size;
1183
1184        let mut params = HashMap::new();
1185
1186        // Embedding
1187        params.insert(
1188            "model.embed_tokens.weight".to_string(),
1189            Tensor::from_vec(vec![0.1; vocab_size * hidden_size], true),
1190        );
1191
1192        // All layers
1193        for layer_idx in 0..config.num_hidden_layers {
1194            let prefix = format!("model.layers.{layer_idx}");
1195            params.insert(
1196                format!("{prefix}.input_layernorm.weight"),
1197                Tensor::from_vec(vec![1.0; hidden_size], true),
1198            );
1199            params.insert(
1200                format!("{prefix}.self_attn.q_proj.weight"),
1201                Tensor::from_vec(vec![0.1; hidden_size * hidden_size], true),
1202            );
1203            params.insert(
1204                format!("{prefix}.self_attn.k_proj.weight"),
1205                Tensor::from_vec(vec![0.1; hidden_size * kv_hidden_size], true),
1206            );
1207            params.insert(
1208                format!("{prefix}.self_attn.v_proj.weight"),
1209                Tensor::from_vec(vec![0.1; hidden_size * kv_hidden_size], true),
1210            );
1211            params.insert(
1212                format!("{prefix}.self_attn.o_proj.weight"),
1213                Tensor::from_vec(vec![0.1; hidden_size * hidden_size], true),
1214            );
1215            params.insert(
1216                format!("{prefix}.post_attention_layernorm.weight"),
1217                Tensor::from_vec(vec![1.0; hidden_size], true),
1218            );
1219            params.insert(
1220                format!("{prefix}.mlp.gate_proj.weight"),
1221                Tensor::from_vec(vec![0.1; hidden_size * intermediate_size], true),
1222            );
1223            params.insert(
1224                format!("{prefix}.mlp.up_proj.weight"),
1225                Tensor::from_vec(vec![0.1; hidden_size * intermediate_size], true),
1226            );
1227            params.insert(
1228                format!("{prefix}.mlp.down_proj.weight"),
1229                Tensor::from_vec(vec![0.1; intermediate_size * hidden_size], true),
1230            );
1231        }
1232
1233        // Final norm - no lm_head
1234        params.insert(
1235            "model.norm.weight".to_string(),
1236            Tensor::from_vec(vec![1.0; hidden_size], true),
1237        );
1238
1239        let transformer = Transformer::from_params(&config, &params);
1240        assert!(transformer.is_some());
1241        let transformer = transformer.expect("operation should succeed");
1242        assert!(transformer.lm_head.is_none()); // Should use tied embeddings
1243    }
1244
1245    #[test]
1246    fn test_transformer_parameters_with_lm_head() {
1247        let config = TransformerConfig::tiny();
1248        let mut transformer = Transformer::new(&config);
1249
1250        // Add a separate lm_head
1251        transformer.lm_head =
1252            Some(Tensor::from_vec(vec![0.1; config.hidden_size * config.vocab_size], true));
1253
1254        let params = transformer.parameters();
1255        // embed_tokens + norm + (layers * 9) + lm_head
1256        // = 1 + 1 + (2 * 9) + 1 = 21
1257        assert_eq!(params.len(), 21);
1258    }
1259
1260    #[test]
1261    fn test_transformer_forward_with_lm_head() {
1262        let config = TransformerConfig::tiny();
1263        let mut transformer = Transformer::new(&config);
1264
1265        // Add a separate lm_head
1266        transformer.lm_head =
1267            Some(Tensor::from_vec(vec![0.1; config.hidden_size * config.vocab_size], true));
1268
1269        let tokens = vec![1, 2, 3];
1270        let logits = transformer.forward(&tokens);
1271        assert_eq!(logits.len(), 3 * config.vocab_size);
1272        assert!(logits.data().iter().all(|&v| v.is_finite()));
1273    }
1274
1275    // =========================================================================
1276    // FALSIFY-L: §2.1.2 LM Head Contract — Five-Whys Gap Analysis (Refs PMAT-329)
1277    //
1278    // Contract: tensor-layout-v1.yaml §tensors.lm_head
1279    //   critical: "true"
1280    //   note: "GH-202 root cause - wrong shape caused [PAD] garbage output"
1281    //
1282    // Five-Whys:
1283    //   Why 1: entrenar-trained model's lm_head could corrupt inference
1284    //   Why 2: lm_head save/load has no shape validation
1285    //   Why 3: from_params accepts ANY tensor for lm_head (like embedding)
1286    //   Why 4: entrenar predates ValidatedWeight contract
1287    //   Why 5: No cross-crate contract enforcement for trained models
1288    //
1289    // Popper (1959): "These tests attempt to falsify the claim that
1290    // entrenar's lm_head handling prevents garbage output after training."
1291    // =========================================================================
1292
1293    /// FALSIFY-L1e: from_params rejects wrong-shape lm_head (PMAT-329 fix)
1294    ///
1295    /// from_params now validates lm_head shape against vocab*hidden.
1296    /// A tensor of 50 elements is rejected when vocab*hidden is expected.
1297    #[test]
1298    fn falsify_l1e_from_params_rejects_wrong_shape_lm_head() {
1299        let config = TransformerConfig::tiny();
1300        let hidden_size = config.hidden_size;
1301        let vocab_size = config.vocab_size;
1302        let kv_hidden_size = config.num_kv_heads * config.head_dim();
1303        let intermediate_size = config.intermediate_size;
1304
1305        let mut params = HashMap::new();
1306
1307        // Valid embedding + layers + norm
1308        params.insert(
1309            "model.embed_tokens.weight".to_string(),
1310            Tensor::from_vec(vec![0.1; vocab_size * hidden_size], true),
1311        );
1312        for layer_idx in 0..config.num_hidden_layers {
1313            let prefix = format!("model.layers.{layer_idx}");
1314            params.insert(
1315                format!("{prefix}.input_layernorm.weight"),
1316                Tensor::from_vec(vec![1.0; hidden_size], true),
1317            );
1318            params.insert(
1319                format!("{prefix}.self_attn.q_proj.weight"),
1320                Tensor::from_vec(vec![0.1; hidden_size * hidden_size], true),
1321            );
1322            params.insert(
1323                format!("{prefix}.self_attn.k_proj.weight"),
1324                Tensor::from_vec(vec![0.1; hidden_size * kv_hidden_size], true),
1325            );
1326            params.insert(
1327                format!("{prefix}.self_attn.v_proj.weight"),
1328                Tensor::from_vec(vec![0.1; hidden_size * kv_hidden_size], true),
1329            );
1330            params.insert(
1331                format!("{prefix}.self_attn.o_proj.weight"),
1332                Tensor::from_vec(vec![0.1; hidden_size * hidden_size], true),
1333            );
1334            params.insert(
1335                format!("{prefix}.post_attention_layernorm.weight"),
1336                Tensor::from_vec(vec![1.0; hidden_size], true),
1337            );
1338            params.insert(
1339                format!("{prefix}.mlp.gate_proj.weight"),
1340                Tensor::from_vec(vec![0.1; hidden_size * intermediate_size], true),
1341            );
1342            params.insert(
1343                format!("{prefix}.mlp.up_proj.weight"),
1344                Tensor::from_vec(vec![0.1; hidden_size * intermediate_size], true),
1345            );
1346            params.insert(
1347                format!("{prefix}.mlp.down_proj.weight"),
1348                Tensor::from_vec(vec![0.1; intermediate_size * hidden_size], true),
1349            );
1350        }
1351        params.insert(
1352            "model.norm.weight".to_string(),
1353            Tensor::from_vec(vec![1.0; hidden_size], true),
1354        );
1355
1356        // WRONG-SHAPE lm_head: 50 elements for hidden*vocab expected
1357        params.insert("lm_head.weight".to_string(), Tensor::from_vec(vec![0.1; 50], true));
1358
1359        let transformer = Transformer::from_params(&config, &params);
1360        // FIXED (PMAT-329): now rejected
1361        assert!(
1362            transformer.is_none(),
1363            "FALSIFY-L1e: PMAT-329 fix — from_params MUST reject wrong-shape lm_head"
1364        );
1365    }
1366
1367    /// FALSIFY-L2e: Tied embeddings produce valid logit dimensions
1368    ///
1369    /// When lm_head is None, the embedding weight [vocab, hidden] is used as lm_head.
1370    /// The matmul must produce [seq_len, vocab_size] logits.
1371    #[test]
1372    fn falsify_l2e_tied_embeddings_produce_correct_logit_dims() {
1373        let config = TransformerConfig::tiny();
1374        let transformer = Transformer::new(&config);
1375        assert!(transformer.lm_head.is_none(), "Default should use tied embeddings");
1376
1377        let tokens = vec![1, 2, 3];
1378        let logits = transformer.forward(&tokens);
1379        assert_eq!(
1380            logits.len(),
1381            3 * config.vocab_size,
1382            "FALSIFY-L2e: Tied embedding logits must be seq_len * vocab_size"
1383        );
1384
1385        // All logits must be finite (not NaN/Inf)
1386        let data = logits.data();
1387        let nan_count = data.iter().filter(|v| v.is_nan()).count();
1388        let inf_count = data.iter().filter(|v| v.is_infinite()).count();
1389        assert_eq!(nan_count, 0, "FALSIFY-L2e: Tied logits must not contain NaN");
1390        assert_eq!(inf_count, 0, "FALSIFY-L2e: Tied logits must not contain Inf");
1391    }
1392
1393    /// FALSIFY-L3e: Separate lm_head produces valid logit dimensions
1394    #[test]
1395    fn falsify_l3e_separate_lm_head_produces_correct_logit_dims() {
1396        let config = TransformerConfig::tiny();
1397        let mut transformer = Transformer::new(&config);
1398        transformer.lm_head =
1399            Some(Tensor::from_vec(vec![0.1; config.hidden_size * config.vocab_size], true));
1400
1401        let tokens = vec![1, 2, 3];
1402        let logits = transformer.forward(&tokens);
1403        assert_eq!(
1404            logits.len(),
1405            3 * config.vocab_size,
1406            "FALSIFY-L3e: Separate lm_head logits must be seq_len * vocab_size"
1407        );
1408        let data = logits.data();
1409        assert!(
1410            data.iter().all(|v| v.is_finite()),
1411            "FALSIFY-L3e: Separate lm_head logits must all be finite"
1412        );
1413    }
1414
1415    /// FALSIFY-L4e: lm_head is included in parameters() and parameters_mut()
1416    ///
1417    /// If lm_head is present but not returned by parameters(), the optimizer
1418    /// won't update it during training → frozen lm_head → garbage after finetuning.
1419    #[test]
1420    fn falsify_l4e_lm_head_in_parameter_list() {
1421        let config = TransformerConfig::tiny();
1422        let mut transformer = Transformer::new(&config);
1423
1424        // Without lm_head: N params
1425        let n_without = transformer.parameters().len();
1426
1427        // With lm_head: N+1 params
1428        transformer.lm_head =
1429            Some(Tensor::from_vec(vec![0.1; config.hidden_size * config.vocab_size], true));
1430        let n_with = transformer.parameters().len();
1431        assert_eq!(
1432            n_with,
1433            n_without + 1,
1434            "FALSIFY-L4e: lm_head must be included in parameters() — optimizer needs it"
1435        );
1436
1437        // Also check parameters_mut
1438        let n_mut = transformer.parameters_mut().len();
1439        assert_eq!(
1440            n_mut, n_with,
1441            "FALSIFY-L4e: parameters_mut() must include lm_head for gradient updates"
1442        );
1443    }
1444
1445    /// FALSIFY-L5e: forward_last returns exactly vocab_size logits
1446    ///
1447    /// The last token's logits are used for next-token prediction.
1448    /// Off-by-one in the slice extraction → wrong token generated.
1449    #[test]
1450    fn falsify_l5e_forward_last_correct_size() {
1451        let config = TransformerConfig::tiny();
1452        let transformer = Transformer::new(&config);
1453
1454        let tokens = vec![1, 2, 3, 4, 5];
1455        let logits = transformer.forward_last(&tokens);
1456        assert_eq!(
1457            logits.len(),
1458            config.vocab_size,
1459            "FALSIFY-L5e: forward_last must return exactly vocab_size logits"
1460        );
1461        let data = logits.data();
1462        assert!(
1463            data.iter().all(|v| v.is_finite()),
1464            "FALSIFY-L5e: forward_last logits must all be finite"
1465        );
1466    }
1467
1468    #[test]
1469    fn test_causal_lm_loss_backward() {
1470        use crate::train::CausalLMLoss;
1471        use crate::train::LossFn;
1472
1473        let vocab_size = 100;
1474        let seq_len = 3;
1475        let loss_fn = CausalLMLoss::new(vocab_size);
1476
1477        // Create some logits
1478        let logits = Tensor::from_vec(
1479            (0..seq_len * vocab_size).map(|i| (i as f32 * 0.01).sin()).collect(),
1480            true,
1481        );
1482
1483        // Target token IDs
1484        let targets = Tensor::from_vec(vec![5.0, 10.0, 15.0], false);
1485
1486        let mut loss = loss_fn.forward(&logits, &targets);
1487
1488        // Backward
1489        crate::autograd::backward(&mut loss, None);
1490
1491        // Loss should be positive
1492        assert!(loss.data()[0] > 0.0);
1493        assert!(loss.data()[0].is_finite());
1494
1495        // Logits should have gradient
1496        assert!(logits.grad().is_some());
1497        let grad = logits.grad().expect("gradient should be available");
1498        assert!(grad.iter().all(|&v| v.is_finite()));
1499    }
1500
1501    // =========================================================================
1502    // FALSIFY-EMB-003 / FALSIFY-TE-001..004: Tied Embeddings Contract
1503    //
1504    // Five-Whys (PMAT-354):
1505    //   Why 1: entrenar had L-series lm_head tests but no EMB-003/TE-* tagged tests
1506    //   Why 2: L-series validates shape, not tied-weight CONTRACT claims
1507    //   Why 3: no mapping from tied-embeddings-v1.yaml to entrenar test names
1508    //   Why 4: entrenar predates the provable-contracts YAML
1509    //   Why 5: tied weights were assumed correct because code path is "just fallback"
1510    //
1511    // References:
1512    //   - provable-contracts/contracts/embedding-algebra-v1.yaml (EMB-003)
1513    //   - provable-contracts/contracts/tied-embeddings-v1.yaml (TE-001..004)
1514    //   - Press & Wolf (2017) "Using the Output Embedding to Improve Language Models"
1515    // =========================================================================
1516
1517    /// FALSIFY-EMB-003: Tied weight sharing — lm_head uses embed_tokens.weight
1518    ///
1519    /// Contract: when lm_head is None, forward() uses embed_tokens.weight directly
1520    /// (pointer/identity sharing, not a copy)
1521    #[test]
1522    fn falsify_emb_003_tied_weight_sharing() {
1523        let config = TransformerConfig::tiny();
1524        let transformer = Transformer::new(&config);
1525
1526        // Default: lm_head is None → tied
1527        assert!(transformer.lm_head.is_none());
1528
1529        // The weight used for lm_head projection IS embed_tokens.weight
1530        let lm_weight = transformer.lm_head.as_ref().unwrap_or(&transformer.embed_tokens.weight);
1531        let embed_weight = &transformer.embed_tokens.weight;
1532
1533        // They must be the same Tensor (same data pointer, not just equal values)
1534        assert!(
1535            std::ptr::eq(lm_weight, embed_weight),
1536            "FALSIFIED EMB-003: tied lm_head must be same object as embed_tokens.weight"
1537        );
1538    }
1539
1540    /// FALSIFY-TE-001: Output shape = (seq_len, vocab_size)
1541    #[test]
1542    fn falsify_te_001_output_shape() {
1543        let config = TransformerConfig::tiny();
1544        let transformer = Transformer::new(&config);
1545
1546        for seq_len in [1, 3, 10] {
1547            let tokens: Vec<u32> = (0..seq_len).collect();
1548            let logits = transformer.forward(&tokens);
1549            assert_eq!(
1550                logits.len(),
1551                seq_len as usize * config.vocab_size,
1552                "FALSIFIED TE-001: output shape for seq_len={seq_len}"
1553            );
1554        }
1555    }
1556
1557    /// FALSIFY-TE-002: Tied equivalence — tied output == explicit matmul with cloned W
1558    ///
1559    /// Contract: forward() with tied lm_head must produce bit-identical output
1560    /// to manually computing matmul(hidden, W_embed) with a separate copy of the
1561    /// embedding weight matrix. If they diverge, the tied path silently aliases
1562    /// or transposes incorrectly.
1563    #[test]
1564    fn falsify_te_002_tied_equivalence() {
1565        let config = TransformerConfig::tiny();
1566        let transformer = Transformer::new(&config);
1567
1568        // Tied path: forward() uses embed_tokens.weight as lm_head
1569        let tokens = vec![0u32, 3, 7, 15, 42];
1570        let tied_logits = transformer.forward(&tokens);
1571
1572        // Explicit path: clone embed weight, run hidden states, matmul manually
1573        let hidden = transformer.forward_hidden(&tokens);
1574        let w_clone = transformer.embed_tokens.weight.clone();
1575        let explicit_logits =
1576            matmul_nt(&hidden, &w_clone, tokens.len(), config.hidden_size, config.vocab_size);
1577
1578        let tied_data = tied_logits.data();
1579        let explicit_data = explicit_logits.data();
1580
1581        assert_eq!(
1582            tied_data.len(),
1583            explicit_data.len(),
1584            "FALSIFIED TE-002: output lengths differ: {} vs {}",
1585            tied_data.len(),
1586            explicit_data.len()
1587        );
1588
1589        for (i, (&t, &e)) in tied_data.iter().zip(explicit_data.iter()).enumerate() {
1590            assert!(
1591                (t - e).abs() < 1e-6,
1592                "FALSIFIED TE-002: tied[{i}] = {t} != explicit[{i}] = {e}"
1593            );
1594        }
1595    }
1596
1597    /// FALSIFY-TE-003: No extra parameters for tied embeddings
1598    ///
1599    /// Contract: tied model has exactly N params, untied has N+1 (the separate lm_head)
1600    #[test]
1601    fn falsify_te_003_no_extra_params() {
1602        let config = TransformerConfig::tiny();
1603        let tied = Transformer::new(&config);
1604        let tied_count = tied.parameters().len();
1605
1606        let mut untied = Transformer::new(&config);
1607        untied.lm_head =
1608            Some(Tensor::from_vec(vec![0.1; config.hidden_size * config.vocab_size], true));
1609        let untied_count = untied.parameters().len();
1610
1611        assert_eq!(
1612            untied_count,
1613            tied_count + 1,
1614            "FALSIFIED TE-003: tied model must have exactly 1 fewer param than untied"
1615        );
1616    }
1617
1618    /// FALSIFY-TE-004: Finite output for tied embeddings
1619    #[test]
1620    fn falsify_te_004_finite_output() {
1621        let config = TransformerConfig::tiny();
1622        let transformer = Transformer::new(&config);
1623        let tokens = vec![0u32, 5, 10, 50, 99];
1624        let logits = transformer.forward(&tokens);
1625        let data = logits.data();
1626
1627        let nan_count = data.iter().filter(|v| v.is_nan()).count();
1628        let inf_count = data.iter().filter(|v| v.is_infinite()).count();
1629
1630        assert_eq!(
1631            nan_count, 0,
1632            "FALSIFIED TE-004: tied embedding output contains {nan_count} NaN values"
1633        );
1634        assert_eq!(
1635            inf_count, 0,
1636            "FALSIFIED TE-004: tied embedding output contains {inf_count} Inf values"
1637        );
1638    }
1639
1640    // =========================================================================
1641    // PROPTEST FALSIFY-TE: Tied embeddings property-based falsification
1642    //
1643    // Five-Whys (PMAT-354, Phase 9):
1644    //   Why 1: YAML tied-embeddings-v1.yaml calls for "proptest with seq_len in [1,128]"
1645    //   Why 2: All 4 TE tests use fixed token sequences
1646    //   Why 3: TE proptest had ZERO coverage across the entire stack
1647    //   Why 4: Transformer construction is expensive, discouraging property testing
1648    //   Why 5: Fixed tokens miss edge cases in arbitrary token→logit paths
1649    //
1650    // References:
1651    //   - tied-embeddings-v1.yaml FALSIFY-TE-001: "proptest with seq_len in [1,128]"
1652    //   - tied-embeddings-v1.yaml FALSIFY-TE-002: "clone W_embed, compare tied vs explicit"
1653    //   - tied-embeddings-v1.yaml FALSIFY-TE-004: "proptest with finite x, check is_finite()"
1654    // =========================================================================
1655
1656    mod te_proptest_falsify {
1657        use super::*;
1658        use proptest::prelude::*;
1659
1660        // TE-001-prop: Output shape for random seq_len
1661        // Construct transformer once per test run, vary only the token sequence
1662        proptest! {
1663            #![proptest_config(ProptestConfig::with_cases(50))]
1664            #[test]
1665            fn falsify_te_001_prop_output_shape(
1666                seq_len in 1_usize..32,
1667            ) {
1668                let config = TransformerConfig::tiny();
1669                let transformer = Transformer::new(&config);
1670                let tokens: Vec<u32> = (0..seq_len).map(|i| (i % config.vocab_size) as u32).collect();
1671                let logits = transformer.forward(&tokens);
1672                prop_assert_eq!(
1673                    logits.len(),
1674                    seq_len * config.vocab_size,
1675                    "FALSIFIED TE-001-prop: seq_len={}, got len={}", seq_len, logits.len()
1676                );
1677            }
1678        }
1679
1680        // TE-002-prop: Tied equivalence for random tokens
1681        proptest! {
1682            #![proptest_config(ProptestConfig::with_cases(20))]
1683            #[test]
1684            fn falsify_te_002_prop_tied_equivalence(
1685                token_ids in proptest::collection::vec(0_u32..999, 1..8),
1686            ) {
1687                let config = TransformerConfig::tiny();
1688                let transformer = Transformer::new(&config);
1689
1690                let tied_logits = transformer.forward(&token_ids);
1691                let hidden = transformer.forward_hidden(&token_ids);
1692                let w_clone = transformer.embed_tokens.weight.clone();
1693                let explicit_logits = matmul_nt(
1694                    &hidden, &w_clone,
1695                    token_ids.len(), config.hidden_size, config.vocab_size,
1696                );
1697
1698                let tied_data = tied_logits.data();
1699                let explicit_data = explicit_logits.data();
1700                prop_assert_eq!(tied_data.len(), explicit_data.len());
1701
1702                for (i, (&t, &e)) in tied_data.iter().zip(explicit_data.iter()).enumerate() {
1703                    prop_assert!(
1704                        (t - e).abs() < 1e-5,
1705                        "FALSIFIED TE-002-prop: tied[{}]={} != explicit[{}]={}",
1706                        i, t, i, e
1707                    );
1708                }
1709            }
1710        }
1711
1712        // TE-004-prop: All outputs finite for random tokens
1713        proptest! {
1714            #![proptest_config(ProptestConfig::with_cases(30))]
1715            #[test]
1716            fn falsify_te_004_prop_finite(
1717                token_ids in proptest::collection::vec(0_u32..999, 1..16),
1718            ) {
1719                let config = TransformerConfig::tiny();
1720                let transformer = Transformer::new(&config);
1721                let logits = transformer.forward(&token_ids);
1722                let data = logits.data();
1723
1724                for (i, &v) in data.iter().enumerate() {
1725                    prop_assert!(
1726                        v.is_finite(),
1727                        "FALSIFIED TE-004-prop: logits[{}]={} non-finite (n_tokens={})",
1728                        i, v, token_ids.len()
1729                    );
1730                }
1731            }
1732        }
1733    }
1734
1735    // =========================================================================
1736    // FALSIFY-PIPE-001: Cross-contract pipeline test
1737    //
1738    // Five-Whys (PMAT-354, Phase 8):
1739    //   Why 1: no test exercises the full §2.1.1 pipeline as a single chain
1740    //   Why 2: EM, TE, SM tests each validate one contract in isolation
1741    //   Why 3: bugs can hide at contract boundaries (shape mismatch between stages)
1742    //   Why 4: the embed→tied_lm_head→softmax chain is the critical inference path
1743    //   Why 5: cross-contract pipeline faults would only show in integration
1744    //
1745    // Pipeline: embed(token_ids) → transformer_layers → norm → tied_matmul → softmax
1746    // Claims verified:
1747    //   EM-001: embed output shape = (seq_len, d_model)
1748    //   TE-001: tied logits shape = (seq_len, vocab_size)
1749    //   SM-001: softmax(logits) sums to 1.0 per row
1750    //   SM-002: all probabilities positive
1751    //   SM-003: argmax preserved through softmax
1752    // =========================================================================
1753
1754    /// FALSIFY-PIPE-001: Full embed → tied_lm_head → softmax pipeline
1755    #[test]
1756    fn falsify_pipe_001_embed_tied_softmax_pipeline() {
1757        let config = TransformerConfig::tiny();
1758        let transformer = Transformer::new(&config);
1759
1760        let tokens = vec![0u32, 3, 7, 15, 42];
1761        let seq_len = tokens.len();
1762        let vocab_size = config.vocab_size;
1763
1764        // Stage 1: Full forward pass (embed → layers → norm → tied matmul)
1765        let logits = transformer.forward(&tokens);
1766        let logits_data = logits.data();
1767
1768        // TE-001: logits shape = (seq_len, vocab_size)
1769        assert_eq!(
1770            logits_data.len(),
1771            seq_len * vocab_size,
1772            "FALSIFIED PIPE-001/TE-001: logits len={} != seq_len({seq_len}) * vocab({vocab_size})",
1773            logits_data.len()
1774        );
1775
1776        // TE-004: all logits finite
1777        for (i, &l) in logits_data.iter().enumerate() {
1778            assert!(l.is_finite(), "FALSIFIED PIPE-001/TE-004: logits[{i}] = {l} not finite");
1779        }
1780
1781        // Stage 2: Apply softmax per row (the sampling step)
1782        let logits_slice = logits_data.as_slice().expect("operation should succeed");
1783        for row in 0..seq_len {
1784            let start = row * vocab_size;
1785            let end = start + vocab_size;
1786            let row_logits = &logits_slice[start..end];
1787
1788            // Compute softmax for this row
1789            let max_val = row_logits.iter().copied().fold(f32::NEG_INFINITY, f32::max);
1790            let exps: Vec<f32> = row_logits.iter().map(|&x| (x - max_val).exp()).collect();
1791            let sum: f32 = exps.iter().sum();
1792            let probs: Vec<f32> = exps.iter().map(|&e| e / sum).collect();
1793
1794            // SM-001: sums to 1.0
1795            let prob_sum: f32 = probs.iter().sum();
1796            assert!(
1797                (prob_sum - 1.0).abs() < 1e-4,
1798                "FALSIFIED PIPE-001/SM-001: row {row} prob sum={prob_sum}"
1799            );
1800
1801            // SM-002: all positive
1802            for (i, &p) in probs.iter().enumerate() {
1803                assert!(p >= 0.0, "FALSIFIED PIPE-001/SM-002: row {row} prob[{i}]={p} negative");
1804            }
1805
1806            // SM-003: argmax preserved
1807            let logit_argmax = row_logits
1808                .iter()
1809                .enumerate()
1810                .max_by(|(_, a), (_, b)| a.partial_cmp(b).expect("operation should succeed"))
1811                .expect("operation should succeed")
1812                .0;
1813            let prob_argmax = probs
1814                .iter()
1815                .enumerate()
1816                .max_by(|(_, a), (_, b)| a.partial_cmp(b).expect("operation should succeed"))
1817                .expect("operation should succeed")
1818                .0;
1819            assert_eq!(
1820                logit_argmax, prob_argmax,
1821                "FALSIFIED PIPE-001/SM-003: row {row} argmax changed {logit_argmax} → {prob_argmax}"
1822            );
1823        }
1824    }
1825
1826    // =========================================================================
1827    // SSC-024: Transformer::from_safetensors() tests
1828    //
1829    // Tests for loading pretrained weights from SafeTensors files.
1830    // Uses synthetic SafeTensors with the tiny config to avoid needing
1831    // real 500MB model files in CI.
1832    // =========================================================================
1833
1834    mod safetensors_tests {
1835        use super::*;
1836        use safetensors::serialize;
1837        use safetensors::tensor::{Dtype, TensorView};
1838        use tempfile::TempDir;
1839
1840        /// Helper: create a synthetic SafeTensors file with all weights
1841        /// matching the tiny config (hidden=64, 2 layers, vocab=1000).
1842        fn create_tiny_safetensors(dir: &std::path::Path) -> std::path::PathBuf {
1843            let config = TransformerConfig::tiny();
1844            let hidden = config.hidden_size;
1845            let kv_hidden = config.num_kv_heads * config.head_dim();
1846            let intermediate = config.intermediate_size;
1847            let vocab = config.vocab_size;
1848
1849            let mut tensors_data: Vec<(String, Vec<u8>, Vec<usize>)> = Vec::new();
1850
1851            // Helper to create f32 bytes
1852            let make_f32 = |n: usize, val: f32| -> Vec<u8> {
1853                std::iter::repeat_n(val, n).flat_map(f32::to_le_bytes).collect()
1854            };
1855
1856            // Embedding
1857            tensors_data.push((
1858                "model.embed_tokens.weight".to_string(),
1859                make_f32(vocab * hidden, 0.01),
1860                vec![vocab, hidden],
1861            ));
1862
1863            // Final norm
1864            tensors_data.push((
1865                "model.norm.weight".to_string(),
1866                make_f32(hidden, 1.0),
1867                vec![hidden],
1868            ));
1869
1870            // Per-layer weights
1871            for i in 0..config.num_hidden_layers {
1872                let p = format!("model.layers.{i}");
1873
1874                // Layer norms
1875                tensors_data.push((
1876                    format!("{p}.input_layernorm.weight"),
1877                    make_f32(hidden, 1.0),
1878                    vec![hidden],
1879                ));
1880                tensors_data.push((
1881                    format!("{p}.post_attention_layernorm.weight"),
1882                    make_f32(hidden, 1.0),
1883                    vec![hidden],
1884                ));
1885
1886                // Attention projections
1887                tensors_data.push((
1888                    format!("{p}.self_attn.q_proj.weight"),
1889                    make_f32(hidden * hidden, 0.01),
1890                    vec![hidden, hidden],
1891                ));
1892                tensors_data.push((
1893                    format!("{p}.self_attn.k_proj.weight"),
1894                    make_f32(hidden * kv_hidden, 0.01),
1895                    vec![kv_hidden, hidden],
1896                ));
1897                tensors_data.push((
1898                    format!("{p}.self_attn.v_proj.weight"),
1899                    make_f32(hidden * kv_hidden, 0.01),
1900                    vec![kv_hidden, hidden],
1901                ));
1902                tensors_data.push((
1903                    format!("{p}.self_attn.o_proj.weight"),
1904                    make_f32(hidden * hidden, 0.01),
1905                    vec![hidden, hidden],
1906                ));
1907
1908                // MLP projections
1909                tensors_data.push((
1910                    format!("{p}.mlp.gate_proj.weight"),
1911                    make_f32(hidden * intermediate, 0.01),
1912                    vec![intermediate, hidden],
1913                ));
1914                tensors_data.push((
1915                    format!("{p}.mlp.up_proj.weight"),
1916                    make_f32(hidden * intermediate, 0.01),
1917                    vec![intermediate, hidden],
1918                ));
1919                tensors_data.push((
1920                    format!("{p}.mlp.down_proj.weight"),
1921                    make_f32(intermediate * hidden, 0.01),
1922                    vec![hidden, intermediate],
1923                ));
1924            }
1925
1926            // Build TensorViews from owned data and serialize
1927            let views: Vec<TensorView<'_>> = tensors_data
1928                .iter()
1929                .map(|(_, bytes, shape)| {
1930                    TensorView::new(Dtype::F32, shape.clone(), bytes).expect("valid tensor view")
1931                })
1932                .collect();
1933
1934            let named_views: Vec<(&str, &TensorView<'_>)> = tensors_data
1935                .iter()
1936                .zip(views.iter())
1937                .map(|((name, _, _), view)| (name.as_str(), view))
1938                .collect();
1939
1940            let file_path = dir.join("model.safetensors");
1941            let serialized =
1942                serialize(named_views, None::<std::collections::HashMap<String, String>>)
1943                    .expect("serialize safetensors");
1944            std::fs::write(&file_path, serialized).expect("write safetensors file");
1945            file_path
1946        }
1947
1948        /// Helper: create a SafeTensors file with bf16 weights (like real HF models)
1949        fn create_tiny_bf16_safetensors(dir: &std::path::Path) -> std::path::PathBuf {
1950            let config = TransformerConfig::tiny();
1951            let hidden = config.hidden_size;
1952            let kv_hidden = config.num_kv_heads * config.head_dim();
1953            let intermediate = config.intermediate_size;
1954            let vocab = config.vocab_size;
1955
1956            let mut tensors_data: Vec<(String, Vec<u8>, Vec<usize>)> = Vec::new();
1957
1958            // Helper to create bf16 bytes
1959            let make_bf16 = |n: usize, val: f32| -> Vec<u8> {
1960                std::iter::repeat_n(half::bf16::from_f32(val), n)
1961                    .flat_map(half::bf16::to_le_bytes)
1962                    .collect()
1963            };
1964
1965            // Embedding
1966            tensors_data.push((
1967                "model.embed_tokens.weight".to_string(),
1968                make_bf16(vocab * hidden, 0.01),
1969                vec![vocab, hidden],
1970            ));
1971
1972            // Final norm
1973            tensors_data.push((
1974                "model.norm.weight".to_string(),
1975                make_bf16(hidden, 1.0),
1976                vec![hidden],
1977            ));
1978
1979            // Per-layer weights
1980            for i in 0..config.num_hidden_layers {
1981                let p = format!("model.layers.{i}");
1982
1983                tensors_data.push((
1984                    format!("{p}.input_layernorm.weight"),
1985                    make_bf16(hidden, 1.0),
1986                    vec![hidden],
1987                ));
1988                tensors_data.push((
1989                    format!("{p}.post_attention_layernorm.weight"),
1990                    make_bf16(hidden, 1.0),
1991                    vec![hidden],
1992                ));
1993                tensors_data.push((
1994                    format!("{p}.self_attn.q_proj.weight"),
1995                    make_bf16(hidden * hidden, 0.01),
1996                    vec![hidden, hidden],
1997                ));
1998                tensors_data.push((
1999                    format!("{p}.self_attn.k_proj.weight"),
2000                    make_bf16(hidden * kv_hidden, 0.01),
2001                    vec![kv_hidden, hidden],
2002                ));
2003                tensors_data.push((
2004                    format!("{p}.self_attn.v_proj.weight"),
2005                    make_bf16(hidden * kv_hidden, 0.01),
2006                    vec![kv_hidden, hidden],
2007                ));
2008                tensors_data.push((
2009                    format!("{p}.self_attn.o_proj.weight"),
2010                    make_bf16(hidden * hidden, 0.01),
2011                    vec![hidden, hidden],
2012                ));
2013                tensors_data.push((
2014                    format!("{p}.mlp.gate_proj.weight"),
2015                    make_bf16(hidden * intermediate, 0.01),
2016                    vec![intermediate, hidden],
2017                ));
2018                tensors_data.push((
2019                    format!("{p}.mlp.up_proj.weight"),
2020                    make_bf16(hidden * intermediate, 0.01),
2021                    vec![intermediate, hidden],
2022                ));
2023                tensors_data.push((
2024                    format!("{p}.mlp.down_proj.weight"),
2025                    make_bf16(intermediate * hidden, 0.01),
2026                    vec![hidden, intermediate],
2027                ));
2028            }
2029
2030            let views: Vec<TensorView<'_>> = tensors_data
2031                .iter()
2032                .map(|(_, bytes, shape)| {
2033                    TensorView::new(Dtype::BF16, shape.clone(), bytes).expect("valid tensor view")
2034                })
2035                .collect();
2036
2037            let named_views: Vec<(&str, &TensorView<'_>)> = tensors_data
2038                .iter()
2039                .zip(views.iter())
2040                .map(|((name, _, _), view)| (name.as_str(), view))
2041                .collect();
2042
2043            let file_path = dir.join("model.safetensors");
2044            let serialized =
2045                serialize(named_views, None::<std::collections::HashMap<String, String>>)
2046                    .expect("serialize safetensors");
2047            std::fs::write(&file_path, serialized).expect("write safetensors file");
2048            file_path
2049        }
2050
2051        // -----------------------------------------------------------------
2052        // Happy path tests
2053        // -----------------------------------------------------------------
2054
2055        #[test]
2056        fn test_ssc024_from_safetensors_f32_success() {
2057            let dir = TempDir::new().expect("create temp dir");
2058            create_tiny_safetensors(dir.path());
2059            let config = TransformerConfig::tiny();
2060
2061            let result = Transformer::from_safetensors(dir.path(), &config);
2062            assert!(
2063                result.is_ok(),
2064                "from_safetensors should succeed: {}",
2065                result.as_ref().err().map_or(String::new(), std::string::ToString::to_string)
2066            );
2067
2068            let transformer = result.expect("validated above");
2069            assert_eq!(transformer.layers.len(), config.num_hidden_layers);
2070            assert!(transformer.lm_head.is_none()); // tiny config has no lm_head
2071        }
2072
2073        #[test]
2074        fn test_ssc024_from_safetensors_bf16_conversion() {
2075            let dir = TempDir::new().expect("create temp dir");
2076            create_tiny_bf16_safetensors(dir.path());
2077            let config = TransformerConfig::tiny();
2078
2079            let result = Transformer::from_safetensors(dir.path(), &config);
2080            assert!(
2081                result.is_ok(),
2082                "BF16 loading should succeed: {}",
2083                result.as_ref().err().map_or(String::new(), std::string::ToString::to_string)
2084            );
2085
2086            let transformer = result.expect("validated above");
2087            assert_eq!(transformer.layers.len(), config.num_hidden_layers);
2088
2089            // Verify forward pass produces finite output
2090            let tokens = vec![1u32, 2, 3];
2091            let logits = transformer.forward(&tokens);
2092            assert_eq!(logits.len(), 3 * config.vocab_size);
2093            assert!(
2094                logits.data().iter().all(|v| v.is_finite()),
2095                "BF16-loaded model should produce finite outputs"
2096            );
2097        }
2098
2099        #[test]
2100        fn test_ssc024_from_safetensors_single_file_path() {
2101            let dir = TempDir::new().expect("create temp dir");
2102            let file_path = create_tiny_safetensors(dir.path());
2103            let config = TransformerConfig::tiny();
2104
2105            // Pass the file path directly, not the directory
2106            let result = Transformer::from_safetensors(&file_path, &config);
2107            assert!(
2108                result.is_ok(),
2109                "Direct file path should work: {}",
2110                result.as_ref().err().map_or(String::new(), std::string::ToString::to_string)
2111            );
2112        }
2113
2114        #[test]
2115        fn test_ssc024_loaded_model_forward_produces_finite() {
2116            let dir = TempDir::new().expect("create temp dir");
2117            create_tiny_safetensors(dir.path());
2118            let config = TransformerConfig::tiny();
2119
2120            let transformer =
2121                Transformer::from_safetensors(dir.path(), &config).expect("loading should succeed");
2122
2123            // Run forward pass
2124            let tokens = vec![0u32, 5, 42, 99];
2125            let logits = transformer.forward(&tokens);
2126
2127            assert_eq!(logits.len(), tokens.len() * config.vocab_size);
2128            let data = logits.data();
2129            let nan_count = data.iter().filter(|v| v.is_nan()).count();
2130            let inf_count = data.iter().filter(|v| v.is_infinite()).count();
2131            assert_eq!(nan_count, 0, "Loaded model output must not contain NaN");
2132            assert_eq!(inf_count, 0, "Loaded model output must not contain Inf");
2133        }
2134
2135        // -----------------------------------------------------------------
2136        // Error case: no SafeTensors files
2137        // -----------------------------------------------------------------
2138
2139        #[test]
2140        fn test_ssc024_from_safetensors_no_files() {
2141            let dir = TempDir::new().expect("create temp dir");
2142            let config = TransformerConfig::tiny();
2143
2144            let result = Transformer::from_safetensors(dir.path(), &config);
2145            assert!(result.is_err());
2146            let err_msg = match result {
2147                Err(e) => e.to_string(),
2148                Ok(_) => panic!("expected error"),
2149            };
2150            assert!(
2151                err_msg.contains("No SafeTensors files"),
2152                "Error should mention missing files: {err_msg}"
2153            );
2154        }
2155
2156        // -----------------------------------------------------------------
2157        // Error case: shape mismatch
2158        // -----------------------------------------------------------------
2159
2160        #[test]
2161        fn test_ssc024_from_safetensors_wrong_embedding_shape() {
2162            let dir = TempDir::new().expect("create temp dir");
2163            let config = TransformerConfig::tiny();
2164            let hidden = config.hidden_size;
2165
2166            // Create a file with wrong embedding shape
2167            let wrong_embed_bytes: Vec<u8> =
2168                std::iter::repeat_n(0.01_f32, 42).flat_map(f32::to_le_bytes).collect();
2169
2170            // We need at least embedding + norm + 2 layers to pass validate_weights.
2171            // But the embedding shape is wrong, so validate_weight_shapes should catch it.
2172            // Actually, we need ALL required keys for validate_weights to pass first.
2173            // Let's create a full set but with wrong embedding size.
2174            let kv_hidden = config.num_kv_heads * config.head_dim();
2175            let intermediate = config.intermediate_size;
2176
2177            let mut td: Vec<(String, Vec<u8>, Vec<usize>)> = Vec::new();
2178
2179            let make_f32 = |n: usize, val: f32| -> Vec<u8> {
2180                std::iter::repeat_n(val, n).flat_map(f32::to_le_bytes).collect()
2181            };
2182
2183            // WRONG: embedding has 42 elements instead of vocab * hidden
2184            td.push(("model.embed_tokens.weight".to_string(), wrong_embed_bytes, vec![42]));
2185            td.push(("model.norm.weight".to_string(), make_f32(hidden, 1.0), vec![hidden]));
2186
2187            for i in 0..config.num_hidden_layers {
2188                let p = format!("model.layers.{i}");
2189                td.push((
2190                    format!("{p}.input_layernorm.weight"),
2191                    make_f32(hidden, 1.0),
2192                    vec![hidden],
2193                ));
2194                td.push((
2195                    format!("{p}.post_attention_layernorm.weight"),
2196                    make_f32(hidden, 1.0),
2197                    vec![hidden],
2198                ));
2199                td.push((
2200                    format!("{p}.self_attn.q_proj.weight"),
2201                    make_f32(hidden * hidden, 0.01),
2202                    vec![hidden, hidden],
2203                ));
2204                td.push((
2205                    format!("{p}.self_attn.k_proj.weight"),
2206                    make_f32(hidden * kv_hidden, 0.01),
2207                    vec![kv_hidden, hidden],
2208                ));
2209                td.push((
2210                    format!("{p}.self_attn.v_proj.weight"),
2211                    make_f32(hidden * kv_hidden, 0.01),
2212                    vec![kv_hidden, hidden],
2213                ));
2214                td.push((
2215                    format!("{p}.self_attn.o_proj.weight"),
2216                    make_f32(hidden * hidden, 0.01),
2217                    vec![hidden, hidden],
2218                ));
2219                td.push((
2220                    format!("{p}.mlp.gate_proj.weight"),
2221                    make_f32(hidden * intermediate, 0.01),
2222                    vec![intermediate, hidden],
2223                ));
2224                td.push((
2225                    format!("{p}.mlp.up_proj.weight"),
2226                    make_f32(hidden * intermediate, 0.01),
2227                    vec![intermediate, hidden],
2228                ));
2229                td.push((
2230                    format!("{p}.mlp.down_proj.weight"),
2231                    make_f32(intermediate * hidden, 0.01),
2232                    vec![hidden, intermediate],
2233                ));
2234            }
2235
2236            let views: Vec<TensorView<'_>> = td
2237                .iter()
2238                .map(|(_, bytes, shape)| {
2239                    TensorView::new(Dtype::F32, shape.clone(), bytes).expect("view")
2240                })
2241                .collect();
2242            let named: Vec<(&str, &TensorView<'_>)> =
2243                td.iter().zip(views.iter()).map(|((n, _, _), v)| (n.as_str(), v)).collect();
2244
2245            let file_path = dir.path().join("model.safetensors");
2246            let serialized =
2247                serialize(named, None::<std::collections::HashMap<String, String>>).expect("ser");
2248            std::fs::write(&file_path, serialized).expect("write");
2249
2250            let result = Transformer::from_safetensors(dir.path(), &config);
2251            assert!(result.is_err(), "Wrong embedding shape should fail");
2252            let err_msg = match result {
2253                Err(e) => e.to_string(),
2254                Ok(_) => panic!("expected error"),
2255            };
2256            assert!(
2257                err_msg.contains("Shape mismatch") || err_msg.contains("embed_tokens"),
2258                "Error should indicate shape issue: {err_msg}"
2259            );
2260        }
2261
2262        // -----------------------------------------------------------------
2263        // Error case: NaN in weights
2264        // -----------------------------------------------------------------
2265
2266        #[test]
2267        fn test_ssc024_from_safetensors_nan_detection() {
2268            let dir = TempDir::new().expect("create temp dir");
2269            let config = TransformerConfig::tiny();
2270            let hidden = config.hidden_size;
2271            let kv_hidden = config.num_kv_heads * config.head_dim();
2272            let intermediate = config.intermediate_size;
2273            let vocab = config.vocab_size;
2274
2275            let mut td: Vec<(String, Vec<u8>, Vec<usize>)> = Vec::new();
2276
2277            let make_f32 = |n: usize, val: f32| -> Vec<u8> {
2278                std::iter::repeat_n(val, n).flat_map(f32::to_le_bytes).collect()
2279            };
2280
2281            // Embedding with NaN injected
2282            let mut embed_vals: Vec<f32> = vec![0.01; vocab * hidden];
2283            embed_vals[42] = f32::NAN;
2284            let embed_bytes: Vec<u8> = embed_vals.iter().flat_map(|v| v.to_le_bytes()).collect();
2285
2286            td.push(("model.embed_tokens.weight".to_string(), embed_bytes, vec![vocab, hidden]));
2287            td.push(("model.norm.weight".to_string(), make_f32(hidden, 1.0), vec![hidden]));
2288
2289            for i in 0..config.num_hidden_layers {
2290                let p = format!("model.layers.{i}");
2291                td.push((
2292                    format!("{p}.input_layernorm.weight"),
2293                    make_f32(hidden, 1.0),
2294                    vec![hidden],
2295                ));
2296                td.push((
2297                    format!("{p}.post_attention_layernorm.weight"),
2298                    make_f32(hidden, 1.0),
2299                    vec![hidden],
2300                ));
2301                td.push((
2302                    format!("{p}.self_attn.q_proj.weight"),
2303                    make_f32(hidden * hidden, 0.01),
2304                    vec![hidden, hidden],
2305                ));
2306                td.push((
2307                    format!("{p}.self_attn.k_proj.weight"),
2308                    make_f32(hidden * kv_hidden, 0.01),
2309                    vec![kv_hidden, hidden],
2310                ));
2311                td.push((
2312                    format!("{p}.self_attn.v_proj.weight"),
2313                    make_f32(hidden * kv_hidden, 0.01),
2314                    vec![kv_hidden, hidden],
2315                ));
2316                td.push((
2317                    format!("{p}.self_attn.o_proj.weight"),
2318                    make_f32(hidden * hidden, 0.01),
2319                    vec![hidden, hidden],
2320                ));
2321                td.push((
2322                    format!("{p}.mlp.gate_proj.weight"),
2323                    make_f32(hidden * intermediate, 0.01),
2324                    vec![intermediate, hidden],
2325                ));
2326                td.push((
2327                    format!("{p}.mlp.up_proj.weight"),
2328                    make_f32(hidden * intermediate, 0.01),
2329                    vec![intermediate, hidden],
2330                ));
2331                td.push((
2332                    format!("{p}.mlp.down_proj.weight"),
2333                    make_f32(intermediate * hidden, 0.01),
2334                    vec![hidden, intermediate],
2335                ));
2336            }
2337
2338            let views: Vec<TensorView<'_>> = td
2339                .iter()
2340                .map(|(_, bytes, shape)| {
2341                    TensorView::new(Dtype::F32, shape.clone(), bytes).expect("view")
2342                })
2343                .collect();
2344            let named: Vec<(&str, &TensorView<'_>)> =
2345                td.iter().zip(views.iter()).map(|((n, _, _), v)| (n.as_str(), v)).collect();
2346
2347            let file_path = dir.path().join("model.safetensors");
2348            let serialized =
2349                serialize(named, None::<std::collections::HashMap<String, String>>).expect("ser");
2350            std::fs::write(&file_path, serialized).expect("write");
2351
2352            let result = Transformer::from_safetensors(dir.path(), &config);
2353            assert!(result.is_err(), "NaN in weights should fail");
2354            let err_msg = match result {
2355                Err(e) => e.to_string(),
2356                Ok(_) => panic!("expected error"),
2357            };
2358            assert!(err_msg.contains("NaN"), "Error should mention NaN: {err_msg}");
2359        }
2360
2361        // -----------------------------------------------------------------
2362        // Error case: Inf in weights
2363        // -----------------------------------------------------------------
2364
2365        #[test]
2366        fn test_ssc024_from_safetensors_inf_detection() {
2367            let dir = TempDir::new().expect("create temp dir");
2368            let config = TransformerConfig::tiny();
2369            let hidden = config.hidden_size;
2370            let kv_hidden = config.num_kv_heads * config.head_dim();
2371            let intermediate = config.intermediate_size;
2372            let vocab = config.vocab_size;
2373
2374            let mut td: Vec<(String, Vec<u8>, Vec<usize>)> = Vec::new();
2375
2376            let make_f32 = |n: usize, val: f32| -> Vec<u8> {
2377                std::iter::repeat_n(val, n).flat_map(f32::to_le_bytes).collect()
2378            };
2379
2380            // norm with Inf injected
2381            let mut norm_vals: Vec<f32> = vec![1.0; hidden];
2382            norm_vals[0] = f32::INFINITY;
2383            let norm_bytes: Vec<u8> = norm_vals.iter().flat_map(|v| v.to_le_bytes()).collect();
2384
2385            td.push((
2386                "model.embed_tokens.weight".to_string(),
2387                make_f32(vocab * hidden, 0.01),
2388                vec![vocab, hidden],
2389            ));
2390            td.push(("model.norm.weight".to_string(), norm_bytes, vec![hidden]));
2391
2392            for i in 0..config.num_hidden_layers {
2393                let p = format!("model.layers.{i}");
2394                td.push((
2395                    format!("{p}.input_layernorm.weight"),
2396                    make_f32(hidden, 1.0),
2397                    vec![hidden],
2398                ));
2399                td.push((
2400                    format!("{p}.post_attention_layernorm.weight"),
2401                    make_f32(hidden, 1.0),
2402                    vec![hidden],
2403                ));
2404                td.push((
2405                    format!("{p}.self_attn.q_proj.weight"),
2406                    make_f32(hidden * hidden, 0.01),
2407                    vec![hidden, hidden],
2408                ));
2409                td.push((
2410                    format!("{p}.self_attn.k_proj.weight"),
2411                    make_f32(hidden * kv_hidden, 0.01),
2412                    vec![kv_hidden, hidden],
2413                ));
2414                td.push((
2415                    format!("{p}.self_attn.v_proj.weight"),
2416                    make_f32(hidden * kv_hidden, 0.01),
2417                    vec![kv_hidden, hidden],
2418                ));
2419                td.push((
2420                    format!("{p}.self_attn.o_proj.weight"),
2421                    make_f32(hidden * hidden, 0.01),
2422                    vec![hidden, hidden],
2423                ));
2424                td.push((
2425                    format!("{p}.mlp.gate_proj.weight"),
2426                    make_f32(hidden * intermediate, 0.01),
2427                    vec![intermediate, hidden],
2428                ));
2429                td.push((
2430                    format!("{p}.mlp.up_proj.weight"),
2431                    make_f32(hidden * intermediate, 0.01),
2432                    vec![intermediate, hidden],
2433                ));
2434                td.push((
2435                    format!("{p}.mlp.down_proj.weight"),
2436                    make_f32(intermediate * hidden, 0.01),
2437                    vec![hidden, intermediate],
2438                ));
2439            }
2440
2441            let views: Vec<TensorView<'_>> = td
2442                .iter()
2443                .map(|(_, bytes, shape)| {
2444                    TensorView::new(Dtype::F32, shape.clone(), bytes).expect("view")
2445                })
2446                .collect();
2447            let named: Vec<(&str, &TensorView<'_>)> =
2448                td.iter().zip(views.iter()).map(|((n, _, _), v)| (n.as_str(), v)).collect();
2449
2450            let file_path = dir.path().join("model.safetensors");
2451            let serialized =
2452                serialize(named, None::<std::collections::HashMap<String, String>>).expect("ser");
2453            std::fs::write(&file_path, serialized).expect("write");
2454
2455            let result = Transformer::from_safetensors(dir.path(), &config);
2456            assert!(result.is_err(), "Inf in weights should fail");
2457            let err_msg = match result {
2458                Err(e) => e.to_string(),
2459                Ok(_) => panic!("expected error"),
2460            };
2461            assert!(err_msg.contains("Inf"), "Error should mention Inf: {err_msg}");
2462        }
2463
2464        // -----------------------------------------------------------------
2465        // Error case: missing layer weights (wrong layer count)
2466        // -----------------------------------------------------------------
2467
2468        #[test]
2469        fn test_ssc024_from_safetensors_missing_layer() {
2470            let dir = TempDir::new().expect("create temp dir");
2471            // Create a file with 2 layers of weights
2472            create_tiny_safetensors(dir.path());
2473
2474            // But try to load with config expecting 3 layers
2475            let mut config = TransformerConfig::tiny();
2476            config.num_hidden_layers = 3;
2477
2478            let result = Transformer::from_safetensors(dir.path(), &config);
2479            assert!(result.is_err(), "Missing layer 2 should fail");
2480            let err_msg = match result {
2481                Err(e) => e.to_string(),
2482                Ok(_) => panic!("expected error"),
2483            };
2484            assert!(
2485                err_msg.contains("Missing") || err_msg.contains("layers.2"),
2486                "Error should mention missing layer: {err_msg}"
2487            );
2488        }
2489
2490        // -----------------------------------------------------------------
2491        // Error case: wrong attention projection shape
2492        // -----------------------------------------------------------------
2493
2494        #[test]
2495        fn test_ssc024_from_safetensors_wrong_q_proj_shape() {
2496            let dir = TempDir::new().expect("create temp dir");
2497            let config = TransformerConfig::tiny();
2498            let hidden = config.hidden_size;
2499            let q_dim = config.q_dim();
2500            let kv_hidden = config.num_kv_heads * config.head_dim();
2501            let intermediate = config.intermediate_size;
2502            let vocab = config.vocab_size;
2503
2504            let mut td: Vec<(String, Vec<u8>, Vec<usize>)> = Vec::new();
2505
2506            let make_f32 = |n: usize, val: f32| -> Vec<u8> {
2507                std::iter::repeat_n(val, n).flat_map(f32::to_le_bytes).collect()
2508            };
2509
2510            td.push((
2511                "model.embed_tokens.weight".to_string(),
2512                make_f32(vocab * hidden, 0.01),
2513                vec![vocab, hidden],
2514            ));
2515            td.push(("model.norm.weight".to_string(), make_f32(hidden, 1.0), vec![hidden]));
2516
2517            for i in 0..config.num_hidden_layers {
2518                let p = format!("model.layers.{i}");
2519                td.push((
2520                    format!("{p}.input_layernorm.weight"),
2521                    make_f32(hidden, 1.0),
2522                    vec![hidden],
2523                ));
2524                td.push((
2525                    format!("{p}.post_attention_layernorm.weight"),
2526                    make_f32(hidden, 1.0),
2527                    vec![hidden],
2528                ));
2529
2530                // WRONG: q_proj has 7 elements instead of q_dim*hidden
2531                if i == 0 {
2532                    td.push((format!("{p}.self_attn.q_proj.weight"), make_f32(7, 0.01), vec![7]));
2533                } else {
2534                    td.push((
2535                        format!("{p}.self_attn.q_proj.weight"),
2536                        make_f32(q_dim * hidden, 0.01),
2537                        vec![q_dim, hidden],
2538                    ));
2539                }
2540                td.push((
2541                    format!("{p}.self_attn.k_proj.weight"),
2542                    make_f32(kv_hidden * hidden, 0.01),
2543                    vec![kv_hidden, hidden],
2544                ));
2545                td.push((
2546                    format!("{p}.self_attn.v_proj.weight"),
2547                    make_f32(kv_hidden * hidden, 0.01),
2548                    vec![kv_hidden, hidden],
2549                ));
2550                td.push((
2551                    format!("{p}.self_attn.o_proj.weight"),
2552                    make_f32(hidden * q_dim, 0.01),
2553                    vec![hidden, q_dim],
2554                ));
2555                td.push((
2556                    format!("{p}.mlp.gate_proj.weight"),
2557                    make_f32(hidden * intermediate, 0.01),
2558                    vec![intermediate, hidden],
2559                ));
2560                td.push((
2561                    format!("{p}.mlp.up_proj.weight"),
2562                    make_f32(hidden * intermediate, 0.01),
2563                    vec![intermediate, hidden],
2564                ));
2565                td.push((
2566                    format!("{p}.mlp.down_proj.weight"),
2567                    make_f32(intermediate * hidden, 0.01),
2568                    vec![hidden, intermediate],
2569                ));
2570            }
2571
2572            let views: Vec<TensorView<'_>> = td
2573                .iter()
2574                .map(|(_, bytes, shape)| {
2575                    TensorView::new(Dtype::F32, shape.clone(), bytes).expect("view")
2576                })
2577                .collect();
2578            let named: Vec<(&str, &TensorView<'_>)> =
2579                td.iter().zip(views.iter()).map(|((n, _, _), v)| (n.as_str(), v)).collect();
2580
2581            let file_path = dir.path().join("model.safetensors");
2582            let serialized =
2583                serialize(named, None::<std::collections::HashMap<String, String>>).expect("ser");
2584            std::fs::write(&file_path, serialized).expect("write");
2585
2586            let result = Transformer::from_safetensors(dir.path(), &config);
2587            assert!(result.is_err(), "Wrong q_proj shape should fail");
2588            let err_msg = match result {
2589                Err(e) => e.to_string(),
2590                Ok(_) => panic!("expected error"),
2591            };
2592            assert!(
2593                err_msg.contains("Shape mismatch") && err_msg.contains("q_proj"),
2594                "Error should mention q_proj shape mismatch: {err_msg}"
2595            );
2596        }
2597
2598        // -----------------------------------------------------------------
2599        // Validate weight_shapes helper independently
2600        // -----------------------------------------------------------------
2601
2602        #[test]
2603        fn test_ssc024_validate_weight_shapes_success() {
2604            let config = TransformerConfig::tiny();
2605            let hidden = config.hidden_size;
2606            let kv_hidden = config.num_kv_heads * config.head_dim();
2607            let intermediate = config.intermediate_size;
2608            let vocab = config.vocab_size;
2609
2610            let mut weights = HashMap::new();
2611            weights.insert(
2612                "model.embed_tokens.weight".to_string(),
2613                Tensor::from_vec(vec![0.1; vocab * hidden], true),
2614            );
2615            weights
2616                .insert("model.norm.weight".to_string(), Tensor::from_vec(vec![1.0; hidden], true));
2617
2618            for i in 0..config.num_hidden_layers {
2619                let p = format!("model.layers.{i}");
2620                weights.insert(
2621                    format!("{p}.input_layernorm.weight"),
2622                    Tensor::from_vec(vec![1.0; hidden], true),
2623                );
2624                weights.insert(
2625                    format!("{p}.post_attention_layernorm.weight"),
2626                    Tensor::from_vec(vec![1.0; hidden], true),
2627                );
2628                weights.insert(
2629                    format!("{p}.self_attn.q_proj.weight"),
2630                    Tensor::from_vec(vec![0.1; hidden * hidden], true),
2631                );
2632                weights.insert(
2633                    format!("{p}.self_attn.k_proj.weight"),
2634                    Tensor::from_vec(vec![0.1; hidden * kv_hidden], true),
2635                );
2636                weights.insert(
2637                    format!("{p}.self_attn.v_proj.weight"),
2638                    Tensor::from_vec(vec![0.1; hidden * kv_hidden], true),
2639                );
2640                weights.insert(
2641                    format!("{p}.self_attn.o_proj.weight"),
2642                    Tensor::from_vec(vec![0.1; hidden * hidden], true),
2643                );
2644                weights.insert(
2645                    format!("{p}.mlp.gate_proj.weight"),
2646                    Tensor::from_vec(vec![0.1; hidden * intermediate], true),
2647                );
2648                weights.insert(
2649                    format!("{p}.mlp.up_proj.weight"),
2650                    Tensor::from_vec(vec![0.1; hidden * intermediate], true),
2651                );
2652                weights.insert(
2653                    format!("{p}.mlp.down_proj.weight"),
2654                    Tensor::from_vec(vec![0.1; intermediate * hidden], true),
2655                );
2656            }
2657
2658            let result = Transformer::validate_weight_shapes(&weights, &config);
2659            assert!(
2660                result.is_ok(),
2661                "Valid shapes should pass: {}",
2662                result.as_ref().err().map_or(String::new(), std::string::ToString::to_string)
2663            );
2664        }
2665
2666        #[test]
2667        fn test_ssc024_validate_weight_shapes_wrong_norm() {
2668            let config = TransformerConfig::tiny();
2669            let hidden = config.hidden_size;
2670            let vocab = config.vocab_size;
2671
2672            let mut weights = HashMap::new();
2673            weights.insert(
2674                "model.embed_tokens.weight".to_string(),
2675                Tensor::from_vec(vec![0.1; vocab * hidden], true),
2676            );
2677            // Wrong norm size: 3 instead of hidden
2678            weights.insert("model.norm.weight".to_string(), Tensor::from_vec(vec![1.0; 3], true));
2679
2680            let result = Transformer::validate_weight_shapes(&weights, &config);
2681            assert!(result.is_err());
2682            let err_msg = match result {
2683                Err(e) => e.to_string(),
2684                Ok(()) => panic!("expected error"),
2685            };
2686            assert!(err_msg.contains("model.norm.weight"));
2687        }
2688
2689        // -----------------------------------------------------------------
2690        // Validate weight_values helper independently
2691        // -----------------------------------------------------------------
2692
2693        #[test]
2694        fn test_ssc024_validate_weight_values_clean() {
2695            let mut weights = HashMap::new();
2696            weights.insert("a".to_string(), Tensor::from_vec(vec![0.1, 0.2, 0.3], true));
2697            weights.insert("b".to_string(), Tensor::from_vec(vec![1.0, -1.0, 0.0], true));
2698
2699            let result = Transformer::validate_weight_values(&weights);
2700            assert!(result.is_ok());
2701        }
2702
2703        #[test]
2704        fn test_ssc024_validate_weight_values_nan() {
2705            let mut weights = HashMap::new();
2706            weights.insert("clean".to_string(), Tensor::from_vec(vec![0.1, 0.2], true));
2707            weights
2708                .insert("poisoned".to_string(), Tensor::from_vec(vec![0.1, f32::NAN, 0.3], true));
2709
2710            let result = Transformer::validate_weight_values(&weights);
2711            assert!(result.is_err());
2712            let err_msg = match result {
2713                Err(e) => e.to_string(),
2714                Ok(()) => panic!("expected error"),
2715            };
2716            assert!(err_msg.contains("NaN"));
2717            assert!(err_msg.contains("poisoned"));
2718        }
2719
2720        #[test]
2721        fn test_ssc024_validate_weight_values_inf() {
2722            let mut weights = HashMap::new();
2723            weights.insert("w".to_string(), Tensor::from_vec(vec![f32::NEG_INFINITY, 0.2], true));
2724
2725            let result = Transformer::validate_weight_values(&weights);
2726            assert!(result.is_err());
2727            let err_msg = match result {
2728                Err(e) => e.to_string(),
2729                Ok(()) => panic!("expected error"),
2730            };
2731            assert!(err_msg.contains("Inf"));
2732        }
2733
2734        /// GH-262: Qwen3-4B shape validation with q_dim != hidden_size.
2735        ///
2736        /// Uses a 1-layer config mimicking Qwen3-4B dimensions to verify
2737        /// validate_weight_shapes accepts the correct q_dim-based shapes.
2738        #[test]
2739        fn test_gh262_qwen3_4b_weight_shapes_q_dim_ne_hidden() {
2740            // Minimal Qwen3-like config: q_dim (128) != hidden_size (80)
2741            let config = TransformerConfig {
2742                hidden_size: 80,
2743                num_attention_heads: 4,
2744                num_kv_heads: 2,
2745                intermediate_size: 128,
2746                num_hidden_layers: 1,
2747                vocab_size: 256,
2748                max_position_embeddings: 512,
2749                rms_norm_eps: 1e-6,
2750                rope_theta: 10000.0,
2751                use_bias: false,
2752                head_dim_override: Some(32), // head_dim=32, so q_dim = 4*32 = 128, hidden=80
2753                architecture: crate::transformer::config::ModelArchitecture::Decoder,
2754                hf_architecture: None,
2755                hf_model_type: None,
2756                tie_word_embeddings: false,
2757            };
2758
2759            let hidden = config.hidden_size; // 80
2760            let q_dim = config.q_dim(); // 4 * 32 = 128
2761            let kv_hidden = config.num_kv_heads * config.head_dim(); // 2 * 32 = 64
2762            let intermediate = config.intermediate_size; // 128
2763            let vocab = config.vocab_size; // 256
2764
2765            // Verify q_dim != hidden (the Qwen3-4B characteristic)
2766            assert_ne!(q_dim, hidden, "test requires q_dim != hidden_size");
2767
2768            let mut weights = HashMap::new();
2769            weights.insert(
2770                "model.embed_tokens.weight".to_string(),
2771                Tensor::from_vec(vec![0.1; vocab * hidden], true),
2772            );
2773            weights
2774                .insert("model.norm.weight".to_string(), Tensor::from_vec(vec![1.0; hidden], true));
2775
2776            let p = "model.layers.0";
2777            weights.insert(
2778                format!("{p}.input_layernorm.weight"),
2779                Tensor::from_vec(vec![1.0; hidden], true),
2780            );
2781            weights.insert(
2782                format!("{p}.post_attention_layernorm.weight"),
2783                Tensor::from_vec(vec![1.0; hidden], true),
2784            );
2785            // Q: [q_dim, hidden] — NOT [hidden, hidden]
2786            weights.insert(
2787                format!("{p}.self_attn.q_proj.weight"),
2788                Tensor::from_vec(vec![0.1; q_dim * hidden], true),
2789            );
2790            // K: [kv_hidden, hidden]
2791            weights.insert(
2792                format!("{p}.self_attn.k_proj.weight"),
2793                Tensor::from_vec(vec![0.1; kv_hidden * hidden], true),
2794            );
2795            // V: [kv_hidden, hidden]
2796            weights.insert(
2797                format!("{p}.self_attn.v_proj.weight"),
2798                Tensor::from_vec(vec![0.1; kv_hidden * hidden], true),
2799            );
2800            // O: [hidden, q_dim] — NOT [hidden, hidden]
2801            weights.insert(
2802                format!("{p}.self_attn.o_proj.weight"),
2803                Tensor::from_vec(vec![0.1; hidden * q_dim], true),
2804            );
2805            weights.insert(
2806                format!("{p}.mlp.gate_proj.weight"),
2807                Tensor::from_vec(vec![0.1; hidden * intermediate], true),
2808            );
2809            weights.insert(
2810                format!("{p}.mlp.up_proj.weight"),
2811                Tensor::from_vec(vec![0.1; hidden * intermediate], true),
2812            );
2813            weights.insert(
2814                format!("{p}.mlp.down_proj.weight"),
2815                Tensor::from_vec(vec![0.1; intermediate * hidden], true),
2816            );
2817
2818            // Should pass: shapes use q_dim for Q/O projections
2819            let result = Transformer::validate_weight_shapes(&weights, &config);
2820            assert!(
2821                result.is_ok(),
2822                "Qwen3-like shapes (q_dim={q_dim} != hidden={hidden}) should validate: {:?}",
2823                result.err()
2824            );
2825
2826            // Should also construct successfully via from_params
2827            let model = Transformer::from_params(&config, &weights);
2828            assert!(model.is_some(), "Qwen3-like model with q_dim != hidden should construct");
2829        }
2830
2831        /// GH-262: Using hidden_size instead of q_dim for q_proj must fail.
2832        #[test]
2833        fn test_gh262_wrong_q_proj_size_hidden_instead_of_q_dim() {
2834            let config = TransformerConfig {
2835                hidden_size: 80,
2836                num_attention_heads: 4,
2837                num_kv_heads: 2,
2838                intermediate_size: 128,
2839                num_hidden_layers: 1,
2840                vocab_size: 256,
2841                max_position_embeddings: 512,
2842                rms_norm_eps: 1e-6,
2843                rope_theta: 10000.0,
2844                use_bias: false,
2845                head_dim_override: Some(32), // q_dim=128, hidden=80
2846                architecture: crate::transformer::config::ModelArchitecture::Decoder,
2847                hf_architecture: None,
2848                hf_model_type: None,
2849                tie_word_embeddings: false,
2850            };
2851
2852            let hidden = config.hidden_size; // 80
2853            let kv_hidden = config.num_kv_heads * config.head_dim(); // 64
2854            let intermediate = config.intermediate_size;
2855            let vocab = config.vocab_size;
2856
2857            let mut weights = HashMap::new();
2858            weights.insert(
2859                "model.embed_tokens.weight".to_string(),
2860                Tensor::from_vec(vec![0.1; vocab * hidden], true),
2861            );
2862            weights
2863                .insert("model.norm.weight".to_string(), Tensor::from_vec(vec![1.0; hidden], true));
2864
2865            let p = "model.layers.0";
2866            weights.insert(
2867                format!("{p}.input_layernorm.weight"),
2868                Tensor::from_vec(vec![1.0; hidden], true),
2869            );
2870            weights.insert(
2871                format!("{p}.post_attention_layernorm.weight"),
2872                Tensor::from_vec(vec![1.0; hidden], true),
2873            );
2874            // BUG: q_proj uses hidden*hidden (6400) instead of q_dim*hidden (10240)
2875            weights.insert(
2876                format!("{p}.self_attn.q_proj.weight"),
2877                Tensor::from_vec(vec![0.1; hidden * hidden], true),
2878            );
2879            weights.insert(
2880                format!("{p}.self_attn.k_proj.weight"),
2881                Tensor::from_vec(vec![0.1; kv_hidden * hidden], true),
2882            );
2883            weights.insert(
2884                format!("{p}.self_attn.v_proj.weight"),
2885                Tensor::from_vec(vec![0.1; kv_hidden * hidden], true),
2886            );
2887            weights.insert(
2888                format!("{p}.self_attn.o_proj.weight"),
2889                Tensor::from_vec(vec![0.1; hidden * hidden], true),
2890            );
2891            weights.insert(
2892                format!("{p}.mlp.gate_proj.weight"),
2893                Tensor::from_vec(vec![0.1; hidden * intermediate], true),
2894            );
2895            weights.insert(
2896                format!("{p}.mlp.up_proj.weight"),
2897                Tensor::from_vec(vec![0.1; hidden * intermediate], true),
2898            );
2899            weights.insert(
2900                format!("{p}.mlp.down_proj.weight"),
2901                Tensor::from_vec(vec![0.1; intermediate * hidden], true),
2902            );
2903
2904            // Must fail: q_proj has wrong size
2905            let result = Transformer::validate_weight_shapes(&weights, &config);
2906            assert!(result.is_err(), "hidden*hidden q_proj should fail when q_dim != hidden");
2907            let err_msg = result.err().map(|e| e.to_string()).unwrap_or_default();
2908            assert!(
2909                err_msg.contains("q_proj") && err_msg.contains("Shape mismatch"),
2910                "Error should mention q_proj shape mismatch, got: {err_msg}"
2911            );
2912        }
2913
2914        // -----------------------------------------------------------------
2915        // Name mapping integration: Qwen2 bias tensors are preserved
2916        // -----------------------------------------------------------------
2917
2918        #[test]
2919        fn test_ssc024_from_safetensors_with_extra_bias_tensors() {
2920            // Qwen2 models have bias tensors that are loaded alongside weights.
2921            // from_params ignores them (doesn't look for bias keys), but they
2922            // should not cause errors.
2923            let dir = TempDir::new().expect("create temp dir");
2924            let config = TransformerConfig::tiny();
2925            let hidden = config.hidden_size;
2926            let kv_hidden = config.num_kv_heads * config.head_dim();
2927            let intermediate = config.intermediate_size;
2928            let vocab = config.vocab_size;
2929
2930            let mut td: Vec<(String, Vec<u8>, Vec<usize>)> = Vec::new();
2931
2932            let make_f32 = |n: usize, val: f32| -> Vec<u8> {
2933                std::iter::repeat_n(val, n).flat_map(f32::to_le_bytes).collect()
2934            };
2935
2936            td.push((
2937                "model.embed_tokens.weight".to_string(),
2938                make_f32(vocab * hidden, 0.01),
2939                vec![vocab, hidden],
2940            ));
2941            td.push(("model.norm.weight".to_string(), make_f32(hidden, 1.0), vec![hidden]));
2942
2943            for i in 0..config.num_hidden_layers {
2944                let p = format!("model.layers.{i}");
2945                td.push((
2946                    format!("{p}.input_layernorm.weight"),
2947                    make_f32(hidden, 1.0),
2948                    vec![hidden],
2949                ));
2950                td.push((
2951                    format!("{p}.post_attention_layernorm.weight"),
2952                    make_f32(hidden, 1.0),
2953                    vec![hidden],
2954                ));
2955                td.push((
2956                    format!("{p}.self_attn.q_proj.weight"),
2957                    make_f32(hidden * hidden, 0.01),
2958                    vec![hidden, hidden],
2959                ));
2960                td.push((
2961                    format!("{p}.self_attn.k_proj.weight"),
2962                    make_f32(hidden * kv_hidden, 0.01),
2963                    vec![kv_hidden, hidden],
2964                ));
2965                td.push((
2966                    format!("{p}.self_attn.v_proj.weight"),
2967                    make_f32(hidden * kv_hidden, 0.01),
2968                    vec![kv_hidden, hidden],
2969                ));
2970                td.push((
2971                    format!("{p}.self_attn.o_proj.weight"),
2972                    make_f32(hidden * hidden, 0.01),
2973                    vec![hidden, hidden],
2974                ));
2975                td.push((
2976                    format!("{p}.mlp.gate_proj.weight"),
2977                    make_f32(hidden * intermediate, 0.01),
2978                    vec![intermediate, hidden],
2979                ));
2980                td.push((
2981                    format!("{p}.mlp.up_proj.weight"),
2982                    make_f32(hidden * intermediate, 0.01),
2983                    vec![intermediate, hidden],
2984                ));
2985                td.push((
2986                    format!("{p}.mlp.down_proj.weight"),
2987                    make_f32(intermediate * hidden, 0.01),
2988                    vec![hidden, intermediate],
2989                ));
2990
2991                // Qwen2-style bias tensors
2992                td.push((
2993                    format!("{p}.self_attn.q_proj.bias"),
2994                    make_f32(hidden, 0.0),
2995                    vec![hidden],
2996                ));
2997                td.push((
2998                    format!("{p}.self_attn.k_proj.bias"),
2999                    make_f32(kv_hidden, 0.0),
3000                    vec![kv_hidden],
3001                ));
3002                td.push((
3003                    format!("{p}.self_attn.v_proj.bias"),
3004                    make_f32(kv_hidden, 0.0),
3005                    vec![kv_hidden],
3006                ));
3007            }
3008
3009            let views: Vec<TensorView<'_>> = td
3010                .iter()
3011                .map(|(_, bytes, shape)| {
3012                    TensorView::new(Dtype::F32, shape.clone(), bytes).expect("view")
3013                })
3014                .collect();
3015            let named: Vec<(&str, &TensorView<'_>)> =
3016                td.iter().zip(views.iter()).map(|((n, _, _), v)| (n.as_str(), v)).collect();
3017
3018            let file_path = dir.path().join("model.safetensors");
3019            let serialized =
3020                serialize(named, None::<std::collections::HashMap<String, String>>).expect("ser");
3021            std::fs::write(&file_path, serialized).expect("write");
3022
3023            // Should succeed even with extra bias tensors
3024            let result = Transformer::from_safetensors(dir.path(), &config);
3025            assert!(
3026                result.is_ok(),
3027                "Extra bias tensors should not cause failure: {}",
3028                result.as_ref().err().map_or(String::new(), std::string::ToString::to_string)
3029            );
3030        }
3031    }
3032}