use std::sync::Arc;
use std::sync::Mutex;
use std::sync::OnceLock;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use anyhow::{Result, anyhow};
use crate::CeraError;
use crate::backend::cpu::RopeType;
use crate::backend::wgpu::{DevicePollExt, GpuContext, GpuTensor, KvShiftParams, shaders};
use crate::gguf::GgufFile;
use crate::kv_cache::{InferenceState, KvCompression, KvPrefixCache, LayerSnapshot, StateSnapshot};
use crate::lora::{LoraAdapterWeights, LoraTarget};
use crate::model::gpu_turboquant::{TqGpuCache, TqMode, describe_kv_mode};
use crate::model::gpu_weight_source::{
GpuWeightSource, MOE_MAX_EXPERT_USED, MOE_MAX_EXPERTS, stacked_expert_layout,
};
use crate::model::transformer::WeightRef;
use crate::model::{BlockType, Model, ModelConfig, ScalarMultipliers};
use crate::tensor::DType;
const MAX_PREFILL_TOKENS: usize = 2048;
const MAX_ALL_LOGITS_TOKENS: usize = 64;
pub(crate) const MUL_MAT_TILE_WG_M: u32 = 16;
const GEMV_F32_ROWS_PER_WG: u32 = 8;
pub(crate) const MUL_MAT_TILE_WG_N: u32 = 16;
pub(crate) const MUL_MAT_TILE_M: u32 = 4;
pub(crate) const MUL_MAT_TILE_N: u32 = 4;
pub(crate) const MUL_MAT_TILE_K: u32 = 16;
const _: () = assert!(
MUL_MAT_TILE_M == 4 && MUL_MAT_TILE_N == 4,
"mul_mat_reg_tile.wgsl hand-unrolls a 4x4 thread tile; re-unroll it before \
changing MUL_MAT_TILE_M/N"
);
const _: () = assert!(
MUL_MAT_TILE_K.is_multiple_of(8),
"the Q4_0 shmem loader stages 8 consecutive k per thread and indexes within \
one 32-element block; TILE_K must be a multiple of 8"
);
const _: () = assert!(
(MUL_MAT_TILE_K * (MUL_MAT_TILE_WG_M * MUL_MAT_TILE_M + 4)
+ MUL_MAT_TILE_K * (MUL_MAT_TILE_WG_N * MUL_MAT_TILE_N + 4))
* 4
<= 16384,
"reg-tile shmem must stay within WebGPU's guaranteed 16 KiB workgroup-storage \
limit so the pipeline builds on a spec-minimum adapter"
);
const _: () = assert!(
MUL_MAT_TILE_WG_M * MUL_MAT_TILE_WG_N <= 256,
"reg-tile workgroup must stay within WebGPU's guaranteed 256 invocations per \
workgroup so the pipeline builds on a spec-minimum adapter"
);
fn build_mul_mat_pipeline(
ctx: &GpuContext,
label: &str,
src0_loader: &str,
src0_inner: &str,
) -> wgpu::ComputePipeline {
let wg_m = format!("{MUL_MAT_TILE_WG_M}u");
let wg_n = format!("{MUL_MAT_TILE_WG_N}u");
let tile_m = format!("{MUL_MAT_TILE_M}u");
let tile_n = format!("{MUL_MAT_TILE_N}u");
let tile_k = format!("{MUL_MAT_TILE_K}u");
ctx.create_pipeline_with_defines(
shaders::MUL_MAT_REG_TILE,
"main",
label,
&[
("SRC0_INNER_TYPE", src0_inner),
(src0_loader, ""),
("WORKGROUP_SIZE_M", &wg_m),
("WORKGROUP_SIZE_N", &wg_n),
("TILE_M", &tile_m),
("TILE_N", &tile_n),
("TILE_K", &tile_k),
],
)
}
fn use_spirv_passthrough(ctx: &GpuContext) -> bool {
std::env::var("CERA_WGPU_SPIRV_PASSTHROUGH").as_deref() != Ok("0")
&& ctx.supports_spirv_passthrough()
}
fn gcd_u64(mut a: u64, mut b: u64) -> u64 {
while b != 0 {
let r = a % b;
a = b;
b = r;
}
a
}
fn lcm_u64(a: u64, b: u64) -> u64 {
(a / gcd_u64(a, b)) * b
}
fn kv_slab_bytes(rows: usize, cols: usize) -> u64 {
rows as u64 * cols as u64 * 4
}
fn f32_binding(buffer: &wgpu::Buffer, len_floats: u64) -> wgpu::BindingResource<'_> {
let bytes = len_floats
.checked_mul(std::mem::size_of::<f32>() as u64)
.expect("f32 storage binding size overflow");
wgpu::BindingResource::Buffer(wgpu::BufferBinding {
buffer,
offset: 0,
size: wgpu::BufferSize::new(bytes.max(4)),
})
}
fn assert_f32_binding_fits(len_floats: u64, max_binding: u64, what: &str) {
let bytes = len_floats.saturating_mul(std::mem::size_of::<f32>() as u64);
assert!(
bytes <= max_binding,
"wgpu {what} binding is {bytes} bytes, exceeding adapter \
max_storage_buffer_binding_size {max_binding}; context paging is required"
);
}
fn gemv_tile_rows(m: u32, k: u32, max_binding: u64, offset_alignment: u64, elem_size: u64) -> u32 {
const ROWS_PER_WG: u64 = 8;
let row_bytes = u64::from(k) * elem_size;
let full_bytes = (u64::from(m) * row_bytes).div_ceil(4) * 4;
if full_bytes <= max_binding {
return m;
}
let max_rows = (max_binding / row_bytes) as u32;
assert!(
max_rows > 0,
"GPU max storage binding size {} is too small for one GEMV row of {} bytes",
max_binding,
row_bytes
);
let offset_alignment = offset_alignment.max(elem_size.max(4));
let row_alignment = (offset_alignment / gcd_u64(row_bytes, offset_alignment)).max(1) as u32;
let tile_alignment = lcm_u64(u64::from(row_alignment), ROWS_PER_WG) as u32;
let tile_rows = if max_rows >= tile_alignment {
max_rows - (max_rows % tile_alignment)
} else if max_rows >= row_alignment {
max_rows - (max_rows % row_alignment)
} else {
max_rows
};
assert!(
tile_rows > 0 && (u64::from(tile_rows) * row_bytes).is_multiple_of(offset_alignment),
"GPU storage binding alignment {} cannot be satisfied for GEMV rows of {} bytes",
offset_alignment,
row_bytes
);
tile_rows
}
struct GpuWeight {
tensor: GpuTensor,
params_buf: wgpu::Buffer,
cached_bg: Option<wgpu::BindGroup>,
}
enum LmHead {
Quantized(GpuWeight),
F16 {
weight: wgpu::Buffer,
params: wgpu::Buffer,
},
}
enum GpuFfn {
Dense(Box<GpuDenseFfn>),
Moe(Box<GpuMoeFfn>),
}
struct GpuDenseFfn {
gate: GpuWeight,
up: GpuWeight,
down: GpuWeight,
}
struct GpuMoeWeight {
buffer: wgpu::Buffer,
m: u32,
k: u32,
expert_stride: u32,
}
struct GpuMoeFfn {
router: wgpu::Buffer,
bias: wgpu::Buffer,
gate: GpuMoeWeight,
up: GpuMoeWeight,
down: GpuMoeWeight,
scratch: Arc<MoeScratch>,
}
struct GpuLayerWeights {
attn_norm: wgpu::Buffer,
ffn_norm: wgpu::Buffer,
ffn: GpuFfn,
conv_in_proj: Option<GpuWeight>,
conv_out_proj: Option<GpuWeight>,
conv_weight: Option<wgpu::Buffer>,
attn_q: Option<GpuWeight>,
attn_k: Option<GpuWeight>,
attn_v: Option<GpuWeight>,
attn_output: Option<GpuWeight>,
attn_q_norm: Option<wgpu::Buffer>,
attn_k_norm: Option<wgpu::Buffer>,
attn_q_bias: Option<wgpu::Buffer>,
attn_k_bias: Option<wgpu::Buffer>,
attn_v_bias: Option<wgpu::Buffer>,
attn_norm_bg: Option<wgpu::BindGroup>,
ffn_norm_bg: Option<wgpu::BindGroup>,
rope_bg: Option<wgpu::BindGroup>,
conv_fused_bg: Option<wgpu::BindGroup>,
conv_add_bg: Option<wgpu::BindGroup>,
attn_out_add_bg: Option<wgpu::BindGroup>,
silu_bg: Option<wgpu::BindGroup>,
ffn_swiglu_bg: Option<wgpu::BindGroup>,
attn_qkv_bg: Option<wgpu::BindGroup>,
attn_qkv_params_buf: Option<wgpu::Buffer>,
attn_bg: Option<wgpu::BindGroup>,
qn_bg: Option<wgpu::BindGroup>,
kn_bg: Option<wgpu::BindGroup>,
qb_bg: Option<wgpu::BindGroup>,
kb_bg: Option<wgpu::BindGroup>,
vb_bg: Option<wgpu::BindGroup>,
ffn_add_bg: Option<wgpu::BindGroup>,
}
struct MoeScratch {
n_expert: u32,
n_expert_used: u32,
expert_ff_len: u32,
max_entries: u32,
logits: wgpu::Buffer,
sel_expert: wgpu::Buffer,
sel_weight: wgpu::Buffer,
gate: wgpu::Buffer,
up: wgpu::Buffer,
z: wgpu::Buffer,
}
struct MoeStep<'a> {
pipeline: &'a wgpu::ComputePipeline,
bind_group: wgpu::BindGroup,
workgroups: (u32, u32, u32),
}
fn upload_moe(
ctx: &GpuContext,
src: &dyn GpuWeightSource,
scratch: Option<&Arc<MoeScratch>>,
hidden_size: usize,
layer: usize,
moe: &crate::model::lfm2::MoeFfnRefs,
) -> Result<GpuMoeFfn> {
use anyhow::Context;
let scratch = scratch
.with_context(|| {
format!(
"layer {layer} has routed expert weights but the model config carries no \
mixture-of-experts parameters to size their scratch with"
)
})?
.clone();
let n_expert = scratch.n_expert;
let expert_ff_len = scratch.expert_ff_len;
anyhow::ensure!(
moe.exp_probs_b.len() == n_expert as usize,
"layer {layer}: selection bias has {} entries for {n_expert} experts",
moe.exp_probs_b.len(),
);
anyhow::ensure!(
moe.gate.len() == n_expert as usize
&& moe.up.len() == n_expert as usize
&& moe.down.len() == n_expert as usize,
"layer {layer}: routed FFN has {n_expert} experts but {} gate / {} up / {} down expert \
tensors; the routing kernel emits ids the GEMV would index past the stacked weights",
moe.gate.len(),
moe.up.len(),
moe.down.len(),
);
anyhow::ensure!(
moe.router.dtype == DType::F32,
"layer {layer}: wgpu MoE routing needs an F32 router projection, found {:?}",
moe.router.dtype,
);
let stack = |refs: &[WeightRef], what: &str| -> Result<GpuMoeWeight> {
let layout = stacked_expert_layout(refs, layer, what, "wgpu")?;
let stride = layout.expert_stride;
let total = layout.total_bytes;
anyhow::ensure!(
u64::from(total) <= ctx.max_storage_buffer_binding_size,
"layer {layer}: {what} stacks {} experts into {total} bytes, over this adapter's \
{} byte storage-binding limit; the expert GEMV binds the whole stack at once",
refs.len(),
ctx.max_storage_buffer_binding_size,
);
let bytes = refs
.iter()
.fold(Vec::with_capacity(total as usize), |mut acc: Vec<u8>, r| {
acc.extend_from_slice(&src.weight_bytes(r));
acc
});
Ok(GpuMoeWeight {
buffer: ctx.upload_storage(&bytes, &format!("l{layer}.{what}")),
m: u32::try_from(layout.rows)
.with_context(|| format!("layer {layer}: {what} has too many rows"))?,
k: u32::try_from(layout.inner)
.with_context(|| format!("layer {layer}: {what} inner dim too large"))?,
expert_stride: stride,
})
};
let gate = stack(&moe.gate, "ffn_gate_exps")?;
let up = stack(&moe.up, "ffn_up_exps")?;
let down = stack(&moe.down, "ffn_down_exps")?;
let hs = u32::try_from(hidden_size)
.with_context(|| format!("layer {layer}: hidden size {hidden_size} too large"))?;
anyhow::ensure!(
gate.m == expert_ff_len && up.m == expert_ff_len && down.k == expert_ff_len,
"layer {layer}: expert width {expert_ff_len} disagrees with the projection shapes \
(gate.m={}, up.m={}, down.k={})",
gate.m,
up.m,
down.k,
);
anyhow::ensure!(
down.m == hs
&& gate.k == hs
&& up.k == hs
&& moe.router.k == hidden_size
&& moe.router.m == n_expert as usize,
"layer {layer}: routed FFN shapes disagree with hidden size {hs} / expert count \
{n_expert} (down.m={}, gate.k={}, up.k={}, router {}x{}); the wgpu expert kernels \
index every one of these against those two numbers",
down.m,
gate.k,
up.k,
moe.router.m,
moe.router.k,
);
anyhow::ensure!(
gate.m.max(down.m) <= crate::backend::wgpu::MAX_WG,
"layer {layer}: expert projections are {} rows, over the {} workgroups-per-dimension \
cap the expert GEMV dispatches one row per workgroup against",
gate.m.max(down.m),
crate::backend::wgpu::MAX_WG,
);
let router_bytes = (moe.router.m as u64)
.checked_mul(moe.router.k as u64)
.and_then(|n| n.checked_mul(4))
.with_context(|| format!("layer {layer}: router projection size overflows u64"))?;
anyhow::ensure!(
router_bytes <= ctx.max_storage_buffer_binding_size,
"layer {layer}: router projection needs {router_bytes} bytes, over this adapter's {} \
byte storage-binding limit",
ctx.max_storage_buffer_binding_size,
);
Ok(GpuMoeFfn {
router: ctx.upload_f32(
&src.dequantize_weight(&moe.router),
&format!("l{layer}.ffn_gate_inp"),
),
bias: ctx.upload_f32(&moe.exp_probs_b, &format!("l{layer}.exp_probs_b")),
gate,
up,
down,
scratch,
})
}
#[allow(dead_code)]
struct GpuPipelines {
gemv_f32: wgpu::ComputePipeline,
gemv_f16: wgpu::ComputePipeline,
gemv_f32_accum: wgpu::ComputePipeline,
gemm_f32_nt: wgpu::ComputePipeline,
gemm_f32_nt_accum: wgpu::ComputePipeline,
gemv_q4_0: wgpu::ComputePipeline,
gemv_q4_0_fast: wgpu::ComputePipeline,
gemv_q4_k: wgpu::ComputePipeline,
gemv_q5_k: wgpu::ComputePipeline,
gemv_q6_k: wgpu::ComputePipeline,
gemv_q8_0: wgpu::ComputePipeline,
add_inplace: wgpu::ComputePipeline,
scaled_add_inplace: wgpu::ComputePipeline,
scale_f32: wgpu::ComputePipeline,
mul_inplace: wgpu::ComputePipeline,
silu_mul_inplace: wgpu::ComputePipeline,
ffn_swiglu_q4_0: wgpu::ComputePipeline,
gemv_q4_0_qkv: wgpu::ComputePipeline,
rmsnorm: wgpu::ComputePipeline,
per_head_rmsnorm: wgpu::ComputePipeline,
rope: wgpu::ComputePipeline,
kv_shift: wgpu::ComputePipeline,
flash_attention: wgpu::ComputePipeline,
conv1d_fused: wgpu::ComputePipeline,
argmax_f32: wgpu::ComputePipeline,
rmsnorm_batch: wgpu::ComputePipeline,
add_rmsnorm_batch: wgpu::ComputePipeline,
qk_norm_rope_batch: wgpu::ComputePipeline,
conv1d_fused_batch: wgpu::ComputePipeline,
bias_add: wgpu::ComputePipeline,
mul_mat_reg_tile_q4_0: wgpu::ComputePipeline,
mul_mat_reg_tile_q8_0: wgpu::ComputePipeline,
mul_mat_reg_tile_q4_k: wgpu::ComputePipeline,
mul_mat_reg_tile_q5_k: wgpu::ComputePipeline,
mul_mat_reg_tile_q6_k: wgpu::ComputePipeline,
mul_mat_reg_tile_f32: wgpu::ComputePipeline,
attention_prefill: wgpu::ComputePipeline,
moe_route: wgpu::ComputePipeline,
moe_gemv_q4_0: wgpu::ComputePipeline,
moe_combine: wgpu::ComputePipeline,
}
struct WgpuLoraTarget {
a: wgpu::Buffer,
b_scaled: wgpu::Buffer,
b_batched: wgpu::Buffer,
a_params: wgpu::Buffer,
b_params: wgpu::Buffer,
rank: u32,
#[allow(dead_code)]
k: u32,
d: u32,
}
struct WgpuLoraAdapter {
layers: Vec<[Option<WgpuLoraTarget>; crate::lora::LORA_TARGET_COUNT]>,
}
impl WgpuLoraAdapter {
fn upload(ctx: &GpuContext, w: &LoraAdapterWeights, residual_mult: f32) -> Self {
let mut layers = Vec::with_capacity(w.n_layers());
debug_assert!(
!w.has_moe_deltas(),
"adapter carries routed-FFN deltas, which this backend has no hooks for; \
Session::attach_lora_adapters is meant to have rejected it"
);
for layer in 0..w.n_layers() {
let mut targets: [Option<WgpuLoraTarget>; crate::lora::LORA_TARGET_COUNT] =
Default::default();
for target in LoraTarget::ALL {
let Some(t) = w.get(layer, target) else {
continue;
};
let rank = t.rank as u32;
let k = t.k as u32;
let d = t.d as u32;
let b_factor = match target {
LoraTarget::AttnOutput | LoraTarget::FfnDown => t.scale * residual_mult,
_ => t.scale,
};
let b_scaled_data: Vec<f32> = t.b.iter().map(|&x| x * b_factor).collect();
let b_scaled = ctx.upload_f32(&b_scaled_data, "lora_b_scaled");
let b_batched = if b_factor == t.scale {
b_scaled.clone()
} else {
ctx.upload_f32(
&t.b.iter().map(|&x| x * t.scale).collect::<Vec<f32>>(),
"lora_b_batched",
)
};
targets[target.index()] = Some(WgpuLoraTarget {
a: ctx.upload_f32(&t.a, "lora_a"),
b_scaled,
b_batched,
a_params: ctx
.upload_storage(bytemuck::cast_slice(&[rank, k, 0, 0]), "lora_a_p"),
b_params: ctx
.upload_storage(bytemuck::cast_slice(&[d, rank, 0, 0]), "lora_b_p"),
rank,
k,
d,
});
}
layers.push(targets);
}
Self { layers }
}
}
#[allow(dead_code)]
struct GpuState {
kv_caches: OnceLock<Vec<Option<(wgpu::Buffer, wgpu::Buffer)>>>,
conv_buffers: Vec<Option<wgpu::Buffer>>,
seq_len: AtomicUsize,
max_seq_len: usize,
embedding_f32: Vec<f32>,
}
struct HsScratch {
kv: Vec<Option<(wgpu::Buffer, wgpu::Buffer)>>,
conv: Vec<Option<wgpu::Buffer>>,
}
struct LoraGuard<'a>(&'a Mutex<Option<Arc<WgpuLoraAdapter>>>);
impl Drop for LoraGuard<'_> {
fn drop(&mut self) {
*self.0.lock().unwrap_or_else(|e| e.into_inner()) = None;
}
}
pub struct GpuLfm2Model {
ctx: GpuContext,
config: ModelConfig,
pipelines: GpuPipelines,
lm_head: LmHead,
output_norm: wgpu::Buffer,
layers: Vec<GpuLayerWeights>,
rope_type: RopeType,
scalars: ScalarMultipliers,
batched_prefill: bool,
batched_fallback_warned: AtomicBool,
moe_lora_dropped_warned: AtomicBool,
rope_freqs_buf: wgpu::Buffer,
has_freq_factors: bool,
hidden_buf: wgpu::Buffer, normed_buf: wgpu::Buffer, ffn_input_buf: wgpu::Buffer, gate_buf: wgpu::Buffer, up_buf: wgpu::Buffer, out_buf: wgpu::Buffer, q_buf: wgpu::Buffer, k_buf: wgpu::Buffer, v_buf: wgpu::Buffer, kv_shift_scratch: wgpu::Buffer,
attn_out_buf: wgpu::Buffer, logits_buf: wgpu::Buffer, argmax_out_buf: wgpu::Buffer,
argmax_readback_buf: wgpu::Buffer,
#[allow(dead_code)]
argmax_params: wgpu::Buffer,
argmax_bg: wgpu::BindGroup,
rmsnorm_hs_params: wgpu::Buffer, elementwise_hs_params: wgpu::Buffer, elementwise_is_params: wgpu::Buffer, elementwise_qdim_params: wgpu::Buffer,
elementwise_kvdim_params: wgpu::Buffer,
residual_add_params: wgpu::Buffer,
logit_scale_params: Option<wgpu::Buffer>,
conv1d_params: wgpu::Buffer, per_head_norm_params: wgpu::Buffer, rope_params: wgpu::Buffer,
attn_params: wgpu::Buffer, gemv_tile_params: Vec<wgpu::Buffer>, conv_proj_buf: wgpu::Buffer, conv_gate_buf: wgpu::Buffer, prefill_batch_buf: wgpu::Buffer,
prefill_normed_buf: wgpu::Buffer,
prefill_proj_buf: wgpu::Buffer,
prefill_gate_buf: wgpu::Buffer,
prefill_up_buf: wgpu::Buffer,
prefill_all_logits_buf: wgpu::Buffer,
gpu_state: GpuState,
infer_lock: Mutex<()>,
hs_scratch: OnceLock<HsScratch>,
use_hs_scratch: AtomicBool,
tq: OnceLock<TqGpuCache>,
kv_mode: OnceLock<Option<TqMode>>,
kv_cache_tag: OnceLock<String>,
model_id: String,
prefix_cache: Mutex<KvPrefixCache>,
lora_lru: Mutex<Vec<(Arc<LoraAdapterWeights>, Arc<WgpuLoraAdapter>)>>,
active_lora: Mutex<Option<Arc<WgpuLoraAdapter>>>,
lora_tmp: wgpu::Buffer,
lora_tmp_batched: wgpu::Buffer,
lora_params_pool: Mutex<(Vec<wgpu::Buffer>, usize)>,
prefill_params_pool: Mutex<(Vec<wgpu::Buffer>, usize)>,
}
#[derive(Clone, Copy)]
enum HiddenSeed<'a> {
Token(u32),
Embedding(&'a [f32]),
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum TailArgmax {
None,
Dispatch,
DispatchAndStage,
}
#[derive(Clone, Copy, PartialEq, Eq)]
enum DecodeTail {
Logits(TailArgmax),
Hidden,
HiddenUnsubmitted,
LogitsUnsubmitted(TailArgmax),
}
impl GpuLfm2Model {
pub fn from_gguf(gguf: GgufFile, context_size: usize) -> Result<Self> {
Self::from_gguf_with_id(gguf, context_size, String::new())
}
pub fn from_gguf_with_id(
gguf: GgufFile,
context_size: usize,
model_id: String,
) -> Result<Self> {
let ctx = GpuContext::new()?;
Self::from_gguf_with_ctx(gguf, context_size, model_id, ctx)
}
pub fn from_gguf_with_ctx(
gguf: GgufFile,
context_size: usize,
model_id: String,
ctx: GpuContext,
) -> Result<Self> {
let arch = gguf.architecture().unwrap_or("").to_lowercase();
match arch.as_str() {
"llama" | "qwen2" | "qwen3" | "granite" => {
let cpu_model = super::llama::LlamaModel::from_gguf_with_id(
gguf,
context_size,
model_id.clone(),
)?;
Self::from_weight_source_with_ctx(&cpu_model, context_size, model_id, ctx)
}
_ => {
let cpu_model = super::lfm2::Lfm2Model::from_gguf_with_id(
gguf,
context_size,
model_id.clone(),
)?;
Self::from_weight_source_with_ctx(&cpu_model, context_size, model_id, ctx)
}
}
}
pub fn ctx(&self) -> &GpuContext {
&self.ctx
}
pub fn from_llama_with_id(
gguf: GgufFile,
context_size: usize,
model_id: String,
) -> Result<Self> {
let cpu_model =
super::llama::LlamaModel::from_gguf_with_id(gguf, context_size, model_id.clone())?;
Self::from_weight_source(&cpu_model, context_size, model_id)
}
fn from_weight_source(
src: &dyn GpuWeightSource,
context_size: usize,
model_id: String,
) -> Result<Self> {
let ctx = GpuContext::new()?;
Self::from_weight_source_with_ctx(src, context_size, model_id, ctx)
}
pub fn from_dspark_with_ctx(
dspark: std::sync::Arc<crate::model::dspark::DSparkDraftModel>,
context_size: usize,
model_id: String,
ctx: GpuContext,
) -> Result<Self> {
let dspark_cfg = dspark.config.to_model_config(context_size);
let dspark_src = crate::model::dspark::DSparkGpuWeightSource {
config: dspark_cfg,
dspark,
};
Self::from_weight_source_with_ctx(&dspark_src, context_size, model_id, ctx)
}
pub fn from_weight_source_with_ctx(
src: &dyn GpuWeightSource,
context_size: usize,
model_id: String,
ctx: GpuContext,
) -> Result<Self> {
let mut config = src.config().clone();
let max_seq_len = context_size.min(config.max_seq_len);
config.max_seq_len = max_seq_len;
let hs = config.hidden_size;
let is = config.intermediate_size;
let head_dim = config.head_dim;
let q_dim = config.n_heads * head_dim;
let max_kv_dim = config.kv_heads_per_layer.iter().copied().max().unwrap_or(0) * head_dim;
let rope_type = src.rope_type();
let scalars = config.scalars;
let batched_prefill = src.supports_batched_prefill();
anyhow::ensure!(
config.moe.is_none() || scalars.residual == 1.0,
"mixture-of-experts with a residual multiplier ({}) is not supported on the wgpu \
backend; the routed FFN combine adds into the residual unscaled",
scalars.residual,
);
tracing::info!(
"GPU model: {} layers, hs={hs}, is={is}, vocab={}",
config.n_layers,
config.vocab_size
);
let pipelines = GpuPipelines {
gemv_f32: ctx.create_pipeline(shaders::GEMV_F32, "gemv_f32", "gemv_f32"),
gemv_f16: ctx.create_pipeline_with_defines(
shaders::GEMV_F32,
"gemv_f32",
"gemv_f16",
&[("F16_A", "1")],
),
gemv_f32_accum: ctx.create_pipeline(
shaders::GEMV_F32,
"gemv_f32_accum",
"gemv_f32_accum",
),
gemm_f32_nt: ctx.create_pipeline(shaders::GEMM_F32, "gemm_f32_nt", "gemm_f32_nt"),
gemm_f32_nt_accum: ctx.create_pipeline(
shaders::GEMM_F32,
"gemm_f32_nt_accum",
"gemm_f32_nt_accum",
),
gemv_q4_0: ctx.create_pipeline(shaders::GEMV_Q4_0, "gemv_q4_0", "gemv_q4_0"),
gemv_q4_0_fast: ctx.create_pipeline(
shaders::GEMV_Q4_0_FAST,
"gemv_q4_0_fast",
"gemv_q4_0_fast",
),
gemv_q4_k: ctx.create_pipeline(shaders::GEMV_Q4_K, "gemv_q4_k", "gemv_q4_k"),
gemv_q5_k: ctx.create_pipeline(shaders::GEMV_Q5_K, "gemv_q5_k", "gemv_q5_k"),
gemv_q6_k: ctx.create_pipeline(shaders::GEMV_Q6_K, "gemv_q6_k", "gemv_q6_k"),
gemv_q8_0: ctx.create_pipeline(shaders::GEMV_Q8_0, "gemv_q8_0", "gemv_q8_0"),
add_inplace: ctx.create_pipeline(shaders::ELEMENTWISE, "add_inplace", "add"),
scaled_add_inplace: ctx.create_pipeline(
shaders::ELEMENTWISE,
"scaled_add_inplace",
"scaled_add",
),
scale_f32: ctx.create_pipeline(shaders::SCALE_F32, "scale_f32", "scale_f32"),
mul_inplace: ctx.create_pipeline(shaders::ELEMENTWISE, "mul_inplace", "mul"),
silu_mul_inplace: ctx.create_pipeline(
shaders::ELEMENTWISE,
"silu_mul_inplace",
"silu_mul",
),
ffn_swiglu_q4_0: ctx.create_pipeline(
shaders::FFN_SWIGLU_Q4_0,
"ffn_swiglu_q4_0",
"ffn_swiglu_q4_0",
),
gemv_q4_0_qkv: ctx.create_pipeline(
shaders::GEMV_Q4_0_QKV,
"gemv_q4_0_qkv",
"gemv_q4_0_qkv",
),
rmsnorm: ctx.create_pipeline(shaders::RMSNORM, "rmsnorm", "rmsnorm"),
per_head_rmsnorm: ctx.create_pipeline(
shaders::PER_HEAD_RMSNORM,
"per_head_rmsnorm",
"per_head_rmsnorm",
),
rope: ctx.create_pipeline(shaders::ROPE, "rope", "rope"),
kv_shift: ctx.create_pipeline(shaders::KV_SHIFT, "kv_shift", "kv_shift"),
flash_attention: ctx.create_pipeline(
shaders::FLASH_ATTENTION,
"flash_attention",
"flash_attention",
),
conv1d_fused: ctx.create_pipeline(
shaders::CONV1D_FUSED,
"conv1d_fused",
"conv1d_fused",
),
argmax_f32: ctx.create_pipeline(shaders::ARGMAX_F32, "argmax_f32", "argmax_f32"),
rmsnorm_batch: ctx.create_pipeline(
shaders::RMSNORM_BATCH,
"rmsnorm_batch",
"rmsnorm_batch",
),
add_rmsnorm_batch: ctx.create_pipeline(
shaders::RMSNORM_BATCH,
"add_rmsnorm_batch",
"add_rmsnorm_batch",
),
qk_norm_rope_batch: ctx.create_pipeline(
shaders::QK_NORM_ROPE_BATCH,
"qk_norm_rope_batch",
"qk_norm_rope_batch",
),
conv1d_fused_batch: ctx.create_pipeline(
shaders::CONV1D_FUSED_BATCH,
"conv1d_fused_batch",
"conv1d_fused_batch",
),
bias_add: ctx.create_pipeline(shaders::BIAS_ADD, "bias_add", "bias_add"),
mul_mat_reg_tile_q4_0: if use_spirv_passthrough(&ctx) {
tracing::debug!("mul_mat_reg_tile_q4_0: SPIR-V passthrough (slang)");
ctx.mul_mat_reg_tile_q4_0_passthrough()
} else {
build_mul_mat_pipeline(&ctx, "mul_mat_q4_0", "INIT_SRC0_SHMEM_Q4_0", "u32")
},
mul_mat_reg_tile_q8_0: if use_spirv_passthrough(&ctx) {
tracing::debug!("mul_mat_reg_tile_q8_0: SPIR-V passthrough (slang)");
ctx.mul_mat_reg_tile_q8_0_passthrough()
} else {
build_mul_mat_pipeline(&ctx, "mul_mat_q8_0", "INIT_SRC0_SHMEM_Q8_0", "u32")
},
mul_mat_reg_tile_q4_k: if use_spirv_passthrough(&ctx) {
tracing::debug!("mul_mat_reg_tile_q4_k: SPIR-V passthrough (slang)");
ctx.mul_mat_reg_tile_q4_k_passthrough()
} else {
build_mul_mat_pipeline(&ctx, "mul_mat_q4_k", "INIT_SRC0_SHMEM_Q4_K", "u32")
},
mul_mat_reg_tile_q5_k: if use_spirv_passthrough(&ctx) {
tracing::debug!("mul_mat_reg_tile_q5_k: SPIR-V passthrough (slang)");
ctx.mul_mat_reg_tile_q5_k_passthrough()
} else {
build_mul_mat_pipeline(&ctx, "mul_mat_q5_k", "INIT_SRC0_SHMEM_Q5_K", "u32")
},
mul_mat_reg_tile_q6_k: if use_spirv_passthrough(&ctx) {
tracing::debug!("mul_mat_reg_tile_q6_k: SPIR-V passthrough (slang)");
ctx.mul_mat_reg_tile_q6_k_passthrough()
} else {
build_mul_mat_pipeline(&ctx, "mul_mat_q6_k", "INIT_SRC0_SHMEM_Q6_K", "u32")
},
mul_mat_reg_tile_f32: build_mul_mat_pipeline(
&ctx,
"mul_mat_f32",
"INIT_SRC0_SHMEM_FLOAT",
"f32",
),
attention_prefill: ctx.create_pipeline(
shaders::ATTENTION_PREFILL,
"attention_prefill",
"attention_prefill",
),
moe_route: ctx.create_pipeline(shaders::MOE_ROUTE, "moe_route", "moe_route"),
moe_gemv_q4_0: ctx.create_pipeline(
shaders::MOE_GEMV_Q4_0,
"moe_gemv_q4_0",
"moe_gemv_q4_0",
),
moe_combine: ctx.create_pipeline(shaders::MOE_COMBINE, "moe_combine", "moe_combine"),
};
let emb_tensor = src.embedding_tensor()?;
let embedding_raw = emb_tensor.to_f32_vec();
let (lm_head_dtype, lm_head_bytes) = match src.output_ref() {
Some(wref) => (wref.dtype, src.weight_bytes(wref)),
None => (emb_tensor.dtype(), src.embedding_tensor_data()?),
};
let lm_head_bound_bytes = (lm_head_bytes.len() as u64).div_ceil(4) * 4;
let lm_head_params = ctx.upload_storage(
bytemuck::cast_slice(&[
config.vocab_size as u32,
config.hidden_size as u32,
0u32,
0u32,
]),
"lm_head.params",
);
let lm_head = if Self::has_quantized_gemv(lm_head_dtype)
&& lm_head_bound_bytes <= ctx.max_storage_buffer_binding_size
{
LmHead::Quantized(GpuWeight {
tensor: GpuTensor {
buffer: ctx.upload_storage(&lm_head_bytes, "lm_head"),
dtype: lm_head_dtype,
shape: vec![config.vocab_size, config.hidden_size],
},
params_buf: lm_head_params,
cached_bg: None,
})
} else {
let f16_weight = match src.output_ref() {
Some(wref) => ctx.upload_f32_as_f16(&src.dequantize_weight(wref), "output.weight"),
None => ctx.upload_f32_as_f16(&embedding_raw, "token_embd.weight"),
};
LmHead::F16 {
weight: f16_weight,
params: lm_head_params,
}
};
let mut embedding_f32 = embedding_raw;
if scalars.embedding != 1.0 {
for v in embedding_f32.iter_mut() {
*v *= scalars.embedding;
}
}
let output_norm = ctx.upload_f32(src.output_norm_weight(), "output_norm");
let upload_weight = |wref: &WeightRef, name: &str| -> GpuWeight {
let (buf, dtype) = if matches!(
wref.dtype,
DType::Q4_0 | DType::Q8_0 | DType::Q4KM | DType::Q5KM | DType::Q6K
) {
let data = src.weight_bytes(wref);
(ctx.upload_storage(&data, name), wref.dtype)
} else {
let f32_data = src.dequantize_weight(wref);
(ctx.upload_f32(&f32_data, name), DType::F32)
};
let params_buf = ctx.upload_storage(
bytemuck::cast_slice(&[wref.m as u32, wref.k as u32, 0u32, 0u32]),
&format!("{name}.params"),
);
GpuWeight {
tensor: GpuTensor {
buffer: buf,
dtype,
shape: vec![wref.m, wref.k],
},
params_buf,
cached_bg: None,
}
};
let upload_opt_f32 = |data: Option<&[f32]>, name: &str| -> Option<wgpu::Buffer> {
data.map(|d| ctx.upload_f32(d, name))
};
let max_pref = max_seq_len.min(MAX_PREFILL_TOKENS);
let moe_scratch = config
.moe
.as_ref()
.map(|m| -> Result<Arc<MoeScratch>> {
use anyhow::Context;
let buf = |n: usize, name: &str| -> Result<wgpu::Buffer> {
let bytes = n as u64 * 4;
anyhow::ensure!(
bytes <= ctx.max_storage_buffer_binding_size,
"mixture-of-experts scratch `{name}` needs {bytes} bytes, over this \
adapter's {} byte storage-binding limit",
ctx.max_storage_buffer_binding_size,
);
Ok(ctx.create_storage_rw(bytes, name))
};
anyhow::ensure!(
(1..=MOE_MAX_EXPERTS as usize).contains(&m.n_expert)
&& (1..=MOE_MAX_EXPERT_USED as usize).contains(&m.n_expert_used)
&& m.n_expert_used <= m.n_expert,
"wgpu MoE routing supports 1..={MOE_MAX_EXPERTS} experts and \
1..={MOE_MAX_EXPERT_USED} active (and no more active than available), \
model declares {} and {}",
m.n_expert,
m.n_expert_used,
);
let entries = max_pref * m.n_expert_used;
let dim = |v: usize, what: &str| -> Result<u32> {
u32::try_from(v).with_context(|| {
format!("mixture-of-experts {what} {v} does not fit the kernels' u32")
})
};
let cells = |rows: usize, cols: usize, what: &str| -> Result<usize> {
rows.checked_mul(cols).with_context(|| {
format!("mixture-of-experts {what} scratch size overflows")
})
};
let max_entries = dim(entries, "entry count")?;
anyhow::ensure!(
max_entries <= crate::backend::wgpu::MAX_WG,
"a {max_pref}-token chunk over {} experts per token is {max_entries} \
entries, over the {} workgroups-per-dimension cap the expert GEMV \
dispatches one entry per workgroup against",
m.n_expert_used,
crate::backend::wgpu::MAX_WG,
);
let silu_groups = cells(entries, m.expert_ff_len, "SwiGLU")?.div_ceil(256);
anyhow::ensure!(
silu_groups <= crate::backend::wgpu::MAX_WG as usize,
"the routed SwiGLU over {max_entries} entries of width {} needs \
{silu_groups} workgroups, over this backend's {} per-dimension cap",
m.expert_ff_len,
crate::backend::wgpu::MAX_WG,
);
Ok(Arc::new(MoeScratch {
n_expert: dim(m.n_expert, "expert count")?,
n_expert_used: dim(m.n_expert_used, "active expert count")?,
expert_ff_len: dim(m.expert_ff_len, "expert width")?,
max_entries,
logits: buf(cells(max_pref, m.n_expert, "logits")?, "moe.logits")?,
sel_expert: buf(entries, "moe.sel_expert")?,
sel_weight: buf(entries, "moe.sel_weight")?,
gate: buf(cells(entries, m.expert_ff_len, "gate")?, "moe.gate")?,
up: buf(cells(entries, m.expert_ff_len, "up")?, "moe.up")?,
z: buf(cells(entries, hs, "output")?, "moe.z")?,
}))
})
.transpose()?;
let mut layers = Vec::with_capacity(config.n_layers);
for i in 0..config.n_layers {
let attn_norm = ctx.upload_f32(src.attn_norm_weight(i), &format!("l{i}.anorm"));
let ffn_norm = ctx.upload_f32(src.ffn_norm_weight(i), &format!("l{i}.fnorm"));
let ffn = match src.moe_refs(i) {
None => GpuFfn::Dense(Box::new(GpuDenseFfn {
gate: upload_weight(src.ffn_gate_ref(i)?, &format!("l{i}.ffn_gate")),
up: upload_weight(src.ffn_up_ref(i)?, &format!("l{i}.ffn_up")),
down: upload_weight(src.ffn_down_ref(i)?, &format!("l{i}.ffn_down")),
})),
Some(m) => GpuFfn::Moe(Box::new(upload_moe(
&ctx,
src,
moe_scratch.as_ref(),
hs,
i,
m,
)?)),
};
let is_conv = config.block_types[i] == BlockType::GatedConv;
let (conv_in_proj, conv_out_proj, conv_weight) = if is_conv {
let ip = src
.conv_in_proj_ref(i)
.ok_or_else(|| anyhow!("conv layer missing in_proj"))?;
let op = src
.conv_out_proj_ref(i)
.ok_or_else(|| anyhow!("conv layer missing out_proj"))?;
(
Some(upload_weight(ip, &format!("l{i}.conv_ip"))),
Some(upload_weight(op, &format!("l{i}.conv_op"))),
Some(
ctx.upload_f32(
src.conv_weight(i)
.ok_or_else(|| anyhow!("conv layer missing conv weight"))?,
&format!("l{i}.conv_w"),
),
),
)
} else {
(None, None, None)
};
let (attn_q, attn_k, attn_v, attn_output, attn_q_norm, attn_k_norm) = if !is_conv {
(
Some(upload_weight(
src.attn_q_ref(i)
.ok_or_else(|| anyhow!("attn layer missing q"))?,
&format!("l{i}.attn_q"),
)),
Some(upload_weight(
src.attn_k_ref(i)
.ok_or_else(|| anyhow!("attn layer missing k"))?,
&format!("l{i}.attn_k"),
)),
Some(upload_weight(
src.attn_v_ref(i)
.ok_or_else(|| anyhow!("attn layer missing v"))?,
&format!("l{i}.attn_v"),
)),
Some(upload_weight(
src.attn_output_ref(i)
.ok_or_else(|| anyhow!("attn layer missing output"))?,
&format!("l{i}.attn_o"),
)),
upload_opt_f32(src.attn_q_norm_weight(i), &format!("l{i}.qn")),
upload_opt_f32(src.attn_k_norm_weight(i), &format!("l{i}.kn")),
)
} else {
(None, None, None, None, None, None)
};
let attn_q_bias = upload_opt_f32(src.attn_q_bias(i), &format!("l{i}.qb"));
let attn_k_bias = upload_opt_f32(src.attn_k_bias(i), &format!("l{i}.kb"));
let attn_v_bias = upload_opt_f32(src.attn_v_bias(i), &format!("l{i}.vb"));
layers.push(GpuLayerWeights {
attn_norm,
ffn_norm,
ffn,
conv_in_proj,
conv_out_proj,
conv_weight,
attn_q,
attn_k,
attn_v,
attn_output,
attn_q_norm,
attn_k_norm,
attn_q_bias,
attn_k_bias,
attn_v_bias,
attn_norm_bg: None,
ffn_norm_bg: None,
rope_bg: None,
conv_fused_bg: None,
conv_add_bg: None,
attn_out_add_bg: None,
silu_bg: None,
ffn_swiglu_bg: None,
attn_qkv_bg: None,
attn_qkv_params_buf: None,
attn_bg: None,
qn_bg: None,
kn_bg: None,
qb_bg: None,
kb_bg: None,
vb_bg: None,
ffn_add_bg: None,
});
}
let f = |size: usize, name: &str| ctx.create_storage_rw((size * 4) as u64, name);
let hidden_buf = f(hs, "hidden");
let normed_buf = f(hs, "normed");
let ffn_input_buf = f(hs, "ffn_input");
let gate_buf = f(is, "gate");
let up_buf = f(is, "up");
let out_buf = f(hs, "out");
let q_buf = f(q_dim, "q");
let k_buf = f(max_kv_dim, "k");
let v_buf = f(max_kv_dim, "v");
let kv_shift_scratch =
ctx.create_storage_rw(kv_slab_bytes(max_seq_len, max_kv_dim), "kv_shift_scratch");
let attn_out_buf = f(q_dim, "attn_out");
let logits_buf = f(config.vocab_size, "logits");
let conv_proj_buf = f(3 * hs, "conv_proj");
let conv_gate_buf = f(hs, "conv_gate");
let prefill_batch_buf = f(hs * max_pref, "prefill_batch");
let prefill_normed_buf = f(hs.max(q_dim) * max_pref, "prefill_normed");
let prefill_proj_buf = f((3 * hs).max(q_dim) * max_pref, "prefill_proj");
let prefill_gate_buf = f(is.max(max_kv_dim).max(hs) * max_pref, "prefill_gate");
let prefill_up_buf = f(is.max(max_kv_dim).max(hs) * max_pref, "prefill_up");
let max_all_logits = max_seq_len.min(MAX_ALL_LOGITS_TOKENS);
let prefill_all_logits_buf = f(config.vocab_size * max_all_logits, "prefill_all_logits");
let kernel_size = config.conv_kernel_size.unwrap_or(3);
let d_conv = kernel_size - 1;
let mut conv_buffers = Vec::with_capacity(config.n_layers);
for i in 0..config.n_layers {
if config.block_types[i] == BlockType::Attention {
conv_buffers.push(None);
} else {
let cb = f(d_conv * hs, &format!("l{i}.conv_buf"));
conv_buffers.push(Some(cb));
}
}
let gpu_state = GpuState {
kv_caches: OnceLock::new(),
conv_buffers,
seq_len: AtomicUsize::new(0),
max_seq_len,
embedding_f32,
};
let rmsnorm_hs_params = ctx.upload_storage(
bytemuck::cast_slice(&[hs as u32, config.rms_norm_eps.to_bits(), 0u32, 0u32]),
"rmsnorm_hs_params",
);
let elementwise_hs_params =
ctx.upload_storage(bytemuck::cast_slice(&[hs as u32, 0u32]), "ew_hs_params");
let elementwise_is_params =
ctx.upload_storage(bytemuck::cast_slice(&[is as u32, 0u32]), "ew_is_params");
let kv_dim_bias = config.n_kv_heads * head_dim;
let elementwise_qdim_params = ctx.upload_storage(
bytemuck::cast_slice(&[q_dim as u32, 0u32]),
"ew_qdim_params",
);
let elementwise_kvdim_params = ctx.upload_storage(
bytemuck::cast_slice(&[kv_dim_bias as u32, 0u32]),
"ew_kvdim_params",
);
let residual_add_params = ctx.upload_storage(
bytemuck::cast_slice(&[hs as u32, scalars.residual.to_bits()]),
"residual_add_params",
);
let logit_scale_params = (scalars.logit != 1.0).then(|| {
ctx.upload_storage(
bytemuck::cast_slice(&[config.vocab_size as u32, (1.0 / scalars.logit).to_bits()]),
"logit_scale_params",
)
});
let kernel_size = config.conv_kernel_size.unwrap_or(3) as u32;
let d_conv = kernel_size - 1;
let head_dim_u32 = head_dim as u32;
let conv1d_params = ctx.upload_storage(
bytemuck::cast_slice(&[hs as u32, kernel_size, d_conv, 0u32]),
"conv1d_params",
);
let per_head_norm_params = ctx.upload_storage(
bytemuck::cast_slice(&[head_dim_u32, config.rms_norm_eps.to_bits(), 0u32, 0u32]),
"ph_norm_params",
);
let rope_params = ctx.create_storage_rw(7 * 4, "rope_params");
let has_freq_factors = src.rope_freqs().is_some();
let rope_freqs_buf = match src.rope_freqs() {
Some(rf) => ctx.upload_f32(rf, "rope_freqs"),
None => ctx.upload_f32(&[1.0f32], "rope_freqs_dummy"),
};
let attn_params = ctx.create_storage_rw(8 * 4, "attn_params");
let gemv_tile_params = if matches!(lm_head, LmHead::F16 { .. }) {
let tile_rows = gemv_tile_rows(
config.vocab_size as u32,
hs as u32,
ctx.max_storage_buffer_binding_size,
ctx.min_storage_buffer_offset_alignment,
2,
);
let tile_count = (config.vocab_size as u32).div_ceil(tile_rows);
(0..tile_count)
.map(|i| ctx.create_storage_rw(4 * 4, &format!("gemv_tile_params.{i}")))
.collect()
} else {
Vec::new()
};
let argmax_out_buf =
ctx.create_storage_rw((MAX_PREFILL_TOKENS * 4).max(64) as u64, "argmax_out");
let argmax_readback_buf =
ctx.create_readback_buffer((MAX_PREFILL_TOKENS * 4).max(64) as u64, "argmax_readback");
let argmax_params = ctx.upload_storage(
bytemuck::cast_slice(&[config.vocab_size as u32, 0u32]),
"argmax_params",
);
let argmax_bg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("argmax_bg"),
layout: &pipelines.argmax_f32.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: logits_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: argmax_out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: argmax_params.as_entire_binding(),
},
],
});
let prefix_cache = Mutex::new(KvPrefixCache::new(
crate::kv_cache::KvCacheConfig::default(),
&config,
&format!("wgpu:{model_id}"),
));
let lora_tmp = ctx.create_storage_rw((crate::lora::MAX_LORA_RANK * 4) as u64, "lora_tmp");
let lora_tmp_batched = ctx.create_storage_rw(
(crate::lora::MAX_LORA_RANK * max_pref * 4) as u64,
"lora_tmp_batched",
);
let mut model = Self {
ctx,
config,
pipelines,
lm_head,
output_norm,
layers,
rope_type,
scalars,
batched_prefill,
batched_fallback_warned: AtomicBool::new(false),
moe_lora_dropped_warned: AtomicBool::new(false),
rope_freqs_buf,
has_freq_factors,
hidden_buf,
normed_buf,
ffn_input_buf,
gate_buf,
up_buf,
out_buf,
q_buf,
k_buf,
v_buf,
kv_shift_scratch,
attn_out_buf,
logits_buf,
argmax_out_buf,
argmax_readback_buf,
argmax_params,
argmax_bg,
rmsnorm_hs_params,
elementwise_hs_params,
elementwise_is_params,
elementwise_qdim_params,
elementwise_kvdim_params,
residual_add_params,
logit_scale_params,
conv1d_params,
per_head_norm_params,
rope_params,
attn_params,
gemv_tile_params,
conv_proj_buf,
conv_gate_buf,
prefill_batch_buf,
prefill_normed_buf,
prefill_proj_buf,
prefill_gate_buf,
prefill_up_buf,
prefill_all_logits_buf,
gpu_state,
infer_lock: Mutex::new(()),
hs_scratch: OnceLock::new(),
use_hs_scratch: AtomicBool::new(false),
tq: OnceLock::new(),
kv_mode: OnceLock::new(),
kv_cache_tag: OnceLock::new(),
prefix_cache,
model_id,
lora_lru: Mutex::new(Vec::new()),
active_lora: Mutex::new(None),
lora_tmp,
lora_tmp_batched,
lora_params_pool: Mutex::new((Vec::new(), 0)),
prefill_params_pool: Mutex::new((Vec::new(), 0)),
};
model.cache_bind_groups();
Ok(model)
}
fn resolve_lora(&self, state: &InferenceState) -> LoraGuard<'_> {
let usable = state.lora.as_ref().filter(|adapter| {
let ok = !adapter.has_moe_deltas();
if !ok && !self.moe_lora_dropped_warned.swap(true, Ordering::Relaxed) {
tracing::error!(
"ignoring a LoRA adapter that carries mixture-of-experts deltas: this \
backend has no routed-FFN hooks. Attach through `Session`, which refuses \
it with `CeraError::LoraUnsupportedByBackend` instead of ignoring it."
);
}
ok
});
let resolved = usable.map(|adapter| {
let mut lru = self.lora_lru.lock().unwrap_or_else(|e| e.into_inner());
if let Some(pos) = lru.iter().position(|(cpu, _)| Arc::ptr_eq(cpu, adapter)) {
let (cpu, gpu) = lru.remove(pos);
lru.push((cpu, gpu.clone()));
gpu
} else {
let gpu = Arc::new(WgpuLoraAdapter::upload(
&self.ctx,
adapter,
self.scalars.residual,
));
lru.push((adapter.clone(), gpu.clone()));
if lru.len() > 3 {
lru.remove(0);
}
gpu
}
});
*self.active_lora.lock().unwrap_or_else(|e| e.into_inner()) = resolved;
LoraGuard(&self.active_lora)
}
fn lora_target_bgs(
&self,
t: &WgpuLoraTarget,
input: &wgpu::Buffer,
output: &wgpu::Buffer,
) -> (wgpu::BindGroup, wgpu::BindGroup) {
let layout_a = self.pipelines.gemv_f32.get_bind_group_layout(0);
let layout_b = self.pipelines.gemv_f32_accum.get_bind_group_layout(0);
let bg_a = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("lora_a"),
layout: &layout_a,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: t.a.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: input.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.lora_tmp.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: t.a_params.as_entire_binding(),
},
],
});
let bg_b = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("lora_b"),
layout: &layout_b,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: t.b_scaled.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.lora_tmp.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: output.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: t.b_params.as_entire_binding(),
},
],
});
(bg_a, bg_b)
}
fn dispatch_lora_into(
&self,
pass: &mut wgpu::ComputePass<'_>,
t: &WgpuLoraTarget,
bg_a: &wgpu::BindGroup,
bg_b: &wgpu::BindGroup,
) {
let a_groups = t.rank.div_ceil(GEMV_F32_ROWS_PER_WG);
self.dispatch_into(
pass,
&self.pipelines.gemv_f32,
bg_a,
crate::backend::wgpu::gemv_row_workgroups(a_groups),
);
let b_groups = t.d.div_ceil(GEMV_F32_ROWS_PER_WG);
self.dispatch_into(
pass,
&self.pipelines.gemv_f32_accum,
bg_b,
crate::backend::wgpu::gemv_row_workgroups(b_groups),
);
}
fn lora_target(
lora: Option<&Arc<WgpuLoraAdapter>>,
layer: usize,
target: LoraTarget,
) -> Option<&WgpuLoraTarget> {
lora?.layers.get(layer)?[target.index()].as_ref()
}
fn next_lora_params(&self, data: &[u32; 4]) -> wgpu::Buffer {
let mut pool = self
.lora_params_pool
.lock()
.unwrap_or_else(|e| e.into_inner());
let (bufs, next) = &mut *pool;
let idx = *next;
*next += 1;
if bufs.len() <= idx {
bufs.push(self.ctx.create_storage_rw(16, "lora_batched_params"));
}
let buf = bufs[idx].clone();
self.ctx
.queue
.write_buffer(&buf, 0, bytemuck::cast_slice(data));
buf
}
fn next_prefill_params(&self, data: &[u8]) -> wgpu::Buffer {
let mut pool = self
.prefill_params_pool
.lock()
.unwrap_or_else(|e| e.into_inner());
let (bufs, next) = &mut *pool;
let idx = *next;
*next += 1;
let aligned_size = (data.len().max(16).div_ceil(4) * 4) as u64;
if bufs.len() <= idx || bufs[idx].size() < aligned_size {
if bufs.len() <= idx {
bufs.push(
self.ctx
.create_storage_rw(aligned_size.max(64), "prefill_batched_params"),
);
} else {
bufs[idx] = self
.ctx
.create_storage_rw(aligned_size.max(64), "prefill_batched_params");
}
}
let buf = bufs[idx].clone();
self.ctx.queue.write_buffer(&buf, 0, data);
buf
}
fn encode_lora_batched(
&self,
enc: &mut wgpu::CommandEncoder,
t: &WgpuLoraTarget,
input: &wgpu::Buffer,
output: &wgpu::Buffer,
n: u32,
) {
let total1 = n * t.rank;
let p1: [u32; 4] = [n, t.rank, t.k, 0];
let p1_buf = self.next_lora_params(&p1);
let bg1 = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("lora_batched_a"),
layout: &self.pipelines.gemm_f32_nt.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: input.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: t.a.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.lora_tmp_batched.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p1_buf.as_entire_binding(),
},
],
});
self.encode(
enc,
&self.pipelines.gemm_f32_nt,
&bg1,
crate::backend::wgpu::gemv_row_workgroups(total1),
"lora_batched_a",
);
let total2 = n * t.d;
let p2: [u32; 4] = [n, t.d, t.rank, 0];
let p2_buf = self.next_lora_params(&p2);
let bg2 = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("lora_batched_b"),
layout: &self.pipelines.gemm_f32_nt_accum.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.lora_tmp_batched.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: t.b_batched.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: output.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p2_buf.as_entire_binding(),
},
],
});
self.encode(
enc,
&self.pipelines.gemm_f32_nt_accum,
&bg2,
crate::backend::wgpu::gemv_row_workgroups(total2),
"lora_batched_b",
);
}
#[allow(clippy::too_many_arguments)]
fn encode_lora_hook_batched(
&self,
enc: &mut wgpu::CommandEncoder,
lora: Option<&Arc<WgpuLoraAdapter>>,
layer: usize,
target: LoraTarget,
input: &wgpu::Buffer,
output: &wgpu::Buffer,
n: u32,
) {
if let Some(t) = Self::lora_target(lora, layer, target) {
self.encode_lora_batched(enc, t, input, output, n);
}
}
fn make_gemv_bg(
&self,
w: &GpuWeight,
input: &wgpu::Buffer,
output: &wgpu::Buffer,
) -> wgpu::BindGroup {
let (pipeline, _, _) = self.gemv_pipeline_rows_label(w);
self.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: w.tensor.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: input.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: output.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: w.params_buf.as_entire_binding(),
},
],
})
}
fn gemv_pipeline_rows_label(
&self,
w: &GpuWeight,
) -> (&wgpu::ComputePipeline, u32, &'static str) {
match w.tensor.dtype {
DType::Q4_0 => (&self.pipelines.gemv_q4_0_fast, 8, "gemv_q4"),
DType::Q8_0 => (&self.pipelines.gemv_q8_0, 8, "gemv_q8"),
DType::Q4KM => (&self.pipelines.gemv_q4_k, 2, "gemv_q4k"),
DType::Q6K => (&self.pipelines.gemv_q6_k, 2, "gemv_q6"),
DType::Q5KM => (&self.pipelines.gemv_q5_k, 2, "gemv_q5k"),
_ => (&self.pipelines.gemv_f32, 8, "gemv_f32"),
}
}
fn has_quantized_gemv(dtype: DType) -> bool {
matches!(
dtype,
DType::Q4_0 | DType::Q8_0 | DType::Q4KM | DType::Q5KM | DType::Q6K
)
}
fn gemv_workgroups(&self, w: &GpuWeight) -> (u32, u32, u32) {
let (_, rows_per_wg, _) = self.gemv_pipeline_rows_label(w);
let row_groups = (w.tensor.shape[0] as u32).div_ceil(rows_per_wg);
crate::backend::wgpu::gemv_row_workgroups(row_groups)
}
fn dispatch_gemv_into(
&self,
pass: &mut wgpu::ComputePass<'_>,
w: &GpuWeight,
bind_group: &wgpu::BindGroup,
) {
let (pipeline, _, _) = self.gemv_pipeline_rows_label(w);
self.dispatch_into(pass, pipeline, bind_group, self.gemv_workgroups(w));
}
fn cache_bind_groups(&mut self) {
let cfg = &self.config;
for i in 0..cfg.n_layers {
let attn_norm_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.rmsnorm.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.normed_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.layers[i].attn_norm.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.rmsnorm_hs_params.as_entire_binding(),
},
],
});
let ffn_norm_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.rmsnorm.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.ffn_input_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.layers[i].ffn_norm.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.rmsnorm_hs_params.as_entire_binding(),
},
],
});
let silu_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.silu_mul_inplace.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.gate_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.up_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.elementwise_is_params.as_entire_binding(),
},
],
});
let ffn_add_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.scaled_add_inplace.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.hidden_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.residual_add_params.as_entire_binding(),
},
],
});
let (conv_fused_bg, conv_add_bg) = if cfg.block_types[i] == BlockType::GatedConv {
let conv_buf = self.active_conv(i);
let conv_p = &self.conv1d_params;
let bg_fused = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.conv1d_fused.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.conv_proj_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: conv_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.layers[i]
.conv_weight
.as_ref()
.unwrap()
.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: self.conv_gate_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: conv_p.as_entire_binding(),
},
],
});
let add_p = &self.elementwise_hs_params;
let bg_add = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.add_inplace.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.hidden_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: add_p.as_entire_binding(),
},
],
});
(Some(bg_fused), Some(bg_add))
} else {
(None, None)
};
let (rope_bg, attn_out_add_bg) = if cfg.block_types[i] != BlockType::GatedConv {
let bg_rope = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.rope.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.q_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.k_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.rope_params.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: self.rope_freqs_buf.as_entire_binding(),
},
],
});
let bg_add = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.scaled_add_inplace.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.hidden_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.residual_add_params.as_entire_binding(),
},
],
});
(Some(bg_rope), Some(bg_add))
} else {
(None, None)
};
let layer = &mut self.layers[i];
layer.attn_norm_bg = Some(attn_norm_bg);
layer.ffn_norm_bg = Some(ffn_norm_bg);
layer.silu_bg = Some(silu_bg);
layer.ffn_add_bg = Some(ffn_add_bg);
layer.conv_fused_bg = conv_fused_bg;
layer.conv_add_bg = conv_add_bg;
layer.rope_bg = rope_bg;
layer.attn_out_add_bg = attn_out_add_bg;
if let GpuFfn::Dense(d) = &self.layers[i].ffn {
let gate_bg = self.make_gemv_bg(&d.gate, &self.ffn_input_buf, &self.gate_buf);
let up_bg = self.make_gemv_bg(&d.up, &self.ffn_input_buf, &self.up_buf);
let down_bg = self.make_gemv_bg(&d.down, &self.gate_buf, &self.out_buf);
let ffn_swiglu_bg = if d.gate.tensor.dtype == DType::Q4_0
&& d.up.tensor.dtype == DType::Q4_0
{
Some(
self.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("ffn_swiglu_q4_0"),
layout: &self.pipelines.ffn_swiglu_q4_0.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: d.gate.tensor.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: d.up.tensor.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.ffn_input_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: self.gate_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: d.gate.params_buf.as_entire_binding(),
},
],
}),
)
} else {
None
};
if let GpuFfn::Dense(d) = &mut self.layers[i].ffn {
d.gate.cached_bg = Some(gate_bg);
d.up.cached_bg = Some(up_bg);
d.down.cached_bg = Some(down_bg);
}
self.layers[i].ffn_swiglu_bg = ffn_swiglu_bg;
}
if cfg.block_types[i] == BlockType::GatedConv {
if let Some(ref w) = self.layers[i].conv_in_proj {
let bg = self.make_gemv_bg(w, &self.normed_buf, &self.conv_proj_buf);
self.layers[i].conv_in_proj.as_mut().unwrap().cached_bg = Some(bg);
}
if let Some(ref w) = self.layers[i].conv_out_proj {
let bg = self.make_gemv_bg(w, &self.conv_gate_buf, &self.out_buf);
self.layers[i].conv_out_proj.as_mut().unwrap().cached_bg = Some(bg);
}
} else {
let can_fuse_qkv = self.layers[i].attn_q.as_ref().is_some_and(|w| {
w.tensor.dtype == DType::Q4_0 && w.tensor.shape[0].is_multiple_of(4)
}) && self.layers[i].attn_k.as_ref().is_some_and(|w| {
w.tensor.dtype == DType::Q4_0 && w.tensor.shape[0].is_multiple_of(4)
}) && self.layers[i].attn_v.as_ref().is_some_and(|w| {
w.tensor.dtype == DType::Q4_0 && w.tensor.shape[0].is_multiple_of(4)
});
let (attn_qkv_bg, attn_qkv_params_buf) = if can_fuse_qkv {
let q_w = self.layers[i].attn_q.as_ref().unwrap();
let k_w = self.layers[i].attn_k.as_ref().unwrap();
let v_w = self.layers[i].attn_v.as_ref().unwrap();
let params = [
q_w.tensor.shape[0] as u32,
k_w.tensor.shape[0] as u32,
q_w.tensor.shape[1] as u32,
0u32,
];
use wgpu::util::DeviceExt;
let params_buf =
self.ctx
.device
.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("attn_qkv_params"),
contents: bytemuck::cast_slice(¶ms),
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_DST,
});
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("gemv_q4_0_qkv"),
layout: &self.pipelines.gemv_q4_0_qkv.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: q_w.tensor.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: k_w.tensor.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: v_w.tensor.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: self.normed_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: self.q_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 5,
resource: self.k_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 6,
resource: self.v_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 7,
resource: params_buf.as_entire_binding(),
},
],
});
(Some(bg), Some(params_buf))
} else {
(None, None)
};
self.layers[i].attn_qkv_bg = attn_qkv_bg;
self.layers[i].attn_qkv_params_buf = attn_qkv_params_buf;
if let Some(ref w) = self.layers[i].attn_q {
let bg = self.make_gemv_bg(w, &self.normed_buf, &self.q_buf);
self.layers[i].attn_q.as_mut().unwrap().cached_bg = Some(bg);
}
if let Some(ref w) = self.layers[i].attn_k {
let bg = self.make_gemv_bg(w, &self.normed_buf, &self.k_buf);
self.layers[i].attn_k.as_mut().unwrap().cached_bg = Some(bg);
}
if let Some(ref w) = self.layers[i].attn_v {
let bg = self.make_gemv_bg(w, &self.normed_buf, &self.v_buf);
self.layers[i].attn_v.as_mut().unwrap().cached_bg = Some(bg);
}
if let Some(ref w) = self.layers[i].attn_output {
let bg = self.make_gemv_bg(w, &self.attn_out_buf, &self.out_buf);
self.layers[i].attn_output.as_mut().unwrap().cached_bg = Some(bg);
}
if let Some(kv) = self.f32_kv().get(i).and_then(|opt| opt.as_ref()) {
let (k_cache, v_cache) = kv;
let attn_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("flash_attention"),
layout: &self.pipelines.flash_attention.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.q_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: k_cache.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: v_cache.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: self.attn_out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: self.attn_params.as_entire_binding(),
},
],
});
self.layers[i].attn_bg = Some(attn_bg);
}
let per_head_norm_bg = |buf: &wgpu::Buffer, norm: &wgpu::Buffer| {
self.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.per_head_rmsnorm.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: norm.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: self.per_head_norm_params.as_entire_binding(),
},
],
})
};
self.layers[i].qn_bg = self.layers[i]
.attn_q_norm
.as_ref()
.map(|w| per_head_norm_bg(&self.q_buf, w));
self.layers[i].kn_bg = self.layers[i]
.attn_k_norm
.as_ref()
.map(|w| per_head_norm_bg(&self.k_buf, w));
let bias_bg = |buf: &wgpu::Buffer, bias: &wgpu::Buffer, params: &wgpu::Buffer| {
self.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.add_inplace.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: bias.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params.as_entire_binding(),
},
],
})
};
self.layers[i].qb_bg = self.layers[i]
.attn_q_bias
.as_ref()
.map(|b| bias_bg(&self.q_buf, b, &self.elementwise_qdim_params));
self.layers[i].kb_bg = self.layers[i]
.attn_k_bias
.as_ref()
.map(|b| bias_bg(&self.k_buf, b, &self.elementwise_kvdim_params));
self.layers[i].vb_bg = self.layers[i]
.attn_v_bias
.as_ref()
.map(|b| bias_bg(&self.v_buf, b, &self.elementwise_kvdim_params));
}
}
let lm_head_bg = match &self.lm_head {
LmHead::Quantized(w) => Some(self.make_gemv_bg(w, &self.hidden_buf, &self.logits_buf)),
LmHead::F16 { .. } => None,
};
if let (Some(bg), LmHead::Quantized(w)) = (lm_head_bg, &mut self.lm_head) {
w.cached_bg = Some(bg);
}
}
fn encode(
&self,
enc: &mut wgpu::CommandEncoder,
pipeline: &wgpu::ComputePipeline,
bind_group: &wgpu::BindGroup,
workgroups: (u32, u32, u32),
label: &str,
) {
{
let mut pass = self.ctx.begin_pass(enc, label);
pass.set_pipeline(pipeline);
pass.set_bind_group(0, bind_group, &[]);
pass.dispatch_workgroups(workgroups.0, workgroups.1, workgroups.2);
}
}
fn dispatch_into(
&self,
pass: &mut wgpu::ComputePass<'_>,
pipeline: &wgpu::ComputePipeline,
bind_group: &wgpu::BindGroup,
workgroups: (u32, u32, u32),
) {
pass.set_pipeline(pipeline);
pass.set_bind_group(0, bind_group, &[]);
pass.dispatch_workgroups(workgroups.0, workgroups.1, workgroups.2);
}
fn submit_and_wait(&self, enc: wgpu::CommandEncoder) {
self.ctx.submit_encoder(enc);
self.ctx.device.poll_wait();
}
fn new_encoder(&self) -> wgpu::CommandEncoder {
self.ctx.device.create_command_encoder(&Default::default())
}
#[allow(dead_code)]
fn encode_gemv_weight(
&self,
enc: &mut wgpu::CommandEncoder,
w: &GpuWeight,
input: &wgpu::Buffer,
output: &wgpu::Buffer,
) {
let (pipeline, _, label) = self.gemv_pipeline_rows_label(w);
let fresh_bg;
let bg = if let Some(ref cached) = w.cached_bg {
cached
} else {
fresh_bg = self.make_gemv_bg(w, input, output);
&fresh_bg
};
self.encode(enc, pipeline, bg, self.gemv_workgroups(w), label);
}
fn moe_bind_group(
&self,
pipeline: &wgpu::ComputePipeline,
label: &str,
buffers: &[&wgpu::Buffer],
) -> wgpu::BindGroup {
let entries: Vec<wgpu::BindGroupEntry<'_>> = buffers
.iter()
.enumerate()
.map(|(i, b)| wgpu::BindGroupEntry {
binding: i as u32,
resource: b.as_entire_binding(),
})
.collect();
self.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some(label),
layout: &pipeline.get_bind_group_layout(0),
entries: &entries,
})
}
#[allow(clippy::too_many_arguments)]
fn moe_gemv_step(
&self,
w: &GpuMoeWeight,
sel_expert: &wgpu::Buffer,
x: &wgpu::Buffer,
y: &wgpu::Buffer,
n_used: u32,
n_entries: u32,
x_by_entry: bool,
label: &'static str,
) -> MoeStep<'_> {
let params = self.ctx.upload_storage(
bytemuck::cast_slice(&[
w.m,
w.k,
n_used,
n_entries,
w.expert_stride,
u32::from(x_by_entry),
0,
0,
]),
label,
);
MoeStep {
pipeline: &self.pipelines.moe_gemv_q4_0,
bind_group: self.moe_bind_group(
&self.pipelines.moe_gemv_q4_0,
label,
&[&w.buffer, x, y, sel_expert, ¶ms],
),
workgroups: (w.m, n_entries, 1),
}
}
fn moe_ffn_steps(
&self,
moe: &GpuMoeFfn,
x: &wgpu::Buffer,
out: &wgpu::Buffer,
n: u32,
accumulate: bool,
) -> Vec<MoeStep<'_>> {
let scratch = &moe.scratch;
let hs = self.config.hidden_size as u32;
let ff = scratch.expert_ff_len;
let n_used = scratch.n_expert_used;
let entries = n * n_used;
debug_assert!(
entries <= scratch.max_entries,
"routed FFN asked for {n} tokens ({entries} entries) against scratch sized for {}; \
every per-entry buffer below would be indexed past its end",
scratch.max_entries,
);
let router_params = self.ctx.upload_storage(
bytemuck::cast_slice(&[n, scratch.n_expert, hs, 0u32]),
"moe_router_params",
);
let route_params = self.ctx.upload_storage(
bytemuck::cast_slice(&[scratch.n_expert, n_used, n, 0u32]),
"moe_route_params",
);
let silu_total = entries * ff;
let silu_params = self
.ctx
.upload_storage(bytemuck::cast_slice(&[silu_total, 0u32]), "moe_silu_params");
let combine_params = self.ctx.upload_storage(
bytemuck::cast_slice(&[hs, n_used, n, u32::from(accumulate)]),
"moe_combine_params",
);
vec![
MoeStep {
pipeline: &self.pipelines.gemm_f32_nt,
bind_group: self.moe_bind_group(
&self.pipelines.gemm_f32_nt,
"moe_router",
&[x, &moe.router, &scratch.logits, &router_params],
),
workgroups: crate::backend::wgpu::gemv_row_workgroups(n * scratch.n_expert),
},
MoeStep {
pipeline: &self.pipelines.moe_route,
bind_group: self.moe_bind_group(
&self.pipelines.moe_route,
"moe_route",
&[
&scratch.logits,
&moe.bias,
&scratch.sel_expert,
&scratch.sel_weight,
&route_params,
],
),
workgroups: (n, 1, 1),
},
self.moe_gemv_step(
&moe.gate,
&scratch.sel_expert,
x,
&scratch.gate,
n_used,
entries,
false,
"moe_gate",
),
self.moe_gemv_step(
&moe.up,
&scratch.sel_expert,
x,
&scratch.up,
n_used,
entries,
false,
"moe_up",
),
MoeStep {
pipeline: &self.pipelines.silu_mul_inplace,
bind_group: self.moe_bind_group(
&self.pipelines.silu_mul_inplace,
"moe_silu",
&[&scratch.gate, &scratch.up, &silu_params],
),
workgroups: (silu_total.div_ceil(256), 1, 1),
},
self.moe_gemv_step(
&moe.down,
&scratch.sel_expert,
&scratch.gate,
&scratch.z,
n_used,
entries,
true,
"moe_down",
),
MoeStep {
pipeline: &self.pipelines.moe_combine,
bind_group: self.moe_bind_group(
&self.pipelines.moe_combine,
"moe_combine",
&[&scratch.z, &scratch.sel_weight, out, &combine_params],
),
workgroups: (hs.div_ceil(256), n, 1),
},
]
}
fn encode_lm_head(
&self,
enc: &mut wgpu::CommandEncoder,
input: &wgpu::Buffer,
output: &wgpu::Buffer,
) {
match &self.lm_head {
LmHead::Quantized(w) => {
let bg_tmp;
let bg = match w.cached_bg.as_ref() {
Some(bg) => bg,
None => {
bg_tmp = self.make_gemv_bg(w, input, output);
&bg_tmp
}
};
let mut pass = self.ctx.begin_pass(enc, "lm_head");
self.dispatch_gemv_into(&mut pass, w, bg);
}
LmHead::F16 { weight, params } => {
self.encode_gemv_f16(enc, weight, params, input, output)
}
}
}
fn encode_gemv_f16(
&self,
enc: &mut wgpu::CommandEncoder,
weight: &wgpu::Buffer,
params: &wgpu::Buffer,
input: &wgpu::Buffer,
output: &wgpu::Buffer,
) {
let m = self.config.vocab_size as u32;
let k = self.config.hidden_size as u32;
let weight_bytes = (u64::from(m) * u64::from(k) * 2).div_ceil(4) * 4;
let max_binding = self.ctx.max_storage_buffer_binding_size;
if weight_bytes > max_binding {
self.encode_gemv_f16_tiled(enc, weight, input, output, m, k);
return;
}
let params_buf = params;
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.gemv_f16.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: weight.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: input.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: output.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let groups = m.div_ceil(8);
self.encode(
enc,
&self.pipelines.gemv_f16,
&bg,
crate::backend::wgpu::gemv_row_workgroups(groups),
"gemv_f16",
);
}
fn encode_gemv_f16_tiled(
&self,
enc: &mut wgpu::CommandEncoder,
weight: &wgpu::Buffer,
input: &wgpu::Buffer,
output: &wgpu::Buffer,
m: u32,
k: u32,
) {
let row_bytes = u64::from(k) * 2;
let max_binding = self.ctx.max_storage_buffer_binding_size;
let tile_rows = gemv_tile_rows(
m,
k,
max_binding,
self.ctx.min_storage_buffer_offset_alignment,
2, );
let layout = self.pipelines.gemv_f16.get_bind_group_layout(0);
let mut row_start = 0u32;
let mut tile_idx = 0usize;
while row_start < m {
let rows = (m - row_start).min(tile_rows);
let weight_offset = u64::from(row_start) * row_bytes;
let Some(params_buf) = self.gemv_tile_params.get(tile_idx) else {
tracing::error!(
"tile_idx {tile_idx} exceeds preallocated LM-head GEMV tile params count"
);
break;
};
self.ctx.queue.write_buffer(
params_buf,
0,
bytemuck::cast_slice(&[rows, k, row_start, 0u32]),
);
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &layout,
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: wgpu::BindingResource::Buffer(wgpu::BufferBinding {
buffer: weight,
offset: weight_offset,
size: wgpu::BufferSize::new(
(u64::from(rows) * row_bytes).div_ceil(4) * 4,
),
}),
},
wgpu::BindGroupEntry {
binding: 1,
resource: input.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: output.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: params_buf.as_entire_binding(),
},
],
});
let groups = rows.div_ceil(8);
self.encode(
enc,
&self.pipelines.gemv_f16,
&bg,
crate::backend::wgpu::gemv_row_workgroups(groups),
"gemv_f16_tiled",
);
row_start += rows;
tile_idx += 1;
}
}
fn encode_rmsnorm(
&self,
enc: &mut wgpu::CommandEncoder,
x: &wgpu::Buffer,
weight: &wgpu::Buffer,
_n: u32,
_eps: f32,
) {
let params_buf = &self.rmsnorm_hs_params;
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.rmsnorm.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: x.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: weight.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params_buf.as_entire_binding(),
},
],
});
self.encode(enc, &self.pipelines.rmsnorm, &bg, (1, 1, 1), "rmsnorm");
}
fn encode_copy(
enc: &mut wgpu::CommandEncoder,
src: &wgpu::Buffer,
src_off_floats: u64,
dst: &wgpu::Buffer,
dst_off_floats: u64,
len_floats: u64,
) {
let f32_bytes = std::mem::size_of::<f32>() as u64;
enc.copy_buffer_to_buffer(
src,
src_off_floats * f32_bytes,
dst,
dst_off_floats * f32_bytes,
len_floats * f32_bytes,
);
}
fn encode_kv_shift_layers(&self, n_keep: usize, shift: usize, retained: usize) {
debug_assert!(retained > 0, "encode_kv_shift_layers requires retained > 0");
let cfg = &self.config;
let head_dim = cfg.head_dim;
let freq_base_bits = cfg.rope_theta.to_bits();
let mut enc = self
.ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("kv_shift"),
});
for layer_idx in 0..cfg.n_layers {
if cfg.block_types[layer_idx] != BlockType::Attention {
continue;
}
let n_kv_heads = cfg.kv_heads_per_layer[layer_idx];
let kv_dim = n_kv_heads * head_dim;
let (k_cache, v_cache) = self.f32_kv()[layer_idx]
.as_ref()
.expect("attention layer missing GPU kv_caches entry");
let params = KvShiftParams {
n_keep: n_keep as u32,
shift: shift as u32,
retained: retained as u32,
n_kv_heads: n_kv_heads as u32,
head_dim: head_dim as u32,
freq_base_bits,
rope_type: self.rope_type as u32,
has_freq_factors: u32::from(self.has_freq_factors),
};
let params_buf = self.ctx.upload_storage(
bytemuck::cast_slice(¶ms.to_u32_array()),
"kv_shift_params",
);
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("kv_shift"),
layout: &self.pipelines.kv_shift.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: k_cache.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.kv_shift_scratch.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: self.rope_freqs_buf.as_entire_binding(),
},
],
});
self.encode(
&mut enc,
&self.pipelines.kv_shift,
&bg,
params.dispatch_dims(),
"kv_shift",
);
let n_floats = (retained * kv_dim) as u64;
Self::encode_copy(
&mut enc,
&self.kv_shift_scratch,
0,
k_cache,
(n_keep * kv_dim) as u64,
n_floats,
);
Self::encode_copy(
&mut enc,
v_cache,
((n_keep + shift) * kv_dim) as u64,
&self.kv_shift_scratch,
0,
n_floats,
);
Self::encode_copy(
&mut enc,
&self.kv_shift_scratch,
0,
v_cache,
(n_keep * kv_dim) as u64,
n_floats,
);
}
self.submit_and_wait(enc);
}
}
impl GpuLfm2Model {
fn forward_inner(&self, tokens: &[u32], pos: usize, state: &mut InferenceState) -> Vec<f32> {
self.forward_inner_compute(tokens, pos, state);
self.ctx
.download_f32(&self.logits_buf, self.config.vocab_size)
}
fn hs_scratch(&self) -> &HsScratch {
self.hs_scratch.get_or_init(|| {
let cfg = &self.config;
let head_dim = cfg.head_dim;
let hs = cfg.hidden_size;
let d_conv = cfg.conv_kernel_size.unwrap_or(3) - 1;
let max_seq_len = self.gpu_state.max_seq_len;
let f = |size: usize, name: &str| self.ctx.create_storage_rw((size * 4) as u64, name);
let mut kv = Vec::with_capacity(cfg.n_layers);
let mut conv = Vec::with_capacity(cfg.n_layers);
for i in 0..cfg.n_layers {
if cfg.block_types[i] == BlockType::Attention {
let kv_dim = cfg.kv_heads_per_layer[i] * head_dim;
let bytes = kv_slab_bytes(max_seq_len, kv_dim);
let k = self.ctx.create_storage_rw(bytes, &format!("hs.l{i}.k"));
let v = self.ctx.create_storage_rw(bytes, &format!("hs.l{i}.v"));
kv.push(Some((k, v)));
conv.push(None);
} else {
kv.push(None);
conv.push(Some(f(d_conv * hs, &format!("hs.l{i}.conv"))));
}
}
HsScratch { kv, conv }
})
}
fn f32_kv(&self) -> &Vec<Option<(wgpu::Buffer, wgpu::Buffer)>> {
assert!(
self.tq.get().is_none(),
"f32 KV cache requested while TurboQuant is active"
);
self.gpu_state.kv_caches.get_or_init(|| {
let cfg = &self.config;
let head_dim = cfg.head_dim;
let max_seq_len = self.gpu_state.max_seq_len;
let mut kv = Vec::with_capacity(cfg.n_layers);
for i in 0..cfg.n_layers {
if cfg.block_types[i] == BlockType::Attention {
let kv_dim = cfg.kv_heads_per_layer[i] * head_dim;
let bytes = kv_slab_bytes(max_seq_len, kv_dim);
kv.push(Some((
self.ctx.create_storage_rw(bytes, &format!("l{i}.k_cache")),
self.ctx.create_storage_rw(bytes, &format!("l{i}.v_cache")),
)));
} else {
kv.push(None);
}
}
kv
})
}
#[inline]
fn active_kv(&self, i: usize) -> &(wgpu::Buffer, wgpu::Buffer) {
let caches = if self.use_hs_scratch.load(Ordering::Relaxed) {
&self
.hs_scratch
.get()
.expect("hs_scratch built before use_hs_scratch is set")
.kv
} else {
self.f32_kv()
};
caches[i].as_ref().unwrap()
}
fn cache_namespace(&self) -> String {
let tag = self.kv_cache_tag.get().map(String::as_str).unwrap_or("");
format!("wgpu:{tag}{}", self.model_id)
}
#[inline]
fn tq_cache(&self) -> Option<&TqGpuCache> {
if self.use_hs_scratch.load(Ordering::Relaxed) {
return None;
}
self.tq.get()
}
#[inline]
fn active_conv(&self, i: usize) -> &wgpu::Buffer {
let bufs = if self.use_hs_scratch.load(Ordering::Relaxed) {
&self
.hs_scratch
.get()
.expect("hs_scratch built before use_hs_scratch is set")
.conv
} else {
&self.gpu_state.conv_buffers
};
bufs[i].as_ref().unwrap()
}
fn forward_inner_compute(&self, tokens: &[u32], pos: usize, state: &mut InferenceState) {
self.forward_inner_compute_tail(tokens, pos, state, DecodeTail::Logits(TailArgmax::None));
}
fn forward_inner_compute_from_embedding(
&self,
embedding: &[f32],
pos: usize,
state: &mut InferenceState,
) {
self.forward_inner_compute_tail_seeded(
HiddenSeed::Embedding(embedding),
pos,
state,
DecodeTail::Logits(TailArgmax::None),
);
}
pub fn seed_embeddings(
&self,
embeddings: &[f32],
n_tokens: usize,
start_pos: usize,
state: &mut InferenceState,
) {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.seed_embeddings_locked(embeddings, n_tokens, start_pos, state);
}
fn seed_embeddings_locked(
&self,
embeddings: &[f32],
n_tokens: usize,
start_pos: usize,
state: &mut InferenceState,
) {
let hidden_size = self.config.hidden_size;
assert!(n_tokens > 0, "seed_embeddings requires at least one frame");
assert_eq!(
embeddings.len(),
n_tokens * hidden_size,
"embeddings.len() ({}) != n_tokens ({}) * hidden_size ({})",
embeddings.len(),
n_tokens,
hidden_size
);
if start_pos == 0 {
self.gpu_state.seq_len.store(0, Ordering::Relaxed);
self.zero_conv_buffers_locked();
}
for i in 0..n_tokens {
let frame = &embeddings[i * hidden_size..(i + 1) * hidden_size];
let pos = state.seq_len;
self.forward_inner_compute_from_embedding(frame, pos, state);
}
}
fn forward_inner_compute_tail(
&self,
tokens: &[u32],
pos: usize,
state: &mut InferenceState,
tail: DecodeTail,
) -> Option<wgpu::CommandEncoder> {
assert_eq!(tokens.len(), 1, "GPU forward expects single token");
self.forward_inner_compute_tail_seeded(HiddenSeed::Token(tokens[0]), pos, state, tail)
}
fn forward_inner_compute_tail_seeded(
&self,
seed: HiddenSeed<'_>,
pos: usize,
state: &mut InferenceState,
tail: DecodeTail,
) -> Option<wgpu::CommandEncoder> {
let cfg = &self.config;
let hs = cfg.hidden_size;
let hs32 = hs as u32;
self.ctx.reset_profiler();
assert!(
self.gpu_state.seq_len.load(Ordering::Relaxed) < self.gpu_state.max_seq_len,
"GPU seq_len {} exceeds max_seq_len {}",
self.gpu_state.seq_len.load(Ordering::Relaxed),
self.gpu_state.max_seq_len,
);
match seed {
HiddenSeed::Token(token) => {
let emb_offset = token as usize * hs;
self.ctx.queue.write_buffer(
&self.hidden_buf,
0,
bytemuck::cast_slice(
&self.gpu_state.embedding_f32[emb_offset..emb_offset + hs],
),
);
}
HiddenSeed::Embedding(embedding) => {
assert_eq!(
embedding.len(),
hs,
"GPU forward_from_embedding expects one hidden-size vector"
);
self.ctx
.queue
.write_buffer(&self.hidden_buf, 0, bytemuck::cast_slice(embedding));
}
}
let lora = self
.active_lora
.lock()
.unwrap_or_else(|e| e.into_inner())
.clone();
let head_dim = cfg.head_dim as u32;
let n_heads = cfg.n_heads as u32;
let n_kv_heads = cfg
.kv_heads_per_layer
.iter()
.copied()
.find(|&h| h > 0)
.unwrap_or(cfg.n_kv_heads) as u32;
let rope_data: [u32; 7] = [
pos as u32,
n_heads,
n_kv_heads,
head_dim,
cfg.rope_theta.to_bits(),
self.rope_type as u32,
self.has_freq_factors as u32,
];
self.ctx
.queue
.write_buffer(&self.rope_params, 0, bytemuck::cast_slice(&rope_data));
let seq_len = self.gpu_state.seq_len.load(Ordering::Relaxed);
let scale = self
.scalars
.attn
.unwrap_or_else(|| 1.0 / (head_dim as f32).sqrt());
let kv_dim = n_kv_heads * head_dim;
let attn_params: [u32; 8] = [
n_heads,
n_kv_heads,
head_dim,
kv_dim,
(seq_len + 1) as u32,
scale.to_bits(),
0,
0,
];
self.ctx
.queue
.write_buffer(&self.attn_params, 0, bytemuck::cast_slice(&attn_params));
if let Some(tq) = self.tq_cache() {
tq.write_params(
&self.ctx,
cfg,
1,
self.gpu_state.seq_len.load(Ordering::Relaxed),
scale,
);
}
let mut enc = self.new_encoder();
for i in 0..cfg.n_layers {
let lw = &self.layers[i];
if cfg.block_types[i] == BlockType::GatedConv {
let kernel_size = cfg.conv_kernel_size.unwrap_or(3) as u32;
let _d_conv = kernel_size - 1;
let norm_bg = lw.attn_norm_bg.as_ref().unwrap();
let in_w = lw.conv_in_proj.as_ref().unwrap();
let in_bg_tmp;
let in_bg = match in_w.cached_bg.as_ref() {
Some(b) => b,
None => {
in_bg_tmp = self.make_gemv_bg(in_w, &self.normed_buf, &self.conv_proj_buf);
&in_bg_tmp
}
};
let in_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::ShortconvInProj);
let in_lora_bgs = in_lora.map(|t| {
(
t,
self.lora_target_bgs(t, &self.normed_buf, &self.conv_proj_buf),
)
});
Self::encode_copy(
&mut enc,
&self.hidden_buf,
0,
&self.normed_buf,
0,
hs as u64,
);
let conv_fused_bg = lw.conv_fused_bg.as_ref().unwrap();
let out_w = lw.conv_out_proj.as_ref().unwrap();
let out_bg_tmp;
let out_bg = match out_w.cached_bg.as_ref() {
Some(b) => b,
None => {
out_bg_tmp = self.make_gemv_bg(out_w, &self.conv_gate_buf, &self.out_buf);
&out_bg_tmp
}
};
let out_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::ShortconvOutProj);
let out_lora_bgs = out_lora.map(|t| {
(
t,
self.lora_target_bgs(t, &self.conv_gate_buf, &self.out_buf),
)
});
let add_bg = lw.conv_add_bg.as_ref().unwrap();
{
let mut pass = self.ctx.begin_pass(&mut enc, "conv");
self.dispatch_into(&mut pass, &self.pipelines.rmsnorm, norm_bg, (1, 1, 1));
self.dispatch_gemv_into(&mut pass, in_w, in_bg);
if let Some((t, (bg_a, bg_b))) = &in_lora_bgs {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
self.dispatch_into(
&mut pass,
&self.pipelines.conv1d_fused,
conv_fused_bg,
(hs32.div_ceil(256), 1, 1),
);
self.dispatch_gemv_into(&mut pass, out_w, out_bg);
if let Some((t, (bg_a, bg_b))) = &out_lora_bgs {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
self.dispatch_into(
&mut pass,
&self.pipelines.add_inplace,
add_bg,
(hs32.div_ceil(256), 1, 1),
);
}
} else {
let q_dim = n_heads * head_dim;
let norm_bg = lw.attn_norm_bg.as_ref().unwrap();
let q_w = lw.attn_q.as_ref().unwrap();
let q_bg_tmp;
let q_bg = match q_w.cached_bg.as_ref() {
Some(b) => b,
None => {
q_bg_tmp = self.make_gemv_bg(q_w, &self.normed_buf, &self.q_buf);
&q_bg_tmp
}
};
let k_w = lw.attn_k.as_ref().unwrap();
let k_bg_tmp;
let k_bg = match k_w.cached_bg.as_ref() {
Some(b) => b,
None => {
k_bg_tmp = self.make_gemv_bg(k_w, &self.normed_buf, &self.k_buf);
&k_bg_tmp
}
};
let v_w = lw.attn_v.as_ref().unwrap();
let v_bg_tmp;
let v_bg = match v_w.cached_bg.as_ref() {
Some(b) => b,
None => {
v_bg_tmp = self.make_gemv_bg(v_w, &self.normed_buf, &self.v_buf);
&v_bg_tmp
}
};
let rope_bg = lw.rope_bg.as_ref().unwrap();
let max_pairs = std::cmp::max(n_heads, n_kv_heads) * (head_dim / 2);
let q_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::AttnQ);
let k_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::AttnK);
let v_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::AttnV);
let q_lora_bgs =
q_lora.map(|t| (t, self.lora_target_bgs(t, &self.normed_buf, &self.q_buf)));
let k_lora_bgs =
k_lora.map(|t| (t, self.lora_target_bgs(t, &self.normed_buf, &self.k_buf)));
let v_lora_bgs =
v_lora.map(|t| (t, self.lora_target_bgs(t, &self.normed_buf, &self.v_buf)));
Self::encode_copy(
&mut enc,
&self.hidden_buf,
0,
&self.normed_buf,
0,
hs as u64,
);
{
let mut pass = self.ctx.begin_pass(&mut enc, "attn_pre");
self.dispatch_into(&mut pass, &self.pipelines.rmsnorm, norm_bg, (1, 1, 1));
if let Some(qkv_bg) = lw.attn_qkv_bg.as_ref() {
let total_rows = q_dim + 2 * kv_dim;
self.dispatch_into(
&mut pass,
&self.pipelines.gemv_q4_0_qkv,
qkv_bg,
(total_rows.div_ceil(4), 1, 1),
);
} else {
self.dispatch_gemv_into(&mut pass, q_w, q_bg);
self.dispatch_gemv_into(&mut pass, k_w, k_bg);
self.dispatch_gemv_into(&mut pass, v_w, v_bg);
}
if let Some((t, (bg_a, bg_b))) = q_lora_bgs.as_ref() {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
if let Some((t, (bg_a, bg_b))) = k_lora_bgs.as_ref() {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
if let Some((t, (bg_a, bg_b))) = v_lora_bgs.as_ref() {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
if let Some(bg) = lw.qb_bg.as_ref() {
self.dispatch_into(
&mut pass,
&self.pipelines.add_inplace,
bg,
(q_dim.div_ceil(256), 1, 1),
);
}
if let Some(bg) = lw.kb_bg.as_ref() {
self.dispatch_into(
&mut pass,
&self.pipelines.add_inplace,
bg,
(kv_dim.div_ceil(256), 1, 1),
);
}
if let Some(bg) = lw.vb_bg.as_ref() {
self.dispatch_into(
&mut pass,
&self.pipelines.add_inplace,
bg,
(kv_dim.div_ceil(256), 1, 1),
);
}
if let Some(bg) = lw.qn_bg.as_ref() {
self.dispatch_into(
&mut pass,
&self.pipelines.per_head_rmsnorm,
bg,
(n_heads, 1, 1),
);
}
if let Some(bg) = lw.kn_bg.as_ref() {
self.dispatch_into(
&mut pass,
&self.pipelines.per_head_rmsnorm,
bg,
(n_kv_heads, 1, 1),
);
}
self.dispatch_into(
&mut pass,
&self.pipelines.rope,
rope_bg,
(max_pairs.div_ceil(256), 1, 1),
);
}
let seq_len = self.gpu_state.seq_len.load(Ordering::Relaxed);
if let Some(tq) = self.tq_cache() {
tq.encode_kv(&self.ctx, &mut enc, i, &self.k_buf, &self.v_buf, 1);
tq.rotate_queries(&self.ctx, &mut enc, i, &self.q_buf, 1, n_heads as usize);
tq.attention(
&self.ctx,
&mut enc,
i,
&self.attn_out_buf,
1,
n_heads as usize,
);
} else {
let (k_cache, v_cache) = self.active_kv(i);
let kv_offset_floats = (seq_len * kv_dim as usize) as u64;
Self::encode_copy(
&mut enc,
&self.k_buf,
0,
k_cache,
kv_offset_floats,
kv_dim as u64,
);
Self::encode_copy(
&mut enc,
&self.v_buf,
0,
v_cache,
kv_offset_floats,
kv_dim as u64,
);
}
let out_w = lw.attn_output.as_ref().unwrap();
let out_bg_tmp;
let out_bg = match out_w.cached_bg.as_ref() {
Some(b) => b,
None => {
out_bg_tmp = self.make_gemv_bg(out_w, &self.attn_out_buf, &self.out_buf);
&out_bg_tmp
}
};
let add_bg = lw.attn_out_add_bg.as_ref().unwrap();
let o_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::AttnOutput);
let o_lora_bgs = o_lora.map(|t| {
(
t,
self.lora_target_bgs(t, &self.attn_out_buf, &self.hidden_buf),
)
});
{
let mut pass = self.ctx.begin_pass(&mut enc, "attn_post");
if self.tq_cache().is_none() {
let attn_bg_tmp;
let attn_bg = if self.use_hs_scratch.load(Ordering::Relaxed) {
let (k_buf, v_buf) = self.active_kv(i);
attn_bg_tmp =
self.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("flash_attention_hs_bg"),
layout: &self
.pipelines
.flash_attention
.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.q_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: k_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: v_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: self.attn_out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: self.attn_params.as_entire_binding(),
},
],
});
&attn_bg_tmp
} else {
lw.attn_bg.as_ref().unwrap()
};
self.dispatch_into(
&mut pass,
&self.pipelines.flash_attention,
attn_bg,
(n_heads, 1, 1),
);
}
self.dispatch_gemv_into(&mut pass, out_w, out_bg);
self.dispatch_into(
&mut pass,
&self.pipelines.scaled_add_inplace,
add_bg,
(hs32.div_ceil(256), 1, 1),
);
if let Some((t, (bg_a, bg_b))) = o_lora_bgs.as_ref() {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
}
}
Self::encode_copy(
&mut enc,
&self.hidden_buf,
0,
&self.ffn_input_buf,
0,
hs as u64,
);
let norm_bg = lw.ffn_norm_bg.as_ref().unwrap();
let dense = match &lw.ffn {
GpuFfn::Moe(moe) => {
let steps =
self.moe_ffn_steps(moe, &self.ffn_input_buf, &self.hidden_buf, 1, true);
{
let mut pass = self.ctx.begin_pass(&mut enc, "ffn_moe");
self.dispatch_into(&mut pass, &self.pipelines.rmsnorm, norm_bg, (1, 1, 1));
steps.iter().for_each(|s| {
self.dispatch_into(&mut pass, s.pipeline, &s.bind_group, s.workgroups);
});
}
continue;
}
GpuFfn::Dense(d) => d,
};
let gate_bg_tmp;
let gate_bg = match dense.gate.cached_bg.as_ref() {
Some(bg) => bg,
None => {
gate_bg_tmp =
self.make_gemv_bg(&dense.gate, &self.ffn_input_buf, &self.gate_buf);
&gate_bg_tmp
}
};
let up_bg_tmp;
let up_bg = match dense.up.cached_bg.as_ref() {
Some(bg) => bg,
None => {
up_bg_tmp = self.make_gemv_bg(&dense.up, &self.ffn_input_buf, &self.up_buf);
&up_bg_tmp
}
};
let silu_bg = lw.silu_bg.as_ref().unwrap();
let down_bg_tmp;
let down_bg = match dense.down.cached_bg.as_ref() {
Some(bg) => bg,
None => {
down_bg_tmp = self.make_gemv_bg(&dense.down, &self.gate_buf, &self.out_buf);
&down_bg_tmp
}
};
let add_bg = lw.ffn_add_bg.as_ref().unwrap();
let gate_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::FfnGate);
let up_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::FfnUp);
let down_lora = Self::lora_target(lora.as_ref(), i, LoraTarget::FfnDown);
let gate_lora_bgs = gate_lora.map(|t| {
(
t,
self.lora_target_bgs(t, &self.ffn_input_buf, &self.gate_buf),
)
});
let up_lora_bgs = up_lora.map(|t| {
(
t,
self.lora_target_bgs(t, &self.ffn_input_buf, &self.up_buf),
)
});
let down_lora_bgs =
down_lora.map(|t| (t, self.lora_target_bgs(t, &self.gate_buf, &self.hidden_buf)));
{
let mut pass = self.ctx.begin_pass(&mut enc, "ffn");
self.dispatch_into(&mut pass, &self.pipelines.rmsnorm, norm_bg, (1, 1, 1));
self.dispatch_gemv_into(&mut pass, &dense.gate, gate_bg);
self.dispatch_gemv_into(&mut pass, &dense.up, up_bg);
if let Some((t, (bg_a, bg_b))) = gate_lora_bgs.as_ref() {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
if let Some((t, (bg_a, bg_b))) = up_lora_bgs.as_ref() {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
self.dispatch_into(
&mut pass,
&self.pipelines.silu_mul_inplace,
silu_bg,
((dense.gate.tensor.shape[0] as u32).div_ceil(256), 1, 1),
);
self.dispatch_gemv_into(&mut pass, &dense.down, down_bg);
self.dispatch_into(
&mut pass,
&self.pipelines.scaled_add_inplace,
add_bg,
(hs32.div_ceil(256), 1, 1),
);
if let Some((t, (bg_a, bg_b))) = down_lora_bgs.as_ref() {
self.dispatch_lora_into(&mut pass, t, bg_a, bg_b);
}
}
}
self.encode_rmsnorm(
&mut enc,
&self.hidden_buf,
&self.output_norm,
hs32,
cfg.rms_norm_eps,
);
match tail {
DecodeTail::Hidden => {
self.submit_and_wait(enc);
self.gpu_state.seq_len.fetch_add(1, Ordering::Relaxed);
state.seq_len += 1;
self.ctx.finish_profiler();
None
}
DecodeTail::HiddenUnsubmitted => {
self.gpu_state.seq_len.fetch_add(1, Ordering::Relaxed);
state.seq_len += 1;
self.ctx.finish_profiler();
Some(enc)
}
DecodeTail::Logits(argmax) | DecodeTail::LogitsUnsubmitted(argmax) => {
self.encode_lm_head(&mut enc, &self.hidden_buf, &self.logits_buf);
if let Some(params) = self.logit_scale_params.as_ref() {
let scale_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("logit_scale_bg"),
layout: &self.pipelines.scale_f32.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.logits_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: params.as_entire_binding(),
},
],
});
let mut pass = self.ctx.begin_pass(&mut enc, "logit_scale");
self.dispatch_into(
&mut pass,
&self.pipelines.scale_f32,
&scale_bg,
((cfg.vocab_size as u32).div_ceil(256), 1, 1),
);
drop(pass);
}
if argmax != TailArgmax::None {
self.encode_argmax_pass(&mut enc);
}
if argmax == TailArgmax::DispatchAndStage {
enc.copy_buffer_to_buffer(
&self.argmax_out_buf,
0,
&self.argmax_readback_buf,
0,
4,
);
}
self.gpu_state.seq_len.fetch_add(1, Ordering::Relaxed);
state.seq_len += 1;
self.ctx.finish_profiler();
if matches!(tail, DecodeTail::LogitsUnsubmitted(_)) {
Some(enc)
} else {
self.submit_and_wait(enc);
None
}
}
}
}
fn encode_argmax_pass(&self, enc: &mut wgpu::CommandEncoder) {
let mut pass = self.ctx.begin_pass(enc, "argmax");
pass.set_pipeline(&self.pipelines.argmax_f32);
pass.set_bind_group(0, &self.argmax_bg, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
fn forward_greedy_inner(&self, tokens: &[u32], pos: usize, state: &mut InferenceState) -> u32 {
self.forward_inner_compute_tail(
tokens,
pos,
state,
DecodeTail::Logits(TailArgmax::DispatchAndStage),
);
self.ctx.read_mapped_u32(&self.argmax_readback_buf, 1)[0]
}
pub fn forward_prefill_step(&self, token: u32, pos: usize, state: &mut InferenceState) {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
if pos == 0 {
self.zero_conv_buffers_locked();
}
self.forward_inner_compute(&[token], pos, state);
}
pub async fn forward_greedy_async(
&self,
token: u32,
pos: usize,
state: &mut InferenceState,
) -> Result<u32> {
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
let enc = self
.forward_inner_compute_tail(
&[token],
pos,
state,
DecodeTail::LogitsUnsubmitted(TailArgmax::Dispatch),
)
.ok_or_else(|| anyhow::anyhow!("LogitsUnsubmitted returned None"))?;
Ok::<_, anyhow::Error>(self.ctx.begin_download_with_encoder(
enc,
&self.argmax_out_buf,
std::mem::size_of::<u32>() as u64,
))
}?;
let bytes = pending.recv().await?;
if bytes.len() < 4 {
anyhow::bail!(
"GPU argmax readback buffer truncated (expected 4 bytes, got {})",
bytes.len()
);
}
let token = bytemuck::pod_read_unaligned::<u32>(&bytes[..4]);
Ok(token)
}
pub async fn forward_logits_async(
&self,
token: u32,
pos: usize,
state: &mut InferenceState,
) -> Result<Vec<f32>> {
let vocab = self.config.vocab_size;
let expected_bytes = vocab * std::mem::size_of::<f32>();
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
let enc = self
.forward_inner_compute_tail(
&[token],
pos,
state,
DecodeTail::LogitsUnsubmitted(TailArgmax::None),
)
.ok_or_else(|| anyhow::anyhow!("LogitsUnsubmitted returned None"))?;
Ok::<_, anyhow::Error>(self.ctx.begin_download_with_encoder(
enc,
&self.logits_buf,
expected_bytes as u64,
))
}?;
let bytes = pending.recv().await?;
if bytes.len() < expected_bytes {
anyhow::bail!(
"GPU logits readback buffer truncated (expected {expected_bytes} bytes, got {})",
bytes.len()
);
}
let mut out = vec![0f32; vocab];
bytemuck::cast_slice_mut(&mut out).copy_from_slice(&bytes[..expected_bytes]);
Ok(out)
}
pub async fn forward_embedding_async(
&self,
token: u32,
pos: usize,
state: &mut InferenceState,
) -> Result<Vec<f32>> {
let hidden_size = self.config.hidden_size;
let expected_bytes = hidden_size * std::mem::size_of::<f32>();
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
let enc = self
.forward_inner_compute_tail(&[token], pos, state, DecodeTail::HiddenUnsubmitted)
.ok_or_else(|| anyhow::anyhow!("HiddenUnsubmitted returned None"))?;
Ok::<_, anyhow::Error>(self.ctx.begin_download_with_encoder(
enc,
&self.hidden_buf,
expected_bytes as u64,
))
}?;
let bytes = pending.recv().await?;
if bytes.len() < expected_bytes {
anyhow::bail!(
"GPU hidden readback buffer truncated (expected {expected_bytes} bytes, got {})",
bytes.len()
);
}
let mut out = vec![0f32; hidden_size];
bytemuck::cast_slice_mut(&mut out).copy_from_slice(&bytes[..expected_bytes]);
Ok(out)
}
pub async fn forward_hidden_from_embedding_async(
&self,
embedding: &[f32],
pos: usize,
state: &mut InferenceState,
) -> Result<Vec<f32>> {
let hidden_size = self.config.hidden_size;
let expected_bytes = hidden_size * std::mem::size_of::<f32>();
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
let enc = self
.forward_inner_compute_tail_seeded(
HiddenSeed::Embedding(embedding),
pos,
state,
DecodeTail::HiddenUnsubmitted,
)
.ok_or_else(|| anyhow::anyhow!("HiddenUnsubmitted returned None"))?;
Ok::<_, anyhow::Error>(self.ctx.begin_download_with_encoder(
enc,
&self.hidden_buf,
expected_bytes as u64,
))
}?;
let bytes = pending.recv().await?;
if bytes.len() < expected_bytes {
anyhow::bail!(
"GPU hidden readback buffer truncated (expected {expected_bytes} bytes, got {})",
bytes.len()
);
}
let mut out = vec![0f32; hidden_size];
bytemuck::cast_slice_mut(&mut out).copy_from_slice(&bytes[..expected_bytes]);
Ok(out)
}
pub fn forward_hidden_gpu(
&self,
token: u32,
pos: usize,
state: &mut InferenceState,
) -> Result<&wgpu::Buffer> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
self.forward_inner_compute_tail(&[token], pos, state, DecodeTail::Hidden);
Ok(&self.hidden_buf)
}
pub fn forward_hidden_from_embedding_gpu(
&self,
embedding: &[f32],
pos: usize,
state: &mut InferenceState,
) -> Result<&wgpu::Buffer> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
self.forward_inner_compute_tail_seeded(
HiddenSeed::Embedding(embedding),
pos,
state,
DecodeTail::Hidden,
);
Ok(&self.hidden_buf)
}
pub fn hidden_buffer(&self) -> &wgpu::Buffer {
&self.hidden_buf
}
pub async fn forward_logits_from_embedding_async(
&self,
embedding: &[f32],
pos: usize,
state: &mut InferenceState,
) -> Result<Vec<f32>> {
let vocab_size = self.config.vocab_size;
let expected_bytes = vocab_size * std::mem::size_of::<f32>();
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
let enc = self
.forward_inner_compute_tail_seeded(
HiddenSeed::Embedding(embedding),
pos,
state,
DecodeTail::LogitsUnsubmitted(TailArgmax::None),
)
.ok_or_else(|| anyhow::anyhow!("LogitsUnsubmitted returned None"))?;
Ok::<_, anyhow::Error>(self.ctx.begin_download_with_encoder(
enc,
&self.logits_buf,
expected_bytes as u64,
))
}?;
let bytes = pending.recv().await?;
if bytes.len() < expected_bytes {
anyhow::bail!(
"GPU logits readback buffer truncated (expected {expected_bytes} bytes, got {})",
bytes.len()
);
}
let mut out = vec![0f32; vocab_size];
bytemuck::cast_slice_mut(&mut out).copy_from_slice(&bytes[..expected_bytes]);
Ok(out)
}
pub async fn forward_greedy_from_embedding_async(
&self,
embedding: &[f32],
pos: usize,
state: &mut InferenceState,
) -> Result<u32> {
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
let enc = self
.forward_inner_compute_tail_seeded(
HiddenSeed::Embedding(embedding),
pos,
state,
DecodeTail::LogitsUnsubmitted(TailArgmax::Dispatch),
)
.ok_or_else(|| anyhow::anyhow!("LogitsUnsubmitted returned None"))?;
Ok::<_, anyhow::Error>(self.ctx.begin_download_with_encoder(
enc,
&self.argmax_out_buf,
std::mem::size_of::<u32>() as u64,
))
}?;
let bytes = pending.recv().await?;
if bytes.len() < 4 {
anyhow::bail!(
"GPU argmax readback buffer truncated (expected 4 bytes, got {})",
bytes.len()
);
}
let token = bytemuck::pod_read_unaligned::<u32>(&bytes[..4]);
Ok(token)
}
pub async fn lm_head_argmax_async(&self, hidden: &[f32]) -> Result<u32> {
let hs = self.config.hidden_size;
anyhow::ensure!(
hidden.len() == hs,
"dspark_hidden length ({}) != hidden_size ({})",
hidden.len(),
hs
);
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
self.ctx
.queue
.write_buffer(&self.hidden_buf, 0, bytemuck::cast_slice(hidden));
let mut enc = self
.ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("dspark_lm_head_argmax"),
});
self.encode_lm_head(&mut enc, &self.hidden_buf, &self.logits_buf);
if let Some(params) = self.logit_scale_params.as_ref() {
let scale_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("logit_scale_bg"),
layout: &self.pipelines.scale_f32.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.logits_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: params.as_entire_binding(),
},
],
});
let mut pass = self.ctx.begin_pass(&mut enc, "logit_scale");
self.dispatch_into(
&mut pass,
&self.pipelines.scale_f32,
&scale_bg,
((self.config.vocab_size as u32).div_ceil(256), 1, 1),
);
}
self.encode_argmax_pass(&mut enc);
self.ctx.begin_download_with_encoder(
enc,
&self.argmax_out_buf,
std::mem::size_of::<u32>() as u64,
)
};
let bytes = pending.recv().await?;
if bytes.len() < std::mem::size_of::<u32>() {
anyhow::bail!("short read ({}) for argmax token readback", bytes.len());
}
let token = bytemuck::pod_read_unaligned::<u32>(&bytes[..4]);
Ok(token)
}
pub fn gpu_info(&self) -> (&str, &str) {
(&self.ctx.adapter_name, &self.ctx.backend)
}
pub fn kv_mode_label(&self) -> String {
describe_kv_mode(self.kv_mode.get().unwrap_or(&None))
}
}
impl GpuLfm2Model {
fn unbatchable_matmul_weight(&self) -> Option<(usize, &'static str, DType)> {
for (li, lw) in self.layers.iter().enumerate() {
let (gate, up, down) = match &lw.ffn {
GpuFfn::Dense(d) => (Some(&d.gate), Some(&d.up), Some(&d.down)),
GpuFfn::Moe(_) => (None, None, None),
};
let weights: [(&'static str, Option<&GpuWeight>); 9] = [
("ffn_gate", gate),
("ffn_up", up),
("ffn_down", down),
("attn_q", lw.attn_q.as_ref()),
("attn_k", lw.attn_k.as_ref()),
("attn_v", lw.attn_v.as_ref()),
("attn_output", lw.attn_output.as_ref()),
("conv_in_proj", lw.conv_in_proj.as_ref()),
("conv_out_proj", lw.conv_out_proj.as_ref()),
];
for (name, w) in weights {
let Some(w) = w else { continue };
let dt = w.tensor.dtype;
if !matches!(
dt,
DType::Q4_0 | DType::Q8_0 | DType::Q4KM | DType::Q5KM | DType::Q6K | DType::F32
) {
return Some((li, name, dt));
}
}
}
None
}
fn encode_rmsnorm_batch(
&self,
enc: &mut wgpu::CommandEncoder,
src: &wgpu::Buffer,
dst: &wgpu::Buffer,
weight: &wgpu::Buffer,
n: u32,
hs: u32,
) {
let params: [u32; 5] = [
hs,
self.config.rms_norm_eps.to_bits(),
hs,
hs,
1.0f32.to_bits(),
];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.rmsnorm_batch.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: src.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: dst.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: weight.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
self.encode(
enc,
&self.pipelines.rmsnorm_batch,
&bg,
(n, 1, 1),
"rmsnorm_batch",
);
}
#[allow(clippy::too_many_arguments)]
fn encode_add_rmsnorm_batch(
&self,
enc: &mut wgpu::CommandEncoder,
src: &wgpu::Buffer,
dst: &wgpu::Buffer,
weight: &wgpu::Buffer,
residual: &wgpu::Buffer,
n: u32,
hs: u32,
) {
let params: [u32; 5] = [
hs,
self.config.rms_norm_eps.to_bits(),
hs,
hs,
self.scalars.residual.to_bits(),
];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.add_rmsnorm_batch.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: src.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: dst.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: weight.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: residual.as_entire_binding(),
},
],
});
self.encode(
enc,
&self.pipelines.add_rmsnorm_batch,
&bg,
(n, 1, 1),
"add_rmsnorm_batch",
);
}
#[allow(clippy::too_many_arguments)] fn encode_mul_mat_reg_tile(
&self,
enc: &mut wgpu::CommandEncoder,
w: &GpuWeight,
x: &wgpu::Buffer,
y: &wgpu::Buffer,
n: u32,
k: u32,
x_stride: u32,
y_stride: u32,
) {
debug_assert!(
matches!(
w.tensor.dtype,
DType::Q4_0 | DType::Q8_0 | DType::Q4KM | DType::Q5KM | DType::Q6K | DType::F32
),
"encode_mul_mat_reg_tile only supports Q4_0/Q8_0/Q4KM/Q5KM/Q6K/F32 weights"
);
let m = w.tensor.shape[0] as u32;
let (pipeline, label) = match w.tensor.dtype {
DType::Q4_0 => (&self.pipelines.mul_mat_reg_tile_q4_0, "mul_mat_tile"),
DType::Q8_0 => (&self.pipelines.mul_mat_reg_tile_q8_0, "mul_mat_q8_0"),
DType::Q4KM => (&self.pipelines.mul_mat_reg_tile_q4_k, "mul_mat_q4k"),
DType::Q5KM => (&self.pipelines.mul_mat_reg_tile_q5_k, "mul_mat_q5k"),
DType::Q6K => (&self.pipelines.mul_mat_reg_tile_q6_k, "mul_mat_q6k"),
DType::F32 => (&self.pipelines.mul_mat_reg_tile_f32, "mul_mat_f32"),
_ => unreachable!("batched prefill only supports Q4_0/Q8_0/Q4KM/Q5KM/Q6K/F32"),
};
let wg_m = m.div_ceil(MUL_MAT_TILE_WG_M * MUL_MAT_TILE_M);
let wg_n = n.div_ceil(MUL_MAT_TILE_WG_N * MUL_MAT_TILE_N);
let params: [u32; 5] = [m, k, n, x_stride, y_stride];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: w.tensor.buffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: x.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: y.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: p_buf.as_entire_binding(),
},
],
});
self.encode(enc, pipeline, &bg, (wg_m, wg_n, 1), label);
}
fn encode_bias_add_batch(
&self,
enc: &mut wgpu::CommandEncoder,
buf: &wgpu::Buffer,
bias: &wgpu::Buffer,
n: u32,
dim: u32,
) {
let total = n * dim;
let params: [u32; 2] = [total, dim];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.bias_add.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: bias.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: p_buf.as_entire_binding(),
},
],
});
self.encode(
enc,
&self.pipelines.bias_add,
&bg,
(total.div_ceil(256), 1, 1),
"bias_add_batch",
);
}
#[allow(clippy::too_many_arguments)]
fn encode_qk_norm_rope_batch(
&self,
enc: &mut wgpu::CommandEncoder,
q_batch: &wgpu::Buffer,
k_batch: &wgpu::Buffer,
q_norm_w: Option<&wgpu::Buffer>,
k_norm_w: Option<&wgpu::Buffer>,
start_pos: u32,
n: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
q_stride: u32,
k_stride: u32,
) {
debug_assert_eq!(
q_norm_w.is_some(),
k_norm_w.is_some(),
"QK-norm weights must be both present or both absent",
);
let has_qk_norm = q_norm_w.is_some() && k_norm_w.is_some();
let q_norm = q_norm_w.unwrap_or(&self.rope_freqs_buf);
let k_norm = k_norm_w.unwrap_or(&self.rope_freqs_buf);
let params: [u32; 12] = [
start_pos,
n,
n_heads,
n_kv_heads,
head_dim,
self.config.rms_norm_eps.to_bits(),
self.config.rope_theta.to_bits(),
self.rope_type as u32,
q_stride,
k_stride,
self.has_freq_factors as u32,
has_qk_norm as u32,
];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.qk_norm_rope_batch.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: q_batch.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: k_batch.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: q_norm.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: k_norm.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 5,
resource: self.rope_freqs_buf.as_entire_binding(),
},
],
});
let tg_count = n * (n_heads + n_kv_heads);
self.encode(
enc,
&self.pipelines.qk_norm_rope_batch,
&bg,
(tg_count, 1, 1),
"qk_norm_rope_batch",
);
}
#[allow(clippy::too_many_arguments)]
fn encode_conv1d_fused_batch(
&self,
enc: &mut wgpu::CommandEncoder,
proj: &wgpu::Buffer,
rbuffer: &wgpu::Buffer,
weight: &wgpu::Buffer,
output: &wgpu::Buffer,
n: u32,
hs: u32,
) {
let kernel_size = self.config.conv_kernel_size.unwrap_or(3) as u32;
let d_conv = kernel_size - 1;
let params: [u32; 6] = [hs, kernel_size, d_conv, n, 3 * hs, hs];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.conv1d_fused_batch.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: proj.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: rbuffer.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: weight.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 3,
resource: output.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
],
});
let groups = hs.div_ceil(256);
self.encode(
enc,
&self.pipelines.conv1d_fused_batch,
&bg,
(groups, 1, 1),
"conv1d_fused_batch",
);
}
#[allow(clippy::too_many_arguments)]
fn encode_attention_prefill(
&self,
enc: &mut wgpu::CommandEncoder,
q_batch: &wgpu::Buffer,
k_cache: &wgpu::Buffer,
v_cache: &wgpu::Buffer,
out_batch: &wgpu::Buffer,
n: u32,
n_heads: u32,
n_kv_heads: u32,
head_dim: u32,
kv_dim: u32,
max_seq: u32,
start_pos: u32,
q_stride: u32,
out_stride: u32,
scale: f32,
) {
if n == 0 {
return;
}
assert!(
head_dim <= 128,
"wgpu attention_prefill supports head_dim <= 128 (q_shared/acc are \
sized 128); got {head_dim}"
);
assert!(
n_kv_heads > 0 && n_heads.is_multiple_of(n_kv_heads),
"wgpu attention_prefill requires n_kv_heads > 0 and n_heads divisible \
by n_kv_heads; got n_heads={n_heads}, n_kv_heads={n_kv_heads}"
);
assert_eq!(
kv_dim,
n_kv_heads * head_dim,
"wgpu attention_prefill requires kv_dim == n_kv_heads * head_dim; got \
kv_dim={kv_dim}, n_kv_heads={n_kv_heads}, head_dim={head_dim}"
);
let kv_live_floats = u64::from(max_seq).saturating_mul(u64::from(kv_dim));
assert_f32_binding_fits(
kv_live_floats,
self.ctx.max_storage_buffer_binding_size,
"attention_prefill live KV",
);
let params: [u32; 12] = [
n_heads,
n_kv_heads,
head_dim,
kv_dim,
max_seq,
scale.to_bits(),
start_pos,
n,
q_stride,
out_stride,
0, 0,
];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.attention_prefill.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: q_batch.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: f32_binding(k_cache, kv_live_floats),
},
wgpu::BindGroupEntry {
binding: 2,
resource: f32_binding(v_cache, kv_live_floats),
},
wgpu::BindGroupEntry {
binding: 3,
resource: out_batch.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 4,
resource: p_buf.as_entire_binding(),
},
],
});
self.encode(
enc,
&self.pipelines.attention_prefill,
&bg,
(n_heads, n, 1),
"attention_prefill",
);
}
fn encode_prefill_batched_locked(
&self,
tokens: &[u32],
start_pos: usize,
_state: &mut InferenceState,
all_logits: bool,
need_logits: bool,
) -> wgpu::CommandEncoder {
debug_assert!(!tokens.is_empty());
let n = tokens.len();
assert!(
start_pos + n <= self.gpu_state.max_seq_len,
"prefill start_pos {start_pos} + n {n} exceeds max_seq_len {}",
self.gpu_state.max_seq_len,
);
debug_assert!(
n <= self.gpu_state.max_seq_len.min(MAX_PREFILL_TOKENS),
"n {n} exceeds chunk capacity (max_seq_len = {}, MAX_PREFILL_TOKENS = {MAX_PREFILL_TOKENS})",
self.gpu_state.max_seq_len,
);
let cfg = &self.config;
let hs = cfg.hidden_size;
let is = cfg.intermediate_size;
self.ctx.reset_profiler();
self.gpu_state.seq_len.store(start_pos, Ordering::Relaxed);
let lora = self
.active_lora
.lock()
.unwrap_or_else(|e| e.into_inner())
.clone();
if let Some(tq) = self.tq_cache() {
let scale = self
.scalars
.attn
.unwrap_or_else(|| 1.0 / (cfg.head_dim as f32).sqrt());
tq.write_params(&self.ctx, cfg, n, start_pos, scale);
}
let mut staged: Vec<f32> = Vec::with_capacity(n * hs);
for &t in tokens {
let off = (t as usize) * hs;
staged.extend_from_slice(&self.gpu_state.embedding_f32[off..off + hs]);
}
self.ctx
.queue
.write_buffer(&self.prefill_batch_buf, 0, bytemuck::cast_slice(&staged));
if lora.is_some() {
self.lora_params_pool
.lock()
.unwrap_or_else(|e| e.into_inner())
.1 = 0;
}
self.prefill_params_pool
.lock()
.unwrap_or_else(|e| e.into_inner())
.1 = 0;
let mut enc = self.new_encoder();
let n_u = n as u32;
let hs_u = hs as u32;
let is_u = is as u32;
for layer in 0..cfg.n_layers {
let lw = &self.layers[layer];
if layer > 0 {
self.encode_add_rmsnorm_batch(
&mut enc,
&self.prefill_batch_buf,
&self.prefill_normed_buf,
&lw.attn_norm,
&self.prefill_up_buf,
n_u,
hs_u,
);
} else {
self.encode_rmsnorm_batch(
&mut enc,
&self.prefill_batch_buf,
&self.prefill_normed_buf,
&lw.attn_norm,
n_u,
hs_u,
);
}
if cfg.block_types[layer] == BlockType::GatedConv {
let conv_buf = self.gpu_state.conv_buffers[layer].as_ref().unwrap();
let w_in = lw.conv_in_proj.as_ref().unwrap();
let w_out = lw.conv_out_proj.as_ref().unwrap();
let conv_weight = lw.conv_weight.as_ref().unwrap();
self.encode_mul_mat_reg_tile(
&mut enc,
w_in,
&self.prefill_normed_buf,
&self.prefill_proj_buf,
n_u,
hs_u,
hs_u,
3 * hs_u,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::ShortconvInProj,
&self.prefill_normed_buf,
&self.prefill_proj_buf,
n_u,
);
self.encode_conv1d_fused_batch(
&mut enc,
&self.prefill_proj_buf,
conv_buf,
conv_weight,
&self.prefill_normed_buf,
n_u,
hs_u,
);
self.encode_mul_mat_reg_tile(
&mut enc,
w_out,
&self.prefill_normed_buf,
&self.prefill_gate_buf,
n_u,
hs_u,
hs_u,
hs_u,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::ShortconvOutProj,
&self.prefill_normed_buf,
&self.prefill_gate_buf,
n_u,
);
} else {
let head_dim = cfg.head_dim as u32;
let n_kv_heads = cfg.kv_heads_per_layer[layer] as u32;
let kv_dim = n_kv_heads * head_dim;
let n_heads = cfg.n_heads as u32;
let q_dim = n_heads * head_dim;
let w_q = lw.attn_q.as_ref().unwrap();
let w_k = lw.attn_k.as_ref().unwrap();
let w_v = lw.attn_v.as_ref().unwrap();
let w_o = lw.attn_output.as_ref().unwrap();
self.encode_mul_mat_reg_tile(
&mut enc,
w_q,
&self.prefill_normed_buf,
&self.prefill_proj_buf,
n_u,
hs_u,
hs_u,
q_dim,
);
self.encode_mul_mat_reg_tile(
&mut enc,
w_k,
&self.prefill_normed_buf,
&self.prefill_gate_buf,
n_u,
hs_u,
hs_u,
kv_dim,
);
self.encode_mul_mat_reg_tile(
&mut enc,
w_v,
&self.prefill_normed_buf,
&self.prefill_up_buf,
n_u,
hs_u,
hs_u,
kv_dim,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::AttnQ,
&self.prefill_normed_buf,
&self.prefill_proj_buf,
n_u,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::AttnK,
&self.prefill_normed_buf,
&self.prefill_gate_buf,
n_u,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::AttnV,
&self.prefill_normed_buf,
&self.prefill_up_buf,
n_u,
);
if let Some(b) = lw.attn_q_bias.as_ref() {
self.encode_bias_add_batch(&mut enc, &self.prefill_proj_buf, b, n_u, q_dim);
}
if let Some(b) = lw.attn_k_bias.as_ref() {
self.encode_bias_add_batch(&mut enc, &self.prefill_gate_buf, b, n_u, kv_dim);
}
if let Some(b) = lw.attn_v_bias.as_ref() {
self.encode_bias_add_batch(&mut enc, &self.prefill_up_buf, b, n_u, kv_dim);
}
self.encode_qk_norm_rope_batch(
&mut enc,
&self.prefill_proj_buf,
&self.prefill_gate_buf,
lw.attn_q_norm.as_ref(),
lw.attn_k_norm.as_ref(),
start_pos as u32,
n_u,
n_heads,
n_kv_heads,
head_dim,
q_dim,
kv_dim,
);
let attn_scale = self
.scalars
.attn
.unwrap_or_else(|| 1.0 / (head_dim as f32).sqrt());
if let Some(tq) = self.tq_cache() {
tq.encode_kv(
&self.ctx,
&mut enc,
layer,
&self.prefill_gate_buf,
&self.prefill_up_buf,
n,
);
tq.rotate_queries(
&self.ctx,
&mut enc,
layer,
&self.prefill_proj_buf,
n,
n_heads as usize,
);
tq.attention(
&self.ctx,
&mut enc,
layer,
&self.prefill_normed_buf,
n,
n_heads as usize,
);
} else {
let (k_cache, v_cache) = self.active_kv(layer);
let kv_off_floats = (start_pos * kv_dim as usize) as u64;
let kv_chunk_floats = (n * kv_dim as usize) as u64;
Self::encode_copy(
&mut enc,
&self.prefill_gate_buf,
0,
k_cache,
kv_off_floats,
kv_chunk_floats,
);
Self::encode_copy(
&mut enc,
&self.prefill_up_buf,
0,
v_cache,
kv_off_floats,
kv_chunk_floats,
);
let max_seq_for_kv = (start_pos + n) as u32;
self.encode_attention_prefill(
&mut enc,
&self.prefill_proj_buf,
k_cache,
v_cache,
&self.prefill_normed_buf,
n_u,
n_heads,
n_kv_heads,
head_dim,
kv_dim,
max_seq_for_kv,
start_pos as u32,
q_dim,
q_dim,
attn_scale,
);
}
self.encode_mul_mat_reg_tile(
&mut enc,
w_o,
&self.prefill_normed_buf,
&self.prefill_gate_buf,
n_u,
q_dim,
q_dim,
hs_u,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::AttnOutput,
&self.prefill_normed_buf,
&self.prefill_gate_buf,
n_u,
);
}
self.encode_add_rmsnorm_batch(
&mut enc,
&self.prefill_batch_buf,
&self.prefill_normed_buf,
&lw.ffn_norm,
&self.prefill_gate_buf,
n_u,
hs_u,
);
let dense = match &lw.ffn {
GpuFfn::Moe(moe) => {
let steps = self.moe_ffn_steps(
moe,
&self.prefill_normed_buf,
&self.prefill_up_buf,
n_u,
false,
);
let mut pass = self.ctx.begin_pass(&mut enc, "ffn_moe_batch");
steps.iter().for_each(|s| {
self.dispatch_into(&mut pass, s.pipeline, &s.bind_group, s.workgroups);
});
continue;
}
GpuFfn::Dense(d) => d,
};
self.encode_mul_mat_reg_tile(
&mut enc,
&dense.gate,
&self.prefill_normed_buf,
&self.prefill_gate_buf,
n_u,
hs_u,
hs_u,
is_u,
);
self.encode_mul_mat_reg_tile(
&mut enc,
&dense.up,
&self.prefill_normed_buf,
&self.prefill_up_buf,
n_u,
hs_u,
hs_u,
is_u,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::FfnGate,
&self.prefill_normed_buf,
&self.prefill_gate_buf,
n_u,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::FfnUp,
&self.prefill_normed_buf,
&self.prefill_up_buf,
n_u,
);
{
let total = n_u * is_u;
let params: [u32; 2] = [total, 0];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.silu_mul_inplace.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.prefill_gate_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.prefill_up_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: p_buf.as_entire_binding(),
},
],
});
self.encode(
&mut enc,
&self.pipelines.silu_mul_inplace,
&bg,
(total.div_ceil(256), 1, 1),
"silu_mul_batch",
);
}
self.encode_mul_mat_reg_tile(
&mut enc,
&dense.down,
&self.prefill_gate_buf,
&self.prefill_up_buf,
n_u,
is_u,
is_u,
hs_u,
);
self.encode_lora_hook_batched(
&mut enc,
lora.as_ref(),
layer,
LoraTarget::FfnDown,
&self.prefill_gate_buf,
&self.prefill_up_buf,
n_u,
);
}
{
let total = n_u * hs_u;
let params: [u32; 2] = [total, self.scalars.residual.to_bits()];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pipelines.scaled_add_inplace.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.prefill_batch_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.prefill_up_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: p_buf.as_entire_binding(),
},
],
});
self.encode(
&mut enc,
&self.pipelines.scaled_add_inplace,
&bg,
(total.div_ceil(256), 1, 1),
"final_add",
);
}
if !need_logits && !all_logits {
} else if !all_logits {
let last_off_floats = ((n - 1) * hs) as u64;
Self::encode_copy(
&mut enc,
&self.prefill_batch_buf,
last_off_floats,
&self.hidden_buf,
0,
hs as u64,
);
self.encode_rmsnorm(
&mut enc,
&self.hidden_buf,
&self.output_norm,
hs_u,
cfg.rms_norm_eps,
);
self.encode_lm_head(&mut enc, &self.hidden_buf, &self.logits_buf);
if let Some(params) = self.logit_scale_params.as_ref() {
let scale_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("logit_scale_bg"),
layout: &self.pipelines.scale_f32.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.logits_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: params.as_entire_binding(),
},
],
});
self.encode(
&mut enc,
&self.pipelines.scale_f32,
&scale_bg,
((cfg.vocab_size as u32).div_ceil(256), 1, 1),
"logit_scale",
);
}
} else {
let vocab = cfg.vocab_size;
self.encode_rmsnorm_batch(
&mut enc,
&self.prefill_batch_buf,
&self.prefill_normed_buf,
&self.output_norm,
n_u,
hs_u,
);
match &self.lm_head {
LmHead::Quantized(w) => {
self.encode_mul_mat_reg_tile(
&mut enc,
w,
&self.prefill_normed_buf,
&self.prefill_all_logits_buf,
n_u,
hs_u,
hs_u,
vocab as u32,
);
}
LmHead::F16 { weight, params } => {
for j in 0..n {
let tok_off_floats = (j * hs) as u64;
Self::encode_copy(
&mut enc,
&self.prefill_normed_buf,
tok_off_floats,
&self.hidden_buf,
0,
hs as u64,
);
self.encode_gemv_f16(
&mut enc,
weight,
params,
&self.hidden_buf,
&self.logits_buf,
);
Self::encode_copy(
&mut enc,
&self.logits_buf,
0,
&self.prefill_all_logits_buf,
(j * vocab) as u64,
vocab as u64,
);
}
}
}
if self.scalars.logit != 1.0 {
let total = n_u * (vocab as u32);
let params_data: [u32; 2] = [total, (1.0 / self.scalars.logit).to_bits()];
let p_buf = self.next_prefill_params(bytemuck::cast_slice(¶ms_data));
let scale_bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("logit_scale_bg"),
layout: &self.pipelines.scale_f32.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.prefill_all_logits_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: p_buf.as_entire_binding(),
},
],
});
self.encode(
&mut enc,
&self.pipelines.scale_f32,
&scale_bg,
(total.div_ceil(256), 1, 1),
"logit_scale",
);
}
}
enc
}
fn forward_prefill_batched_locked(
&self,
tokens: &[u32],
start_pos: usize,
state: &mut InferenceState,
all_logits: bool,
need_logits: bool,
) -> Vec<f32> {
let n = tokens.len();
let enc =
self.encode_prefill_batched_locked(tokens, start_pos, state, all_logits, need_logits);
self.submit_and_wait(enc);
self.gpu_state
.seq_len
.store(start_pos + n, Ordering::Relaxed);
state.seq_len = start_pos + n;
self.ctx.finish_profiler();
if !need_logits && !all_logits {
Vec::new()
} else if !all_logits {
self.ctx
.download_f32(&self.logits_buf, self.config.vocab_size)
} else {
self.ctx
.download_f32(&self.prefill_all_logits_buf, n * self.config.vocab_size)
}
}
pub async fn forward_prefill_logits_all_async(
&self,
tokens: &[u32],
start_pos: usize,
state: &mut InferenceState,
) -> Result<Vec<f32>> {
let n = tokens.len();
if n == 0 {
return Ok(Vec::new());
}
if n > MAX_ALL_LOGITS_TOKENS {
anyhow::bail!("batch size {n} exceeds MAX_ALL_LOGITS_TOKENS ({MAX_ALL_LOGITS_TOKENS})");
}
if !self.batched_prefill || self.unbatchable_matmul_weight().is_some() {
anyhow::bail!("batched prefill verification not supported for this model");
}
let vocab = self.config.vocab_size;
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
let enc = self.encode_prefill_batched_locked(tokens, start_pos, state, true, true);
self.gpu_state
.seq_len
.store(start_pos + n, Ordering::Relaxed);
state.seq_len = start_pos + n;
self.ctx.begin_download_with_encoder(
enc,
&self.prefill_all_logits_buf,
(n * vocab * std::mem::size_of::<f32>()) as u64,
)
};
let bytes = pending.recv().await?;
let expected_bytes = n * vocab * std::mem::size_of::<f32>();
if bytes.len() < expected_bytes {
anyhow::bail!(
"GPU prefill logits readback buffer truncated (expected {expected_bytes} bytes, got {})",
bytes.len()
);
}
let mut out = vec![0.0f32; n * vocab];
bytemuck::cast_slice_mut(&mut out).copy_from_slice(&bytes[..expected_bytes]);
Ok(out)
}
pub async fn forward_prefill_argmax_all_async(
&self,
tokens: &[u32],
start_pos: usize,
state: &mut InferenceState,
) -> Result<Vec<u32>> {
let n = tokens.len();
if n == 0 {
return Ok(Vec::new());
}
if n > MAX_ALL_LOGITS_TOKENS {
anyhow::bail!("batch size {n} exceeds MAX_ALL_LOGITS_TOKENS ({MAX_ALL_LOGITS_TOKENS})");
}
if !self.batched_prefill || self.unbatchable_matmul_weight().is_some() {
anyhow::bail!("batched prefill verification not supported for this model");
}
let vocab = self.config.vocab_size;
let pending = {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
let mut enc = self.encode_prefill_batched_locked(tokens, start_pos, state, true, true);
self.gpu_state
.seq_len
.store(start_pos + n, Ordering::Relaxed);
state.seq_len = start_pos + n;
for j in 0..n {
Self::encode_copy(
&mut enc,
&self.prefill_all_logits_buf,
(j * vocab) as u64,
&self.logits_buf,
0,
vocab as u64,
);
let params_buf =
self.next_prefill_params(bytemuck::cast_slice(&[vocab as u32, j as u32]));
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: Some("argmax_batch_row_bg"),
layout: &self.pipelines.argmax_f32.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry {
binding: 0,
resource: self.logits_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 1,
resource: self.argmax_out_buf.as_entire_binding(),
},
wgpu::BindGroupEntry {
binding: 2,
resource: params_buf.as_entire_binding(),
},
],
});
let mut pass = self.ctx.begin_pass(&mut enc, "argmax_batch_row");
pass.set_pipeline(&self.pipelines.argmax_f32);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(1, 1, 1);
}
self.ctx.begin_download_with_encoder(
enc,
&self.argmax_out_buf,
std::mem::size_of_val(tokens) as u64,
)
};
let bytes = pending.recv().await?;
let expected_bytes = std::mem::size_of_val(tokens);
if bytes.len() < expected_bytes {
anyhow::bail!(
"GPU prefill argmax readback buffer truncated (expected {expected_bytes} bytes, got {})",
bytes.len()
);
}
let mut out = vec![0u32; n];
bytemuck::cast_slice_mut(&mut out).copy_from_slice(&bytes[..expected_bytes]);
Ok(out)
}
pub fn truncate_kv_direct(&self, state: &mut InferenceState, len: usize) {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
self.gpu_state.seq_len.store(len, Ordering::Relaxed);
state.seq_len = len;
}
}
impl GpuLfm2Model {
fn snapshot_state_locked(&self) -> StateSnapshot {
let seq_len = self.gpu_state.seq_len.load(Ordering::Relaxed);
let cfg = &self.config;
let head_dim = cfg.head_dim;
let kernel_size = cfg.conv_kernel_size.unwrap_or(3);
let d_conv = kernel_size - 1;
let download_exact =
|buf: &wgpu::Buffer, count: usize| -> Vec<f32> { self.ctx.download_f32(buf, count) };
let tq = self.tq_cache();
let mut layers = Vec::with_capacity(cfg.n_layers);
for i in 0..cfg.n_layers {
if cfg.block_types[i] == BlockType::Attention {
if let Some(tq) = tq {
let (keys, values) = tq.snapshot_layer(&self.ctx, i, seq_len);
layers.push(LayerSnapshot::AttentionCompressed { keys, values });
continue;
}
let kv_dim = cfg.kv_heads_per_layer[i] * head_dim;
let count = seq_len * kv_dim;
let (k_buf, v_buf) = self.active_kv(i);
let k_floats = download_exact(k_buf, count);
let v_floats = download_exact(v_buf, count);
layers.push(LayerSnapshot::Attention {
k_data: bytemuck::cast_slice(&k_floats).to_vec(),
v_data: bytemuck::cast_slice(&v_floats).to_vec(),
});
} else {
let count = d_conv * cfg.hidden_size;
let conv_buf = self.gpu_state.conv_buffers[i]
.as_ref()
.expect("conv layer must have rolling buffer");
let floats = download_exact(conv_buf, count);
layers.push(LayerSnapshot::Conv {
buffer: bytemuck::cast_slice(&floats).to_vec(),
});
}
}
StateSnapshot::new(layers, seq_len)
}
fn restore_state_locked(&self, snapshot: &StateSnapshot) {
let cfg = &self.config;
for (i, layer_snap) in snapshot.layers.iter().enumerate() {
match layer_snap {
LayerSnapshot::Attention { k_data, v_data } => {
assert_eq!(
cfg.block_types[i],
BlockType::Attention,
"snapshot layer {i} attention vs state config"
);
assert!(
self.tq_cache().is_none(),
"f32 Attention snapshot restored into a TurboQuant-configured \
wgpu model at layer {i}; the lookup gate in forward_prefill \
must reject a mode-mismatched snapshot"
);
let (k_buf, v_buf) = self.active_kv(i);
self.ctx.queue.write_buffer(k_buf, 0, k_data);
self.ctx.queue.write_buffer(v_buf, 0, v_data);
}
LayerSnapshot::Conv { buffer } => {
assert_eq!(
cfg.block_types[i],
BlockType::GatedConv,
"snapshot layer {i} conv vs state config"
);
let conv_buf = self.gpu_state.conv_buffers[i]
.as_ref()
.expect("conv layer must have rolling buffer");
self.ctx.queue.write_buffer(conv_buf, 0, buffer);
}
LayerSnapshot::AttentionCompressed { keys, values } => {
assert_eq!(
cfg.block_types[i],
BlockType::Attention,
"snapshot layer {i} attention vs state config"
);
let tq = self.tq_cache().unwrap_or_else(|| {
panic!(
"GpuLfm2Model::restore_state_locked received a \
TurboQuant-compressed snapshot at layer {i} but this \
model is not TurboQuant-configured; callers must gate \
on `StateSnapshot::is_compressed`"
)
});
let restored =
tq.restore_layer(&self.ctx, i, keys, values)
.unwrap_or_else(|| {
panic!(
"invalid or shape-mismatched TurboQuant blob in \
snapshot at layer {i}"
)
});
assert_eq!(
restored, snapshot.seq_len,
"layer {i}: restored TurboQuant seq_len {restored} disagrees \
with the snapshot's {}",
snapshot.seq_len
);
}
LayerSnapshot::AttentionF16 { .. } => {
panic!(
"GpuLfm2Model::restore_state_locked received an f16 \
snapshot at layer {i}; wgpu uses f32 KV. This indicates \
a cross-backend cache-namespace leak."
);
}
}
}
self.gpu_state
.seq_len
.store(snapshot.seq_len, Ordering::Relaxed);
}
fn zero_conv_buffers_locked(&self) {
let cfg = &self.config;
let mut enc = self
.ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor {
label: Some("zero_conv_buffers"),
});
for i in 0..cfg.n_layers {
if cfg.block_types[i] == BlockType::GatedConv
&& let Some(conv_buf) = self.gpu_state.conv_buffers[i].as_ref()
{
enc.clear_buffer(conv_buf, 0, None);
}
}
self.ctx.submit_encoder(enc);
}
}
impl Model for GpuLfm2Model {
fn supports_all_logits(&self) -> bool {
self.batched_prefill && self.unbatchable_matmul_weight().is_none()
}
fn forward_prefill_logits_all(
&self,
tokens: &[u32],
start_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
let n = tokens.len();
if n == 0 {
return Vec::new();
}
if n == 1 {
self.gpu_state.seq_len.store(start_pos, Ordering::Relaxed);
return self.forward_inner(tokens, start_pos, state);
}
assert!(
n <= MAX_ALL_LOGITS_TOKENS,
"forward_prefill_logits_all token count ({n}) exceeds MAX_ALL_LOGITS_TOKENS ({MAX_ALL_LOGITS_TOKENS})"
);
self.forward_prefill_batched_locked(tokens, start_pos, state, true, true)
}
fn truncate_kv(&self, state: &mut InferenceState, len: usize) {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
self.gpu_state.seq_len.store(len, Ordering::Relaxed);
state.seq_len = len;
}
fn supports_hidden_states(&self) -> bool {
true
}
fn hidden_states(&self, tokens: &[u32], state: &mut InferenceState) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
assert!(
!tokens.is_empty(),
"hidden_states requires at least one token"
);
let hs = self.config.hidden_size;
let vocab = self.config.vocab_size;
assert!(
tokens.len() <= self.gpu_state.max_seq_len,
"hidden_states chunk ({}) exceeds max_seq_len ({})",
tokens.len(),
self.gpu_state.max_seq_len
);
let scratch = self.hs_scratch();
let mut enc = self.new_encoder();
for buf in scratch.conv.iter().flatten() {
enc.clear_buffer(buf, 0, None);
}
self.submit_and_wait(enc);
let saved_seq = self.gpu_state.seq_len.load(Ordering::Relaxed);
struct HsGuard<'a> {
flag: &'a AtomicBool,
seq: &'a AtomicUsize,
saved: usize,
}
impl Drop for HsGuard<'_> {
fn drop(&mut self) {
self.seq.store(self.saved, Ordering::Relaxed);
self.flag.store(false, Ordering::Relaxed);
}
}
self.gpu_state.seq_len.store(0, Ordering::Relaxed);
self.use_hs_scratch.store(true, Ordering::Relaxed);
let _hs_guard = HsGuard {
flag: &self.use_hs_scratch,
seq: &self.gpu_state.seq_len,
saved: saved_seq,
};
let mut dummy = InferenceState::for_prefill(&self.config, 1)
.expect("hidden_states: 1-token scratch InferenceState allocation failed");
let mut out = Vec::with_capacity(tokens.len() * hs);
for (pos, &token) in tokens.iter().enumerate() {
let token_id = token as usize;
assert!(
token_id < vocab,
"token_id {token_id} out of range (vocab_size={vocab})"
);
self.forward_inner_compute(&[token], pos, &mut dummy);
out.extend_from_slice(&self.ctx.download_f32(&self.hidden_buf, hs));
}
out
}
fn forward(&self, tokens: &[u32], pos: usize, state: &mut InferenceState) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
self.forward_inner(tokens, pos, state)
}
fn forward_greedy(&self, tokens: &[u32], pos: usize, state: &mut InferenceState) -> u32 {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.gpu_state.seq_len.store(pos, Ordering::Relaxed);
self.forward_greedy_inner(tokens, pos, state)
}
fn supports_embedding_input(&self) -> bool {
true
}
fn forward_from_embedding(
&self,
embedding: &[f32],
_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
let pos = state.seq_len;
self.forward_inner_compute_from_embedding(embedding, pos, state);
self.ctx
.download_f32(&self.logits_buf, self.config.vocab_size)
}
fn forward_embedding(
&self,
tokens: &[u32],
_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
let pos = state.seq_len;
self.forward_inner_compute_tail(tokens, pos, state, DecodeTail::Hidden);
self.ctx
.download_f32(&self.hidden_buf, self.config.hidden_size)
}
fn forward_hidden_from_embedding(
&self,
embedding: &[f32],
_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
let pos = state.seq_len;
self.forward_inner_compute_tail_seeded(
HiddenSeed::Embedding(embedding),
pos,
state,
DecodeTail::Hidden,
);
self.ctx
.download_f32(&self.hidden_buf, self.config.hidden_size)
}
fn forward_prefill_from_embeddings(
&self,
embeddings: &[f32],
n_tokens: usize,
start_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
self.seed_embeddings_locked(embeddings, n_tokens, start_pos, state);
self.ctx
.download_f32(&self.logits_buf, self.config.vocab_size)
}
fn forward_prefill(
&self,
tokens: &[u32],
start_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let _lora_guard = self.resolve_lora(state);
let lora_active = self
.active_lora
.lock()
.unwrap_or_else(|e| e.into_inner())
.is_some();
self.gpu_state.seq_len.store(start_pos, Ordering::Relaxed);
if start_pos == 0 {
let hit = (!lora_active)
.then(|| {
self.prefix_cache
.lock()
.unwrap_or_else(|e| e.into_inner())
.find_longest_prefix(tokens)
})
.flatten()
.filter(|(snapshot, _)| snapshot.is_compressed() == self.tq_cache().is_some());
if let Some((snapshot, prefix_len)) = hit {
if prefix_len < tokens.len() && prefix_len > 0 {
let use_len = prefix_len;
self.restore_state_locked(&snapshot);
self.gpu_state.seq_len.store(use_len, Ordering::Relaxed);
state.seq_len = use_len;
let remaining = &tokens[use_len..];
let last = remaining.len() - 1;
let mut logits = Vec::new();
for (j, &token) in remaining.iter().enumerate() {
if j == last {
logits = self.forward_inner(&[token], use_len + j, state);
} else {
self.forward_inner_compute(&[token], use_len + j, state);
}
}
self.prefix_cache
.lock()
.unwrap_or_else(|e| e.into_inner())
.insert(tokens, self.snapshot_state_locked());
return logits;
}
}
self.zero_conv_buffers_locked();
}
let unbatchable = self.unbatchable_matmul_weight();
if let Some((layer, name, dtype)) = unbatchable
&& !tokens.is_empty()
&& self.batched_prefill
&& !self.batched_fallback_warned.swap(true, Ordering::Relaxed)
{
tracing::warn!(
layer,
tensor = name,
?dtype,
"no batched prefill GEMM for this dtype — falling back to the \
per-token loop, which issues ~340x the GPU submits and makes \
prefill no faster than decode. Add a batched kernel for {dtype:?} \
to put this model back on the fast path.",
);
}
if !tokens.is_empty() && self.batched_prefill && unbatchable.is_none() {
let chunk_size = self.gpu_state.max_seq_len.min(MAX_PREFILL_TOKENS);
let mut logits = Vec::new();
let mut pos = 0usize;
while pos < tokens.len() {
let end = (pos + chunk_size).min(tokens.len());
let is_last = end >= tokens.len();
logits = self.forward_prefill_batched_locked(
&tokens[pos..end],
start_pos + pos,
state,
false,
is_last,
);
pos = end;
}
if start_pos == 0 && !lora_active {
self.prefix_cache
.lock()
.unwrap_or_else(|e| e.into_inner())
.insert(tokens, self.snapshot_state_locked());
}
return logits;
}
let mut logits = Vec::new();
if !tokens.is_empty() {
let last = tokens.len() - 1;
for (i, &token) in tokens.iter().enumerate() {
if i == last {
logits = self.forward_inner(&[token], start_pos + i, state);
} else {
self.forward_inner_compute(&[token], start_pos + i, state);
}
}
}
if start_pos == 0 && !lora_active {
self.prefix_cache
.lock()
.unwrap_or_else(|e| e.into_inner())
.insert(tokens, self.snapshot_state_locked());
}
logits
}
fn configure_cache(&self, config: crate::kv_cache::KvCacheConfig) {
let id = self.cache_namespace();
*self.prefix_cache.lock().unwrap_or_else(|e| e.into_inner()) =
KvPrefixCache::new(config, &self.config, &id);
}
fn clear_warm_cache(&self) {
self.prefix_cache
.lock()
.unwrap_or_else(|e| e.into_inner())
.clear_warm();
}
fn clear_cache(&self) {
self.prefix_cache
.lock()
.unwrap_or_else(|e| e.into_inner())
.clear();
}
fn snapshot_state(&self) -> StateSnapshot {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
self.snapshot_state_locked()
}
fn restore_state(&self, snapshot: &StateSnapshot) {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
self.restore_state_locked(snapshot);
}
fn supports_moe_lora(&self) -> bool {
false
}
fn turboquant_supported(&self) -> bool {
crate::model::gpu_turboquant::head_dim_supported(self.config.head_dim)
}
fn configure_kv_compression(&self, compression: &KvCompression) -> Result<(), CeraError> {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
let want = TqMode::from_compression(compression, self.config.head_dim);
if want.is_none() && matches!(compression, KvCompression::TurboQuant { .. }) {
tracing::warn!(
target: "cera::gpu",
head_dim = self.config.head_dim,
"TurboQuant requested but not supported for this configuration on \
the wgpu backend (needs keys+values compression and a \
power-of-two head_dim <= 128 that is a multiple of 32); \
falling back to f32 KV"
);
}
if let Some(&configured) = self.kv_mode.get() {
if configured != want {
return Err(CeraError::KvCompressionConflict {
configured: describe_kv_mode(&configured),
requested: describe_kv_mode(&want),
});
}
debug_assert_eq!(
self.tq.get().map(|t| t.mode),
configured,
"kv_mode and the built TurboQuant cache disagree"
);
return Ok(());
}
if let Some(mode) = want {
let q_cap = self.gpu_state.max_seq_len.min(MAX_PREFILL_TOKENS);
let cache = TqGpuCache::new(
&self.ctx,
&self.config,
self.gpu_state.max_seq_len,
q_cap,
mode,
)?;
assert!(self.tq.set(cache).is_ok(), "tq cache set race");
}
let _ = self.kv_mode.set(want);
let _ = self.kv_cache_tag.set(if want.is_some() {
compression.resolved_for(&self.config).cache_tag()
} else {
KvCompression::None.cache_tag()
});
let tag_changed = self.kv_cache_tag.get().is_some_and(|t| !t.is_empty());
if tag_changed {
let id = self.cache_namespace();
let mut cache = self.prefix_cache.lock().unwrap_or_else(|e| e.into_inner());
let cache_config = cache.config.clone();
*cache = KvPrefixCache::new(cache_config, &self.config, &id);
}
Ok(())
}
fn supports_kv_shift(&self) -> bool {
self.tq.get().is_none()
}
fn shift_kv(&self, state: &mut InferenceState, n_keep: usize, shift: usize) {
let _guard = self.infer_lock.lock().unwrap_or_else(|e| e.into_inner());
assert!(shift > 0, "shift must be > 0");
let cur_len = self.gpu_state.seq_len.load(Ordering::Relaxed);
debug_assert_eq!(
state.seq_len, cur_len,
"seq_len mirrors out of sync: state.seq_len={} gpu_state.seq_len={cur_len}",
state.seq_len,
);
assert!(
n_keep + shift <= cur_len,
"shift range out of bounds: n_keep={n_keep} + shift={shift} > seq_len={cur_len}",
);
assert!(
!state.is_compressed(),
"shift_kv called on a TurboQuant-compressed state; \
shifting compressed caches is not supported on the wgpu backend"
);
let new_seq_len = cur_len - shift;
let retained = new_seq_len - n_keep;
if retained > 0 {
self.encode_kv_shift_layers(n_keep, shift, retained);
}
self.gpu_state.seq_len.store(new_seq_len, Ordering::Relaxed);
state.seq_len = new_seq_len;
}
fn config(&self) -> &ModelConfig {
&self.config
}
}
#[cfg(test)]
#[cfg(not(target_arch = "wasm32"))]
mod tests {
use crate::backend::wgpu::GpuContext;
fn gpu_ctx_or_skip() -> Option<GpuContext> {
match GpuContext::new() {
Ok(ctx) => Some(ctx),
Err(e) => {
let required = std::env::var("CERA_REQUIRE_GPU").unwrap_or_default();
assert!(
required.is_empty(),
"CERA_REQUIRE_GPU is set but no GPU adapter is available: {e}"
);
eprintln!("skipping: no GPU adapter ({e})");
None
}
}
}
#[test]
fn gemv_tile_rows_fits_and_aligns() {
use super::gemv_tile_rows;
assert_eq!(gemv_tile_rows(1000, 512, 1 << 30, 256, 2), 1000);
assert_eq!(gemv_tile_rows(1000, 512, 1 << 30, 256, 4), 1000);
let m = 131_072u32;
let align = 256u64;
for &k in &[100u32, 99u32] {
for &elem in &[2u64, 4u64] {
let row_bytes = u64::from(k) * elem;
for &max_binding in &[1u64 << 20, 4 << 20, 512 << 10] {
let rows = gemv_tile_rows(m, k, max_binding, align, elem);
assert!(rows > 0, "k={k} elem={elem} max={max_binding}");
assert!(
u64::from(rows) * row_bytes <= max_binding,
"tile exceeds max_binding (k={k}, elem={elem}, max={max_binding})",
);
assert_eq!(
(u64::from(rows) * row_bytes) % align,
0,
"tile byte size not offset-aligned (k={k}, elem={elem}, max={max_binding})",
);
let final_rows = m % rows;
if final_rows > 0 {
let offset = u64::from(m - final_rows) * row_bytes;
let bound = (u64::from(final_rows) * row_bytes).div_ceil(4) * 4;
let padded_buf = (u64::from(m) * row_bytes).div_ceil(4) * 4;
assert_eq!(offset % 4, 0, "offset not u32-aligned (k={k}, elem={elem})");
assert!(
offset + bound <= padded_buf,
"final tile rounded binding overruns padded buffer (k={k}, elem={elem})",
);
}
}
}
}
let mb = 1u64 << 20;
assert!(gemv_tile_rows(m, 100, mb, align, 2) >= gemv_tile_rows(m, 100, mb, align, 4));
}
#[test]
fn encode_copy_scales_float_offsets_to_bytes() {
let Some(ctx) = gpu_ctx_or_skip() else {
return;
};
let src: Vec<f32> = (0..16).map(|x| x as f32).collect();
let src_buf = ctx.upload_f32(&src, "encode_copy_src");
let dst_buf = ctx.create_storage_rw((16 * 4) as u64, "encode_copy_dst");
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
super::GpuLfm2Model::encode_copy(&mut enc, &src_buf, 2, &dst_buf, 5, 4);
ctx.queue.submit(Some(enc.finish()));
let got = ctx.download_f32(&dst_buf, 16);
let mut want = vec![0.0f32; 16];
want[5..9].copy_from_slice(&src[2..6]); assert_eq!(
got, want,
"encode_copy must scale float offsets/length to bytes \
(src_off=2, dst_off=5, len=4 → dst[5..9] == src[2..6])"
);
}
}