spark_model/mistral_loader/loader_impl/
mod.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! `impl ModelWeightLoader for MistralWeightLoader` — the trait body
4//! delegates to a per-layer phased pipeline split across sibling files
5//! to stay under the ≤500 LoC cap.
6//!
7//! Layer-loading phases:
8//! - `phase_lora_qkv`     — wq_a/b, wkv_a/b LoRA + NVFP4 + TP shard
9//! - `phase_per_head`     — W_UK_T / W_UV per-head transpose, wq_b_rope
10//! - `phase_qk_absorbed`  — fused W_QK_absorbed CPU compute
11//! - `phase_block_diag`   — block-diagonal W_UK_BD / W_UV_BD
12//! - `phase_o_proj`       — output projection NVFP4
13//! - `yarn`               — YaRN inv_freq table (computed once)
14//! - `phase_assemble`     — MlaWeights + MoE + TransformerLayer
15
16use 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
28// `ctx` + the three name-independent transform phases are `pub(crate)` so the
29// LongCat MLA loader reuses the SAME per-head transpose / absorbed-QK /
30// block-diagonal math instead of duplicating it (only the two name-bound
31// phases — lora_qkv and o_proj — are Mistral-specific).
32pub(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        // MLA TP: wq_b and wkv_b are sharded ColumnParallel on the
44        // head-output axis. wq_a / wkv_a (LoRA down-projections to
45        // latent) stay replicated since the latent dim isn't
46        // head-dependent. q_a_norm / kv_a_norm (latent norms) also
47        // replicated. The CPU-side wkv_b transpose for absorbed MLA
48        // (W_UK / W_UV) iterates over `config.num_key_value_heads`
49        // which is already TP-local after main.rs's head split, so
50        // it naturally produces per-rank absorbed weights.
51        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
124/// Inherent helpers — outside the trait impl block.
125impl 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}