Skip to main content

lattice_embed/
vision.rs

1//! Image (and image+text) embedding through the Qwen3.5 vision-language
2//! pooled-embedding pipeline (ADR-069 S5, #1007).
3//!
4//! This is a wire-through: [`VisionEmbeddingModel::embed_image`] and
5//! [`VisionEmbeddingModel::embed_text`] call straight into
6//! `lattice_inference::vision::embed_image_from_bytes_f16` /
7//! `lattice_inference::forward::cpu_f16::embed_text_vlm_f16`, the same
8//! pooling + L2-normalization contract #1007 established. No new math lives
9//! here — only checkpoint loading (mirroring the directory-loading pattern
10//! `service::native` uses for the BERT/Qwen text models) and error mapping.
11//!
12//! [`VisionEmbeddingModel::from_directory`] accepts either a
13//! `model.safetensors.index.json` naming exactly one decoder shard (the
14//! Qwen3.5-0.8B layout), or an unindexed directory containing exactly one
15//! safetensors file named `model.safetensors`. The indexed path is
16//! authoritative when both are present. The unindexed path parses the
17//! safetensors header and validates the exact `model.visual.*` inventory
18//! before tensor payloads are materialized. Decoder and visual tensors are
19//! materialized from the same mmap-backed open file, so a path replacement
20//! during loading cannot mix two checkpoint versions. For an indexed
21//! one-shard layout, the authoritative `weight_map` must exactly match that
22//! opened shard's header inventory. `quantize_index.json` is rejected: this
23//! loader constructs an f16 decoder and cannot bind a quantized visual file
24//! set coherently. Callers with pre-loaded components, or a multi-shard
25//! decoder checkpoint, can assemble their own weights and call
26//! [`VisionEmbeddingModel::new`] directly.
27
28use 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
85/// A loaded Qwen3.5 vision-language checkpoint, ready to pool image (and
86/// image+text) embeddings.
87///
88/// See [`docs/model.md`](../docs/model.md) for the general model-loading design; this type
89/// follows the same "load once, reuse" shape as `NativeEmbeddingService`'s wrapped models.
90pub struct VisionEmbeddingModel {
91    weights: F16ModelWeights,
92    config: Qwen35Config,
93    vision_weights: Qwen35VisionWeights,
94    tokenizer: BpeTokenizer,
95}
96
97impl VisionEmbeddingModel {
98    /// Compose a model from already-loaded components (no I/O).
99    ///
100    /// Use this when the checkpoint spans multiple safetensors shards (not
101    /// supported by [`Self::from_directory`]) or when components are shared
102    /// across other in-process model instances.
103    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    /// Load a Qwen3.5 vision-language checkpoint directory: `config.json`,
118    /// `tokenizer.json`, the `model.visual.*` vision-encoder tensors, and a
119    /// single-shard f16 decoder checkpoint. An existing
120    /// `model.safetensors.index.json` is authoritative. Without an index, the
121    /// directory must contain exactly one `*.safetensors` file, named
122    /// `model.safetensors`; its header must carry the exact `model.visual.*`
123    /// inventory implied by `vision_config`. A `quantize_index.json` is
124    /// rejected because this constructor has no quantized decoder loader. The
125    /// chosen safetensors file is opened once and that same reader supplies
126    /// both visual and decoder tensors.
127    ///
128    /// # Errors
129    ///
130    /// Returns [`EmbedError::ModelInitialization`] if `config.json` is
131    /// missing or invalid, if the checkpoint has no `vision_config`, if a
132    /// quantized checkpoint is present, if the unindexed safetensors layout is
133    /// missing or ambiguous, if the decoder weights are sharded across more
134    /// than one file, or if any component tensor fails to load.
135    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        // Tokenizer parsing is independent of tensor payloads, so reject a
164        // missing or malformed tokenizer before materializing multi-GB model
165        // weights. The checkpoint itself is still opened only once below so
166        // visual and decoder tensors stay bound to one file descriptor.
167        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    /// Pool an image (plus an optional text prompt) into a single
186    /// L2-normalized `[dimensions()]` embedding vector.
187    ///
188    /// Same scaffold and pooling contract as
189    /// [`lattice_inference::vision::embed_image_from_bytes_f16`] (see that
190    /// function's docs for the exact prompt-assembly layout).
191    ///
192    /// # Errors
193    ///
194    /// Returns [`EmbedError::InvalidInput`] if `image_bytes` cannot be
195    /// decoded, its dimensions are not compatible with the checkpoint's
196    /// patch/merge geometry, or the assembled request otherwise fails
197    /// validation (the error message names the offending field). Returns
198    /// [`EmbedError::InferenceFailed`] for every other underlying failure —
199    /// e.g. the prompt plus image tokens exceeding the checkpoint's context
200    /// window.
201    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    /// Metal-dispatching sibling of [`Self::embed_image`]: runs the ViT
220    /// forward pass on the Metal GPU instead of the CPU (see
221    /// [`lattice_inference::vision::embed_image_from_bytes_f16_metal`]).
222    /// Same scaffold, pooling contract, and error semantics as
223    /// [`Self::embed_image`], with one addition: on a build/platform with no
224    /// Metal support, this returns [`EmbedError::InferenceFailed`] (Metal
225    /// unavailability is a runtime-backend failure, not caller-input
226    /// validation) rather than silently falling back to
227    /// [`Self::embed_image`]'s CPU path — callers that want a fallback must
228    /// call [`Self::embed_image`] themselves.
229    ///
230    /// # Errors
231    ///
232    /// See [`Self::embed_image`]'s docs.
233    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    /// Pool a text-only prompt through the same decoder + pooling path as
252    /// [`Self::embed_image`], landing in the same vector space.
253    ///
254    /// # Errors
255    ///
256    /// Returns [`EmbedError::InvalidInput`] if the prompt is empty or
257    /// tokenizes to an out-of-vocabulary id. Returns
258    /// [`EmbedError::InferenceFailed`] for every other underlying failure —
259    /// e.g. the prompt exceeding the checkpoint's context window.
260    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    /// Output embedding dimension (the checkpoint's decoder hidden size).
272    pub fn dimensions(&self) -> usize {
273        self.config.hidden_size
274    }
275}
276
277/// Map an inference-layer error to the embed crate's two-variant contract:
278/// caller-supplied-input problems stay distinguishable from every other
279/// (model/runtime) failure, so callers can tell "fix your request" apart
280/// from "retry or report a bug" (see `embed_image`/`embed_text` docs).
281fn 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    /// Deterministic pseudo-random f32 fill (xorshift LCG), matching the
301    /// fixture builder in `lattice_inference::vision::pooled_embed`'s own
302    /// unit tests, so this crate's wrapper is exercised against
303    /// non-trivial weights without needing a real checkpoint.
304    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, // must equal decoder hidden_size below
323            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    /// A minimal one-layer full-attention decoder + vision config wired
385    /// together: small enough to hand-construct, non-trivial (pseudo-random)
386    /// projections so the pipeline is actually exercised end to end.
387    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    /// Single-character vocab: with no merges, a byte-level BPE tokenizer
496    /// falls back to per-character tokens, so (unlike `tiny_tokenizer`'s
497    /// whole-word entries) this actually produces non-empty `real_length`
498    /// output — required by `embed_text_vlm_f16`'s empty-prompt guard.
499    /// Mirrors the tokenizer `cpu_f16.rs`'s own `embed_text_vlm_f16` tests use.
500    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    /// The embed-crate wrapper must return the exact same vector as calling
731    /// the raw inference-crate primitive directly: wiring adds no numerical
732    /// difference. This is the core claim of this module (a wire-through,
733    /// not a reimplementation).
734    #[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    /// Same wire-through claim as
772    /// `embed_image_matches_raw_inference_primitive`, for the Metal entry
773    /// point: the crate wrapper must add no numerical difference relative to
774    /// calling `embed_image_from_bytes_f16_metal` directly.
775    #[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    /// Off this cfg gate, the wrapper must surface the CPU/Metal
816    /// distinction the inference crate makes (`InferenceError::
817    /// UnsupportedModel`) as `EmbedError::InferenceFailed`, never silently
818    /// substituting `embed_image`'s CPU output.
819    #[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        // factor = patch_size(2) * merge(2) = 4; 6 is not a multiple of 4.
897        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    /// A runtime (non-input) failure -- the prompt exceeding the checkpoint's
937    /// context window, surfaced as `InferenceError::Inference` from the
938    /// shared prefill path (cpu_f16.rs) -- must map to
939    /// `EmbedError::InferenceFailed`, not `EmbedError::InvalidInput`: the
940    /// prompt itself is well-formed, the checkpoint just can't fit it.
941    #[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}