spark_model/weight_loader/
qwen3_vl.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3use anyhow::{Context, Result};
4use atlas_core::config::ModelConfig;
5use spark_runtime::gpu::GpuBackend;
6use spark_runtime::kv_cache::KvCacheDtype;
7use spark_runtime::weights::WeightStore;
8
9use super::{ModelWeightLoader, WeightFormat};
10use crate::layer::TransformerLayer;
11use crate::layers::vision_encoder::{MergerLayer, ViTBlock};
12use crate::layers::{FfnComponent, MoeLayer, Qwen3AttentionLayer, VisionEncoder};
13use crate::tp_shard::{TpShardKind, load_qkvo_tp, shard_quantized_nvfp4};
14use crate::weight_map::{
15    AttentionWeights, DenseWeight, MtpWeights, dense, detect_nvfp4_variant, load_kv_scales,
16    load_moe_no_shared, quantize_to_nvfp4, quantized_auto,
17};
18
19pub struct Qwen3VLWeightLoader;
20
21impl ModelWeightLoader for Qwen3VLWeightLoader {
22    fn supports_tp(&self) -> bool {
23        // Q/K/V column-parallel + O row-parallel via shard_quantized_nvfp4
24        // (single quant path: NVFP4 from disk). Per-head q_norm/k_norm
25        // are replicated naturally — TP just shards heads, the norm
26        // weights duplicate. MoE and vision encoder remain full-replica.
27        true
28    }
29
30    fn load_layers(
31        &self,
32        store: &WeightStore,
33        config: &ModelConfig,
34        gpu: &dyn GpuBackend,
35        layer_kv_dtypes: &[KvCacheDtype],
36    ) -> Result<Vec<Box<dyn TransformerLayer>>> {
37        let mut layers: Vec<Box<dyn TransformerLayer>> =
38            Vec::with_capacity(config.num_hidden_layers);
39
40        let variant = detect_nvfp4_variant(store, config);
41        let weight_format = WeightFormat::detect(store, config);
42        tracing::info!(
43            "Weight format: {:?}, NVFP4 variant: {:?}",
44            weight_format,
45            variant
46        );
47
48        // MoE prefill copies. This loader used to build NONE of them, so every
49        // prefill took the fallback `moe_w4a16_grouped_gemm` — measured 695 ms in
50        // `grouped_gate_up` + 351 ms in `grouped_silu_down` on
51        // ig1/Qwen3-VL-30B-A3B-Instruct-NVFP4, i.e. 59 % of a 1798 ms cold TTFT,
52        // against 98.9 / 69.5 ms for a model that does build them. Decided once
53        // here rather than per layer: `free_memory()` shrinks as layers load, so
54        // a per-layer probe transposes the early layers and silently skips the
55        // late ones, leaving prefill on a mix of both paths.
56        let moe_prefill_copies = super::moe_prefill_copies_fit(config, gpu);
57        let absmax_k = gpu.kernel("quantize_nvfp4", "nvfp4_global_absmax")?;
58        let quantize_k = gpu.kernel("quantize_nvfp4", "quantize_bf16_to_nvfp4")?;
59        let stream = gpu.default_stream();
60        let h = config.hidden_size;
61
62        for i in 0..config.num_hidden_layers {
63            let lp = config.layer_prefix(i);
64            let input_norm = dense(store, &format!("{lp}.input_layernorm.weight"))?;
65            let post_attn_norm = dense(store, &format!("{lp}.post_attention_layernorm.weight"))?;
66
67            // MoE without shared experts
68            let moe_weights =
69                load_moe_no_shared(store, &lp, config.num_experts, gpu, config, variant)?;
70            let gate_nvfp4 = quantize_to_nvfp4(
71                &moe_weights.gate,
72                config.num_experts,
73                h,
74                gpu,
75                absmax_k,
76                quantize_k,
77                stream,
78            )?;
79            let mut moe_layer = MoeLayer::new(
80                moe_weights,
81                config.num_experts,
82                Some(gate_nvfp4),
83                gpu,
84                config,
85            )?;
86            if moe_prefill_copies {
87                moe_layer.transpose_for_prefill(gpu, config)?;
88                moe_layer.predequant_for_prefill(gpu, config, stream)?;
89            }
90            let ffn = FfnComponent::Moe(moe_layer);
91
92            // All layers are FullAttention with ungated Q projection.
93            //
94            // TP: Q/K/V column-parallel + O row-parallel via the shared
95            // `load_qkvo_tp` helper. Each of the four projections is loaded
96            // as NVFP4 from disk, then sliced on the matching axis.
97            // NVFP4 group_size is 16 for Qwen3-VL.
98            let p = format!("{lp}.self_attn");
99            let tp_rank = config.tp_rank;
100            let tp_size = config.tp_world_size.max(1);
101            let group_size = 16usize;
102            let load_proj = |name: &str,
103                             full_n: usize,
104                             full_k: usize,
105                             kind: TpShardKind|
106             -> Result<crate::weight_map::QuantizedWeight> {
107                let src = quantized_auto(store, &format!("{p}.{name}"), gpu, variant)?;
108                if tp_size == 1 {
109                    return Ok(src);
110                }
111                let sharded = shard_quantized_nvfp4(
112                    &src, full_n, full_k, kind, tp_rank, tp_size, group_size, gpu,
113                )?;
114                gpu.free(src.weight)?;
115                gpu.free(src.weight_scale)?;
116                Ok(sharded)
117            };
118            let [q, k, v, o] = load_qkvo_tp(config, load_proj)?;
119            let dummy = DenseWeight {
120                weight: spark_runtime::gpu::DevicePtr::NULL,
121            };
122            let (k_scale, v_scale) = load_kv_scales(store, &p, gpu);
123            let attn = AttentionWeights {
124                q_proj: dummy,
125                k_proj: dummy,
126                v_proj: dummy,
127                o_proj: o,
128                q_norm: dense(store, &format!("{p}.q_norm.weight"))?,
129                k_norm: dense(store, &format!("{p}.k_norm.weight"))?,
130                q_norm_full: None,
131                k_norm_full: None,
132                k_scale,
133                v_scale,
134            };
135
136            layers.push(Box::new(Qwen3AttentionLayer::new_ungated(
137                input_norm,
138                attn,
139                post_attn_norm,
140                ffn,
141                i, // attn_layer_idx = layer index (all layers are attention)
142                Some(q),
143                Some(k),
144                Some(v),
145                gpu,
146                layer_kv_dtypes[i],
147                config.fp8_kv_calibration_tokens,
148                config,
149            )?));
150
151            if (i + 1) % 10 == 0 {
152                tracing::info!("Loaded layers 0..{}", i + 1);
153            }
154        }
155
156        tracing::info!(
157            "Qwen3-VL weight loader: {} layers (all attention, ungated)",
158            layers.len(),
159        );
160
161        Ok(layers)
162    }
163
164    fn load_embedding(
165        &self,
166        store: &WeightStore,
167        config: &ModelConfig,
168        _gpu: &dyn GpuBackend,
169    ) -> Result<DenseWeight> {
170        let prefix = &config.weight_prefix;
171        dense(store, &format!("{prefix}.embed_tokens.weight"))
172    }
173
174    fn load_final_norm(
175        &self,
176        store: &WeightStore,
177        config: &ModelConfig,
178        _gpu: &dyn GpuBackend,
179    ) -> Result<DenseWeight> {
180        let prefix = &config.weight_prefix;
181        dense(store, &format!("{prefix}.norm.weight"))
182    }
183
184    fn load_lm_head(
185        &self,
186        store: &WeightStore,
187        config: &ModelConfig,
188        _gpu: &dyn GpuBackend,
189    ) -> Result<DenseWeight> {
190        for pattern in &[
191            "lm_head.weight",
192            "language_model.lm_head.weight",
193            "model.lm_head.weight",
194        ] {
195            if store.contains(pattern) {
196                return dense(store, pattern);
197            }
198        }
199        let prefix = &config.weight_prefix;
200        dense(store, &format!("{prefix}.embed_tokens.weight"))
201    }
202
203    fn load_mtp_weights(
204        &self,
205        _store: &WeightStore,
206        _config: &ModelConfig,
207        _gpu: &dyn GpuBackend,
208    ) -> Result<Option<MtpWeights>> {
209        Ok(None) // VL model has no MTP
210    }
211
212    fn load_vision_encoder(
213        &self,
214        store: &WeightStore,
215        config: &ModelConfig,
216        gpu: &dyn GpuBackend,
217    ) -> Result<Option<VisionEncoder>> {
218        let vcfg = match &config.vision {
219            Some(v) => v.clone(),
220            None => return Ok(None),
221        };
222        let vp = "model.visual";
223
224        let patch_embed_w = dense(store, &format!("{vp}.patch_embed.proj.weight"))?;
225        let patch_embed_b = dense(store, &format!("{vp}.patch_embed.proj.bias"))?;
226        let pos_embed = dense(store, &format!("{vp}.pos_embed.weight"))?;
227        let pos_embed_shape = store.get(&format!("{vp}.pos_embed.weight"))?.shape.clone();
228        let num_position_embeddings = pos_embed_shape
229            .first()
230            .copied()
231            .context("pos_embed shape missing rows")?;
232
233        let mut blocks = Vec::with_capacity(vcfg.depth);
234        for i in 0..vcfg.depth {
235            let bp = format!("{vp}.blocks.{i}");
236            blocks.push(ViTBlock {
237                norm1_w: dense(store, &format!("{bp}.norm1.weight"))?.weight,
238                norm1_b: dense(store, &format!("{bp}.norm1.bias"))?.weight,
239                qkv_w: dense(store, &format!("{bp}.attn.qkv.weight"))?.weight,
240                qkv_b: dense(store, &format!("{bp}.attn.qkv.bias"))?.weight,
241                proj_w: dense(store, &format!("{bp}.attn.proj.weight"))?.weight,
242                proj_b: dense(store, &format!("{bp}.attn.proj.bias"))?.weight,
243                norm2_w: dense(store, &format!("{bp}.norm2.weight"))?.weight,
244                norm2_b: dense(store, &format!("{bp}.norm2.bias"))?.weight,
245                fc1_w: dense(store, &format!("{bp}.mlp.linear_fc1.weight"))?.weight,
246                fc1_b: dense(store, &format!("{bp}.mlp.linear_fc1.bias"))?.weight,
247                fc2_w: dense(store, &format!("{bp}.mlp.linear_fc2.weight"))?.weight,
248                fc2_b: dense(store, &format!("{bp}.mlp.linear_fc2.bias"))?.weight,
249            });
250        }
251
252        let mut deepstack = Vec::with_capacity(vcfg.deepstack_visual_indexes.len());
253        for i in 0..vcfg.deepstack_visual_indexes.len() {
254            let mp = format!("{vp}.deepstack_merger_list.{i}");
255            deepstack.push(MergerLayer {
256                norm_w: dense(store, &format!("{mp}.norm.weight"))?.weight,
257                norm_b: dense(store, &format!("{mp}.norm.bias"))?.weight,
258                fc1_w: dense(store, &format!("{mp}.linear_fc1.weight"))?.weight,
259                fc1_b: dense(store, &format!("{mp}.linear_fc1.bias"))?.weight,
260                fc2_w: dense(store, &format!("{mp}.linear_fc2.weight"))?.weight,
261                fc2_b: dense(store, &format!("{mp}.linear_fc2.bias"))?.weight,
262            });
263        }
264
265        let mp = format!("{vp}.merger");
266        let merger = MergerLayer {
267            norm_w: dense(store, &format!("{mp}.norm.weight"))?.weight,
268            norm_b: dense(store, &format!("{mp}.norm.bias"))?.weight,
269            fc1_w: dense(store, &format!("{mp}.linear_fc1.weight"))?.weight,
270            fc1_b: dense(store, &format!("{mp}.linear_fc1.bias"))?.weight,
271            fc2_w: dense(store, &format!("{mp}.linear_fc2.weight"))?.weight,
272            fc2_b: dense(store, &format!("{mp}.linear_fc2.bias"))?.weight,
273        };
274
275        let deepstack_indexes = vcfg.deepstack_visual_indexes.clone();
276        let ve = VisionEncoder::new(
277            patch_embed_w.weight,
278            patch_embed_b.weight,
279            pos_embed.weight,
280            num_position_embeddings,
281            blocks,
282            deepstack,
283            deepstack_indexes,
284            merger,
285            vcfg.hidden_size,
286            vcfg.num_heads,
287            vcfg.spatial_merge_size,
288            vcfg.out_hidden_size,
289            vcfg.intermediate_size,
290            vcfg.patch_size,
291            vcfg.max_pixels,
292            gpu,
293        )?;
294        tracing::info!(
295            "Vision encoder loaded: depth={}, hidden={}, heads={}, deepstack={:?}",
296            vcfg.depth,
297            vcfg.hidden_size,
298            vcfg.num_heads,
299            vcfg.deepstack_visual_indexes,
300        );
301        Ok(Some(ve))
302    }
303}