1use crate::gpu::{DevicePtr, GpuBackend};
9use anyhow::Result;
10use atlas_core::config::ModelConfig;
11
12mod accessors;
13pub mod decode_meta;
14mod sizes;
15mod sizes_q12;
16mod sizes_q2;
17pub use decode_meta::{DECODE_META_MAX_ROWS, DECODE_META_MIN_ROWS, DecodeMetaLayout};
18pub use sizes::BufferSizes;
19pub use sizes_q2::q2_dequant_scratch_bytes;
20pub use sizes_q12::{
21 Q12_SIZING_STREAMS, q12_batched_scratch_bytes, q12_batched_scratch_bytes_varlen,
22};
23
24pub struct BufferArena {
36 hidden_states: DevicePtr,
38 residual: DevicePtr,
40 norm_output: DevicePtr,
42 qkv_output: DevicePtr,
44 attn_output: DevicePtr,
46 gate_logits: DevicePtr,
48 gate_logits_f32: DevicePtr,
50 moe_router_in_f32: DevicePtr,
52 moe_output: DevicePtr,
54 logits: DevicePtr,
56 ssm_qkvz: DevicePtr,
58 ssm_ba: DevicePtr,
60 ssm_deinterleaved: DevicePtr,
62 ssm_gates: DevicePtr,
64 ssm_conv_out_f32: DevicePtr,
69 scratch: DevicePtr,
71 expert_gate_out: DevicePtr,
73 expert_up_out: DevicePtr,
75 expert_down_out: DevicePtr,
77 splitk_workspace: DevicePtr,
79 o_latent: DevicePtr,
81 norm_unit_w: DevicePtr,
84 hc_streams: DevicePtr,
86 hc_post: DevicePtr,
88 hc_comb: DevicePtr,
90 hc_lowrank_scratch: DevicePtr,
91 qsa_select_scratch: DevicePtr,
92 gdn_fla_scratch: DevicePtr,
95 ssd_scratch: DevicePtr,
98 token_ids: DevicePtr,
101 ffn_act_q8: DevicePtr,
107 ffn_act_a: DevicePtr,
108 ffn_act_scale: DevicePtr,
109 fp8_act: DevicePtr,
111 fp8_act_scale: DevicePtr,
113 q2_dequant_scratch: DevicePtr,
116 lora_xa: DevicePtr,
119 lora_delta: DevicePtr,
122 lora_hact: DevicePtr,
125 lora_seq_slot: DevicePtr,
128 q2_act_q8: DevicePtr,
132 max_batch_tokens: usize,
134 decode_meta: DecodeMetaLayout,
138 sizes: BufferSizes,
140}
141
142impl BufferArena {
143 pub fn new(
145 config: &ModelConfig,
146 max_batch_tokens: usize,
147 max_seq_len: usize,
148 kv_block_size: usize,
149 max_batch_size: usize,
150 gpu: &dyn GpuBackend,
151 ) -> Result<Self> {
152 let decode_meta = DecodeMetaLayout::for_max_batch_size(max_batch_size);
153 let sizes = BufferSizes::from_config(
154 config,
155 max_batch_tokens,
156 max_seq_len,
157 kv_block_size,
158 max_batch_size,
159 );
160
161 let hidden_states = gpu.alloc(sizes.hidden_states)?;
162 let residual = gpu.alloc(sizes.residual)?;
163 let norm_output = gpu.alloc(sizes.norm_output)?;
164 let qkv_output = gpu.alloc(sizes.qkv_output)?;
165 let attn_output = gpu.alloc(sizes.attn_output)?;
166 let gate_logits = gpu.alloc(sizes.gate_logits)?;
167 let gate_logits_f32 = gpu.alloc(sizes.gate_logits_f32)?;
168 let moe_router_in_f32 = gpu.alloc(sizes.moe_router_in_f32)?;
169 let moe_output = gpu.alloc(sizes.moe_output)?;
170 let logits = gpu.alloc(sizes.logits)?;
171 let ssm_qkvz = gpu.alloc(sizes.ssm_qkvz)?;
172 let ssm_ba = gpu.alloc(sizes.ssm_ba)?;
173 let ssm_deinterleaved = gpu.alloc(sizes.ssm_deinterleaved)?;
174 let ssm_gates = gpu.alloc(sizes.ssm_gates)?;
175 let ssm_conv_out_f32 = gpu.alloc(sizes.ssm_conv_out_f32)?;
176 let scratch = gpu.alloc(sizes.scratch)?;
177 let expert_gate_out = gpu.alloc(sizes.expert_gate_out)?;
178 let expert_up_out = gpu.alloc(sizes.expert_up_out)?;
179 let expert_down_out = gpu.alloc(sizes.expert_down_out)?;
180 let splitk_workspace = gpu.alloc(sizes.splitk_workspace)?;
181 let o_latent = gpu.alloc(sizes.o_latent)?;
182 let norm_unit_w = gpu.alloc(sizes.norm_unit_w)?;
186 gpu.memset(norm_unit_w, 0, sizes.norm_unit_w)?;
187 let hc_streams = gpu.alloc(sizes.hc_streams)?;
188 let hc_post = gpu.alloc(sizes.hc_post)?;
189 let hc_comb = gpu.alloc(sizes.hc_comb)?;
190 let hc_lowrank_scratch = gpu.alloc(sizes.hc_lowrank_scratch)?;
191 let qsa_select_scratch = gpu.alloc(sizes.qsa_select_scratch)?;
192 let ssd_scratch = if sizes.ssd_scratch > 0 {
195 gpu.alloc(sizes.ssd_scratch)?
196 } else {
197 DevicePtr::NULL
198 };
199 let gdn_fla_scratch = if sizes.gdn_fla_scratch > 0 {
200 gpu.alloc(sizes.gdn_fla_scratch)?
201 } else {
202 DevicePtr::NULL
203 };
204 let token_ids = gpu.alloc(sizes.token_ids)?;
205 let ffn_act_q8 = if sizes.ffn_act_q8 > 0 {
208 gpu.alloc(sizes.ffn_act_q8)?
209 } else {
210 DevicePtr::NULL
211 };
212 let ffn_act_a = if sizes.ffn_act_a > 0 {
213 gpu.alloc(sizes.ffn_act_a)?
214 } else {
215 DevicePtr::NULL
216 };
217 let ffn_act_scale = if sizes.ffn_act_scale > 0 {
218 gpu.alloc(sizes.ffn_act_scale)?
219 } else {
220 DevicePtr::NULL
221 };
222 let fp8_act = gpu.alloc(sizes.fp8_act)?;
223 let fp8_act_scale = gpu.alloc(sizes.fp8_act_scale)?;
224 let q2_dequant_scratch = if sizes.q2_dequant_scratch > 0 {
226 gpu.alloc(sizes.q2_dequant_scratch)?
227 } else {
228 DevicePtr::NULL
229 };
230 let lora_xa = if sizes.lora_xa > 0 {
233 gpu.alloc(sizes.lora_xa)?
234 } else {
235 DevicePtr::NULL
236 };
237 let lora_delta = if sizes.lora_delta > 0 {
238 gpu.alloc(sizes.lora_delta)?
239 } else {
240 DevicePtr::NULL
241 };
242 let lora_hact = if sizes.lora_hact > 0 {
243 gpu.alloc(sizes.lora_hact)?
244 } else {
245 DevicePtr::NULL
246 };
247 let lora_seq_slot = if sizes.lora_seq_slot > 0 {
248 gpu.alloc(sizes.lora_seq_slot)?
249 } else {
250 DevicePtr::NULL
251 };
252 let q2_act_q8 = if sizes.q2_act_q8 > 0 {
254 gpu.alloc(sizes.q2_act_q8)?
255 } else {
256 DevicePtr::NULL
257 };
258
259 tracing::info!(
260 "Buffer arena: {} tokens × {:.1} MB total (attn_out={:.1}MB, ssm_deint={:.1}MB, kv_lora_rank={})",
261 max_batch_tokens,
262 sizes.total_bytes() as f64 / (1024.0 * 1024.0),
263 sizes.attn_output as f64 / (1024.0 * 1024.0),
264 sizes.ssm_deinterleaved as f64 / (1024.0 * 1024.0),
265 config.kv_lora_rank,
266 );
267
268 Ok(Self {
269 hidden_states,
270 residual,
271 norm_output,
272 qkv_output,
273 attn_output,
274 gate_logits,
275 gate_logits_f32,
276 moe_router_in_f32,
277 moe_output,
278 logits,
279 ssm_qkvz,
280 ssm_ba,
281 ssm_deinterleaved,
282 ssm_gates,
283 ssm_conv_out_f32,
284 scratch,
285 expert_gate_out,
286 expert_up_out,
287 expert_down_out,
288 splitk_workspace,
289 o_latent,
290 norm_unit_w,
291 hc_streams,
292 hc_post,
293 hc_comb,
294 hc_lowrank_scratch,
295 qsa_select_scratch,
296 gdn_fla_scratch,
297 ssd_scratch,
298 token_ids,
299 ffn_act_q8,
300 ffn_act_a,
301 ffn_act_scale,
302 fp8_act,
303 fp8_act_scale,
304 q2_dequant_scratch,
305 lora_xa,
306 lora_delta,
307 lora_hact,
308 lora_seq_slot,
309 q2_act_q8,
310 max_batch_tokens,
311 decode_meta,
312 sizes,
313 })
314 }
315}
316
317impl atlas_core::scope::ModelResource<dyn GpuBackend> for BufferArena {
325 fn label(&self) -> &'static str {
326 "buffer arena"
327 }
328
329 fn release(&mut self, gpu: &dyn GpuBackend) -> anyhow::Result<()> {
330 let Self {
331 sizes: _,
334 max_batch_tokens: _,
335 decode_meta: _,
337 hidden_states,
338 residual,
339 norm_output,
340 qkv_output,
341 attn_output,
342 gate_logits,
343 gate_logits_f32,
344 moe_router_in_f32,
345 moe_output,
346 logits,
347 ssm_qkvz,
348 ssm_ba,
349 ssm_deinterleaved,
350 ssm_gates,
351 ssm_conv_out_f32,
352 scratch,
353 expert_gate_out,
354 expert_up_out,
355 expert_down_out,
356 splitk_workspace,
357 o_latent,
358 norm_unit_w,
359 hc_streams,
360 hc_post,
361 hc_comb,
362 hc_lowrank_scratch,
363 qsa_select_scratch,
364 gdn_fla_scratch,
365 ssd_scratch,
366 token_ids,
367 ffn_act_q8,
368 ffn_act_a,
369 ffn_act_scale,
370 fp8_act,
371 fp8_act_scale,
372 lora_xa,
373 lora_delta,
374 lora_hact,
375 lora_seq_slot,
376 q2_dequant_scratch,
377 q2_act_q8,
378 } = self;
379 let owned = [
382 *hidden_states,
383 *residual,
384 *norm_output,
385 *qkv_output,
386 *attn_output,
387 *gate_logits,
388 *gate_logits_f32,
389 *moe_router_in_f32,
390 *moe_output,
391 *logits,
392 *ssm_qkvz,
393 *ssm_ba,
394 *ssm_deinterleaved,
395 *ssm_gates,
396 *ssm_conv_out_f32,
397 *scratch,
398 *expert_gate_out,
399 *expert_up_out,
400 *expert_down_out,
401 *splitk_workspace,
402 *o_latent,
403 *norm_unit_w,
404 *hc_streams,
405 *hc_lowrank_scratch,
406 *qsa_select_scratch,
407 *hc_post,
408 *hc_comb,
409 *gdn_fla_scratch,
410 *ssd_scratch,
411 *token_ids,
412 *ffn_act_q8,
413 *ffn_act_a,
414 *ffn_act_scale,
415 *fp8_act,
416 *fp8_act_scale,
417 *lora_xa,
418 *lora_delta,
419 *lora_hact,
420 *lora_seq_slot,
421 *q2_dequant_scratch,
422 *q2_act_q8,
423 ];
424 let mut first_error = None;
425 for ptr in owned {
426 if let Err(e) = gpu.free(ptr)
427 && first_error.is_none()
428 {
429 first_error = Some(e);
430 }
431 }
432 *hidden_states = DevicePtr::NULL;
433 *residual = DevicePtr::NULL;
434 *norm_output = DevicePtr::NULL;
435 *qkv_output = DevicePtr::NULL;
436 *attn_output = DevicePtr::NULL;
437 *gate_logits = DevicePtr::NULL;
438 *gate_logits_f32 = DevicePtr::NULL;
439 *moe_router_in_f32 = DevicePtr::NULL;
440 *moe_output = DevicePtr::NULL;
441 *logits = DevicePtr::NULL;
442 *ssm_qkvz = DevicePtr::NULL;
443 *ssm_ba = DevicePtr::NULL;
444 *ssm_deinterleaved = DevicePtr::NULL;
445 *ssm_gates = DevicePtr::NULL;
446 *ssm_conv_out_f32 = DevicePtr::NULL;
447 *scratch = DevicePtr::NULL;
448 *expert_gate_out = DevicePtr::NULL;
449 *expert_up_out = DevicePtr::NULL;
450 *expert_down_out = DevicePtr::NULL;
451 *splitk_workspace = DevicePtr::NULL;
452 *o_latent = DevicePtr::NULL;
453 *norm_unit_w = DevicePtr::NULL;
454 *hc_streams = DevicePtr::NULL;
455 *hc_lowrank_scratch = DevicePtr::NULL;
456 *qsa_select_scratch = DevicePtr::NULL;
457 *hc_post = DevicePtr::NULL;
458 *hc_comb = DevicePtr::NULL;
459 *gdn_fla_scratch = DevicePtr::NULL;
460 *ssd_scratch = DevicePtr::NULL;
461 *token_ids = DevicePtr::NULL;
462 *ffn_act_q8 = DevicePtr::NULL;
463 *ffn_act_a = DevicePtr::NULL;
464 *ffn_act_scale = DevicePtr::NULL;
465 *fp8_act = DevicePtr::NULL;
466 *fp8_act_scale = DevicePtr::NULL;
467 *lora_xa = DevicePtr::NULL;
468 *lora_delta = DevicePtr::NULL;
469 *lora_hact = DevicePtr::NULL;
470 *lora_seq_slot = DevicePtr::NULL;
471 *q2_dequant_scratch = DevicePtr::NULL;
472 *q2_act_q8 = DevicePtr::NULL;
473 match first_error {
474 Some(e) => Err(e),
475 None => Ok(()),
476 }
477 }
478}
479
480#[cfg(test)]
481mod tests;