1use crate::error::{EmbedError, Result};
29use lattice_inference::InferenceError;
30use lattice_inference::model::qwen35_config::Qwen35Config;
31use lattice_inference::tokenizer::bpe::BpeTokenizer;
32use lattice_inference::vision::checkpoint::{
33 Qwen35VisionWeights, load_qwen35_vision_weights_from_safetensors,
34 open_qwen35_single_decoder_safetensors,
35};
36use lattice_inference::vision::{embed_image_from_bytes_f16, embed_image_from_bytes_f16_metal};
37use lattice_inference::weights::f16_weights::{F16ModelWeights, load_f16_weights};
38use std::path::Path;
39
40pub use lattice_inference::forward::cpu_f16::PoolingStrategy;
41
42#[cfg(test)]
43thread_local! {
44 static AFTER_VISUAL_LOAD_HOOK: std::cell::RefCell<Option<Box<dyn FnOnce()>>> =
45 std::cell::RefCell::new(None);
46}
47
48#[cfg(test)]
49fn run_after_visual_load_hook() {
50 let hook = AFTER_VISUAL_LOAD_HOOK.with(|slot| slot.borrow_mut().take());
51 if let Some(hook) = hook {
52 hook();
53 }
54}
55
56#[cfg(test)]
57fn with_after_visual_load_hook<T>(hook: impl FnOnce() + 'static, action: impl FnOnce() -> T) -> T {
58 struct ClearHookOnDrop;
59 impl Drop for ClearHookOnDrop {
60 fn drop(&mut self) {
61 AFTER_VISUAL_LOAD_HOOK.with(|slot| {
62 slot.borrow_mut().take();
63 });
64 }
65 }
66
67 AFTER_VISUAL_LOAD_HOOK.with(|slot| {
68 let previous = slot.borrow_mut().replace(Box::new(hook));
69 assert!(
70 previous.is_none(),
71 "visual-load test hook already installed"
72 );
73 });
74 let _clear_on_drop = ClearHookOnDrop;
75 let result = action();
76 AFTER_VISUAL_LOAD_HOOK.with(|slot| {
77 assert!(
78 slot.borrow().is_none(),
79 "VisionEmbeddingModel::from_directory did not traverse the visual-load test hook"
80 );
81 });
82 result
83}
84
85pub struct VisionEmbeddingModel {
91 weights: F16ModelWeights,
92 config: Qwen35Config,
93 vision_weights: Qwen35VisionWeights,
94 tokenizer: BpeTokenizer,
95}
96
97impl VisionEmbeddingModel {
98 pub fn new(
104 weights: F16ModelWeights,
105 config: Qwen35Config,
106 vision_weights: Qwen35VisionWeights,
107 tokenizer: BpeTokenizer,
108 ) -> Self {
109 Self {
110 weights,
111 config,
112 vision_weights,
113 tokenizer,
114 }
115 }
116
117 pub fn from_directory(dir: &Path) -> Result<Self> {
136 let quantized_index = dir.join("quantize_index.json");
137 match std::fs::symlink_metadata(&quantized_index) {
138 Ok(_) => {
139 return Err(EmbedError::ModelInitialization(format!(
140 "{} is present, but quantized checkpoints are not supported by \
141 VisionEmbeddingModel::from_directory's f16 decoder loader",
142 quantized_index.display()
143 )));
144 }
145 Err(err) if err.kind() == std::io::ErrorKind::NotFound => {}
146 Err(err) => {
147 return Err(EmbedError::ModelInitialization(format!(
148 "failed to inspect {}: {err}",
149 quantized_index.display()
150 )));
151 }
152 }
153
154 let config = Qwen35Config::from_model_dir(dir)
155 .map_err(|e| EmbedError::ModelInitialization(format!("config.json: {e}")))?;
156 let vision_cfg = config.vision_config.clone().ok_or_else(|| {
157 EmbedError::ModelInitialization(format!(
158 "{} has no vision_config; not a vision-language checkpoint",
159 dir.display()
160 ))
161 })?;
162
163 let tokenizer_path = dir.join("tokenizer.json");
168 let tokenizer = BpeTokenizer::from_tokenizer_json(&tokenizer_path).map_err(|e| {
169 EmbedError::ModelInitialization(format!("{}: {e}", tokenizer_path.display()))
170 })?;
171
172 let (mut sf, shard_path) = open_qwen35_single_decoder_safetensors(dir)
173 .map_err(|e| EmbedError::ModelInitialization(format!("decoder checkpoint: {e}")))?;
174 let vision_weights =
175 load_qwen35_vision_weights_from_safetensors(&mut sf, &shard_path, &vision_cfg)
176 .map_err(|e| EmbedError::ModelInitialization(format!("vision weights: {e}")))?;
177 #[cfg(test)]
178 run_after_visual_load_hook();
179 let weights = load_f16_weights(&sf, &config)
180 .map_err(|e| EmbedError::ModelInitialization(format!("decoder weights: {e}")))?;
181
182 Ok(Self::new(weights, config, vision_weights, tokenizer))
183 }
184
185 pub fn embed_image(
202 &self,
203 image_bytes: &[u8],
204 prompt: &str,
205 pooling: PoolingStrategy,
206 ) -> Result<Vec<f32>> {
207 embed_image_from_bytes_f16(
208 &self.weights,
209 &self.config,
210 &self.vision_weights,
211 &self.tokenizer,
212 image_bytes,
213 prompt,
214 pooling,
215 )
216 .map_err(map_inference_error)
217 }
218
219 pub fn embed_image_metal(
234 &self,
235 image_bytes: &[u8],
236 prompt: &str,
237 pooling: PoolingStrategy,
238 ) -> Result<Vec<f32>> {
239 embed_image_from_bytes_f16_metal(
240 &self.weights,
241 &self.config,
242 &self.vision_weights,
243 &self.tokenizer,
244 image_bytes,
245 prompt,
246 pooling,
247 )
248 .map_err(map_inference_error)
249 }
250
251 pub fn embed_text(&self, prompt: &str, pooling: PoolingStrategy) -> Result<Vec<f32>> {
261 lattice_inference::forward::cpu_f16::embed_text_vlm_f16(
262 &self.weights,
263 &self.config,
264 &self.tokenizer,
265 prompt,
266 pooling,
267 )
268 .map_err(map_inference_error)
269 }
270
271 pub fn dimensions(&self) -> usize {
273 self.config.hidden_size
274 }
275}
276
277fn map_inference_error(e: InferenceError) -> EmbedError {
282 match e {
283 InferenceError::InvalidInput(msg) => EmbedError::InvalidInput(msg),
284 other => EmbedError::InferenceFailed(other.to_string()),
285 }
286}
287
288#[cfg(test)]
289mod tests {
290 use super::*;
291 use lattice_inference::model::qwen35_config::{LayerType, RopeParams, VisionModelConfig};
292 use lattice_inference::vision::checkpoint::{
293 VisualBlockWeights, VisualMergerWeights, resolve_qwen35_single_decoder_safetensors,
294 };
295 use lattice_inference::weights::f16_weights::{
296 F16AttentionWeights, F16CommonLayerWeights, F16FeedForwardWeights,
297 F16FullAttentionLayerWeights, f32_to_f16_slice,
298 };
299
300 fn pseudo_random_fill(seed: u32, n: usize) -> Vec<f32> {
305 let mut state = seed | 1;
306 let mut next = move || {
307 state ^= state << 13;
308 state ^= state >> 17;
309 state ^= state << 5;
310 (state as f32 / u32::MAX as f32) * 0.2 - 0.1
311 };
312 (0..n).map(|_| next()).collect()
313 }
314
315 fn tiny_vision_cfg() -> VisionModelConfig {
316 VisionModelConfig {
317 depth: 1,
318 hidden_size: 8,
319 num_heads: 2,
320 patch_size: 2,
321 spatial_merge_size: 2,
322 out_hidden_size: 8, temporal_patch_size: 1,
324 num_position_embeddings: 16,
325 in_channels: 1,
326 deepstack_visual_indexes: vec![],
327 intermediate_size: None,
328 }
329 }
330
331 fn tiny_vision_weights(vision_cfg: &VisionModelConfig, seed: u32) -> Qwen35VisionWeights {
332 let hidden = vision_cfg.hidden_size;
333 let patch_len = vision_cfg.in_channels
334 * vision_cfg.temporal_patch_size
335 * vision_cfg.patch_size
336 * vision_cfg.patch_size;
337 let mlp_dim = 2 * hidden;
338 let merge_in = vision_cfg.spatial_merge_size * vision_cfg.spatial_merge_size * hidden;
339
340 let block = VisualBlockWeights {
341 qkv_weight: pseudo_random_fill(seed, 3 * hidden * hidden),
342 qkv_bias: pseudo_random_fill(seed.wrapping_add(1), 3 * hidden),
343 proj_weight: pseudo_random_fill(seed.wrapping_add(2), hidden * hidden),
344 proj_bias: pseudo_random_fill(seed.wrapping_add(3), hidden),
345 fc1_weight: pseudo_random_fill(seed.wrapping_add(4), mlp_dim * hidden),
346 fc1_bias: pseudo_random_fill(seed.wrapping_add(5), mlp_dim),
347 fc2_weight: pseudo_random_fill(seed.wrapping_add(6), hidden * mlp_dim),
348 fc2_bias: pseudo_random_fill(seed.wrapping_add(7), hidden),
349 norm1_weight: vec![1.0; hidden],
350 norm1_bias: vec![0.0; hidden],
351 norm2_weight: vec![1.0; hidden],
352 norm2_bias: vec![0.0; hidden],
353 };
354
355 Qwen35VisionWeights {
356 patch_embed_weight: pseudo_random_fill(seed.wrapping_add(8), hidden * patch_len),
357 patch_embed_weight_shape: vec![
358 hidden,
359 vision_cfg.in_channels,
360 vision_cfg.temporal_patch_size,
361 vision_cfg.patch_size,
362 vision_cfg.patch_size,
363 ],
364 patch_embed_bias: pseudo_random_fill(seed.wrapping_add(9), hidden),
365 pos_embed: pseudo_random_fill(
366 seed.wrapping_add(10),
367 vision_cfg.num_position_embeddings * hidden,
368 ),
369 blocks: vec![block],
370 merger: VisualMergerWeights {
371 fc1_weight: pseudo_random_fill(seed.wrapping_add(11), merge_in * merge_in),
372 fc1_bias: pseudo_random_fill(seed.wrapping_add(12), merge_in),
373 fc2_weight: pseudo_random_fill(
374 seed.wrapping_add(13),
375 vision_cfg.out_hidden_size * merge_in,
376 ),
377 fc2_bias: pseudo_random_fill(seed.wrapping_add(14), vision_cfg.out_hidden_size),
378 norm_weight: vec![1.0; hidden],
379 norm_bias: vec![0.0; hidden],
380 },
381 }
382 }
383
384 fn tiny_vlm_fixture() -> (Qwen35Config, F16ModelWeights, Qwen35VisionWeights) {
388 let hidden = 8usize;
389 let vocab = 16usize;
390 let vision_cfg = tiny_vision_cfg();
391
392 let cfg = Qwen35Config {
393 hidden_size: hidden,
394 num_hidden_layers: 1,
395 vocab_size: vocab,
396 intermediate_size: 4,
397 rms_norm_eps: 1e-6,
398 num_attention_heads: 1,
399 num_key_value_heads: 1,
400 head_dim: hidden,
401 rope_theta: 1.0e7,
402 partial_rotary_factor: 1.0,
403 rope_parameters: Some(RopeParams {
404 rope_theta: 1.0e7,
405 partial_rotary_factor: Some(1.0),
406 mrope_section: Some(vec![2, 1, 1]),
407 mrope_interleaved: Some(true),
408 }),
409 linear_num_key_heads: 2,
410 linear_num_value_heads: Some(2),
411 linear_key_head_dim: 32,
412 linear_value_head_dim: 32,
413 linear_conv_kernel_dim: 4,
414 num_experts: None,
415 num_experts_per_tok: None,
416 moe_intermediate_size: None,
417 shared_expert_intermediate_size: None,
418 output_router_logits: false,
419 router_aux_loss_coef: None,
420 tie_word_embeddings: true,
421 full_attention_interval: 1,
422 layer_types: vec![LayerType::FullAttention],
423 layer_mask: vec![true],
424 eos_token_id: 999,
425 max_position_embeddings: 512,
426 mtp_num_hidden_layers: 0,
427 mtp_use_dedicated_embeddings: false,
428 quarot_rotation_seed: None,
429 vision_config: Some(vision_cfg.clone()),
430 image_token_id: Some(9),
431 video_token_id: None,
432 vision_start_token_id: Some(10),
433 vision_end_token_id: Some(11),
434 };
435
436 let to_f16 = |src: &[f32]| -> Vec<u16> {
437 let mut dst = vec![0u16; src.len()];
438 f32_to_f16_slice(src, &mut dst);
439 dst
440 };
441
442 let embed_tokens_f32 = pseudo_random_fill(777, vocab * hidden);
443 let q_dim = cfg.full_q_dim();
444 let kv_dim = cfg.full_kv_dim();
445 let full_weights = F16FullAttentionLayerWeights {
446 q_proj: to_f16(&pseudo_random_fill(101, 2 * q_dim * hidden)),
447 k_proj: to_f16(&pseudo_random_fill(102, kv_dim * hidden)),
448 v_proj: to_f16(&pseudo_random_fill(103, kv_dim * hidden)),
449 o_proj: to_f16(&pseudo_random_fill(104, hidden * q_dim)),
450 q_norm: vec![0.0f32; hidden],
451 k_norm: vec![0.0f32; hidden],
452 };
453 let common = F16CommonLayerWeights {
454 input_layernorm: vec![0.0f32; hidden],
455 post_attention_layernorm: vec![0.0f32; hidden],
456 ffn: F16FeedForwardWeights::Dense {
457 gate_proj: to_f16(&vec![0.0f32; 4 * hidden]),
458 up_proj: to_f16(&vec![0.0f32; 4 * hidden]),
459 down_proj: to_f16(&vec![0.0f32; hidden * 4]),
460 },
461 };
462 let weights = F16ModelWeights {
463 embed_tokens: to_f16(&embed_tokens_f32),
464 final_norm: vec![0.0f32; hidden],
465 layers: vec![(F16AttentionWeights::Full(full_weights), common)],
466 };
467
468 let vision_weights = tiny_vision_weights(&vision_cfg, 555);
469 (cfg, weights, vision_weights)
470 }
471
472 fn make_test_png(w: u32, h: u32, seed: u8) -> Vec<u8> {
473 use image::RgbImage;
474 let mut img = RgbImage::new(w, h);
475 for y in 0..h {
476 for x in 0..w {
477 let v = ((x + y + seed as u32) % 256) as u8;
478 img.put_pixel(x, y, image::Rgb([v, v, v]));
479 }
480 }
481 let mut buf = Vec::new();
482 img.write_to(&mut std::io::Cursor::new(&mut buf), image::ImageFormat::Png)
483 .unwrap();
484 buf
485 }
486
487 fn tiny_tokenizer() -> BpeTokenizer {
488 let mut vocab_map = std::collections::HashMap::new();
489 for (i, c) in ["describe", "this", "image"].iter().enumerate() {
490 vocab_map.insert((*c).to_string(), i as u32);
491 }
492 BpeTokenizer::from_vocab_and_merges(vocab_map, vec![]).expect("tokenizer constructs")
493 }
494
495 fn single_char_tokenizer() -> BpeTokenizer {
501 let mut vocab_map = std::collections::HashMap::new();
502 for (i, c) in ["a", "b", "c"].iter().enumerate() {
503 vocab_map.insert((*c).to_string(), i as u32);
504 }
505 BpeTokenizer::from_vocab_and_merges(vocab_map, vec![]).expect("tokenizer constructs")
506 }
507
508 fn tiny_vlm_checkpoint_shapes() -> Vec<(String, Vec<usize>)> {
509 let hidden = 8usize;
510 let mut shapes = vec![
511 (
512 "model.language_model.embed_tokens.weight".to_string(),
513 vec![16, hidden],
514 ),
515 ("model.language_model.norm.weight".to_string(), vec![hidden]),
516 (
517 "model.language_model.layers.0.input_layernorm.weight".to_string(),
518 vec![hidden],
519 ),
520 (
521 "model.language_model.layers.0.post_attention_layernorm.weight".to_string(),
522 vec![hidden],
523 ),
524 (
525 "model.language_model.layers.0.mlp.gate_proj.weight".to_string(),
526 vec![4, hidden],
527 ),
528 (
529 "model.language_model.layers.0.mlp.up_proj.weight".to_string(),
530 vec![4, hidden],
531 ),
532 (
533 "model.language_model.layers.0.mlp.down_proj.weight".to_string(),
534 vec![hidden, 4],
535 ),
536 (
537 "model.language_model.layers.0.self_attn.q_proj.weight".to_string(),
538 vec![16, hidden],
539 ),
540 (
541 "model.language_model.layers.0.self_attn.k_proj.weight".to_string(),
542 vec![hidden, hidden],
543 ),
544 (
545 "model.language_model.layers.0.self_attn.v_proj.weight".to_string(),
546 vec![hidden, hidden],
547 ),
548 (
549 "model.language_model.layers.0.self_attn.o_proj.weight".to_string(),
550 vec![hidden, hidden],
551 ),
552 (
553 "model.language_model.layers.0.self_attn.q_norm.weight".to_string(),
554 vec![hidden],
555 ),
556 (
557 "model.language_model.layers.0.self_attn.k_norm.weight".to_string(),
558 vec![hidden],
559 ),
560 (
561 "model.visual.patch_embed.proj.weight".to_string(),
562 vec![hidden, 3, 1, 2, 2],
563 ),
564 (
565 "model.visual.patch_embed.proj.bias".to_string(),
566 vec![hidden],
567 ),
568 (
569 "model.visual.pos_embed.weight".to_string(),
570 vec![16, hidden],
571 ),
572 (
573 "model.visual.merger.linear_fc1.weight".to_string(),
574 vec![32, 32],
575 ),
576 ("model.visual.merger.linear_fc1.bias".to_string(), vec![32]),
577 (
578 "model.visual.merger.linear_fc2.weight".to_string(),
579 vec![hidden, 32],
580 ),
581 (
582 "model.visual.merger.linear_fc2.bias".to_string(),
583 vec![hidden],
584 ),
585 ("model.visual.merger.norm.weight".to_string(), vec![hidden]),
586 ("model.visual.merger.norm.bias".to_string(), vec![hidden]),
587 ];
588 for (suffix, shape) in [
589 ("attn.qkv.weight", vec![24, hidden]),
590 ("attn.qkv.bias", vec![24]),
591 ("attn.proj.weight", vec![hidden, hidden]),
592 ("attn.proj.bias", vec![hidden]),
593 ("mlp.linear_fc1.weight", vec![32, hidden]),
594 ("mlp.linear_fc1.bias", vec![32]),
595 ("mlp.linear_fc2.weight", vec![hidden, 32]),
596 ("mlp.linear_fc2.bias", vec![hidden]),
597 ("norm1.weight", vec![hidden]),
598 ("norm1.bias", vec![hidden]),
599 ("norm2.weight", vec![hidden]),
600 ("norm2.bias", vec![hidden]),
601 ] {
602 shapes.push((format!("model.visual.blocks.0.{suffix}"), shape));
603 }
604 shapes
605 }
606
607 fn write_f32_safetensors(path: &Path, shapes: &[(String, Vec<usize>)]) {
608 write_f32_safetensors_with_offset(path, shapes, 0.0);
609 }
610
611 fn write_f32_safetensors_with_offset(
612 path: &Path,
613 shapes: &[(String, Vec<usize>)],
614 offset: f32,
615 ) {
616 let mut header_parts = Vec::with_capacity(shapes.len());
617 let mut data = Vec::new();
618 for (i, (name, shape)) in shapes.iter().enumerate() {
619 let start = data.len();
620 let numel: usize = shape.iter().product();
621 for _ in 0..numel {
622 data.extend_from_slice(&(offset + (i + 1) as f32 / 100.0).to_le_bytes());
623 }
624 let end = data.len();
625 let shape = shape
626 .iter()
627 .map(usize::to_string)
628 .collect::<Vec<_>>()
629 .join(",");
630 header_parts.push(format!(
631 r#""{name}":{{"dtype":"F32","shape":[{shape}],"data_offsets":[{start},{end}]}}"#
632 ));
633 }
634 let header = format!("{{{}}}", header_parts.join(","));
635 let mut bytes = Vec::with_capacity(8 + header.len() + data.len());
636 bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
637 bytes.extend_from_slice(header.as_bytes());
638 bytes.extend_from_slice(&data);
639 std::fs::write(path, bytes).expect("write safetensors fixture");
640 }
641
642 fn write_tiny_tokenizer_json(dir: &Path) {
643 let tokenizer = r#"{
644 "model": {
645 "type": "BPE",
646 "vocab": {
647 "a": 0, "b": 1, "c": 2, "d": 3,
648 "e": 4, "f": 5, "g": 6, "h": 7,
649 "i": 8, "j": 9, "k": 10, "l": 11,
650 "m": 12, "n": 13, "o": 14, "p": 15
651 },
652 "merges": []
653 }
654 }"#;
655 std::fs::write(dir.join("tokenizer.json"), tokenizer).expect("write tokenizer.json");
656 }
657
658 fn write_tiny_vlm_checkpoint(dir: &Path, indexed: bool) {
659 let config = r#"{
660 "text_config": {
661 "hidden_size": 8,
662 "num_hidden_layers": 1,
663 "vocab_size": 16,
664 "intermediate_size": 4,
665 "rms_norm_eps": 0.000001,
666 "num_attention_heads": 1,
667 "num_key_value_heads": 1,
668 "head_dim": 8,
669 "rope_theta": 10000000.0,
670 "partial_rotary_factor": 1.0,
671 "rope_parameters": {
672 "rope_theta": 10000000.0,
673 "partial_rotary_factor": 1.0,
674 "mrope_section": [2, 1, 1],
675 "mrope_interleaved": true
676 },
677 "linear_num_key_heads": 2,
678 "linear_num_value_heads": 2,
679 "linear_key_head_dim": 32,
680 "linear_value_head_dim": 32,
681 "linear_conv_kernel_dim": 4,
682 "tie_word_embeddings": true,
683 "full_attention_interval": 1,
684 "layer_types": ["full_attention"],
685 "layer_mask": [true],
686 "eos_token_id": 15,
687 "max_position_embeddings": 512
688 },
689 "vision_config": {
690 "depth": 1,
691 "hidden_size": 8,
692 "num_heads": 2,
693 "patch_size": 2,
694 "spatial_merge_size": 2,
695 "out_hidden_size": 8,
696 "temporal_patch_size": 1,
697 "num_position_embeddings": 16,
698 "in_channels": 3,
699 "deepstack_visual_indexes": []
700 },
701 "image_token_id": 9,
702 "vision_start_token_id": 10,
703 "vision_end_token_id": 11,
704 "tie_word_embeddings": true
705 }"#;
706 std::fs::write(dir.join("config.json"), config).expect("write config.json");
707 write_tiny_tokenizer_json(dir);
708
709 let shapes = tiny_vlm_checkpoint_shapes();
710 let shard_name = if indexed {
711 "model-00001-of-00001.safetensors"
712 } else {
713 "model.safetensors"
714 };
715 write_f32_safetensors(&dir.join(shard_name), &shapes);
716 if indexed {
717 let weight_map = shapes
718 .iter()
719 .map(|(name, _)| format!(r#""{name}":"{shard_name}""#))
720 .collect::<Vec<_>>()
721 .join(",");
722 std::fs::write(
723 dir.join("model.safetensors.index.json"),
724 format!(r#"{{"weight_map":{{{weight_map}}}}}"#),
725 )
726 .expect("write one-shard index");
727 }
728 }
729
730 #[test]
735 fn embed_image_matches_raw_inference_primitive() {
736 let (cfg, weights, vision_weights) = tiny_vlm_fixture();
737 let tokenizer = tiny_tokenizer();
738 let png = make_test_png(8, 8, 0);
739
740 let model = VisionEmbeddingModel::new(
741 weights.clone(),
742 cfg.clone(),
743 vision_weights.clone(),
744 tokenizer.clone(),
745 );
746 let via_wrapper = model
747 .embed_image(
748 &png,
749 "describe this image",
750 PoolingStrategy::MeanVisualTokens,
751 )
752 .expect("wrapper embed_image succeeds");
753
754 let via_raw = embed_image_from_bytes_f16(
755 &weights,
756 &cfg,
757 &vision_weights,
758 &tokenizer,
759 &png,
760 "describe this image",
761 PoolingStrategy::MeanVisualTokens,
762 )
763 .expect("raw primitive succeeds");
764
765 assert_eq!(
766 via_wrapper, via_raw,
767 "embed-crate wrapper must return the identical vector to the raw primitive"
768 );
769 }
770
771 #[cfg(all(target_os = "macos", feature = "metal-gpu"))]
776 #[test]
777 fn embed_image_metal_matches_raw_inference_primitive() {
778 use lattice_inference::vision::embed_image_from_bytes_f16_metal;
779
780 let (cfg, weights, vision_weights) = tiny_vlm_fixture();
781 let tokenizer = tiny_tokenizer();
782 let png = make_test_png(8, 8, 0);
783
784 let model = VisionEmbeddingModel::new(
785 weights.clone(),
786 cfg.clone(),
787 vision_weights.clone(),
788 tokenizer.clone(),
789 );
790 let via_wrapper = model
791 .embed_image_metal(
792 &png,
793 "describe this image",
794 PoolingStrategy::MeanVisualTokens,
795 )
796 .expect("wrapper embed_image_metal succeeds");
797
798 let via_raw = embed_image_from_bytes_f16_metal(
799 &weights,
800 &cfg,
801 &vision_weights,
802 &tokenizer,
803 &png,
804 "describe this image",
805 PoolingStrategy::MeanVisualTokens,
806 )
807 .expect("raw metal primitive succeeds");
808
809 assert_eq!(
810 via_wrapper, via_raw,
811 "embed-crate Metal wrapper must return the identical vector to the raw primitive"
812 );
813 }
814
815 #[cfg(not(all(target_os = "macos", feature = "metal-gpu")))]
820 #[test]
821 fn embed_image_metal_fails_closed_without_metal_gpu() {
822 let (cfg, weights, vision_weights) = tiny_vlm_fixture();
823 let tokenizer = tiny_tokenizer();
824 let png = make_test_png(8, 8, 0);
825 let model = VisionEmbeddingModel::new(weights, cfg, vision_weights, tokenizer);
826
827 let err = model
828 .embed_image_metal(
829 &png,
830 "describe this image",
831 PoolingStrategy::MeanVisualTokens,
832 )
833 .expect_err("Metal wrapper must fail without the metal-gpu feature");
834 assert!(matches!(err, EmbedError::InferenceFailed(_)));
835 }
836
837 #[test]
838 fn embed_image_is_deterministic_and_normalized() {
839 let (cfg, weights, vision_weights) = tiny_vlm_fixture();
840 let tokenizer = tiny_tokenizer();
841 let png = make_test_png(8, 8, 0);
842 let model = VisionEmbeddingModel::new(weights, cfg.clone(), vision_weights, tokenizer);
843
844 let v1 = model
845 .embed_image(
846 &png,
847 "describe this image",
848 PoolingStrategy::MeanVisualTokens,
849 )
850 .expect("embed succeeds");
851 let v2 = model
852 .embed_image(
853 &png,
854 "describe this image",
855 PoolingStrategy::MeanVisualTokens,
856 )
857 .expect("embed succeeds");
858
859 assert_eq!(
860 v1, v2,
861 "same image + prompt must produce an identical vector"
862 );
863 assert_eq!(v1.len(), model.dimensions());
864 assert!(v1.iter().all(|x| x.is_finite()));
865 let norm: f32 = v1.iter().map(|x| x * x).sum::<f32>().sqrt();
866 assert!((norm - 1.0).abs() < 1e-4, "expected unit norm, got {norm}");
867 }
868
869 #[test]
870 fn embed_image_rejects_non_vlm_checkpoint() {
871 let (mut cfg, weights, vision_weights) = tiny_vlm_fixture();
872 cfg.vision_config = None;
873 let tokenizer = tiny_tokenizer();
874 let png = make_test_png(8, 8, 0);
875 let model = VisionEmbeddingModel::new(weights, cfg, vision_weights, tokenizer);
876
877 let err = model
878 .embed_image(
879 &png,
880 "describe this image",
881 PoolingStrategy::MeanVisualTokens,
882 )
883 .expect_err("a checkpoint with no vision_config must be rejected");
884 let msg = err.to_string();
885 assert!(matches!(err, EmbedError::InvalidInput(_)));
886 assert!(
887 msg.contains("vision_config"),
888 "error must name the missing field, got: {msg}"
889 );
890 }
891
892 #[test]
893 fn embed_image_rejects_misaligned_image() {
894 let (cfg, weights, vision_weights) = tiny_vlm_fixture();
895 let tokenizer = tiny_tokenizer();
896 let png = make_test_png(6, 4, 0);
898 let model = VisionEmbeddingModel::new(weights, cfg, vision_weights, tokenizer);
899
900 let err = model
901 .embed_image(
902 &png,
903 "describe this image",
904 PoolingStrategy::MeanVisualTokens,
905 )
906 .expect_err("a misaligned image must be rejected, not panic");
907 assert!(matches!(err, EmbedError::InvalidInput(_)));
908 }
909
910 #[test]
911 fn embed_text_matches_raw_inference_primitive() {
912 let (cfg, weights, vision_weights) = tiny_vlm_fixture();
913 let tokenizer = single_char_tokenizer();
914 let model = VisionEmbeddingModel::new(
915 weights.clone(),
916 cfg.clone(),
917 vision_weights,
918 tokenizer.clone(),
919 );
920
921 let via_wrapper = model
922 .embed_text("abc", PoolingStrategy::LastToken)
923 .expect("wrapper embed_text succeeds");
924 let via_raw = lattice_inference::forward::cpu_f16::embed_text_vlm_f16(
925 &weights,
926 &cfg,
927 &tokenizer,
928 "abc",
929 PoolingStrategy::LastToken,
930 )
931 .expect("raw primitive succeeds");
932
933 assert_eq!(via_wrapper, via_raw);
934 }
935
936 #[test]
942 fn embed_text_maps_context_overflow_to_inference_failed() {
943 let (mut cfg, weights, vision_weights) = tiny_vlm_fixture();
944 cfg.max_position_embeddings = 1;
945 let tokenizer = single_char_tokenizer();
946 let model = VisionEmbeddingModel::new(weights, cfg, vision_weights, tokenizer);
947
948 let err = model
949 .embed_text("abc", PoolingStrategy::LastToken)
950 .expect_err("a prompt longer than max_position_embeddings must fail");
951 assert!(
952 matches!(err, EmbedError::InferenceFailed(_)),
953 "context-window overflow is a runtime failure, not caller-input validation, got: {err:?}"
954 );
955 assert!(
956 err.to_string().contains("context window"),
957 "error should retain the underlying context-window detail, got: {err}"
958 );
959 }
960
961 #[test]
962 fn resolve_single_shard_rejects_multi_shard_index() {
963 let tmp = tempfile::tempdir().expect("tempdir");
964 let index_path = tmp.path().join("model.safetensors.index.json");
965 std::fs::write(
966 &index_path,
967 r#"{"metadata":{},"weight_map":{"a":"shard1.safetensors","b":"shard2.safetensors"}}"#,
968 )
969 .expect("write index");
970
971 let err = resolve_qwen35_single_decoder_safetensors(tmp.path())
972 .expect_err("multi-shard must be rejected");
973 let msg = err.to_string();
974 assert!(msg.contains("sharded across 2 files"), "got: {msg}");
975 }
976
977 #[test]
978 fn resolve_single_shard_rejects_missing_checkpoint() {
979 let tmp = tempfile::tempdir().expect("tempdir");
980 let err = resolve_qwen35_single_decoder_safetensors(tmp.path())
981 .expect_err("missing checkpoint must be rejected");
982 assert!(matches!(err, InferenceError::ModelNotFound(_)));
983 }
984
985 #[test]
986 fn from_directory_without_checkpoint_reports_actionable_error() {
987 let tmp = tempfile::tempdir().expect("tempdir");
988 let config_json = include_str!(concat!(
989 env!("CARGO_MANIFEST_DIR"),
990 "/../inference/tests/fixtures/qwen35_0_8b_config.json"
991 ));
992 std::fs::write(tmp.path().join("config.json"), config_json).expect("write config.json");
993 write_tiny_tokenizer_json(tmp.path());
994
995 let Err(err) = VisionEmbeddingModel::from_directory(tmp.path()) else {
996 panic!("a directory with no checkpoint must be rejected")
997 };
998 assert!(matches!(err, EmbedError::ModelInitialization(_)));
999 let msg = err.to_string();
1000 assert!(
1001 msg.contains("model.safetensors") && msg.contains("model.safetensors.index.json"),
1002 "error must name the supported checkpoint layouts, got: {msg}"
1003 );
1004 }
1005
1006 #[test]
1007 fn from_directory_rejects_missing_tokenizer_before_checkpoint_materialization() {
1008 let tmp = tempfile::tempdir().expect("tempdir");
1009 write_tiny_vlm_checkpoint(tmp.path(), false);
1010 std::fs::remove_file(tmp.path().join("tokenizer.json")).expect("remove tokenizer fixture");
1011 std::fs::write(tmp.path().join("model.safetensors"), u64::MAX.to_le_bytes())
1012 .expect("replace checkpoint with a corrupt header");
1013
1014 let Err(err) = VisionEmbeddingModel::from_directory(tmp.path()) else {
1015 panic!("a checkpoint without tokenizer.json must be rejected")
1016 };
1017 assert!(matches!(err, EmbedError::ModelInitialization(_)));
1018 let msg = err.to_string();
1019 assert!(msg.contains("tokenizer.json"), "got: {msg}");
1020 assert!(
1021 !msg.contains("vision weights") && !msg.contains("decoder weights"),
1022 "tokenizer admission must fail before tensor materialization, got: {msg}"
1023 );
1024 }
1025
1026 #[test]
1027 fn from_directory_rejects_quantized_checkpoint_before_tensor_loading() {
1028 let tmp = tempfile::tempdir().expect("tempdir");
1029 std::fs::write(tmp.path().join("quantize_index.json"), b"not valid json")
1030 .expect("write quantized checkpoint sentinel");
1031
1032 let Err(err) = VisionEmbeddingModel::from_directory(tmp.path()) else {
1033 panic!("the f16 pooled decoder loader must reject quantized checkpoints")
1034 };
1035 assert!(matches!(err, EmbedError::ModelInitialization(_)));
1036 let msg = err.to_string();
1037 assert!(msg.contains("quantize_index.json"), "got: {msg}");
1038 assert!(msg.contains("not supported"), "got: {msg}");
1039 assert!(
1040 !msg.contains("config.json"),
1041 "the unsupported file-set must fail before unrelated component loading, got: {msg}"
1042 );
1043 }
1044
1045 #[test]
1046 fn from_directory_loads_single_model_safetensors_without_index() {
1047 let tmp = tempfile::tempdir().expect("tempdir");
1048 write_tiny_vlm_checkpoint(tmp.path(), false);
1049
1050 let model = VisionEmbeddingModel::from_directory(tmp.path())
1051 .expect("single-file VLM checkpoint must load without a synthetic index");
1052 assert_eq!(model.dimensions(), 8);
1053 }
1054
1055 #[test]
1056 fn single_file_and_one_shard_index_produce_identical_image_embeddings() {
1057 let single = tempfile::tempdir().expect("single tempdir");
1058 let indexed = tempfile::tempdir().expect("indexed tempdir");
1059 write_tiny_vlm_checkpoint(single.path(), false);
1060 write_tiny_vlm_checkpoint(indexed.path(), true);
1061
1062 let single_model = VisionEmbeddingModel::from_directory(single.path())
1063 .expect("single-file VLM checkpoint loads");
1064 let indexed_model = VisionEmbeddingModel::from_directory(indexed.path())
1065 .expect("equivalent one-shard indexed VLM checkpoint loads");
1066 let image = make_test_png(4, 4, 17);
1067 let from_single = single_model
1068 .embed_image(&image, "a", PoolingStrategy::MeanVisualTokens)
1069 .expect("single-file image embedding succeeds");
1070 let from_index = indexed_model
1071 .embed_image(&image, "a", PoolingStrategy::MeanVisualTokens)
1072 .expect("indexed image embedding succeeds");
1073
1074 assert_eq!(
1075 from_single, from_index,
1076 "equivalent single-file and indexed layouts must produce parity embeddings"
1077 );
1078 }
1079
1080 #[test]
1081 fn from_directory_rejects_index_map_that_contradicts_opened_shard_header() {
1082 let tmp = tempfile::tempdir().expect("tempdir");
1083 write_tiny_vlm_checkpoint(tmp.path(), true);
1084 std::fs::write(
1085 tmp.path().join("model.safetensors.index.json"),
1086 r#"{"weight_map":{"not.a.real.tensor":"model-00001-of-00001.safetensors"}}"#,
1087 )
1088 .expect("replace index with contradictory weight_map");
1089
1090 let Err(err) = VisionEmbeddingModel::from_directory(tmp.path()) else {
1091 panic!("an authoritative index that omits the physical tensors must be rejected")
1092 };
1093 let msg = err.to_string();
1094 assert!(
1095 msg.contains("weight_map/header inventory mismatch"),
1096 "got: {msg}"
1097 );
1098 }
1099
1100 #[cfg(unix)]
1101 #[test]
1102 fn from_directory_binds_visual_and_decoder_weights_across_path_replacement() {
1103 let tmp = tempfile::tempdir().expect("tempdir");
1104 write_tiny_vlm_checkpoint(tmp.path(), false);
1105 let checkpoint_path = tmp.path().join("model.safetensors");
1106 let replacement = tmp.path().join("replacement-checkpoint");
1107 write_f32_safetensors_with_offset(&replacement, &tiny_vlm_checkpoint_shapes(), 10.0);
1108
1109 let model = with_after_visual_load_hook(
1110 move || {
1111 std::fs::rename(&replacement, &checkpoint_path)
1112 .expect("atomically replace checkpoint pathname with checkpoint B");
1113 },
1114 || {
1115 VisionEmbeddingModel::from_directory(tmp.path())
1116 .expect("constructor keeps both components on checkpoint A")
1117 },
1118 );
1119
1120 assert_eq!(
1121 model.vision_weights.patch_embed_weight[0], 0.14,
1122 "visual weights must remain bound to checkpoint A"
1123 );
1124 let mut expected_embed = [0u16];
1125 f32_to_f16_slice(&[0.01], &mut expected_embed);
1126 assert_eq!(
1127 model.weights.embed_tokens[0], expected_embed[0],
1128 "decoder weights must remain bound to checkpoint A"
1129 );
1130 }
1131
1132 #[test]
1133 fn resolve_single_shard_prefers_existing_index_over_plain_file() {
1134 let tmp = tempfile::tempdir().expect("tempdir");
1135 std::fs::write(tmp.path().join("model.safetensors"), b"plain")
1136 .expect("write convenience file");
1137 std::fs::write(
1138 tmp.path().join("model.safetensors.index.json"),
1139 r#"{"weight_map":{"tensor":"indexed.safetensors"}}"#,
1140 )
1141 .expect("write index");
1142
1143 let resolved = resolve_qwen35_single_decoder_safetensors(tmp.path())
1144 .expect("single-shard index resolves");
1145 assert_eq!(resolved, tmp.path().join("indexed.safetensors"));
1146 }
1147
1148 #[test]
1149 fn resolve_single_shard_rejects_index_entry_escaping_model_directory() {
1150 let tmp = tempfile::tempdir().expect("tempdir");
1151 std::fs::write(
1152 tmp.path().join("model.safetensors.index.json"),
1153 r#"{"weight_map":{"tensor":"../outside.safetensors"}}"#,
1154 )
1155 .expect("write index");
1156
1157 let err = resolve_qwen35_single_decoder_safetensors(tmp.path())
1158 .expect_err("an index entry must not escape the checkpoint directory");
1159 assert!(
1160 err.to_string().contains("escapes the model directory"),
1161 "got: {err}"
1162 );
1163 }
1164}