use std::path::Path;
use crate::mamba_ssm::gpu::blas::{
TiedLmDims, TypedPtr, gpu_gemm_ex_forward_raw, gpu_gemm_ex_tied_lm_head_raw,
gpu_sgemm_forward_ptr, gpu_sgemm_tied_lm_head_raw,
};
use crate::mamba_ssm::gpu::buffers::{GpuBuffer, GpuByteBuffer};
use crate::mamba_ssm::gpu::dtype::WeightDtype;
use crate::mamba_ssm::gpu::inference::GpuMambaBackbone;
const PREFILL_PARALLEL_THRESHOLD: usize = 256;
use crate::hf::embed::embed_lookup;
use crate::hf::load::{HfModel, load_hf};
use super::sample::{SampleParams, Xoshiro256PlusPlus, sample_token};
use rayon::prelude::*;
const SAMPLE_PARALLEL_THRESHOLD: usize = 8;
enum EmbedStorage {
F32 {
embed: GpuBuffer,
lm_head: Option<GpuBuffer>,
},
Half {
embed: GpuByteBuffer,
lm_head: Option<GpuByteBuffer>,
dtype: WeightDtype,
},
}
pub struct GpuMambaLM {
backbone: GpuMambaBackbone,
embed_storage: EmbedStorage,
embed_cpu: Vec<f32>,
input_cpu: Vec<f32>,
gpu_logits: GpuBuffer,
logits_padded_cpu: Vec<f32>,
logits_cpu: Vec<f32>,
pub vocab_size: usize,
vocab_size_padded: usize,
pub d_model: usize,
pub batch: usize,
}
impl GpuMambaLM {
pub fn ctx(&self) -> &crate::mamba_ssm::gpu::context::GpuCtx {
self.backbone.ctx()
}
pub fn last_logits(&self, b: usize) -> &[f32] {
&self.logits_cpu[b * self.vocab_size..(b + 1) * self.vocab_size]
}
}
impl GpuMambaLM {
pub fn from_hf(dir: &Path, gpu_ordinal: usize) -> Result<Self, String> {
Self::from_hf_with_dtype_batch(dir, gpu_ordinal, WeightDtype::F32, 1)
}
pub fn from_hf_with_dtype(
dir: &Path,
gpu_ordinal: usize,
dtype: WeightDtype,
) -> Result<Self, String> {
Self::from_hf_with_dtype_batch(dir, gpu_ordinal, dtype, 1)
}
pub fn from_hf_with_dtype_batch(
dir: &Path,
gpu_ordinal: usize,
dtype: WeightDtype,
batch: usize,
) -> Result<Self, String> {
let HfModel {
backbone: cpu_backbone,
embed,
lm_head,
vocab_size,
vocab_size_padded,
d_model,
} = load_hf(dir)?;
let cfg = *cpu_backbone.config();
let backbone = GpuMambaBackbone::new_with_dtype(
gpu_ordinal,
cpu_backbone.weights(),
cfg,
d_model,
batch,
dtype,
)?;
let stream = backbone.stream();
let lm_head_padded: Option<Vec<f32>> = lm_head.as_ref().map(|lm| {
if vocab_size == vocab_size_padded {
lm.clone()
} else {
let mut padded = vec![0.0f32; vocab_size_padded * d_model];
for row in 0..d_model {
let src = &lm[row * vocab_size..(row + 1) * vocab_size];
let dst =
&mut padded[row * vocab_size_padded..row * vocab_size_padded + vocab_size];
dst.copy_from_slice(src);
}
padded
}
});
let embed_storage = match dtype {
WeightDtype::F32 => {
let mut e = GpuBuffer::zeros(stream, vocab_size_padded * d_model)?;
e.upload(stream, &embed)?;
let lm = if let Some(ref lm_w) = lm_head_padded {
let mut b = GpuBuffer::zeros(stream, lm_w.len())?;
b.upload(stream, lm_w)?;
Some(b)
} else {
None
};
EmbedStorage::F32 {
embed: e,
lm_head: lm,
}
}
WeightDtype::Bf16 | WeightDtype::F16 => {
let embed_bytes = embed.len() * dtype.size_bytes();
let e = GpuByteBuffer::zeros(stream, embed_bytes)?;
upload_f32_as_dtype(stream, &e, 0, &embed, embed.len(), dtype)?;
let lm = if let Some(ref lm_w) = lm_head_padded {
let lm_bytes = lm_w.len() * dtype.size_bytes();
let b = GpuByteBuffer::zeros(stream, lm_bytes)?;
upload_f32_as_dtype(stream, &b, 0, lm_w, lm_w.len(), dtype)?;
Some(b)
} else {
None
};
EmbedStorage::Half {
embed: e,
lm_head: lm,
dtype,
}
}
};
let gpu_logits = GpuBuffer::zeros(stream, batch * vocab_size_padded)?;
Ok(Self {
backbone,
embed_storage,
embed_cpu: embed,
input_cpu: vec![0.0; batch * d_model],
gpu_logits,
logits_padded_cpu: vec![0.0; batch * vocab_size_padded],
logits_cpu: vec![0.0; batch * vocab_size],
vocab_size,
vocab_size_padded,
d_model,
batch,
})
}
pub fn dtype(&self) -> WeightDtype {
match &self.embed_storage {
EmbedStorage::F32 { .. } => WeightDtype::F32,
EmbedStorage::Half { dtype, .. } => *dtype,
}
}
pub fn capture_graph(&mut self) -> Result<(), String> {
self.backbone.capture_graph()
}
#[doc(hidden)]
pub fn debug_download_temporal(&self, out: &mut [f32]) -> Result<(), String> {
self.backbone.download_temporal(out)
}
#[doc(hidden)]
pub fn debug_step_one_token(
&mut self,
token: u32,
layer_limit: usize,
out: &mut [f32],
) -> Result<(), String> {
self.backbone.reset()?;
let emb =
crate::hf::embed::embed_lookup(&self.embed_cpu, token, self.d_model, self.vocab_size);
self.input_cpu[..self.d_model].copy_from_slice(emb);
self.backbone
.debug_step_partial(&self.input_cpu[..self.d_model], layer_limit, out)
}
pub fn reset(&mut self) -> Result<(), String> {
self.backbone.reset()
}
pub fn generate(&mut self, prompt: &[u32], params: &SampleParams) -> Result<Vec<u32>, String> {
let mut tokens = Vec::with_capacity(params.max_tokens);
self.generate_streaming(prompt, params, |tok, _| {
tokens.push(tok);
})?;
Ok(tokens)
}
pub fn generate_streaming(
&mut self,
prompt: &[u32],
params: &SampleParams,
mut cb: impl FnMut(u32, &str),
) -> Result<(), String> {
assert_eq!(
self.batch, 1,
"generate_streaming requires batch=1; use generate_batch for batch>1"
);
self.backbone.reset()?;
let mut rng = Xoshiro256PlusPlus::new(params.seed);
if prompt.len() >= PREFILL_PARALLEL_THRESHOLD {
self.prefill_parallel(prompt)?;
} else {
for &token_id in prompt {
let emb = embed_lookup(&self.embed_cpu, token_id, self.d_model, self.vocab_size);
self.input_cpu[..self.d_model].copy_from_slice(emb);
self.backbone
.step_gpu_only(&self.input_cpu[..self.d_model])?;
}
}
self.compute_logits()?;
let mut seen: Vec<u32> = prompt.to_vec();
for _ in 0..params.max_tokens {
let next = sample_token(&mut self.logits_cpu, params, &seen, &mut rng);
if params.eos_token_ids.contains(&next) {
break;
}
seen.push(next);
cb(next, "");
let emb = embed_lookup(&self.embed_cpu, next, self.d_model, self.vocab_size);
self.input_cpu[..self.d_model].copy_from_slice(emb);
self.backbone
.step_gpu_only(&self.input_cpu[..self.d_model])?;
self.compute_logits()?;
}
Ok(())
}
pub fn generate_batch(
&mut self,
prompts: &[&[u32]],
params: &[SampleParams],
) -> Result<Vec<Vec<u32>>, String> {
assert_eq!(prompts.len(), self.batch, "prompts.len() != batch");
assert_eq!(params.len(), self.batch, "params.len() != batch");
self.backbone.reset()?;
let mut rngs: Vec<Xoshiro256PlusPlus> = params
.iter()
.map(|p| Xoshiro256PlusPlus::new(p.seed))
.collect();
let b = self.batch;
let d = self.d_model;
let vocab_size = self.vocab_size;
let max_prompt = prompts.iter().map(|p| p.len()).max().unwrap_or(0);
let max_tokens = params.iter().map(|p| p.max_tokens).max().unwrap_or(0);
let mut prompt_pos = vec![0usize; b]; let mut finished = vec![false; b];
let mut outputs: Vec<Vec<u32>> = (0..b).map(|_| Vec::new()).collect();
let mut last_token = vec![0u32; b];
for i in 0..b {
if prompts[i].is_empty() {
finished[i] = true;
continue;
}
last_token[i] = prompts[i][0];
prompt_pos[i] = 1; }
let total_steps = max_prompt + max_tokens;
for _step in 0..total_steps {
if finished.iter().all(|&f| f) {
break;
}
for i in 0..b {
if finished[i] {
for v in &mut self.input_cpu[i * d..(i + 1) * d] {
*v = 0.0;
}
} else {
let emb = embed_lookup(&self.embed_cpu, last_token[i], d, vocab_size);
self.input_cpu[i * d..(i + 1) * d].copy_from_slice(emb);
}
}
self.backbone.step_gpu_only(&self.input_cpu)?;
self.compute_logits()?;
let need_decode_for_slot: Vec<bool> = (0..b)
.map(|i| !finished[i] && prompt_pos[i] >= prompts[i].len())
.collect();
let new_tokens: Vec<Option<u32>> = if b >= SAMPLE_PARALLEL_THRESHOLD {
self.logits_cpu
.par_chunks_mut(vocab_size)
.zip(rngs.par_iter_mut())
.enumerate()
.map(|(i, (slot_logits, rng))| {
if need_decode_for_slot[i] {
Some(sample_token(slot_logits, ¶ms[i], &outputs[i], rng))
} else {
None
}
})
.collect()
} else {
(0..b)
.map(|i| {
if need_decode_for_slot[i] {
let slot_logits =
&mut self.logits_cpu[i * vocab_size..(i + 1) * vocab_size];
Some(sample_token(
slot_logits,
¶ms[i],
&outputs[i],
&mut rngs[i],
))
} else {
None
}
})
.collect()
};
for i in 0..b {
if finished[i] {
continue;
}
if prompt_pos[i] < prompts[i].len() {
last_token[i] = prompts[i][prompt_pos[i]];
prompt_pos[i] += 1;
continue;
}
let next = new_tokens[i].expect("decode-phase slot must have a sampled token");
if params[i].eos_token_ids.contains(&next)
|| outputs[i].len() >= params[i].max_tokens
{
finished[i] = true;
continue;
}
outputs[i].push(next);
last_token[i] = next;
}
}
Ok(outputs)
}
fn prefill_parallel(&mut self, prompt: &[u32]) -> Result<(), String> {
let t = prompt.len();
let d = self.d_model;
let b = self.batch;
let stream = self.backbone.stream().clone();
let mut embed_flat = vec![0.0f32; b * t * d];
for ti in 0..t {
let emb = embed_lookup(&self.embed_cpu, prompt[ti], d, self.vocab_size);
embed_flat[ti * d..(ti + 1) * d].copy_from_slice(emb);
}
let mut ip_out_flat = GpuBuffer::zeros(&stream, b * t * d)?;
ip_out_flat.upload(&stream, &embed_flat)?;
match self.backbone.dtype() {
WeightDtype::F32 => {
let mut prefill_scratch = self.backbone.alloc_prefill_scratch(t)?;
self.backbone
.prefill_sequence(&ip_out_flat, &mut prefill_scratch)?;
}
WeightDtype::Bf16 | WeightDtype::F16 => {
let mut prefill_scratch = self.backbone.alloc_prefill_mixed_scratch(t)?;
self.backbone
.prefill_sequence_mixed(&ip_out_flat, &mut prefill_scratch)?;
}
}
Ok(())
}
fn compute_logits(&mut self) -> Result<(), String> {
let ctx = self.backbone.ctx();
let stream = self.backbone.stream().clone();
let temporal_ptr = self.backbone.temporal_ptr();
let b = self.batch;
let d = self.d_model;
match &self.embed_storage {
EmbedStorage::F32 { embed, lm_head } => {
if let Some(lm) = lm_head {
gpu_sgemm_forward_ptr(
ctx,
&mut self.gpu_logits,
temporal_ptr,
lm.cached_ptr(),
None,
(b, d, self.vocab_size_padded),
)?;
} else {
gpu_sgemm_tied_lm_head_raw(
ctx,
self.gpu_logits.cached_ptr(),
temporal_ptr,
embed.cached_ptr(),
b,
d,
self.vocab_size_padded,
)?;
}
}
EmbedStorage::Half {
embed,
lm_head,
dtype,
} => {
let backbone_dtype = self.backbone.temporal_dtype();
let temporal_half_ptr = if backbone_dtype == *dtype {
temporal_ptr
} else {
let half_bytes = b * d * dtype.size_bytes();
ctx.ensure_half_staging(half_bytes)?;
let staging_ptr = ctx.half_staging_ptr();
use cudarc::driver::PushKernelArg;
let n = (b * d) as i32;
let kernel = match *dtype {
WeightDtype::Bf16 => &ctx.kernels.cast_f32_to_bf16,
WeightDtype::F16 => &ctx.kernels.cast_f32_to_f16,
WeightDtype::F32 => unreachable!(),
};
let mut builder = stream.launch_builder(kernel);
builder.arg(&staging_ptr);
builder.arg(&temporal_ptr);
builder.arg(&n);
use crate::mamba_ssm::gpu::launch::grid_1d;
unsafe { builder.launch(grid_1d(b * d)) }
.map_err(|e| format!("cast temporal: {e:?}"))?;
staging_ptr
};
if let Some(lm) = lm_head {
gpu_gemm_ex_forward_raw(
ctx,
&mut self.gpu_logits,
TypedPtr {
ptr: temporal_half_ptr,
dtype: *dtype,
},
TypedPtr {
ptr: lm.cached_ptr(),
dtype: *dtype,
},
None,
(b, d, self.vocab_size_padded),
)?;
} else {
gpu_gemm_ex_tied_lm_head_raw(
ctx,
self.gpu_logits.cached_ptr(),
temporal_half_ptr,
embed.cached_ptr(),
*dtype,
TiedLmDims {
batch: b,
d_model: d,
vocab_padded: self.vocab_size_padded,
},
)?;
}
}
}
stream
.synchronize()
.map_err(|e| format!("logits sync: {e:?}"))?;
self.gpu_logits
.download(&stream, &mut self.logits_padded_cpu)?;
for bi in 0..self.batch {
let src = &self.logits_padded_cpu
[bi * self.vocab_size_padded..bi * self.vocab_size_padded + self.vocab_size];
let dst = &mut self.logits_cpu[bi * self.vocab_size..(bi + 1) * self.vocab_size];
dst.copy_from_slice(src);
}
Ok(())
}
}
fn upload_f32_as_dtype(
stream: &std::sync::Arc<cudarc::driver::CudaStream>,
dst: &GpuByteBuffer,
elem_offset: usize,
src: &[f32],
src_elems: usize,
dtype: WeightDtype,
) -> Result<(), String> {
use crate::mamba_ssm::gpu::buffers::cu_memcpy_htod_raw;
assert_eq!(src.len(), src_elems, "src size mismatch");
let byte_off = elem_offset * dtype.size_bytes();
let byte_count = src_elems * dtype.size_bytes();
let dst_ptr = dst.cached_ptr() + byte_off as u64;
match dtype {
WeightDtype::F32 => {
let bytes: &[u8] = bytemuck::cast_slice(src);
assert_eq!(bytes.len(), byte_count);
cu_memcpy_htod_raw(stream, dst_ptr, bytes)
}
WeightDtype::Bf16 => {
let buf: Vec<half::bf16> = src.iter().map(|&v| half::bf16::from_f32(v)).collect();
let bytes: &[u8] = bytemuck::cast_slice(&buf);
assert_eq!(bytes.len(), byte_count);
cu_memcpy_htod_raw(stream, dst_ptr, bytes)
}
WeightDtype::F16 => {
let buf: Vec<half::f16> = src.iter().map(|&v| half::f16::from_f32(v)).collect();
let bytes: &[u8] = bytemuck::cast_slice(&buf);
assert_eq!(bytes.len(), byte_count);
cu_memcpy_htod_raw(stream, dst_ptr, bytes)
}
}
}