spark_model/layers/glm5next_dsa/
build.rs1use anyhow::{Result, bail};
29use spark_runtime::gpu::{DevicePtr, GpuBackend};
30
31use super::Glm5NextDsaConfig;
32use super::layer::Glm5NextDsaWeights;
33use super::tp::{DsaShard, DsaTpPlan};
34
35pub type LoadFn<'a> = &'a dyn Fn(&str) -> Result<Vec<f32>>;
37
38pub fn absorb_q(
47 cfg: &Glm5NextDsaConfig,
48 q_b: &[f32],
49 kv_b: &[f32],
50 full_heads: usize,
51) -> Result<Vec<f32>> {
52 let (nope, vd, kvl, ql) = (
53 cfg.qk_nope_head_dim,
54 cfg.v_head_dim,
55 cfg.kv_lora_rank,
56 cfg.q_lora_rank,
57 );
58 let qk = cfg.qk_head_dim();
59 if q_b.len() != full_heads * qk * ql {
60 bail!(
61 "absorb_q: q_b_proj has {} elems, expected {}",
62 q_b.len(),
63 full_heads * qk * ql
64 );
65 }
66 if kv_b.len() != full_heads * (nope + vd) * kvl {
67 bail!(
68 "absorb_q: kv_b_proj has {} elems, expected {}",
69 kv_b.len(),
70 full_heads * (nope + vd) * kvl
71 );
72 }
73 if cfg.qk_rope_head_dim != 0 {
76 bail!(
77 "absorb_q: NoPE only; qk_rope_head_dim is {}",
78 cfg.qk_rope_head_dim
79 );
80 }
81
82 let mut out = vec![0f32; full_heads * kvl * ql];
83 for h in 0..full_heads {
84 let kv_base = h * (nope + vd);
85 let qb_base = h * nope;
86 for c in 0..kvl {
87 for k in 0..ql {
88 let mut acc = 0f32;
89 for r in 0..nope {
90 acc += kv_b[(kv_base + r) * kvl + c] * q_b[(qb_base + r) * ql + k];
91 }
92 out[(h * kvl + c) * ql + k] = acc;
93 }
94 }
95 }
96 Ok(out)
97}
98
99pub fn absorb_o(
116 cfg: &Glm5NextDsaConfig,
117 o_local: &[f32],
118 kv_b_local: &[f32],
119 local_heads: usize,
120) -> Result<Vec<f32>> {
121 let (nope, vd, kvl, hidden) = (
122 cfg.qk_nope_head_dim,
123 cfg.v_head_dim,
124 cfg.kv_lora_rank,
125 cfg.hidden,
126 );
127 if cfg.qk_rope_head_dim != 0 {
128 bail!(
129 "absorb_o: NoPE only; qk_rope_head_dim is {}",
130 cfg.qk_rope_head_dim
131 );
132 }
133 let in_w = local_heads * vd;
134 let out_w = local_heads * kvl;
135 if o_local.len() != hidden * in_w {
136 bail!(
137 "absorb_o: o_proj has {} elems, expected {} ([{hidden}, {in_w}])",
138 o_local.len(),
139 hidden * in_w
140 );
141 }
142 if kv_b_local.len() != local_heads * (nope + vd) * kvl {
143 bail!(
144 "absorb_o: kv_b_proj slice has {} elems, expected {}",
145 kv_b_local.len(),
146 local_heads * (nope + vd) * kvl
147 );
148 }
149
150 let mut out = vec![0f32; hidden * out_w];
151 let threads = std::thread::available_parallelism()
156 .map(|n| n.get())
157 .unwrap_or(1)
158 .clamp(1, hidden);
159 let rows = hidden.div_ceil(threads);
160 std::thread::scope(|sc| {
161 for (o_chunk, i_chunk) in out
162 .chunks_mut(rows * out_w)
163 .zip(o_local.chunks(rows * in_w))
164 {
165 sc.spawn(move || {
166 for (dst_row, src_row) in o_chunk.chunks_mut(out_w).zip(i_chunk.chunks(in_w)) {
167 for h in 0..local_heads {
168 let dst = &mut dst_row[h * kvl..(h + 1) * kvl];
169 let v_base = h * (nope + vd) + nope;
170 for r in 0..vd {
171 let w = src_row[h * vd + r];
172 let src = &kv_b_local[(v_base + r) * kvl..(v_base + r) * kvl + kvl];
173 for (d, k) in dst.iter_mut().zip(src) {
174 *d += w * k;
175 }
176 }
177 }
178 }
179 });
180 }
181 });
182 Ok(out)
183}
184
185fn row_slice(v: &[f32], row_elems: usize, start: usize, end: usize) -> Vec<f32> {
187 v[start * row_elems..end * row_elems].to_vec()
188}
189
190fn col_slice(v: &[f32], row_elems: usize, start: usize, end: usize) -> Vec<f32> {
192 v.chunks(row_elems)
193 .flat_map(|r| r[start..end].iter().copied())
194 .collect()
195}
196
197pub fn shard_host(plan: &super::tp::DsaTensorPlan, full: &[f32]) -> Vec<f32> {
199 match plan.kind {
200 DsaShard::Replicated => full.to_vec(),
201 DsaShard::HeadRows => row_slice(
202 full,
203 plan.full_row_elems,
204 plan.src_row_offset,
205 plan.src_row_offset + plan.local_rows,
206 ),
207 DsaShard::HeadCols => col_slice(
208 full,
209 plan.full_row_elems,
210 plan.src_col_offset,
211 plan.src_col_offset + plan.local_row_elems,
212 ),
213 }
214}
215
216fn up_bf16(gpu: &dyn GpuBackend, v: &[f32]) -> Result<DevicePtr> {
217 let b: Vec<u8> = v
218 .iter()
219 .flat_map(|x| half::bf16::from_f32(*x).to_le_bytes())
220 .collect();
221 let p = gpu.alloc(b.len().max(1))?;
222 gpu.copy_h2d(&b, p)?;
223 Ok(p)
224}
225fn up_f32(gpu: &dyn GpuBackend, v: &[f32]) -> Result<DevicePtr> {
226 let b: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
227 let p = gpu.alloc(b.len().max(1))?;
228 gpu.copy_h2d(&b, p)?;
229 Ok(p)
230}
231
232pub fn build_dsa_weights(
234 gpu: &dyn GpuBackend,
235 cfg: &Glm5NextDsaConfig,
236 plan: &DsaTpPlan,
237 load: LoadFn<'_>,
238) -> Result<Glm5NextDsaWeights> {
239 let get = |n: &str| -> Result<Vec<f32>> { load(&format!("self_attn.{n}")) };
240 let shard = |n: &'static str, full: Vec<f32>| -> Result<Vec<f32>> {
241 let p = plan
242 .get(n)
243 .ok_or_else(|| anyhow::anyhow!("no shard plan for {n}"))?;
244 Ok(shard_host(p, &full))
245 };
246
247 let kv_b = get("kv_b_proj.weight")?;
248
249 let q_absorb_full = absorb_q(cfg, &get("q_b_proj.weight")?, &kv_b, plan.full_heads)?;
253 let per_head = cfg.kv_lora_rank;
254 let start = plan.tp_rank * plan.local_heads * per_head;
255 let len = plan.local_heads * per_head;
256 let q_absorb = row_slice(&q_absorb_full, cfg.q_lora_rank, start, start + len);
257
258 let kv_b_rows = cfg.qk_nope_head_dim + cfg.v_head_dim;
262 let kv_b_local = row_slice(
263 &kv_b,
264 cfg.kv_lora_rank,
265 plan.tp_rank * plan.local_heads * kv_b_rows,
266 (plan.tp_rank + 1) * plan.local_heads * kv_b_rows,
267 );
268 let o_absorb = absorb_o(
269 cfg,
270 &shard("o_proj", get("o_proj.weight")?)?,
271 &kv_b_local,
272 plan.local_heads,
273 )?;
274
275 let scale = (cfg.index_heads as f32).powf(-0.5);
277 let weights_proj: Vec<f32> = get("indexer.weights_proj.weight")?
278 .iter()
279 .map(|x| x * scale)
280 .collect();
281
282 let ape = get("indexer.index_kpool_compress_ape")?;
284
285 Ok(Glm5NextDsaWeights {
286 q_a_proj: up_bf16(gpu, &get("q_a_proj.weight")?)?,
287 q_a_layernorm: up_bf16(gpu, &get("q_a_layernorm.weight")?)?,
288 q_absorb: up_bf16(gpu, &q_absorb)?,
289 kv_a_proj: up_bf16(gpu, &get("kv_a_proj_with_mqa.weight")?)?,
290 kv_a_layernorm: up_bf16(gpu, &get("kv_a_layernorm.weight")?)?,
291 o_absorb: up_bf16(gpu, &o_absorb)?,
292 wk: up_bf16(gpu, &get("indexer.wk.weight")?)?,
293 k_norm_weight: up_bf16(gpu, &get("indexer.k_norm.weight")?)?,
294 k_norm_bias: up_bf16(gpu, &get("indexer.k_norm.bias")?)?,
296 compress_gate: up_bf16(gpu, &get("indexer.index_kpool_compress_gate")?)?,
297 wq_b: up_bf16(gpu, &get("indexer.wq_b.weight")?)?,
298 weights_proj: up_bf16(gpu, &weights_proj)?,
299 ape: up_f32(gpu, &ape)?,
300 })
301}
302
303#[cfg(test)]
304mod tests;