spark_model/mistral_loader/loader_impl/
mod.rs1use anyhow::{Context, Result};
17use atlas_core::config::ModelConfig;
18use spark_runtime::gpu::GpuBackend;
19use spark_runtime::kv_cache::KvCacheDtype;
20use spark_runtime::weights::WeightStore;
21
22use super::MistralWeightLoader;
23use crate::layer::TransformerLayer;
24use crate::layers::vision_encoder::VisionEncoder;
25use crate::weight_loader::ModelWeightLoader;
26use crate::weight_map::{DenseWeight, MtpWeights, dense};
27
28pub(crate) mod ctx;
33mod phase_assemble;
34pub(crate) mod phase_block_diag;
35mod phase_lora_qkv;
36mod phase_o_proj;
37pub(crate) mod phase_per_head;
38pub(crate) mod phase_qk_absorbed;
39mod yarn;
40
41impl ModelWeightLoader for MistralWeightLoader {
42 fn supports_tp(&self) -> bool {
43 true
52 }
53
54 fn load_layers(
55 &self,
56 store: &WeightStore,
57 config: &ModelConfig,
58 gpu: &dyn GpuBackend,
59 layer_kv_dtypes: &[KvCacheDtype],
60 ) -> Result<Vec<Box<dyn TransformerLayer>>> {
61 self.load_layers_inner(store, config, gpu, layer_kv_dtypes)
62 }
63
64 fn load_embedding(
65 &self,
66 store: &WeightStore,
67 _config: &ModelConfig,
68 _gpu: &dyn GpuBackend,
69 ) -> Result<DenseWeight> {
70 dense(store, "tok_embeddings.weight")
71 .or_else(|_| dense(store, "model.embed_tokens.weight"))
72 .context("Mistral: embedding not found")
73 }
74
75 fn load_final_norm(
76 &self,
77 store: &WeightStore,
78 _config: &ModelConfig,
79 _gpu: &dyn GpuBackend,
80 ) -> Result<DenseWeight> {
81 dense(store, "norm.weight")
82 .or_else(|_| dense(store, "model.norm.weight"))
83 .context("Mistral: final norm not found")
84 }
85
86 fn load_lm_head(
87 &self,
88 store: &WeightStore,
89 config: &ModelConfig,
90 _gpu: &dyn GpuBackend,
91 ) -> Result<DenseWeight> {
92 if store.contains("output.weight") {
93 dense(store, "output.weight")
94 } else if store.contains("lm_head.weight") {
95 dense(store, "lm_head.weight")
96 } else if config.tie_word_embeddings {
97 dense(store, "tok_embeddings.weight")
98 .or_else(|_| dense(store, "model.embed_tokens.weight"))
99 .context("Mistral: tied embedding lm_head not found")
100 } else {
101 anyhow::bail!("Mistral: lm_head/output weight not found")
102 }
103 }
104
105 fn load_mtp_weights(
106 &self,
107 _store: &WeightStore,
108 _config: &ModelConfig,
109 _gpu: &dyn GpuBackend,
110 ) -> Result<Option<MtpWeights>> {
111 Ok(None)
112 }
113
114 fn load_vision_encoder(
115 &self,
116 _store: &WeightStore,
117 _config: &ModelConfig,
118 _gpu: &dyn GpuBackend,
119 ) -> Result<Option<VisionEncoder>> {
120 Ok(None)
121 }
122}
123
124impl MistralWeightLoader {
126 pub(crate) fn load_layers_inner(
127 &self,
128 store: &WeightStore,
129 config: &ModelConfig,
130 gpu: &dyn GpuBackend,
131 layer_kv_dtypes: &[KvCacheDtype],
132 ) -> Result<Vec<Box<dyn TransformerLayer>>> {
133 let n = config.num_hidden_layers;
134 let q_lora = config.q_lora_rank;
135 let kv_lora = config.kv_lora_rank;
136 let nope = config.qk_nope_head_dim;
137 let rope = config.qk_rope_head_dim;
138 let v_dim = config.v_head_dim;
139
140 tracing::info!(
141 "Mistral MLA→GQA: expanding LoRA on GPU (q_lora={q_lora}, kv_lora={kv_lora}, \
142 nope={nope}, rope={rope}, v_dim={v_dim})"
143 );
144
145 let absmax_k = gpu.kernel("quantize_nvfp4", "nvfp4_global_absmax")?;
146 let quantize_k = gpu.kernel("quantize_nvfp4", "quantize_bf16_to_nvfp4")?;
147 let stream = gpu.default_stream();
148
149 let mut layers: Vec<Box<dyn TransformerLayer>> = Vec::with_capacity(n);
150 let mut yarn_inv_freq_shared = spark_runtime::gpu::DevicePtr::NULL;
151
152 for i in 0..n {
153 let mut ctx =
154 ctx::MistralLayerCtx::new(store, config, gpu, absmax_k, quantize_k, stream, i);
155 phase_lora_qkv::load_lora_qkv(&mut ctx)?;
156 phase_per_head::build_per_head_views(&mut ctx)?;
157 phase_qk_absorbed::build_w_qk_absorbed(&mut ctx)?;
158 phase_block_diag::build_block_diagonals(&mut ctx)?;
159 phase_o_proj::load_o_proj(&mut ctx)?;
160 let yarn_inv_freq =
161 ctx::ensure_yarn_inv_freq(&mut yarn_inv_freq_shared, config, rope, gpu)?;
162 let layer = phase_assemble::assemble_layer(ctx, yarn_inv_freq, layer_kv_dtypes)?;
163 layers.push(layer);
164
165 if (i + 1) % 6 == 0 || i == n - 1 {
166 let free = gpu.free_memory().unwrap_or(0);
167 tracing::info!("L{}/{n} done — {:.1} GB free", i + 1, free as f64 / 1e9);
168 }
169 }
170 Ok(layers)
171 }
172}