1use 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 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 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 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 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, 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) }
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}