use super::super::{RealizarError, Result};
use super::{
gpu_err, CudaLayer, CudaQuantWeight, Qwen35CudaDims, Qwen35CudaModel, Qwen35CudaState,
};
use crate::cuda::types::WeightQuantType;
use crate::gguf::forward_qwen35::{Qwen35Model, Qwen35OwnedLayer};
use trueno_gpu::driver::GpuBuffer;
pub const PREFILL_MAX_CHUNK_ROWS: usize = 512;
pub const UNIFIED_PREFILL_CHUNK_ROWS: usize = 2048;
pub const PREFILL_SCORES_BUDGET_BYTES: usize = 1 << 30;
struct View(std::mem::ManuallyDrop<GpuBuffer<f32>>);
impl View {
fn at(ptr: u64, elems: usize) -> Self {
Self(std::mem::ManuallyDrop::new(unsafe {
GpuBuffer::<f32>::from_raw_parts(ptr, elems)
}))
}
}
impl std::ops::Deref for View {
type Target = GpuBuffer<f32>;
fn deref(&self) -> &GpuBuffer<f32> {
&self.0
}
}
struct PrefillBuffers {
rows: usize,
attention: PrefillAttention,
attn_rows: usize,
x: GpuBuffer<f32>,
normed: GpuBuffer<f32>,
post_normed: GpuBuffer<f32>,
proj: GpuBuffer<f32>,
conv_in: GpuBuffer<f32>,
conv_out: GpuBuffer<f32>,
alpha_raw: GpuBuffer<f32>,
beta_raw: GpuBuffer<f32>,
dt: GpuBuffer<f32>,
beta: GpuBuffer<f32>,
gate: GpuBuffer<f32>,
out_h: GpuBuffer<f32>,
ssm_out_in: GpuBuffer<f32>,
q_full: GpuBuffer<f32>,
q: GpuBuffer<f32>,
q_normed: GpuBuffer<f32>,
attn_gate: GpuBuffer<f32>,
k_raw: GpuBuffer<f32>,
attn_out_in: GpuBuffer<f32>,
scores: GpuBuffer<f32>,
ffn_gate: GpuBuffer<f32>,
ffn_up: GpuBuffer<f32>,
ffn_act: GpuBuffer<f32>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PrefillAttention {
FlashF16In,
CublasF32,
}
impl PrefillAttention {
#[must_use]
pub const fn as_str(self) -> &'static str {
match self {
Self::FlashF16In => "flash (f16 inputs, f32 accumulation)",
Self::CublasF32 => "cuBLAS f32",
}
}
}
pub const PREFILL_ATTENTION_ENV: &str = "APR_QWEN35_PREFILL_ATTENTION";
pub(crate) fn f16_prewarm_fits(bytes: usize, free: usize) -> bool {
bytes.checked_add(1 << 30).is_some_and(|need| need <= free)
}
#[cfg(test)]
thread_local! {
static ATTENTION_OVERRIDE: std::cell::Cell<Option<PrefillAttention>> = const { std::cell::Cell::new(None) };
}
#[must_use]
pub fn attention_candidates(forced: Option<&str>, flash_supported: bool) -> Vec<PrefillAttention> {
let default = if flash_supported {
vec![PrefillAttention::CublasF32, PrefillAttention::FlashF16In]
} else {
vec![PrefillAttention::CublasF32]
};
match forced {
None => default,
Some("f32") => vec![PrefillAttention::CublasF32],
Some("flash") if flash_supported => vec![PrefillAttention::FlashF16In],
Some("flash") => {
eprintln!(
"warning: {PREFILL_ATTENTION_ENV}=flash, but this device or these heads cannot run the flash kernel; using cuBLAS f32"
);
vec![PrefillAttention::CublasF32]
},
Some(other) => {
eprintln!(
"warning: {PREFILL_ATTENTION_ENV}={other:?} is not one of \"f32\" / \"flash\"; using the default"
);
default
},
}
}
pub(crate) fn prefill_attention_candidates(
ex: &crate::cuda::CudaExecutor,
d: Qwen35CudaDims,
) -> Vec<PrefillAttention> {
#[cfg(test)]
if let Some(a) = ATTENTION_OVERRIDE.with(std::cell::Cell::get) {
return vec![a];
}
attention_candidates(
std::env::var(PREFILL_ATTENTION_ENV).ok().as_deref(),
ex.qwen35_flash_attention_supported(d.num_heads, d.num_kv_heads, d.attn_head_dim),
)
}
pub(crate) fn default_prefill_attention(
ex: &crate::cuda::CudaExecutor,
d: Qwen35CudaDims,
) -> PrefillAttention {
prefill_attention_candidates(ex, d)
.first()
.copied()
.unwrap_or(PrefillAttention::CublasF32)
}
fn chunk_rows_for(max_rows: usize, total_positions: usize) -> usize {
max_rows.max(1).min(total_positions.max(1))
}
fn attention_rows_for(d: Qwen35CudaDims, total_positions: usize, max_rows: usize) -> usize {
let hpk = (d.num_heads / d.num_kv_heads.max(1)).max(1) as usize;
let by_scores = scores_budget_bytes() / (4 * hpk * total_positions.max(1));
by_scores.clamp(1, chunk_rows_for(max_rows, total_positions))
}
#[cfg(test)]
thread_local! {
static SCORES_BUDGET_OVERRIDE: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
}
fn scores_budget_bytes() -> usize {
#[cfg(test)]
if let Some(b) = SCORES_BUDGET_OVERRIDE.with(std::cell::Cell::get) {
return b;
}
PREFILL_SCORES_BUDGET_BYTES
}
fn workspace_bytes_for(
d: Qwen35CudaDims,
largest_projection: usize,
total_positions: usize,
attention: PrefillAttention,
max_rows: usize,
) -> usize {
let rows = chunk_rows_for(max_rows, total_positions);
let hpk = (d.num_heads / d.num_kv_heads.max(1)).max(1) as usize;
let q_dim = (d.num_heads * d.attn_head_dim) as usize;
let kv_dim = (d.num_kv_heads * d.attn_head_dim) as usize;
let (hidden, inter) = (d.hidden_dim as usize, d.intermediate_dim as usize);
let (conv, v, nv) = (
d.conv_dim as usize,
d.v_dim as usize,
d.num_v_heads as usize,
);
let per_row = 4 * hidden
+ 2 * conv
+ 4 * nv
+ 3 * v
+ 2 * q_dim + 4 * q_dim + kv_dim
+ 3 * inter;
let scores = match attention {
PrefillAttention::CublasF32 => {
hpk * attention_rows_for(d, total_positions, max_rows) * total_positions
},
PrefillAttention::FlashF16In => 0,
};
4 * (rows * per_row + scores + largest_projection)
}
fn quant_bytes(t: &crate::gguf::OwnedQuantizedTensor) -> (u64, usize) {
(t.data.len() as u64, t.out_dim * t.in_dim)
}
impl Qwen35CudaModel<'_> {
#[must_use]
pub fn prefill_chunk_rows(&self, total_positions: usize) -> usize {
chunk_rows_for(self.prefill_rows, total_positions)
}
pub fn set_prefill_chunk_rows(&mut self, rows: usize) {
self.prefill_rows = rows.max(1);
}
pub fn set_prefill_attention(&mut self, attention: PrefillAttention) {
self.prefill_attention = attention;
}
#[must_use]
pub fn prefill_workspace_bytes(&self, total_positions: usize) -> usize {
workspace_bytes_for(
self.dims,
self.largest_projection_elems(),
total_positions,
self.prefill_attention,
self.prefill_rows,
)
}
#[must_use]
pub fn capacity_inputs(
model: &Qwen35Model<'_>,
seq_len: usize,
gpu_free: u64,
gpu_total: u64,
attention: PrefillAttention,
chunk_rows: usize,
) -> crate::capacity::CapacityInputs {
let d = Self::dims_of(model);
let f32s = |v: &[f32]| 4 * v.len() as u64;
let mut weights = 0u64;
let mut largest = 0usize;
let mut attn_layers = 0u64;
for layer in &model.layers {
let (vecs, quants): (Vec<&[f32]>, Vec<&crate::gguf::OwnedQuantizedTensor>) = match layer
{
Qwen35OwnedLayer::DeltaNet(l) => (
vec![
&l.attn_norm,
&l.ssm_a,
&l.ssm_dt_bias,
&l.ssm_conv1d_weight,
&l.ssm_norm_weight,
&l.post_attention_norm,
],
vec![
&l.attn_qkv,
&l.attn_gate,
&l.ssm_alpha,
&l.ssm_beta,
&l.ssm_out,
&l.ffn_gate,
&l.ffn_up,
&l.ffn_down,
],
),
Qwen35OwnedLayer::Attention(l) => {
attn_layers += 1;
(
vec![
&l.attn_norm,
&l.attn_q_norm,
&l.attn_k_norm,
&l.post_attention_norm,
],
vec![
&l.attn_q,
&l.attn_k,
&l.attn_v,
&l.attn_output,
&l.ffn_gate,
&l.ffn_up,
&l.ffn_down,
],
)
},
};
weights += vecs.iter().map(|v| f32s(v)).sum::<u64>();
for q in quants {
let (bytes, elems) = quant_bytes(q);
weights += bytes;
largest = largest.max(elems);
}
}
weights +=
2 * f32s(model.base.output_norm_weight()) + quant_bytes(model.base.lm_head_weight()).0;
let kv_row = u64::from(d.num_kv_heads * d.attn_head_dim);
let kv_bytes_per_token_f32 = attn_layers * 2 * kv_row * 4;
let recurrent_per_state = model.layers.len() as u64
* 4
* u64::from(
d.conv_dim * (d.conv_kernel - 1) + d.num_v_heads * d.head_v_dim * d.head_k_dim,
);
let own_state = recurrent_per_state
+ kv_bytes_per_token_f32 * seq_len.min(super::DEFAULT_MAX_SEQ_LEN) as u64;
let per_token_scratch = 4 * u64::from(
8 * d.hidden_dim
+ 2 * d.conv_dim
+ 4 * d.v_dim
+ 8 * d.num_heads * d.attn_head_dim
+ 4 * d.intermediate_dim
+ d.vocab_size,
);
let workspace = workspace_bytes_for(d, largest, seq_len, attention, chunk_rows) as u64
+ recurrent_per_state + own_state
+ per_token_scratch;
crate::capacity::CapacityInputs {
weights_bytes: weights,
kv_bytes_per_token_f32,
seq_len: seq_len as u64,
workspace_bytes: workspace,
overhead_bytes: crate::capacity::OVERHEAD_BYTES,
gpu_free_bytes: gpu_free,
gpu_total_bytes: gpu_total,
f16_kv_decode_available: false,
memory: None,
}
}
#[must_use]
pub fn prefill_attention_candidates_for(
model: &Qwen35Model<'_>,
executor: &crate::cuda::CudaExecutor,
) -> Vec<PrefillAttention> {
prefill_attention_candidates(executor, Self::dims_of(model))
}
#[must_use]
pub fn prefill_attention_mode(&self) -> PrefillAttention {
self.prefill_attention
}
fn projection_weights(&self) -> Vec<&CudaQuantWeight> {
self.layers
.iter()
.flat_map(|l| match l {
CudaLayer::DeltaNet(w) => vec![
&w.attn_qkv,
&w.ssm_alpha,
&w.ssm_beta,
&w.attn_gate,
&w.ssm_out,
&w.ffn_gate,
&w.ffn_up,
&w.ffn_down,
],
CudaLayer::Attention(w) => vec![
&w.attn_q,
&w.attn_k,
&w.attn_v,
&w.attn_output,
&w.ffn_gate,
&w.ffn_up,
&w.ffn_down,
],
})
.collect()
}
fn largest_projection_elems(&self) -> usize {
self.projection_weights()
.iter()
.map(|q| q.n as usize * q.k as usize)
.max()
.unwrap_or(0)
}
pub(crate) fn warm_prefill_weights(&mut self) -> usize {
use crate::cuda::{qwen35_prefill_gemm_mode, Qwen35PrefillGemm};
self.executor.set_qwen35_prefill_f16(false);
if qwen35_prefill_gemm_mode() != Qwen35PrefillGemm::F16 {
return 0;
}
let weights: Vec<(WeightQuantType, u64, u32, u32)> = self
.projection_weights()
.iter()
.map(|w| (w.qtype, w.ptr, w.n, w.k))
.collect();
let bytes: usize = weights
.iter()
.map(|w| w.2 as usize * w.3 as usize * 2)
.sum();
let free = match self.executor.context().memory_info() {
Ok((free, _)) => free,
Err(e) => {
eprintln!("[qwen35] fp16 prefill weights NOT prewarmed ({e}); prefill uses f32");
return 0;
},
};
if !f16_prewarm_fits(bytes, free) {
eprintln!(
"[qwen35] fp16 prefill weights NOT prewarmed: {} MiB needed, {} MiB free; prefill uses f32",
bytes >> 20,
free >> 20
);
return 0;
}
let t0 = std::time::Instant::now();
for &(qtype, ptr, n, k) in &weights {
if let Err(e) = self.executor.qwen35_fp16_weight(qtype, ptr, n, k) {
let ptrs: Vec<u64> = weights.iter().map(|w| w.1).collect();
self.executor.drop_fp16_weights(&ptrs);
eprintln!("[qwen35] fp16 prefill prewarm failed ({e}); prefill uses f32");
return 0;
}
}
if let Err(e) = self.executor.synchronize() {
let ptrs: Vec<u64> = weights.iter().map(|w| w.1).collect();
self.executor.drop_fp16_weights(&ptrs);
eprintln!("[qwen35] fp16 prefill prewarm failed ({e}); prefill uses f32");
return 0;
}
self.executor.set_qwen35_prefill_f16(true);
eprintln!(
"[qwen35] fp16 prefill weights prewarmed: {} MiB in {} ms (#4313)",
bytes >> 20,
t0.elapsed().as_millis()
);
bytes
}
fn alloc_prefill(&self, rows: usize, total_positions: usize) -> Result<PrefillBuffers> {
let d = self.dims;
let attention = self.prefill_attention_mode();
let ctx = self.executor.context();
let z = |n: usize| {
GpuBuffer::<f32>::new(ctx, n.max(1))
.map_err(|e| gpu_err("qwen35_cuda_prefill_alloc", &e))
};
let hpk = (d.num_heads / d.num_kv_heads.max(1)).max(1) as usize;
let q_dim = (d.num_heads * d.attn_head_dim) as usize;
let kv_dim = (d.num_kv_heads * d.attn_head_dim) as usize;
let (hidden, inter) = (d.hidden_dim as usize, d.intermediate_dim as usize);
let (conv, v, nv) = (
d.conv_dim as usize,
d.v_dim as usize,
d.num_v_heads as usize,
);
Ok(PrefillBuffers {
rows,
attention,
attn_rows: attention_rows_for(d, total_positions, self.prefill_rows),
x: z(rows * hidden)?,
normed: z(rows * hidden)?,
post_normed: z(rows * hidden)?,
proj: z(rows * hidden)?,
conv_in: z(rows * conv)?,
conv_out: z(rows * conv)?,
alpha_raw: z(rows * nv)?,
beta_raw: z(rows * nv)?,
dt: z(rows * nv)?,
beta: z(rows * nv)?,
gate: z(rows * v)?,
out_h: z(rows * v)?,
ssm_out_in: z(rows * v)?,
q_full: z(rows * 2 * q_dim)?,
q: z(rows * q_dim)?,
q_normed: z(rows * q_dim)?,
attn_gate: z(rows * q_dim)?,
k_raw: z(rows * kv_dim)?,
attn_out_in: z(rows * q_dim)?,
scores: z(match attention {
PrefillAttention::CublasF32 => {
hpk * attention_rows_for(d, total_positions, self.prefill_rows)
* total_positions
},
PrefillAttention::FlashF16In => 1,
})?,
ffn_gate: z(rows * inter)?,
ffn_up: z(rows * inter)?,
ffn_act: z(rows * inter)?,
})
}
pub fn prefill(
&mut self,
tokens: &[u32],
state: &mut Qwen35CudaState,
pos0: usize,
) -> Result<Vec<f32>> {
self.prefill_check(tokens, state, pos0)?;
let end = pos0 + tokens.len();
let rows = self.prefill_chunk_rows(end);
let bufs = self.alloc_prefill(rows, end)?;
let mut pos = pos0;
let mut last_rows = 0;
for chunk in tokens.chunks(rows) {
self.prefill_chunk(&bufs, chunk, state, pos)?;
pos += chunk.len();
last_rows = chunk.len();
}
self.prefill_tail(&bufs, last_rows - 1)
}
pub fn prefill_logits_at(
&mut self,
tokens: &[u32],
state: &mut Qwen35CudaState,
pos0: usize,
positions: &[usize],
) -> Result<Vec<Vec<f32>>> {
self.prefill_check(tokens, state, pos0)?;
let end = pos0 + tokens.len();
if positions.windows(2).any(|w| w[0] >= w[1])
|| positions.iter().any(|&p| p < pos0 || p >= end)
{
return Err(RealizarError::InvalidShape {
reason: format!(
"qwen35_cuda prefill: requested positions must ascend within {pos0}..{end}"
),
});
}
let rows = self.prefill_chunk_rows(end);
let bufs = self.alloc_prefill(rows, end)?;
let mut pos = pos0;
let mut want = positions.iter().copied().peekable();
let mut out = Vec::with_capacity(positions.len());
for chunk in tokens.chunks(rows) {
self.prefill_chunk(&bufs, chunk, state, pos)?;
while let Some(&p) = want.peek() {
if p >= pos + chunk.len() {
break;
}
out.push(self.prefill_tail(&bufs, p - pos)?);
want.next();
}
pos += chunk.len();
}
Ok(out)
}
fn prefill_check(&self, tokens: &[u32], state: &Qwen35CudaState, pos0: usize) -> Result<()> {
if tokens.is_empty() {
return Err(RealizarError::InvalidShape {
reason: "qwen35_cuda prefill: the prompt is empty".to_string(),
});
}
if pos0 != state.kv_len {
return Err(RealizarError::InvalidShape {
reason: format!(
"qwen35_cuda prefill: pos0 {pos0} is not where the state stands ({} KV rows \
written) — a prefill continues the state, it cannot skip or rewind it",
state.kv_len
),
});
}
let end = pos0 + tokens.len();
if end > state.max_seq_len {
return Err(RealizarError::InvalidShape {
reason: format!(
"qwen35_cuda prefill: positions {pos0}..{end} are past the KV cache ({} rows)",
state.max_seq_len
),
});
}
let hidden = self.dims.hidden_dim as usize;
let vocab_rows = self.model.base.token_embedding().len() / hidden;
if let Some(bad) = tokens.iter().find(|&&t| t as usize >= vocab_rows) {
return Err(RealizarError::InvalidShape {
reason: format!(
"qwen35_cuda prefill: token {bad} is outside the {vocab_rows}-row embedding table"
),
});
}
Ok(())
}
fn prefill_tail(&mut self, bufs: &PrefillBuffers, row: usize) -> Result<Vec<f32>> {
let d = self.dims;
let hidden = d.hidden_dim as usize;
let last = View::at(bufs.x.as_ptr() + (row * hidden * 4) as u64, hidden);
let err = |e: trueno_gpu::GpuError| gpu_err("qwen35_cuda_prefill_tail", &e);
self.executor
.rmsnorm_into(
&last,
&self.output_norm,
&self.out_normed,
d.hidden_dim,
d.eps,
)
.map_err(err)?;
self.executor
.gemv_dispatch(
self.lm_head.qtype,
self.lm_head.ptr,
&self.out_normed,
&self.logits_buf,
self.lm_head.n,
self.lm_head.k,
)
.map_err(err)?;
self.executor.sync_stream().map_err(err)?;
let mut logits = vec![0.0f32; d.vocab_size as usize];
self.logits_buf.copy_to_host(&mut logits).map_err(err)?;
Ok(logits)
}
fn prefill_chunk(
&mut self,
bufs: &PrefillBuffers,
chunk: &[u32],
state: &mut Qwen35CudaState,
pos: usize,
) -> Result<()> {
let hidden = self.dims.hidden_dim as usize;
let n = chunk.len();
debug_assert!(n <= bufs.rows);
let err = |e: trueno_gpu::GpuError| gpu_err("qwen35_cuda_prefill", &e);
self.executor.sync_stream().map_err(err)?;
let table = self.model.base.token_embedding();
let mut host = Vec::with_capacity(n * hidden);
for &t in chunk {
let start = t as usize * hidden;
host.extend_from_slice(&table[start..start + hidden]);
}
let mut x = View::at(bufs.x.as_ptr(), bufs.rows * hidden);
x.0.copy_from_host_at(&host, 0).map_err(err)?;
for il in 0..self.layers.len() {
match self.layers[il] {
CudaLayer::DeltaNet(_) => self
.prefill_deltanet(bufs, state, il, n)
.map_err(|e| gpu_err("qwen35_cuda_prefill_deltanet", &e))?,
CudaLayer::Attention(_) => self
.prefill_attention(bufs, state, il, n, pos)
.map_err(|e| gpu_err("qwen35_cuda_prefill_attention", &e))?,
}
}
state.kv_len = state.kv_len.max(pos + n);
Ok(())
}
#[allow(clippy::too_many_lines)]
fn prefill_deltanet(
&mut self,
b: &PrefillBuffers,
state: &Qwen35CudaState,
il: usize,
n: usize,
) -> std::result::Result<(), trueno_gpu::GpuError> {
let d = self.dims;
let CudaLayer::DeltaNet(w) = &self.layers[il] else {
unreachable!("the caller matched a DeltaNet layer")
};
let ex = &mut self.executor;
let rows = n as u32;
let f = 4u64;
ex.batched_rmsnorm_into(&b.x, &w.attn_norm, &b.normed, d.hidden_dim, rows, d.eps)?;
let p = &w.attn_qkv;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.normed.as_ptr(),
b.conv_in.as_ptr(),
rows,
p.n,
p.k,
p.n,
)?;
ex.qwen35_conv1d_rows(
b.conv_in.as_ptr(),
state.conv[il].as_ptr(),
w.conv1d_weight.as_ptr(),
b.conv_out.as_ptr(),
d.conv_dim,
d.conv_kernel,
rows,
)?;
let conv_base = b.conv_out.as_ptr();
ex.qwen35_l2_norm_rows(
conv_base,
d.head_k_dim,
d.num_k_heads,
d.eps,
d.conv_dim,
rows,
)?;
ex.qwen35_l2_norm_rows(
conv_base + u64::from(d.k_dim) * f,
d.head_k_dim,
d.num_k_heads,
d.eps,
d.conv_dim,
rows,
)?;
let p = &w.ssm_alpha;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.normed.as_ptr(),
b.alpha_raw.as_ptr(),
rows,
p.n,
p.k,
p.n,
)?;
let p = &w.ssm_beta;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.normed.as_ptr(),
b.beta_raw.as_ptr(),
rows,
p.n,
p.k,
p.n,
)?;
ex.qwen35_gates_rows(
b.alpha_raw.as_ptr(),
w.ssm_dt_bias.as_ptr(),
w.ssm_a.as_ptr(),
b.beta_raw.as_ptr(),
b.dt.as_ptr(),
b.beta.as_ptr(),
d.num_v_heads,
rows,
)?;
let p = &w.attn_gate;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.normed.as_ptr(),
b.gate.as_ptr(),
rows,
p.n,
p.k,
p.n,
)?;
ex.qwen35_delta_rule_scan(
conv_base,
conv_base + u64::from(d.k_dim) * f,
conv_base + u64::from(2 * d.k_dim) * f,
b.beta.as_ptr(),
b.dt.as_ptr(),
state.ssm[il].as_ptr(),
b.out_h.as_ptr(),
(d.num_k_heads, d.head_k_dim, d.num_v_heads, d.head_v_dim),
d.conv_dim,
rows,
)?;
ex.gdn_gated_rmsnorm_into(
&b.out_h,
&b.gate,
&w.ssm_norm_weight,
&b.ssm_out_in,
d.head_v_dim,
rows * d.num_v_heads,
d.eps,
)?;
let p = &w.ssm_out;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.ssm_out_in.as_ptr(),
b.proj.as_ptr(),
rows,
p.n,
p.k,
p.n,
)?;
ex.residual_add_into(&b.x, &b.proj, &b.x, rows * d.hidden_dim)?;
Self::prefill_ffn_rows(
ex,
b,
d,
rows,
&w.post_attention_norm,
[&w.ffn_gate, &w.ffn_up, &w.ffn_down],
)
}
#[allow(clippy::too_many_lines)]
fn prefill_attention(
&mut self,
b: &PrefillBuffers,
state: &Qwen35CudaState,
il: usize,
n: usize,
pos: usize,
) -> std::result::Result<(), trueno_gpu::GpuError> {
let d = self.dims;
let CudaLayer::Attention(w) = &self.layers[il] else {
unreachable!("the caller matched an attention layer")
};
let ex = &mut self.executor;
let rows = n as u32;
let q_dim = d.num_heads * d.attn_head_dim;
let kv_dim = d.num_kv_heads * d.attn_head_dim;
let (k_cache, v_cache) = state.kv[il]
.as_ref()
.expect("an attention layer owns a KV cache");
let row_off = (pos as u64) * u64::from(kv_dim) * 4;
let k_rows = View::at(k_cache.as_ptr() + row_off, n * kv_dim as usize);
let v_rows_ptr = v_cache.as_ptr() + row_off;
ex.batched_rmsnorm_into(&b.x, &w.attn_norm, &b.normed, d.hidden_dim, rows, d.eps)?;
let p = &w.attn_q;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.normed.as_ptr(),
b.q_full.as_ptr(),
rows,
p.n,
p.k,
p.n,
)?;
let p = &w.attn_k;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.normed.as_ptr(),
b.k_raw.as_ptr(),
rows,
p.n,
p.k,
p.n,
)?;
let p = &w.attn_v;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.normed.as_ptr(),
v_rows_ptr,
rows,
p.n,
p.k,
kv_dim,
)?;
ex.gdn_split_interleaved_into(
&b.q_full,
&b.q,
&b.attn_gate,
rows * d.num_heads,
d.attn_head_dim,
)?;
ex.per_head_rmsnorm_into(
&b.q,
&w.attn_q_norm,
&b.q_normed,
d.attn_head_dim,
rows * d.num_heads,
d.eps,
)?;
ex.per_head_rmsnorm_into(
&b.k_raw,
&w.attn_k_norm,
&k_rows,
d.attn_head_dim,
rows * d.num_kv_heads,
d.eps,
)?;
let pos32 = u32::try_from(pos).unwrap_or(u32::MAX);
ex.qwen35_rope_rows(
b.q_normed.as_ptr(),
d.num_heads,
d.attn_head_dim,
d.n_rot,
q_dim,
rows,
pos32,
d.theta_scale,
)?;
ex.qwen35_rope_rows(
k_rows.as_ptr(),
d.num_kv_heads,
d.attn_head_dim,
d.n_rot,
kv_dim,
rows,
pos32,
d.theta_scale,
)?;
match b.attention {
PrefillAttention::FlashF16In => ex.qwen35_flash_prefill_attention(
b.q_normed.as_ptr(),
k_cache.as_ptr(),
v_cache.as_ptr(),
b.attn_out_in.as_ptr(),
rows,
pos32,
d.num_heads,
d.num_kv_heads,
)?,
PrefillAttention::CublasF32 => ex.qwen35_prefill_attention(
b.q_normed.as_ptr(),
k_cache.as_ptr(),
v_cache.as_ptr(),
b.attn_out_in.as_ptr(),
b.scores.as_ptr(),
rows,
b.attn_rows as u32,
pos32,
d.num_heads,
d.num_kv_heads,
d.attn_head_dim,
)?,
}
ex.gdn_sigmoid_gate_into(&b.attn_out_in, &b.attn_gate, rows * q_dim)?;
let p = &w.attn_output;
ex.qwen35_project_rows(
p.qtype,
p.ptr,
b.attn_out_in.as_ptr(),
b.proj.as_ptr(),
rows,
p.n,
p.k,
p.n,
)?;
ex.residual_add_into(&b.x, &b.proj, &b.x, rows * d.hidden_dim)?;
Self::prefill_ffn_rows(
ex,
b,
d,
rows,
&w.post_attention_norm,
[&w.ffn_gate, &w.ffn_up, &w.ffn_down],
)
}
fn prefill_ffn_rows(
ex: &mut crate::cuda::CudaExecutor,
b: &PrefillBuffers,
d: super::Qwen35CudaDims,
rows: u32,
post_norm: &GpuBuffer<f32>,
[gate, up, down]: [&super::CudaQuantWeight; 3],
) -> std::result::Result<(), trueno_gpu::GpuError> {
ex.batched_rmsnorm_into(&b.x, post_norm, &b.post_normed, d.hidden_dim, rows, d.eps)?;
ex.qwen35_project_rows(
gate.qtype,
gate.ptr,
b.post_normed.as_ptr(),
b.ffn_gate.as_ptr(),
rows,
gate.n,
gate.k,
gate.n,
)?;
ex.qwen35_project_rows(
up.qtype,
up.ptr,
b.post_normed.as_ptr(),
b.ffn_up.as_ptr(),
rows,
up.n,
up.k,
up.n,
)?;
ex.fused_swiglu_into(
&b.ffn_gate,
&b.ffn_up,
&b.ffn_act,
rows * d.intermediate_dim,
)?;
ex.qwen35_project_rows(
down.qtype,
down.ptr,
b.ffn_act.as_ptr(),
b.proj.as_ptr(),
rows,
down.n,
down.k,
down.n,
)?;
ex.residual_add_into(&b.x, &b.proj, &b.x, rows * d.hidden_dim)
}
}
#[cfg(test)]
#[path = "forward_qwen35_cuda_prefill_tests.rs"]
mod prefill_tests;