use super::config::DFlashConfig;
use anyhow::{anyhow, Result};
pub struct PrefillCapture<'a> {
pub target_layer_ids: &'a [usize],
pub hidden_output: &'a mut [f32],
pub per_position_argmaxes: Option<&'a mut [u32]>,
}
pub fn trim_capture_to(
session: &mut super::hidden_capture::DFlashCaptureSession,
new_seq_len: usize,
) {
if new_seq_len >= session.seq_len {
return;
}
let n_layers = session.target_layer_ids.len();
let hs = session.hidden_size;
let old_seq_len = session.seq_len;
let mut new_buf = vec![0.0f32; n_layers * new_seq_len * hs];
for l in 0..n_layers {
for t in 0..new_seq_len {
let src = (l * old_seq_len + t) * hs;
let dst = (l * new_seq_len + t) * hs;
new_buf[dst..dst + hs].copy_from_slice(&session.hidden_output[src..src + hs]);
}
}
session.hidden_output = new_buf;
session.seq_len = new_seq_len;
if let Some(pa) = session.per_position_argmaxes.as_mut() {
pa.truncate(new_seq_len);
}
}
impl<'a> PrefillCapture<'a> {
pub fn validate(&self, seq_len: usize, hidden_size: usize) -> Result<()> {
let expected = self.target_layer_ids.len() * seq_len * hidden_size;
if self.hidden_output.len() != expected {
return Err(anyhow!(
"PrefillCapture: hidden_output len {} != target_layer_ids({}) * seq_len({}) * hidden_size({}) = {}",
self.hidden_output.len(),
self.target_layer_ids.len(),
seq_len,
hidden_size,
expected
));
}
if let Some(ref pa) = self.per_position_argmaxes {
if pa.len() != seq_len {
return Err(anyhow!(
"PrefillCapture: per_position_argmaxes len {} != seq_len {}",
pa.len(),
seq_len
));
}
}
for w in self.target_layer_ids.windows(2) {
if w[0] >= w[1] {
return Err(anyhow!(
"PrefillCapture: target_layer_ids must be strictly increasing; got {:?}",
self.target_layer_ids
));
}
}
Ok(())
}
pub fn offset_for(
capture_layer_idx: usize,
token_pos: usize,
seq_len: usize,
hidden_size: usize,
) -> usize {
(capture_layer_idx * seq_len + token_pos) * hidden_size
}
pub fn write_layer_slab(
&mut self,
capture_layer_idx: usize,
pf_hidden_data: &[f32],
seq_len: usize,
hidden_size: usize,
) -> Result<()> {
let slab_len = seq_len * hidden_size;
if pf_hidden_data.len() != slab_len {
return Err(anyhow!(
"PrefillCapture::write_layer_slab: pf_hidden_data len {} != seq_len({}) * hidden_size({}) = {}",
pf_hidden_data.len(), seq_len, hidden_size, slab_len
));
}
let start = Self::offset_for(capture_layer_idx, 0, seq_len, hidden_size);
let end = start + slab_len;
if end > self.hidden_output.len() {
return Err(anyhow!(
"PrefillCapture::write_layer_slab: write out of bounds (start {}, end {}, buf {})",
start,
end,
self.hidden_output.len()
));
}
self.hidden_output[start..end].copy_from_slice(pf_hidden_data);
Ok(())
}
pub fn permute_to_concat(&self, seq_len: usize, hidden_size: usize) -> Vec<f32> {
let n_layers = self.target_layer_ids.len();
let mut out = vec![0.0f32; seq_len * n_layers * hidden_size];
for layer_idx in 0..n_layers {
for t in 0..seq_len {
let src = (layer_idx * seq_len + t) * hidden_size;
let dst = (t * n_layers + layer_idx) * hidden_size;
out[dst..dst + hidden_size]
.copy_from_slice(&self.hidden_output[src..src + hidden_size]);
}
}
out
}
}
pub fn capture_buffer_len(cfg: &DFlashConfig, seq_len: usize) -> usize {
cfg.target_layer_ids.len() * seq_len * cfg.hidden_size
}
pub fn extract_drafter_concat(
hidden_output: &[f32],
combined_capture_layer_ids: &[usize],
drafter_target_layer_ids: &[usize],
seq_len: usize,
hidden_size: usize,
) -> Result<Vec<f32>> {
let drafter_n = drafter_target_layer_ids.len();
let total_in = combined_capture_layer_ids.len() * seq_len * hidden_size;
if hidden_output.len() != total_in {
return Err(anyhow!(
"extract_drafter_concat: hidden_output len {} != combined_layers({}) * seq_len({}) * hs({}) = {}",
hidden_output.len(), combined_capture_layer_ids.len(), seq_len, hidden_size, total_in
));
}
let mut drafter_to_combined: Vec<usize> = Vec::with_capacity(drafter_n);
for &dl in drafter_target_layer_ids {
match combined_capture_layer_ids.iter().position(|&c| c == dl) {
Some(idx) => drafter_to_combined.push(idx),
None => {
return Err(anyhow!(
"extract_drafter_concat: drafter target_layer_id {} not in combined capture set {:?}",
dl, combined_capture_layer_ids
));
}
}
}
let mut out = vec![0.0f32; seq_len * drafter_n * hidden_size];
for (drafter_l, &combined_idx) in drafter_to_combined.iter().enumerate() {
for t in 0..seq_len {
let src = (combined_idx * seq_len + t) * hidden_size;
let dst = (t * drafter_n + drafter_l) * hidden_size;
out[dst..dst + hidden_size].copy_from_slice(&hidden_output[src..src + hidden_size]);
}
}
Ok(out)
}
pub fn extract_final_layer_slab(
hidden_output: &[f32],
combined_capture_layer_ids: &[usize],
final_layer_idx: usize,
seq_len: usize,
hidden_size: usize,
) -> Result<Vec<f32>> {
let final_combined_idx = combined_capture_layer_ids
.iter()
.position(|&c| c == final_layer_idx)
.ok_or_else(|| {
anyhow!(
"extract_final_layer_slab: final_layer_idx {} not in combined capture set {:?}",
final_layer_idx,
combined_capture_layer_ids
)
})?;
let start = final_combined_idx * seq_len * hidden_size;
let end = start + seq_len * hidden_size;
if end > hidden_output.len() {
return Err(anyhow!(
"extract_final_layer_slab: end offset {} > buffer len {}",
end,
hidden_output.len()
));
}
Ok(hidden_output[start..end].to_vec())
}
#[derive(Debug, Default, Clone)]
pub struct DFlashCaptureSession {
pub target_layer_ids: Vec<usize>,
pub hidden_output: Vec<f32>,
pub per_position_argmaxes: Option<Vec<u32>>,
pub seq_len: usize,
pub hidden_size: usize,
}
impl DFlashCaptureSession {
pub fn new(
target_layer_ids: Vec<usize>,
seq_len: usize,
hidden_size: usize,
with_argmaxes: bool,
) -> Self {
let hidden_output = vec![0.0f32; target_layer_ids.len() * seq_len * hidden_size];
let per_position_argmaxes = if with_argmaxes {
Some(vec![0u32; seq_len])
} else {
None
};
Self {
target_layer_ids,
hidden_output,
per_position_argmaxes,
seq_len,
hidden_size,
}
}
pub fn capture_index_for(&self, target_layer_idx: usize) -> Option<usize> {
self.target_layer_ids
.iter()
.position(|&i| i == target_layer_idx)
}
pub fn write_layer_slab(
&mut self,
capture_layer_idx: usize,
pf_hidden_data: &[f32],
) -> Result<()> {
let slab_len = self.seq_len * self.hidden_size;
if pf_hidden_data.len() != slab_len {
return Err(anyhow!(
"DFlashCaptureSession::write_layer_slab: pf_hidden_data len {} != seq_len({}) * hidden_size({}) = {}",
pf_hidden_data.len(), self.seq_len, self.hidden_size, slab_len
));
}
let start =
PrefillCapture::offset_for(capture_layer_idx, 0, self.seq_len, self.hidden_size);
let end = start + slab_len;
if end > self.hidden_output.len() {
return Err(anyhow!(
"DFlashCaptureSession::write_layer_slab: write out of bounds (start {}, end {}, buf {})",
start, end, self.hidden_output.len()
));
}
self.hidden_output[start..end].copy_from_slice(pf_hidden_data);
Ok(())
}
pub fn as_view(&mut self) -> PrefillCapture<'_> {
PrefillCapture {
target_layer_ids: &self.target_layer_ids,
hidden_output: &mut self.hidden_output,
per_position_argmaxes: self.per_position_argmaxes.as_deref_mut(),
}
}
}
pub fn append_capture_positions(
prior: &DFlashCaptureSession,
verify_captured: &DFlashCaptureSession,
n_committed: usize,
) -> Result<DFlashCaptureSession> {
debug_assert_eq!(prior.target_layer_ids, verify_captured.target_layer_ids);
debug_assert_eq!(prior.hidden_size, verify_captured.hidden_size);
if n_committed > verify_captured.seq_len {
return Err(anyhow!(
"append_capture_positions: n_committed ({}) > verify_captured.seq_len ({})",
n_committed,
verify_captured.seq_len
));
}
let hs = prior.hidden_size;
let num_layers = prior.target_layer_ids.len();
let prior_seq = prior.seq_len;
let new_seq = prior_seq + n_committed;
let mut new_hidden = vec![0.0f32; num_layers * new_seq * hs];
for layer in 0..num_layers {
let new_layer_base = layer * new_seq * hs;
let prior_layer_base = layer * prior_seq * hs;
new_hidden[new_layer_base..new_layer_base + prior_seq * hs].copy_from_slice(
&prior.hidden_output[prior_layer_base..prior_layer_base + prior_seq * hs],
);
let verify_layer_base = layer * verify_captured.seq_len * hs;
let new_extended = new_layer_base + prior_seq * hs;
new_hidden[new_extended..new_extended + n_committed * hs].copy_from_slice(
&verify_captured.hidden_output[verify_layer_base..verify_layer_base + n_committed * hs],
);
}
Ok(DFlashCaptureSession {
target_layer_ids: prior.target_layer_ids.clone(),
hidden_output: new_hidden,
per_position_argmaxes: None,
seq_len: new_seq,
hidden_size: hs,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::spec_decode::dflash::config::DFlashConfig;
fn gemma4_cfg() -> DFlashConfig {
DFlashConfig::from_json_str(super::super::config::tests::GEMMA4_26B_A4B_DFLASH_CONFIG)
.expect("config parse")
}
#[test]
fn extract_drafter_concat_picks_right_slabs() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let combined = vec![1, 6, 11, 17, 22, 27, 29];
let drafter_ids = vec![1, 6, 11, 17, 22, 27];
let seq_len = 3usize;
let hs = 4usize;
let mut hidden = vec![0.0f32; 7 * 3 * 4];
for cli in 0..7 {
for t in 0..3 {
for d in 0..4 {
hidden[(cli * 3 + t) * 4 + d] = (cli * 100 + t * 10 + d) as f32;
}
}
}
let out = extract_drafter_concat(&hidden, &combined, &drafter_ids, seq_len, hs).unwrap();
for t in 0..seq_len {
for drafter_l in 0..drafter_ids.len() {
let combined_idx = drafter_l; for d in 0..hs {
let expected = (combined_idx * 100 + t * 10 + d) as f32;
let actual = out[(t * drafter_ids.len() + drafter_l) * hs + d];
assert_eq!(actual, expected, "t={t} drafter_l={drafter_l} d={d}");
}
}
}
}
#[test]
fn extract_final_layer_slab_picks_correct_layer() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let combined = vec![1, 6, 11, 17, 22, 27, 29];
let seq_len = 3usize;
let hs = 4usize;
let mut hidden = vec![0.0f32; 7 * 3 * 4];
for cli in 0..7 {
for t in 0..3 {
for d in 0..4 {
hidden[(cli * 3 + t) * 4 + d] = if cli == 6 { 9.0 } else { 1.0 };
}
}
}
let slab = extract_final_layer_slab(&hidden, &combined, 29, seq_len, hs).unwrap();
assert_eq!(slab.len(), seq_len * hs);
for v in slab.iter() {
assert_eq!(*v, 9.0);
}
}
#[test]
fn trim_capture_to_compacts_buffer_and_keeps_first_positions() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut sess = DFlashCaptureSession::new(vec![1, 6, 11], 5, 4, true);
for l in 0..3 {
for t in 0..5 {
for d in 0..4 {
sess.hidden_output[(l * 5 + t) * 4 + d] = (l * 100 + t * 10 + d) as f32;
}
}
}
if let Some(pa) = sess.per_position_argmaxes.as_mut() {
for (i, v) in pa.iter_mut().enumerate() {
*v = i as u32;
}
}
trim_capture_to(&mut sess, 2);
assert_eq!(sess.seq_len, 2);
assert_eq!(sess.hidden_output.len(), 3 * 2 * 4);
for l in 0..3 {
for t in 0..2 {
for d in 0..4 {
let expected = (l * 100 + t * 10 + d) as f32;
let actual = sess.hidden_output[(l * 2 + t) * 4 + d];
assert_eq!(actual, expected, "l={l} t={t} d={d}");
}
}
}
assert_eq!(
sess.per_position_argmaxes.as_ref().map(|v| v.len()),
Some(2)
);
}
#[test]
fn trim_capture_to_no_op_when_new_geq_old() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut sess = DFlashCaptureSession::new(vec![1, 6], 3, 4, false);
let orig_len = sess.hidden_output.len();
for v in sess.hidden_output.iter_mut() {
*v = 7.0;
}
trim_capture_to(&mut sess, 5); assert_eq!(sess.seq_len, 3);
assert_eq!(sess.hidden_output.len(), orig_len);
assert!(sess.hidden_output.iter().all(|&v| v == 7.0));
}
#[test]
fn extract_drafter_concat_errors_on_missing_layer() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let combined = vec![1, 6, 29];
let drafter_ids = vec![1, 100];
let mut hidden = vec![0.0f32; 3 * 3 * 4];
for v in hidden.iter_mut() {
*v = 1.0;
}
let err = extract_drafter_concat(&hidden, &combined, &drafter_ids, 3, 4).unwrap_err();
assert!(format!("{err}").contains("not in combined"));
}
#[test]
fn capture_buffer_len_matches_drafter_fc_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = gemma4_cfg();
assert_eq!(capture_buffer_len(&cfg, 4), 6 * 4 * 2816);
}
#[test]
fn validate_catches_wrong_buffer_size() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let target_layer_ids = vec![1, 6, 11, 17, 22, 27];
let mut hidden = vec![0.0f32; 100]; let cap = PrefillCapture {
target_layer_ids: &target_layer_ids,
hidden_output: &mut hidden,
per_position_argmaxes: None,
};
let err = cap.validate(4, 2816).unwrap_err();
assert!(format!("{err}").contains("hidden_output len"));
}
#[test]
fn validate_catches_non_monotonic_layer_ids() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let target_layer_ids = vec![1, 11, 6, 17, 22, 27]; let mut hidden = vec![0.0f32; 6 * 4 * 2816];
let cap = PrefillCapture {
target_layer_ids: &target_layer_ids,
hidden_output: &mut hidden,
per_position_argmaxes: None,
};
let err = cap.validate(4, 2816).unwrap_err();
assert!(format!("{err}").contains("strictly increasing"));
}
#[test]
fn validate_catches_argmax_size_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let target_layer_ids = vec![1, 6, 11, 17, 22, 27];
let mut hidden = vec![0.0f32; 6 * 4 * 2816];
let mut argmaxes = vec![0u32; 99]; let cap = PrefillCapture {
target_layer_ids: &target_layer_ids,
hidden_output: &mut hidden,
per_position_argmaxes: Some(&mut argmaxes),
};
let err = cap.validate(4, 2816).unwrap_err();
assert!(format!("{err}").contains("per_position_argmaxes"));
}
#[test]
fn offset_for_matches_row_major_layout() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
assert_eq!(PrefillCapture::offset_for(0, 0, 4, 2816), 0);
assert_eq!(PrefillCapture::offset_for(0, 1, 4, 2816), 2816);
assert_eq!(PrefillCapture::offset_for(1, 0, 4, 2816), 4 * 2816);
assert_eq!(
PrefillCapture::offset_for(5, 3, 4, 2816),
(5 * 4 + 3) * 2816
);
}
#[test]
fn write_layer_slab_places_data_correctly() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let target_layer_ids = vec![1, 6, 11, 17, 22, 27];
let mut hidden = vec![0.0f32; 6 * 4 * 2816];
let mut cap = PrefillCapture {
target_layer_ids: &target_layer_ids,
hidden_output: &mut hidden,
per_position_argmaxes: None,
};
let slab = vec![1.0f32; 4 * 2816];
cap.write_layer_slab(2, &slab, 4, 2816)
.expect("write_layer_slab");
let layer2_start = 2 * 4 * 2816;
let layer2_end = 3 * 4 * 2816;
for i in 0..hidden.len() {
let expected = if i >= layer2_start && i < layer2_end {
1.0
} else {
0.0
};
assert_eq!(hidden[i], expected, "at index {i}");
}
}
#[test]
fn session_new_allocates_correct_sizes() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let sess = DFlashCaptureSession::new(vec![1, 6, 11, 17, 22, 27], 4, 2816, true);
assert_eq!(sess.hidden_output.len(), 6 * 4 * 2816);
assert_eq!(
sess.per_position_argmaxes.as_ref().map(|v| v.len()),
Some(4)
);
assert_eq!(sess.seq_len, 4);
assert_eq!(sess.hidden_size, 2816);
let sess2 = DFlashCaptureSession::new(vec![1, 6], 8, 128, false);
assert_eq!(sess2.hidden_output.len(), 2 * 8 * 128);
assert!(sess2.per_position_argmaxes.is_none());
}
#[test]
fn session_capture_index_for_finds_target_layers() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let sess = DFlashCaptureSession::new(vec![1, 6, 11, 17, 22, 27], 4, 2816, false);
assert_eq!(sess.capture_index_for(1), Some(0));
assert_eq!(sess.capture_index_for(11), Some(2));
assert_eq!(sess.capture_index_for(27), Some(5));
assert_eq!(sess.capture_index_for(0), None);
assert_eq!(sess.capture_index_for(5), None);
assert_eq!(sess.capture_index_for(28), None);
}
#[test]
fn session_write_layer_slab_places_data_at_correct_offset() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut sess = DFlashCaptureSession::new(vec![1, 6, 11, 17, 22, 27], 3, 4, false);
let slab = vec![5.0f32; 3 * 4];
sess.write_layer_slab(2, &slab).expect("write");
for i in 0..sess.hidden_output.len() {
let expected = if (24..36).contains(&i) { 5.0 } else { 0.0 };
assert_eq!(sess.hidden_output[i], expected, "i={i}");
}
}
#[test]
fn session_as_view_borrows_buffers_consistently() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let mut sess = DFlashCaptureSession::new(vec![1, 6], 2, 4, true);
let slab = vec![3.0f32; 2 * 4];
sess.write_layer_slab(0, &slab).expect("write");
let view = sess.as_view();
for i in 0..8 {
assert_eq!(view.hidden_output[i], 3.0);
}
assert!(view.per_position_argmaxes.is_some());
}
#[test]
fn option_a_round2_prior_captured_delivers_correct_new_row() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafter_target_layer_ids: Vec<usize> = vec![1, 6, 11, 17, 22, 27];
let final_layer_idx = 29usize;
let mut combined_ids: Vec<usize> = drafter_target_layer_ids.clone();
combined_ids.push(final_layer_idx);
combined_ids.sort_unstable();
combined_ids.dedup();
let prompt_len = 10usize;
let block_size = 8usize;
let hs = 16usize;
let mut prior_captured =
DFlashCaptureSession::new(combined_ids.clone(), prompt_len, hs, false);
for (cli, &_clid) in combined_ids.iter().enumerate() {
for t in 0..prompt_len {
for d in 0..hs {
let off = (cli * prompt_len + t) * hs + d;
prior_captured.hidden_output[off] = (cli * 1000 + t * 10 + d) as f32;
}
}
}
let mut verify_captured =
DFlashCaptureSession::new(combined_ids.clone(), block_size, hs, false);
for (cli, &_clid) in combined_ids.iter().enumerate() {
for t in 0..block_size {
for d in 0..hs {
let off = (cli * block_size + t) * hs + d;
verify_captured.hidden_output[off] =
100_000.0 + (cli * 1000 + t * 10 + d) as f32;
}
}
}
let n_committed = 1usize;
let next_prior = append_capture_positions(&prior_captured, &verify_captured, n_committed)
.expect("append_capture_positions");
assert_eq!(next_prior.seq_len, prompt_len + n_committed);
let drafter_concat = extract_drafter_concat(
&next_prior.hidden_output,
&combined_ids,
&drafter_target_layer_ids,
next_prior.seq_len,
hs,
)
.expect("extract_drafter_concat");
let drafter_n = drafter_target_layer_ids.len();
let row_stride = drafter_n * hs;
let drafter_cached_seq_len = prompt_len;
let new_rows_start = drafter_cached_seq_len * row_stride;
let drafter_concat_new: &[f32] = &drafter_concat[new_rows_start..];
assert_eq!(drafter_concat_new.len(), n_committed * row_stride);
for (drafter_l, &dl_id) in drafter_target_layer_ids.iter().enumerate() {
let combined_idx = combined_ids
.iter()
.position(|&c| c == dl_id)
.expect("drafter layer in combined");
let verify_layer_base = combined_idx * block_size * hs;
let verify_pos0_base = verify_layer_base; let expected = &verify_captured.hidden_output[verify_pos0_base..verify_pos0_base + hs];
let actual = &drafter_concat_new[drafter_l * hs..(drafter_l + 1) * hs];
assert_eq!(
actual, expected,
"drafter_l={} (target_layer_id={}, combined_idx={}): \
drafter_concat_new row 0 must equal verify_captured[combined={}, pos=0, :]",
drafter_l, dl_id, combined_idx, combined_idx,
);
}
}
#[test]
fn option_a_round2_prior_captured_multi_accept_plumbing() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let drafter_target_layer_ids: Vec<usize> = vec![1, 6, 11, 17, 22, 27];
let final_layer_idx = 29usize;
let mut combined_ids: Vec<usize> = drafter_target_layer_ids.clone();
combined_ids.push(final_layer_idx);
combined_ids.sort_unstable();
combined_ids.dedup();
let prompt_len = 8usize;
let block_size = 8usize;
let hs = 16usize;
let mut prior_captured =
DFlashCaptureSession::new(combined_ids.clone(), prompt_len, hs, false);
for cli in 0..combined_ids.len() {
for t in 0..prompt_len {
for d in 0..hs {
let off = (cli * prompt_len + t) * hs + d;
prior_captured.hidden_output[off] = (cli * 1000 + t * 10 + d) as f32;
}
}
}
let mut verify_captured =
DFlashCaptureSession::new(combined_ids.clone(), block_size, hs, false);
for cli in 0..combined_ids.len() {
for t in 0..block_size {
for d in 0..hs {
let off = (cli * block_size + t) * hs + d;
verify_captured.hidden_output[off] =
100_000.0 + (cli * 1000 + t * 10 + d) as f32;
}
}
}
let n_committed = 5usize; let next_prior = append_capture_positions(&prior_captured, &verify_captured, n_committed)
.expect("append");
assert_eq!(next_prior.seq_len, prompt_len + n_committed);
let drafter_concat = extract_drafter_concat(
&next_prior.hidden_output,
&combined_ids,
&drafter_target_layer_ids,
next_prior.seq_len,
hs,
)
.expect("extract");
let drafter_n = drafter_target_layer_ids.len();
let row_stride = drafter_n * hs;
let drafter_cached_seq_len = prompt_len;
let new_rows_start = drafter_cached_seq_len * row_stride;
let drafter_concat_new: &[f32] = &drafter_concat[new_rows_start..];
assert_eq!(drafter_concat_new.len(), n_committed * row_stride);
for t in 0..n_committed {
for (drafter_l, &dl_id) in drafter_target_layer_ids.iter().enumerate() {
let combined_idx = combined_ids
.iter()
.position(|&c| c == dl_id)
.expect("drafter layer in combined");
let verify_pos_t_base = (combined_idx * block_size + t) * hs;
let expected =
&verify_captured.hidden_output[verify_pos_t_base..verify_pos_t_base + hs];
let actual_row_base = (t * drafter_n + drafter_l) * hs;
let actual = &drafter_concat_new[actual_row_base..actual_row_base + hs];
assert_eq!(
actual, expected,
"t={t} drafter_l={drafter_l} (target_layer_id={dl_id}, combined_idx={combined_idx}): \
drafter_concat_new row {t} must equal verify_captured[combined={combined_idx}, pos={t}, :]",
);
}
}
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn smoke_capture_to_drafter_forward_pipeline() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use crate::inference::spec_decode::dflash::{
forward::dispatch_dflash_model_forward,
kv_cache::DFlashKvCache,
tensors::DFlashModelTensors,
weights::{DFlashWeights, DFlashWeightsFile},
};
use mlx_native::{DType, KernelRegistry, MlxDevice};
let cfg = gemma4_cfg();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
let mut registry = KernelRegistry::new();
let home = std::env::var("HOME").expect("HOME set");
let path = format!(
"{home}/.cache/huggingface/hub/models--z-lab--gemma-4-26B-A4B-it-DFlash/snapshots/77d4202772dfe50b2396ec7bac9cfffc7b9e7057/model.safetensors"
);
let file = DFlashWeightsFile::open(&path).expect("file open");
let weights = DFlashWeights::load(file.bytes(), &cfg).expect("validated load");
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let mut cache = DFlashKvCache::new(&device, &cfg, 128).expect("cache");
let ctx_chunk = 4usize;
let hidden = cfg.hidden_size;
let block_size = 8u32;
let mut hidden_buf = vec![0.0f32; capture_buffer_len(&cfg, ctx_chunk)];
let target_layer_ids: Vec<usize> = cfg.target_layer_ids.clone();
for cli in 0..target_layer_ids.len() {
for t in 0..ctx_chunk {
for d in 0..hidden {
let off = (cli * ctx_chunk + t) * hidden + d;
hidden_buf[off] = 0.1 + ((off % 31) as f32) / 310.0;
}
}
}
let cap = PrefillCapture {
target_layer_ids: &target_layer_ids,
hidden_output: &mut hidden_buf,
per_position_argmaxes: None,
};
cap.validate(ctx_chunk, hidden).expect("validate");
let concat = cap.permute_to_concat(ctx_chunk, hidden);
assert_eq!(concat.len(), ctx_chunk * target_layer_ids.len() * hidden);
let mut target_hidden = device
.alloc_buffer(
concat.len() * 4,
DType::F32,
vec![ctx_chunk, target_layer_ids.len() * hidden],
)
.expect("alloc target_hidden");
target_hidden
.as_mut_slice::<f32>()
.expect("target_hidden slice")
.copy_from_slice(&concat);
let h_elem = (block_size as usize) * hidden;
let mut h = device
.alloc_buffer(h_elem * 4, DType::F32, vec![block_size as usize, hidden])
.expect("alloc h");
{
let s = h.as_mut_slice::<f32>().expect("h slice");
for v in s.iter_mut() {
*v = 1.0;
}
}
let h_final = dispatch_dflash_model_forward(
&mut registry,
&device,
&h,
&target_hidden,
&tensors,
&mut cache,
&cfg,
block_size,
ctx_chunk as u32,
)
.expect("drafter forward on permuted capture");
assert_eq!(h_final.element_count(), (block_size as usize) * hidden);
let host: &[f32] = h_final.as_slice::<f32>().expect("h_final slice");
let n_finite = host.iter().filter(|v| v.is_finite()).count();
assert_eq!(
n_finite,
host.len(),
"drafter output on permuted capture must be all finite"
);
}
#[test]
fn permute_to_concat_round_trips_layout() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let target_layer_ids = vec![1, 6];
let n_layers = target_layer_ids.len();
let seq_len = 3;
let hidden_size = 4;
let mut hidden = vec![0.0f32; n_layers * seq_len * hidden_size];
for layer_idx in 0..n_layers {
for t in 0..seq_len {
for d in 0..hidden_size {
let src = (layer_idx * seq_len + t) * hidden_size + d;
hidden[src] = (layer_idx * 10) as f32 + t as f32 + (d as f32) * 0.1;
}
}
}
let cap = PrefillCapture {
target_layer_ids: &target_layer_ids,
hidden_output: &mut hidden,
per_position_argmaxes: None,
};
let concat = cap.permute_to_concat(seq_len, hidden_size);
assert_eq!(concat.len(), seq_len * n_layers * hidden_size);
for layer_idx in 0..n_layers {
for t in 0..seq_len {
for d in 0..hidden_size {
let dst = (t * n_layers + layer_idx) * hidden_size + d;
let expected = (layer_idx * 10) as f32 + t as f32 + (d as f32) * 0.1;
assert_eq!(
concat[dst], expected,
"concat mismatch t={t} layer={layer_idx} d={d}"
);
}
}
}
}
}