Skip to main content

entrenar/transformer/
embedding.rs

1//! Embedding layer module
2//!
3//! This module provides token embedding layers for transformer models.
4
5use crate::Tensor;
6use std::collections::HashMap;
7
8/// Embedding layer
9pub struct Embedding {
10    /// Embedding weight (vocab_size x hidden_size)
11    pub weight: Tensor,
12    /// Vocabulary size
13    vocab_size: usize,
14    /// Hidden dimension
15    hidden_size: usize,
16}
17
18impl Embedding {
19    /// Create new embedding layer with initialized weights
20    pub fn new(vocab_size: usize, hidden_size: usize) -> Self {
21        use super::init::{get_init_seed, rand_normal_seeded};
22        // C-INIT-001: normal(0, 0.02) matching HuggingFace LLaMA
23        Self {
24            weight: Tensor::from_vec(
25                rand_normal_seeded(vocab_size * hidden_size, get_init_seed(), "embed_tokens"),
26                true,
27            ),
28            vocab_size,
29            hidden_size,
30        }
31    }
32
33    /// Create from parameters
34    ///
35    /// # Contract (PMAT-326)
36    /// Validates weight.len() == vocab_size * hidden_size.
37    /// Returns None if key is missing or shape is wrong.
38    pub fn from_params(
39        params: &HashMap<String, Tensor>,
40        name: &str,
41        vocab_size: usize,
42        hidden_size: usize,
43    ) -> Option<Self> {
44        let weight = params.get(name)?.clone();
45        let expected = vocab_size * hidden_size;
46        if weight.len() != expected {
47            eprintln!(
48                "[PMAT-326] Embedding '{name}': shape mismatch — got {} elements, expected {expected} ({vocab_size}x{hidden_size})",
49                weight.len()
50            );
51            return None;
52        }
53        Some(Self { weight, vocab_size, hidden_size })
54    }
55
56    /// Forward pass - lookup embeddings for token IDs
57    ///
58    /// # Arguments
59    /// * `token_ids` - Token IDs to look up
60    ///
61    /// # Returns
62    /// Embedded vectors (seq_len * hidden_size, flattened)
63    pub fn forward(&self, token_ids: &[u32]) -> Tensor {
64        contract_pre_embedding_lookup!(token_ids);
65        let mut output = Vec::with_capacity(token_ids.len() * self.hidden_size);
66
67        for &token_id in token_ids {
68            let idx = token_id as usize;
69            if idx >= self.vocab_size {
70                // N-09: OOB token → zeros. Contract: embedding-lookup-v1.yaml
71                eprintln!(
72                    "Warning: Embedding::forward token_id {} >= vocab_size {}. N-09 OOB escape.",
73                    token_id, self.vocab_size
74                );
75                output.extend(std::iter::repeat_n(0.0, self.hidden_size));
76            } else {
77                let start = idx * self.hidden_size;
78                let end = start + self.hidden_size;
79                output.extend_from_slice(
80                    &self.weight.data().as_slice().expect("embedding weight must be contiguous")
81                        [start..end],
82                );
83            }
84        }
85
86        let result = Tensor::from_vec(output, true);
87        contract_post_embedding_lookup!(result.data().as_slice().unwrap_or(&[]));
88        result
89    }
90
91    /// Get vocabulary size
92    pub fn vocab_size(&self) -> usize {
93        self.vocab_size
94    }
95
96    /// Get hidden dimension
97    pub fn hidden_size(&self) -> usize {
98        self.hidden_size
99    }
100}
101
102/// Learned absolute position embedding for encoder models (BERT, RoBERTa, CodeBERT).
103///
104/// Unlike RoPE (used by decoders), encoder models learn a position embedding table
105/// of shape (max_position_embeddings × hidden_size) that is added to token embeddings.
106///
107/// # Contract (ENC-003)
108/// - Output shape: seq_len × hidden_size (same as token embedding)
109/// - Positions beyond max_position_embeddings are clamped to max-1
110/// - Output is added element-wise to token embeddings
111pub struct LearnedPositionEmbedding {
112    /// Position embedding weight (max_positions × hidden_size)
113    pub weight: Tensor,
114    /// Maximum number of positions
115    max_positions: usize,
116    /// Hidden dimension
117    hidden_size: usize,
118}
119
120impl LearnedPositionEmbedding {
121    /// Create new learned position embedding with deterministic initialization
122    pub fn new(max_positions: usize, hidden_size: usize) -> Self {
123        let scale = (1.0 / hidden_size as f32).sqrt();
124        Self {
125            weight: Tensor::from_vec(
126                (0..max_positions * hidden_size)
127                    .map(|i| (i as f32 * 0.0731).sin() * scale)
128                    .collect(),
129                true,
130            ),
131            max_positions,
132            hidden_size,
133        }
134    }
135
136    /// Create from pre-trained parameters
137    pub fn from_params(
138        params: &HashMap<String, Tensor>,
139        name: &str,
140        max_positions: usize,
141        hidden_size: usize,
142    ) -> Option<Self> {
143        let weight = params.get(name)?.clone();
144        let expected = max_positions * hidden_size;
145        if weight.len() != expected {
146            eprintln!(
147                "[ENC-003] LearnedPositionEmbedding '{name}': shape mismatch — \
148                 got {} elements, expected {expected} ({max_positions}×{hidden_size})",
149                weight.len()
150            );
151            return None;
152        }
153        Some(Self { weight, max_positions, hidden_size })
154    }
155
156    /// Forward pass: return position embeddings for positions 0..seq_len
157    ///
158    /// Output is (seq_len × hidden_size, flattened) — add element-wise to token embeddings.
159    pub fn forward(&self, seq_len: usize) -> Tensor {
160        let clamped_len = seq_len.min(self.max_positions);
161        let weight_slice = &self.weight.data().as_slice().expect("position weight contiguous")
162            [..clamped_len * self.hidden_size];
163        // For positions beyond max, repeat the last position embedding
164        if seq_len <= self.max_positions {
165            Tensor::from_vec(weight_slice.to_vec(), true)
166        } else {
167            let mut output = weight_slice.to_vec();
168            let last_start = (self.max_positions - 1) * self.hidden_size;
169            let last_end = last_start + self.hidden_size;
170            let last_pos = &self.weight.data().as_slice().expect("position weight contiguous")
171                [last_start..last_end];
172            for _ in self.max_positions..seq_len {
173                output.extend_from_slice(last_pos);
174            }
175            Tensor::from_vec(output, true)
176        }
177    }
178
179    /// Get maximum positions
180    pub fn max_positions(&self) -> usize {
181        self.max_positions
182    }
183
184    /// Get hidden dimension
185    pub fn hidden_size(&self) -> usize {
186        self.hidden_size
187    }
188}
189
190#[cfg(test)]
191mod tests {
192    use super::*;
193
194    #[test]
195    fn test_embedding_forward() {
196        let embed = Embedding::new(100, 8);
197        let tokens = vec![0, 5, 10];
198        let output = embed.forward(&tokens);
199        assert_eq!(output.len(), 3 * 8);
200    }
201
202    #[test]
203    fn test_embedding_out_of_vocab() {
204        let embed = Embedding::new(100, 8);
205        let tokens = vec![0, 200]; // 200 is out of vocab
206        let output = embed.forward(&tokens);
207        assert_eq!(output.len(), 2 * 8);
208        // Out of vocab should be zeros
209        let data = output.data();
210        for i in 8..16 {
211            assert_eq!(data[i], 0.0);
212        }
213    }
214
215    #[test]
216    fn test_embedding_vocab_and_hidden_size() {
217        let embed = Embedding::new(500, 16);
218        assert_eq!(embed.vocab_size(), 500);
219        assert_eq!(embed.hidden_size(), 16);
220    }
221
222    #[test]
223    fn test_embedding_single_token() {
224        let embed = Embedding::new(100, 8);
225        let tokens = vec![42];
226        let output = embed.forward(&tokens);
227        assert_eq!(output.len(), 8);
228        assert!(output.requires_grad());
229    }
230
231    #[test]
232    fn test_embedding_requires_grad() {
233        let embed = Embedding::new(100, 8);
234        assert!(embed.weight.requires_grad());
235    }
236
237    #[test]
238    fn test_embedding_from_params() {
239        let mut params = HashMap::new();
240        params.insert("embed.weight".to_string(), Tensor::from_vec(vec![0.1; 100 * 8], true));
241        let embed = Embedding::from_params(&params, "embed.weight", 100, 8);
242        assert!(embed.is_some());
243        let embed = embed.expect("operation should succeed");
244        assert_eq!(embed.vocab_size(), 100);
245        assert_eq!(embed.hidden_size(), 8);
246    }
247
248    #[test]
249    fn test_embedding_from_params_missing() {
250        let params: HashMap<String, Tensor> = HashMap::new();
251        let embed = Embedding::from_params(&params, "missing.weight", 100, 8);
252        assert!(embed.is_none());
253    }
254
255    // =========================================================================
256    // ENC-003: LearnedPositionEmbedding tests
257    // =========================================================================
258
259    #[test]
260    fn enc_003_learned_position_embedding_shape() {
261        let pos_embed = LearnedPositionEmbedding::new(514, 768);
262        assert_eq!(pos_embed.max_positions(), 514);
263        assert_eq!(pos_embed.hidden_size(), 768);
264        let output = pos_embed.forward(10);
265        assert_eq!(output.len(), 10 * 768);
266    }
267
268    #[test]
269    fn enc_003_learned_position_embedding_deterministic() {
270        // Hold the global init-seed lock across both constructions so concurrent
271        // seed-setting tests can't change INIT_SEED between them (see falsify_e7e).
272        let _seed_guard = crate::transformer::init::lock_init_seed(42);
273        let pe1 = LearnedPositionEmbedding::new(128, 32);
274        let pe2 = LearnedPositionEmbedding::new(128, 32);
275        let o1 = pe1.forward(10);
276        let o2 = pe2.forward(10);
277        assert_eq!(
278            o1.data().as_slice().expect("contiguous"),
279            o2.data().as_slice().expect("contiguous"),
280        );
281    }
282
283    #[test]
284    fn enc_003_learned_position_embedding_clamp_beyond_max() {
285        let pe = LearnedPositionEmbedding::new(4, 8);
286        let output = pe.forward(6); // 6 > max_positions=4
287        assert_eq!(output.len(), 6 * 8);
288        // Positions 4 and 5 should equal position 3 (clamped)
289        let data = output.data();
290        let slice = data.as_slice().expect("contiguous");
291        let pos3 = &slice[3 * 8..4 * 8];
292        let pos4 = &slice[4 * 8..5 * 8];
293        let pos5 = &slice[5 * 8..6 * 8];
294        assert_eq!(pos3, pos4);
295        assert_eq!(pos3, pos5);
296    }
297
298    #[test]
299    fn enc_003_learned_position_from_params() {
300        let mut params = HashMap::new();
301        params.insert("pos.weight".to_string(), Tensor::from_vec(vec![0.1; 128 * 32], true));
302        let pe = LearnedPositionEmbedding::from_params(&params, "pos.weight", 128, 32);
303        assert!(pe.is_some());
304    }
305
306    #[test]
307    fn enc_003_learned_position_from_params_rejects_wrong_shape() {
308        let mut params = HashMap::new();
309        params.insert("pos.weight".to_string(), Tensor::from_vec(vec![0.1; 50], true));
310        let pe = LearnedPositionEmbedding::from_params(&params, "pos.weight", 128, 32);
311        assert!(pe.is_none());
312    }
313
314    // =========================================================================
315    // FALSIFY-E7: Entrenar embedding contract gap analysis (Refs PMAT-326)
316    //
317    // Five-Whys: §2.1.1 "What Are Embeddings" falsification sweep
318    //   Why 1: Trained model could have garbage embeddings
319    //   Why 2: No data quality validation during training
320    //   Why 3: Embedding uses raw Tensor, not ValidatedEmbedding
321    //   Why 4: entrenar predates the ValidatedEmbedding contract
322    //   Why 5: No cross-crate contract enforcement test existed
323    //
324    // Popper (1959): "These tests try to break the claim that
325    // entrenar's embedding pipeline prevents degenerate models."
326    // =========================================================================
327
328    /// FALSIFY-E7a: Embedding initialization produces non-degenerate values
329    ///
330    /// The init formula `(i * 0.111).sin() * scale` MUST produce varied,
331    /// finite values. If it doesn't, freshly-initialized models are DOA.
332    #[test]
333    fn falsify_e7a_init_produces_valid_embedding() {
334        let embed = Embedding::new(100, 64);
335        let data = embed.weight.data();
336        let slice = data.as_slice().expect("data as slice");
337
338        // No NaN
339        let nan_count = slice.iter().filter(|v| v.is_nan()).count();
340        assert_eq!(nan_count, 0, "FALSIFY-E7a: Init must not produce NaN");
341
342        // No Inf
343        let inf_count = slice.iter().filter(|v| v.is_infinite()).count();
344        assert_eq!(inf_count, 0, "FALSIFY-E7a: Init must not produce Inf");
345
346        // Not all zeros (<50% zeros per embedding contract)
347        let zero_count = slice.iter().filter(|v| v.abs() < 1e-10).count();
348        let zero_pct = 100.0 * zero_count as f64 / slice.len() as f64;
349        assert!(zero_pct < 50.0,
350            "FALSIFY-E7a: Init has {zero_pct:.1}% zeros — exceeds embedding contract threshold (50%)");
351
352        // Values vary (not constant)
353        let min = slice.iter().copied().fold(f32::INFINITY, f32::min);
354        let max = slice.iter().copied().fold(f32::NEG_INFINITY, f32::max);
355        assert!(
356            (max - min).abs() > 1e-6,
357            "FALSIFY-E7a: Init values are constant ({min}..{max}) — degenerate embedding"
358        );
359    }
360
361    /// FALSIFY-E7b: Embedding shape matches vocab * hidden
362    #[test]
363    fn falsify_e7b_shape_matches_dimensions() {
364        let vocab_size = 151;
365        let hidden_size = 32;
366        let embed = Embedding::new(vocab_size, hidden_size);
367        assert_eq!(
368            embed.weight.len(),
369            vocab_size * hidden_size,
370            "FALSIFY-E7b: Embedding length must be vocab_size * hidden_size"
371        );
372    }
373
374    /// FALSIFY-E7c: from_params rejects wrong-shape tensor (PMAT-326 fix)
375    ///
376    /// from_params now validates weight.len() == vocab_size * hidden_size.
377    /// A tensor of 50 elements is rejected when 100*8=800 is expected.
378    #[test]
379    fn falsify_e7c_from_params_rejects_wrong_shape() {
380        let mut params = HashMap::new();
381        // Intentionally wrong size: 50 elements for 100*8=800 expected
382        params.insert("embed.weight".to_string(), Tensor::from_vec(vec![0.1; 50], true));
383        let embed = Embedding::from_params(&params, "embed.weight", 100, 8);
384        // FIXED (PMAT-326): now rejected
385        assert!(
386            embed.is_none(),
387            "FALSIFY-E7c: PMAT-326 fix — from_params MUST reject wrong-shape embedding"
388        );
389    }
390
391    /// FALSIFY-E7d: OOB token_id produces zeros (not panic)
392    ///
393    /// Contract divergence: aprender skips OOB tokens, realizar/entrenar zero-fill.
394    /// This test documents entrenar's behavior.
395    #[test]
396    fn falsify_e7d_oob_token_produces_zeros_not_panic() {
397        let embed = Embedding::new(100, 8);
398        let tokens = vec![0, 999]; // 999 is way OOB
399        let output = embed.forward(&tokens);
400        assert_eq!(output.len(), 2 * 8);
401        // Token 0 should have non-zero values
402        let data = output.data();
403        let token0_l2: f32 = (0..8).map(|i| data[i] * data[i]).sum::<f32>().sqrt();
404        assert!(token0_l2 > 1e-6, "Token 0 should have non-zero embedding");
405        // Token 999 should be all zeros
406        let token999_l2: f32 = (8..16).map(|i| data[i] * data[i]).sum::<f32>().sqrt();
407        assert!(token999_l2 < 1e-10, "OOB token should be zero-filled");
408    }
409
410    /// FALSIFY-E7e: Embedding init is deterministic (reproducible)
411    #[test]
412    fn falsify_e7e_init_deterministic() {
413        // INIT_SEED is process-global mutable state. Parallel tests that set it (via
414        // `lock_init_seed`) would otherwise change it between these two constructions,
415        // making this (correct) determinism assertion flaky. Hold the seed lock for the
416        // whole comparison so both `Embedding::new` calls observe the same seed.
417        let _seed_guard = crate::transformer::init::lock_init_seed(42);
418        let embed1 = Embedding::new(100, 64);
419        let embed2 = Embedding::new(100, 64);
420        let d1 = embed1.weight.data();
421        let d2 = embed2.weight.data();
422        assert_eq!(
423            d1.as_slice().expect("operation should succeed"),
424            d2.as_slice().expect("operation should succeed"),
425            "FALSIFY-E7e: Same vocab+hidden must produce identical initialization"
426        );
427    }
428
429    // =========================================================================
430    // FALSIFY-EM-001..004: embedding-lookup-v1.yaml contract mapping
431    //
432    // Five-Whys (PMAT-354):
433    //   Why 1: entrenar has E7a-e init tests but no forward-path EM-* tests
434    //   Why 2: E7 tests validate initialization, not the lookup/forward contract
435    //   Why 3: no mapping from embedding-lookup-v1.yaml to entrenar test names
436    //   Why 4: entrenar predates the provable-contracts YAML
437    //   Why 5: forward() was assumed correct because it's "just slicing"
438    //
439    // References:
440    //   - provable-contracts/contracts/embedding-lookup-v1.yaml
441    //   - src/transformer/embedding.rs::forward()
442    // =========================================================================
443
444    /// FALSIFY-EM-001: forward output shape = seq_len * hidden_size
445    #[test]
446    fn falsify_em_001_forward_output_shape() {
447        let embed = Embedding::new(100, 32);
448
449        for seq_len in [1, 3, 10, 50] {
450            let tokens: Vec<u32> = (0..seq_len).collect();
451            let output = embed.forward(&tokens);
452            assert_eq!(
453                output.len(),
454                seq_len as usize * 32,
455                "FALSIFIED EM-001: forward({seq_len} tokens) produced {} elements, expected {}",
456                output.len(),
457                seq_len as usize * 32
458            );
459        }
460    }
461
462    /// FALSIFY-EM-001b: empty input produces empty output
463    #[test]
464    fn falsify_em_001b_forward_empty_input() {
465        let embed = Embedding::new(100, 32);
466        let output = embed.forward(&[]);
467        assert_eq!(output.len(), 0, "FALSIFIED EM-001b: empty input should produce 0 elements");
468    }
469
470    /// FALSIFY-EM-002: OOB token → zeros, no panic (N-09 escape)
471    ///
472    /// Contract: token_id >= vocab_size produces zero-filled output, not a panic.
473    /// Valid tokens alongside OOB tokens must still produce correct results.
474    #[test]
475    fn falsify_em_002_oob_safety() {
476        let vocab_size = 50;
477        let hidden = 8;
478        let embed = Embedding::new(vocab_size, hidden);
479
480        // Pure OOB tokens
481        let oob_output = embed.forward(&[999, 50, 100]);
482        let oob_data = oob_output.data();
483        for (i, &v) in oob_data.iter().enumerate() {
484            assert!(v.abs() < 1e-10, "FALSIFIED EM-002: OOB output[{i}] = {v}, expected 0.0");
485        }
486
487        // Mixed valid + OOB: valid tokens must still be correct
488        let mixed_output = embed.forward(&[0, 999, 49]);
489        let mixed_data = mixed_output.data();
490        let weight_data = embed.weight.data();
491
492        // Token 0 (valid): should match weight row 0
493        for d in 0..hidden {
494            assert_eq!(
495                mixed_data[d], weight_data[d],
496                "FALSIFIED EM-002: valid token 0 corrupted at dim {d}"
497            );
498        }
499
500        // Token 999 (OOB): should be zeros
501        for d in 0..hidden {
502            assert!(
503                mixed_data[hidden + d].abs() < 1e-10,
504                "FALSIFIED EM-002: OOB token 999 at dim {d} = {}, expected 0.0",
505                mixed_data[hidden + d]
506            );
507        }
508
509        // Token 49 (valid boundary): should match weight row 49
510        for d in 0..hidden {
511            assert_eq!(
512                mixed_data[2 * hidden + d],
513                weight_data[49 * hidden + d],
514                "FALSIFIED EM-002: valid boundary token 49 corrupted at dim {d}"
515            );
516        }
517    }
518
519    /// FALSIFY-EM-003: forward determinism (same tokens → bit-identical output)
520    #[test]
521    fn falsify_em_003_forward_determinism() {
522        let embed = Embedding::new(100, 64);
523        let tokens = vec![5u32, 42, 0, 99, 17];
524
525        let o1 = embed.forward(&tokens);
526        let o2 = embed.forward(&tokens);
527
528        assert_eq!(
529            o1.data().as_slice().expect("operation should succeed"),
530            o2.data().as_slice().expect("operation should succeed"),
531            "FALSIFIED EM-003: forward() is non-deterministic"
532        );
533    }
534
535    /// FALSIFY-EM-004: forward output is finite (no NaN, no Inf)
536    #[test]
537    fn falsify_em_004_forward_finite_output() {
538        let embed = Embedding::new(200, 16);
539        let tokens: Vec<u32> = (0..200).collect();
540        let output = embed.forward(&tokens);
541        let data = output.data();
542
543        let nan_count = data.iter().filter(|v| v.is_nan()).count();
544        let inf_count = data.iter().filter(|v| v.is_infinite()).count();
545
546        assert_eq!(
547            nan_count, 0,
548            "FALSIFIED EM-004: forward output contains {nan_count} NaN values"
549        );
550        assert_eq!(
551            inf_count, 0,
552            "FALSIFIED EM-004: forward output contains {inf_count} Inf values"
553        );
554    }
555
556    /// FALSIFY-EM-005: forward value correctness (extractive — output[i] = W[token_id])
557    #[test]
558    fn falsify_em_005_forward_value_correctness() {
559        let embed = Embedding::new(50, 8);
560        let tokens = vec![0u32, 10, 49];
561        let output = embed.forward(&tokens);
562        let out_data = output.data();
563        let weight_data = embed.weight.data();
564
565        // Token 0: output[0..8] == weight[0..8]
566        for i in 0..8 {
567            assert_eq!(
568                out_data[i], weight_data[i],
569                "FALSIFIED EM-005: output[{i}] != weight[{i}] for token 0"
570            );
571        }
572        // Token 10: output[8..16] == weight[80..88]
573        for i in 0..8 {
574            assert_eq!(
575                out_data[8 + i],
576                weight_data[80 + i],
577                "FALSIFIED EM-005: output[{}] != weight[{}] for token 10",
578                8 + i,
579                80 + i
580            );
581        }
582    }
583
584    // =========================================================================
585    // FALSIFY-EMB-005: Non-zero embeddings (embedding-algebra-v1.yaml)
586    //
587    // Five-Whys (PMAT-354):
588    //   Why 1: entrenar had E7a init tests but no FALSIFY-EMB-005 tagged test
589    //   Why 2: E7a covers init validity, not the EMB "non-zero" algebra claim
590    //   Why 3: no mapping from embedding-algebra-v1.yaml to entrenar test names
591    //   Why 4: entrenar predates the provable-contracts YAML
592    //   Why 5: forward output non-zero was assumed from init non-zero
593    // =========================================================================
594
595    // =========================================================================
596    // FALSIFY-EMB-001: Lookup determinism (embedding-algebra-v1.yaml)
597    //
598    // Five-Whys (PMAT-354, Phase 8):
599    //   Why 1: entrenar had EM-003 (determinism) but not EMB-001 (algebra contract)
600    //   Why 2: EM-003 tests forward() determinism, EMB-001 tests per-token lookup identity
601    //   Why 3: EMB-001 YAML says "proptest: embed(t) == embed(t) for random t"
602    //   Why 4: no mapping from embedding-algebra-v1.yaml EMB-001 to entrenar tests
603    //   Why 5: lookup determinism assumed from EM-003 but never isolated per-token
604    // =========================================================================
605
606    /// FALSIFY-EMB-001: same token always returns same vector
607    #[test]
608    fn falsify_emb_001_lookup_determinism() {
609        let embed = Embedding::new(200, 48);
610        for t in [0u32, 1, 42, 100, 199] {
611            let v1 = embed.forward(&[t]);
612            let v2 = embed.forward(&[t]);
613            assert_eq!(
614                v1.data(),
615                v2.data(),
616                "FALSIFIED EMB-001: embed({t}) != embed({t}) — non-deterministic lookup"
617            );
618        }
619    }
620
621    // =========================================================================
622    // FALSIFY-EMB-002: Shape preservation (embedding-algebra-v1.yaml)
623    //
624    // Five-Whys (PMAT-354, Phase 8):
625    //   Why 1: entrenar EM-001 tests output length but not EMB-002 per-token dimension
626    //   Why 2: EMB-002 YAML says "embedding output is d_model-dimensional"
627    //   Why 3: shape preservation for different hidden sizes never parametrically tested
628    //   Why 4: entrenar only used hidden_size=64 in EM-001 tests
629    //   Why 5: no systematic d_model variation in embedding tests
630    // =========================================================================
631
632    /// FALSIFY-EMB-002: embedding output dimension matches hidden_size
633    #[test]
634    fn falsify_emb_002_shape_preservation() {
635        for (v, d) in [(100, 32), (200, 64), (500, 128), (50, 16)] {
636            let embed = Embedding::new(v, d);
637            let output = embed.forward(&[0, 1, 2]);
638            assert_eq!(
639                output.data().len(),
640                3 * d,
641                "FALSIFIED EMB-002: vocab={v}, d_model={d}, output len={} != 3*{d}",
642                output.data().len()
643            );
644        }
645    }
646
647    // =========================================================================
648    // FALSIFY-EMB-004: Vocabulary bounds (embedding-algebra-v1.yaml)
649    //
650    // Five-Whys (PMAT-354, Phase 8):
651    //   Why 1: entrenar EM-002 tests OOB safety but not EMB-004 (algebra perspective)
652    //   Why 2: EMB-004 YAML says "out-of-range IDs rejected"
653    //   Why 3: entrenar silently zeros OOB (N-09) — need explicit boundary test
654    //   Why 4: boundary between valid and OOB never tested at exact vocab_size edge
655    //   Why 5: no EMB-004 tagged test existed in entrenar
656    // =========================================================================
657
658    /// FALSIFY-EMB-004: valid tokens non-zero, OOB tokens zero
659    #[test]
660    fn falsify_emb_004_vocabulary_bounds() {
661        let vocab = 50;
662        let d = 16;
663        let embed = Embedding::new(vocab, d);
664
665        // Last valid token must be non-zero
666        let valid_output = embed.forward(&[vocab as u32 - 1]);
667        let valid_norm: f32 = valid_output.data().iter().map(|v| v * v).sum();
668        assert!(
669            valid_norm > 0.0,
670            "FALSIFIED EMB-004: valid token {} produced zero embedding",
671            vocab - 1
672        );
673
674        // First OOB token must be zero (N-09 escape)
675        let oob_output = embed.forward(&[vocab as u32]);
676        let oob_norm: f32 = oob_output.data().iter().map(|v| v * v).sum();
677        assert!(
678            oob_norm == 0.0,
679            "FALSIFIED EMB-004: OOB token {vocab} produced non-zero (norm={oob_norm})"
680        );
681    }
682
683    /// FALSIFY-EMB-005: forward output is non-zero for valid tokens
684    #[test]
685    fn falsify_emb_005_forward_non_zero() {
686        let embed = Embedding::new(100, 64);
687        let tokens = vec![0u32, 42, 99];
688        let output = embed.forward(&tokens);
689        let data = output.data();
690
691        let l2_norm: f32 = data.iter().map(|v| v * v).sum::<f32>().sqrt();
692        assert!(l2_norm > 1e-6, "FALSIFIED EMB-005: forward output is all-zero (L2={l2_norm})");
693    }
694
695    // =========================================================================
696    // PROPTEST FALSIFY: Embedding property-based falsification
697    //
698    // Five-Whys (PMAT-354, Phase 9):
699    //   Why 1: EM/EMB tests used fixed vocab=100, hidden=32/48/64
700    //   Why 2: embedding forward() could have off-by-one at edge vocab sizes
701    //   Why 3: proptest explores vocab/hidden/seq_len combos humans don't anticipate
702    //   Why 4: determinism (EM-003, EMB-001) could break under certain init patterns
703    //   Why 5: YAML contracts explicitly call for "proptest with random..."
704    // =========================================================================
705
706    mod em_proptest_falsify {
707        use super::*;
708        use proptest::prelude::*;
709
710        // EM-001-prop: output shape for random seq_len and hidden_size
711        proptest! {
712            #![proptest_config(ProptestConfig::with_cases(100))]
713            #[test]
714            fn falsify_em_001_prop_output_shape(
715                vocab_size in prop::sample::select(vec![50_usize, 100, 200, 500]),
716                hidden_size in prop::sample::select(vec![16_usize, 32, 48, 64]),
717                seq_len in 1_usize..32,
718            ) {
719                let embed = Embedding::new(vocab_size, hidden_size);
720                let tokens: Vec<u32> = (0..seq_len).map(|i| (i % vocab_size) as u32).collect();
721                let output = embed.forward(&tokens);
722                prop_assert_eq!(
723                    output.len(), seq_len * hidden_size,
724                    "FALSIFIED EM-001-prop: len={} != {}*{}={} (v={})",
725                    output.len(), seq_len, hidden_size, seq_len * hidden_size, vocab_size
726                );
727            }
728        }
729
730        // EM-003-prop: determinism for random tokens
731        proptest! {
732            #![proptest_config(ProptestConfig::with_cases(50))]
733            #[test]
734            fn falsify_em_003_prop_determinism(
735                vocab_size in prop::sample::select(vec![50_usize, 100, 200]),
736                hidden_size in prop::sample::select(vec![16_usize, 32, 64]),
737                token_ids in proptest::collection::vec(0_u32..49, 1..16),
738            ) {
739                let embed = Embedding::new(vocab_size, hidden_size);
740                let out1 = embed.forward(&token_ids);
741                let out2 = embed.forward(&token_ids);
742                prop_assert_eq!(
743                    out1.data(), out2.data(),
744                    "FALSIFIED EM-003-prop: two calls differ (v={}, h={})",
745                    vocab_size, hidden_size
746                );
747            }
748        }
749
750        // EM-004-prop: finite output for random tokens
751        proptest! {
752            #![proptest_config(ProptestConfig::with_cases(100))]
753            #[test]
754            fn falsify_em_004_prop_finite(
755                vocab_size in prop::sample::select(vec![50_usize, 100, 200]),
756                hidden_size in prop::sample::select(vec![16_usize, 32, 64]),
757                token_ids in proptest::collection::vec(0_u32..49, 1..16),
758            ) {
759                let embed = Embedding::new(vocab_size, hidden_size);
760                let output = embed.forward(&token_ids);
761                for (i, v) in output.data().iter().enumerate() {
762                    prop_assert!(
763                        v.is_finite(),
764                        "FALSIFIED EM-004-prop: output[{}]={} not finite (v={}, h={})",
765                        i, v, vocab_size, hidden_size
766                    );
767                }
768            }
769        }
770    }
771
772    // =========================================================================
773    // PROPTEST FALSIFY: EMB algebra property-based falsification
774    //
775    // Five-Whys (PMAT-354, Phase 9):
776    //   Why 1: EMB-001/002/004/005 had zero proptest coverage in entrenar
777    //   Why 2: Determinism (EMB-001) only tested 5 fixed token IDs
778    //   Why 3: Shape preservation (EMB-002) only tested 4 (vocab, d) pairs
779    //   Why 4: Vocabulary bounds (EMB-004) only tested vocab=50
780    //   Why 5: proptest explores random token/dim combos at scale
781    // =========================================================================
782
783    mod emb_proptest_falsify {
784        use super::*;
785        use proptest::prelude::*;
786
787        // EMB-001-prop: lookup determinism for random tokens
788        proptest! {
789            #![proptest_config(ProptestConfig::with_cases(100))]
790            #[test]
791            fn falsify_emb_001_prop_determinism(
792                vocab_size in prop::sample::select(vec![50_usize, 100, 200]),
793                hidden_size in prop::sample::select(vec![16_usize, 32, 64]),
794                token_id in 0_u32..49,
795            ) {
796                let embed = Embedding::new(vocab_size, hidden_size);
797                let v1 = embed.forward(&[token_id]);
798                let v2 = embed.forward(&[token_id]);
799                prop_assert_eq!(
800                    v1.data(), v2.data(),
801                    "FALSIFIED EMB-001-prop: embed({}) non-deterministic (v={}, h={})",
802                    token_id, vocab_size, hidden_size
803                );
804            }
805        }
806
807        // EMB-002-prop: shape preservation for random dimensions
808        proptest! {
809            #![proptest_config(ProptestConfig::with_cases(100))]
810            #[test]
811            fn falsify_emb_002_prop_shape(
812                vocab_size in prop::sample::select(vec![50_usize, 100, 200, 500]),
813                hidden_size in prop::sample::select(vec![16_usize, 32, 48, 64, 128]),
814                seq_len in 1_usize..16,
815            ) {
816                let embed = Embedding::new(vocab_size, hidden_size);
817                let tokens: Vec<u32> = (0..seq_len).map(|i| (i % vocab_size) as u32).collect();
818                let output = embed.forward(&tokens);
819                prop_assert_eq!(
820                    output.data().len(), seq_len * hidden_size,
821                    "FALSIFIED EMB-002-prop: data len={} != {}*{}={} (v={})",
822                    output.data().len(), seq_len, hidden_size, seq_len * hidden_size, vocab_size
823                );
824            }
825        }
826
827        // EMB-005-prop: non-zero output for random valid tokens
828        proptest! {
829            #![proptest_config(ProptestConfig::with_cases(100))]
830            #[test]
831            fn falsify_emb_005_prop_non_zero(
832                vocab_size in prop::sample::select(vec![50_usize, 100, 200]),
833                hidden_size in prop::sample::select(vec![16_usize, 32, 64]),
834                token_ids in proptest::collection::vec(0_u32..49, 1..8),
835            ) {
836                let embed = Embedding::new(vocab_size, hidden_size);
837                let output = embed.forward(&token_ids);
838                let l2_norm: f32 = output.data().iter().map(|v| v * v).sum::<f32>().sqrt();
839                prop_assert!(
840                    l2_norm > 1e-6,
841                    "FALSIFIED EMB-005-prop: output all-zero (L2={}, v={}, h={})",
842                    l2_norm, vocab_size, hidden_size
843                );
844            }
845        }
846    }
847}