1use crate::Tensor;
6use std::collections::HashMap;
7
8pub struct Embedding {
10 pub weight: Tensor,
12 vocab_size: usize,
14 hidden_size: usize,
16}
17
18impl Embedding {
19 pub fn new(vocab_size: usize, hidden_size: usize) -> Self {
21 use super::init::{get_init_seed, rand_normal_seeded};
22 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 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 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 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 pub fn vocab_size(&self) -> usize {
93 self.vocab_size
94 }
95
96 pub fn hidden_size(&self) -> usize {
98 self.hidden_size
99 }
100}
101
102pub struct LearnedPositionEmbedding {
112 pub weight: Tensor,
114 max_positions: usize,
116 hidden_size: usize,
118}
119
120impl LearnedPositionEmbedding {
121 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 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 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 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 pub fn max_positions(&self) -> usize {
181 self.max_positions
182 }
183
184 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]; let output = embed.forward(&tokens);
207 assert_eq!(output.len(), 2 * 8);
208 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(¶ms, "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(¶ms, "missing.weight", 100, 8);
252 assert!(embed.is_none());
253 }
254
255 #[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 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); assert_eq!(output.len(), 6 * 8);
288 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(¶ms, "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(¶ms, "pos.weight", 128, 32);
311 assert!(pe.is_none());
312 }
313
314 #[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 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 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 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 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 #[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 #[test]
379 fn falsify_e7c_from_params_rejects_wrong_shape() {
380 let mut params = HashMap::new();
381 params.insert("embed.weight".to_string(), Tensor::from_vec(vec![0.1; 50], true));
383 let embed = Embedding::from_params(¶ms, "embed.weight", 100, 8);
384 assert!(
386 embed.is_none(),
387 "FALSIFY-E7c: PMAT-326 fix — from_params MUST reject wrong-shape embedding"
388 );
389 }
390
391 #[test]
396 fn falsify_e7d_oob_token_produces_zeros_not_panic() {
397 let embed = Embedding::new(100, 8);
398 let tokens = vec![0, 999]; let output = embed.forward(&tokens);
400 assert_eq!(output.len(), 2 * 8);
401 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 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 #[test]
412 fn falsify_e7e_init_deterministic() {
413 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 #[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 #[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 #[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 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 let mixed_output = embed.forward(&[0, 999, 49]);
489 let mixed_data = mixed_output.data();
490 let weight_data = embed.weight.data();
491
492 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 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 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 #[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 #[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 #[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 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 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 #[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 #[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 #[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 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 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 #[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 mod em_proptest_falsify {
707 use super::*;
708 use proptest::prelude::*;
709
710 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 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 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 mod emb_proptest_falsify {
784 use super::*;
785 use proptest::prelude::*;
786
787 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 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 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}