2498 lines
78 KiB
Rust
2498 lines
78 KiB
Rust
use super::checkpoint::{read_buffer, read_u32, write_buffer, write_u32};
|
|
use super::*;
|
|
use crate::engine::qwen::{QwenModel, QwenTensor};
|
|
use std::ffi::CStr;
|
|
|
|
const HIDDEN: u32 = 2_560;
|
|
const HC: u32 = 4;
|
|
const HC_WIDTH: u32 = HIDDEN * HC;
|
|
const HC_RANK: u32 = 320;
|
|
const LAYERS: usize = 48;
|
|
const GDN_HEADS_K: u32 = 16;
|
|
const GDN_HEADS_V: u32 = 48;
|
|
const HEAD_DIM: u32 = 128;
|
|
const GDN_QKV: u32 = (GDN_HEADS_K * 2 + GDN_HEADS_V) * HEAD_DIM;
|
|
const GDN_VALUE: u32 = GDN_HEADS_V * HEAD_DIM;
|
|
const GDN_CONTROLS: u32 = GDN_VALUE + GDN_HEADS_V * 2;
|
|
const ATTN_HEADS: u32 = 24;
|
|
const ATTN_KV_HEADS: u32 = 2;
|
|
const ATTN_DIM: u32 = 256;
|
|
const ATTN_WIDTH: u32 = ATTN_HEADS * ATTN_DIM;
|
|
const ATTN_KV_WIDTH: u32 = ATTN_KV_HEADS * ATTN_DIM;
|
|
const EXPERTS: u32 = 512;
|
|
const EXPERTS_USED: usize = 10;
|
|
const EXPERT_WIDTH: u32 = 640;
|
|
const VOCAB: u32 = 248_320;
|
|
const DENSE_BUDGET: u32 = 2_048;
|
|
const PLE_HEADS: usize = 16;
|
|
const PLE_HEAD_DIM: u32 = 160;
|
|
const PLE_HISTORY: usize = 2;
|
|
const PLE_CONV_STATE: u32 = 9;
|
|
const EOS_TOKEN: i32 = 248_044;
|
|
const PLE_ROW_BYTES: usize = 100;
|
|
const CHECKPOINT_MAGIC: &[u8; 8] = b"DS4QWN01";
|
|
const CHECKPOINT_VERSION: u32 = 2;
|
|
const CHECKPOINT_CHUNK: usize = 8 * 1024 * 1024;
|
|
|
|
#[derive(Clone, Copy)]
|
|
struct Weight<'a> {
|
|
tensor: &'a QwenTensor,
|
|
expert: Option<u32>,
|
|
}
|
|
|
|
struct Affine<'a> {
|
|
packed: Weight<'a>,
|
|
scales: Weight<'a>,
|
|
biases: Weight<'a>,
|
|
bits: u32,
|
|
group: u32,
|
|
}
|
|
|
|
enum LayerState {
|
|
Gdn { conv: Buffer, recurrent: Buffer },
|
|
Attention { kv: Buffer },
|
|
}
|
|
|
|
struct Scratch {
|
|
hidden: Buffer,
|
|
hc: Buffer,
|
|
hc_norm: Buffer,
|
|
hc_mix: Buffer,
|
|
rank: Buffer,
|
|
block: Buffer,
|
|
injection: Buffer,
|
|
qkv: Buffer,
|
|
controls: Buffer,
|
|
gdn_raw: Buffer,
|
|
gdn_out: Buffer,
|
|
router: Buffer,
|
|
route_ids: Buffer,
|
|
route_weights: Buffer,
|
|
gate: Buffer,
|
|
up: Buffer,
|
|
mid: Buffer,
|
|
expert: Buffer,
|
|
moe: Buffer,
|
|
shared: Buffer,
|
|
q_packed: Buffer,
|
|
q: Buffer,
|
|
q_gate: Buffer,
|
|
k: Buffer,
|
|
k_rope: Buffer,
|
|
v: Buffer,
|
|
attention: Buffer,
|
|
ple_packed: Buffer,
|
|
ple_scales: Buffer,
|
|
ple_biases: Buffer,
|
|
ple_embedding: Buffer,
|
|
ple_key: Buffer,
|
|
ple_value: Buffer,
|
|
ple_gated: Buffer,
|
|
ple_norm: Buffer,
|
|
ple_output: Buffer,
|
|
logits: Buffer,
|
|
}
|
|
|
|
impl Scratch {
|
|
fn new() -> Result<Self, String> {
|
|
Ok(Self {
|
|
hidden: Buffer::floats(HIDDEN.into())?,
|
|
hc: Buffer::floats(HC_WIDTH.into())?,
|
|
hc_norm: Buffer::floats(HC_WIDTH.into())?,
|
|
hc_mix: Buffer::floats(HC_WIDTH.into())?,
|
|
rank: Buffer::floats(HC_RANK.into())?,
|
|
block: Buffer::floats(HIDDEN.into())?,
|
|
injection: Buffer::floats(HC.into())?,
|
|
qkv: Buffer::floats(GDN_QKV.into())?,
|
|
controls: Buffer::floats(GDN_CONTROLS.into())?,
|
|
gdn_raw: Buffer::floats(GDN_VALUE.into())?,
|
|
gdn_out: Buffer::floats(GDN_VALUE.into())?,
|
|
router: Buffer::floats(EXPERTS.into())?,
|
|
route_ids: Buffer::bytes((EXPERTS_USED * 4) as u64)?,
|
|
route_weights: Buffer::floats(EXPERTS_USED as u64)?,
|
|
gate: Buffer::floats(EXPERT_WIDTH.into())?,
|
|
up: Buffer::floats(EXPERT_WIDTH.into())?,
|
|
mid: Buffer::floats(EXPERT_WIDTH.into())?,
|
|
expert: Buffer::floats(HIDDEN.into())?,
|
|
moe: Buffer::floats(HIDDEN.into())?,
|
|
shared: Buffer::floats(HIDDEN.into())?,
|
|
q_packed: Buffer::floats((ATTN_WIDTH * 2).into())?,
|
|
q: Buffer::floats(ATTN_WIDTH.into())?,
|
|
q_gate: Buffer::floats(ATTN_WIDTH.into())?,
|
|
k: Buffer::floats(ATTN_KV_WIDTH.into())?,
|
|
k_rope: Buffer::floats(ATTN_KV_WIDTH.into())?,
|
|
v: Buffer::floats(ATTN_KV_WIDTH.into())?,
|
|
attention: Buffer::floats(ATTN_WIDTH.into())?,
|
|
ple_packed: Buffer::bytes((PLE_HEADS as u64) * 80)?,
|
|
ple_scales: Buffer::bytes((PLE_HEADS as u64) * 10)?,
|
|
ple_biases: Buffer::bytes((PLE_HEADS as u64) * 10)?,
|
|
ple_embedding: Buffer::floats(HIDDEN.into())?,
|
|
ple_key: Buffer::floats(HC_WIDTH.into())?,
|
|
ple_value: Buffer::floats(HIDDEN.into())?,
|
|
ple_gated: Buffer::floats(HC_WIDTH.into())?,
|
|
ple_norm: Buffer::floats(HC_WIDTH.into())?,
|
|
ple_output: Buffer::floats(HC_WIDTH.into())?,
|
|
logits: Buffer::floats(VOCAB.into())?,
|
|
})
|
|
}
|
|
}
|
|
|
|
struct PleContract {
|
|
multipliers: [i64; 3],
|
|
sizes: [i64; PLE_HEADS],
|
|
offsets: [i64; PLE_HEADS],
|
|
}
|
|
|
|
struct PleState {
|
|
history: [i32; PLE_HISTORY],
|
|
conv: Buffer,
|
|
}
|
|
|
|
pub(in crate::engine) struct QwenExecutor {
|
|
model: QwenModel,
|
|
states: Vec<LayerState>,
|
|
ple_contract: PleContract,
|
|
ple_state: PleState,
|
|
scratch: Scratch,
|
|
logits: Vec<f32>,
|
|
tokens: Vec<i32>,
|
|
position: u32,
|
|
context: u32,
|
|
checkpoint_tag: [u8; 32],
|
|
_context: Context,
|
|
}
|
|
|
|
pub(in crate::engine) struct QwenResidentState {
|
|
states: Vec<LayerState>,
|
|
ple_state: PleState,
|
|
logits: Vec<f32>,
|
|
tokens: Vec<i32>,
|
|
position: u32,
|
|
checkpoint_tag: [u8; 32],
|
|
}
|
|
|
|
impl QwenExecutor {
|
|
pub(super) fn open(model: QwenModel, context: u32) -> Result<Self, String> {
|
|
let ple_contract = ple_contract(&model)?;
|
|
let native = Context::open_qwen(model.memory().admission)?;
|
|
let states = allocate_states(context)?;
|
|
Ok(Self {
|
|
model,
|
|
states,
|
|
ple_contract,
|
|
ple_state: allocate_ple_state()?,
|
|
scratch: Scratch::new()?,
|
|
logits: vec![0.0; VOCAB as usize],
|
|
tokens: Vec::new(),
|
|
position: 0,
|
|
context,
|
|
checkpoint_tag: [0; 32],
|
|
_context: native,
|
|
})
|
|
}
|
|
|
|
pub(super) fn eval(&mut self, token: i32) -> Result<(), String> {
|
|
if token < 0 || token as u32 >= VOCAB {
|
|
return Err(format!("token {token} is outside the Qwen vocabulary"));
|
|
}
|
|
if self.position >= self.context {
|
|
return Err(format!(
|
|
"the Qwen executor supports {} tokens per session",
|
|
self.context
|
|
));
|
|
}
|
|
if self.position + 1 > DENSE_BUDGET {
|
|
return Err(
|
|
"Qwen sparse QSA selection is required beyond 2048 tokens; native QSA belongs to issue #97"
|
|
.into(),
|
|
);
|
|
}
|
|
|
|
self.begin_token(token)?;
|
|
|
|
for layer in 0..LAYERS {
|
|
if layer == 1 {
|
|
self.ple(token)?;
|
|
}
|
|
self.encode_layer(layer)?;
|
|
}
|
|
self.final_output()?;
|
|
self.tokens.push(token);
|
|
self.position += 1;
|
|
Ok(())
|
|
}
|
|
|
|
fn begin_token(&self, token: i32) -> Result<(), String> {
|
|
let embedding = self.affine("language_model.model.embed_tokens", HIDDEN, VOCAB, None)?;
|
|
let mut args = args();
|
|
args.u[0] = HIDDEN;
|
|
args.u[2] = embedding.bits;
|
|
args.u[3] = embedding.group;
|
|
args.u[4] = token as u32;
|
|
let commands = Commands::begin()?;
|
|
self.dispatch(
|
|
c"kernel_qwen_affine_embedding",
|
|
&self.scratch.hidden,
|
|
None,
|
|
None,
|
|
None,
|
|
&[
|
|
self.view(embedding.packed),
|
|
self.view(embedding.scales),
|
|
self.view(embedding.biases),
|
|
],
|
|
&args,
|
|
HIDDEN,
|
|
1,
|
|
)?;
|
|
args.u[0] = HIDDEN;
|
|
self.dispatch(
|
|
c"kernel_qwen_repeat4",
|
|
&self.scratch.hc,
|
|
Some(&self.scratch.hidden),
|
|
None,
|
|
None,
|
|
&[],
|
|
&args,
|
|
HC_WIDTH,
|
|
1,
|
|
)?;
|
|
commands.finish()
|
|
}
|
|
|
|
fn ple(&mut self, token: i32) -> Result<(), String> {
|
|
let rows = ple_rows(&self.ple_contract, self.ple_state.history, token)?;
|
|
let packed = self.model.tensor("ngram.weight")?;
|
|
let scales = self.model.tensor("ngram.scales")?;
|
|
let biases = self.model.tensor("ngram.biases")?;
|
|
let mut packed_stage = [0_u8; PLE_HEADS * 80];
|
|
let mut scales_stage = [0_u8; PLE_HEADS * 10];
|
|
let mut biases_stage = [0_u8; PLE_HEADS * 10];
|
|
let packed_bytes = self.model.tensor_bytes(packed)?;
|
|
let scales_bytes = self.model.tensor_bytes(scales)?;
|
|
let biases_bytes = self.model.tensor_bytes(biases)?;
|
|
for (head, &row) in rows.iter().enumerate() {
|
|
let bytes = gather_ple_row(row, packed_bytes, scales_bytes, biases_bytes)?;
|
|
packed_stage[head * 80..][..80].copy_from_slice(&bytes[..80]);
|
|
scales_stage[head * 10..][..10].copy_from_slice(&bytes[80..90]);
|
|
biases_stage[head * 10..][..10].copy_from_slice(&bytes[90..]);
|
|
}
|
|
self.scratch.ple_packed.write(0, &packed_stage)?;
|
|
self.scratch.ple_scales.write(0, &scales_stage)?;
|
|
self.scratch.ple_biases.write(0, &biases_stage)?;
|
|
|
|
let commands = Commands::begin()?;
|
|
let mut dequant = args();
|
|
dequant.u[0] = PLE_HEAD_DIM;
|
|
dequant.u[1] = PLE_HEADS as u32;
|
|
dequant.u[2] = 4;
|
|
dequant.u[3] = 32;
|
|
self.dispatch(
|
|
c"kernel_qwen_ple_dequant",
|
|
&self.scratch.ple_embedding,
|
|
Some(&self.scratch.ple_packed),
|
|
Some(&self.scratch.ple_scales),
|
|
Some(&self.scratch.ple_biases),
|
|
&[],
|
|
&dequant,
|
|
HIDDEN,
|
|
1,
|
|
)?;
|
|
let prefix = "language_model.model.layers.1.ple";
|
|
self.bf16_mv(
|
|
self.weight(&format!("{prefix}.key_proj.weight"))?,
|
|
&self.scratch.ple_embedding,
|
|
&self.scratch.ple_key,
|
|
HIDDEN,
|
|
HC_WIDTH,
|
|
)?;
|
|
self.bf16_mv(
|
|
self.weight(&format!("{prefix}.value_proj.weight"))?,
|
|
&self.scratch.ple_embedding,
|
|
&self.scratch.ple_value,
|
|
HIDDEN,
|
|
HIDDEN,
|
|
)?;
|
|
for (input, output, name) in [
|
|
(&self.scratch.ple_key, &self.scratch.ple_key, "norm_key"),
|
|
(&self.scratch.hc, &self.scratch.hc_norm, "norm_query"),
|
|
] {
|
|
let mut norm = args();
|
|
norm.u[0] = HC_WIDTH;
|
|
norm.u[1] = HIDDEN;
|
|
norm.f[0] = 1.0e-6;
|
|
self.dispatch(
|
|
c"kernel_qwen_zero_rms",
|
|
output,
|
|
Some(input),
|
|
None,
|
|
None,
|
|
&[self.view(self.weight(&format!("{prefix}.{name}.weight"))?)],
|
|
&norm,
|
|
HC,
|
|
1,
|
|
)?;
|
|
}
|
|
let mut gate = args();
|
|
gate.u[0] = HIDDEN;
|
|
self.dispatch(
|
|
c"kernel_qwen_ple_gate",
|
|
&self.scratch.ple_gated,
|
|
Some(&self.scratch.ple_key),
|
|
Some(&self.scratch.hc_norm),
|
|
Some(&self.scratch.ple_value),
|
|
&[],
|
|
&gate,
|
|
HC,
|
|
1,
|
|
)?;
|
|
let mut norm = args();
|
|
norm.u[0] = HC_WIDTH;
|
|
norm.u[1] = HIDDEN;
|
|
norm.f[0] = 1.0e-6;
|
|
self.dispatch(
|
|
c"kernel_qwen_zero_rms",
|
|
&self.scratch.ple_norm,
|
|
Some(&self.scratch.ple_gated),
|
|
None,
|
|
None,
|
|
&[self.view(self.weight(&format!("{prefix}.norm_conv.weight"))?)],
|
|
&norm,
|
|
HC,
|
|
1,
|
|
)?;
|
|
let mut conv = args();
|
|
conv.u[0] = HC_WIDTH;
|
|
self.dispatch(
|
|
c"kernel_qwen_ple_conv",
|
|
&self.scratch.ple_output,
|
|
Some(&self.scratch.ple_gated),
|
|
Some(&self.scratch.ple_norm),
|
|
Some(&self.ple_state.conv),
|
|
&[self.view(self.weight(&format!("{prefix}.conv_weight"))?)],
|
|
&conv,
|
|
HC_WIDTH,
|
|
1,
|
|
)?;
|
|
let mut add = args();
|
|
add.u[0] = HC_WIDTH;
|
|
self.dispatch(
|
|
c"kernel_qwen_add",
|
|
&self.scratch.hc_norm,
|
|
Some(&self.scratch.hc),
|
|
Some(&self.scratch.ple_output),
|
|
None,
|
|
&[],
|
|
&add,
|
|
HC_WIDTH,
|
|
1,
|
|
)?;
|
|
self.scratch.hc.copy_from(
|
|
0,
|
|
&self.scratch.hc_norm,
|
|
0,
|
|
u64::from(HC_WIDTH) * 4,
|
|
"committing Qwen PLE injection",
|
|
)?;
|
|
commands.finish()?;
|
|
self.ple_state.history = advance_ple_history(self.ple_state.history, token);
|
|
Ok(())
|
|
}
|
|
|
|
fn encode_layer(&mut self, layer: usize) -> Result<(), String> {
|
|
let prefix = format!("language_model.model.layers.{layer}");
|
|
let commands = Commands::begin()?;
|
|
self.hyper_read(&format!("{prefix}.attn_hyper_connection"))?;
|
|
match &self.states[layer] {
|
|
LayerState::Gdn { .. } => self.gdn(&prefix, layer)?,
|
|
LayerState::Attention { .. } => self.attention(&prefix, layer)?,
|
|
}
|
|
self.hyper_write()?;
|
|
self.hyper_read(&format!("{prefix}.mlp_hyper_connection"))?;
|
|
self.affine_mv_into(
|
|
&self.affine(&format!("{prefix}.mlp.gate"), HIDDEN, EXPERTS, None)?,
|
|
&self.scratch.block,
|
|
&self.scratch.router,
|
|
HIDDEN,
|
|
EXPERTS,
|
|
)?;
|
|
let mut route_args = args();
|
|
route_args.u[0] = EXPERTS;
|
|
self.dispatch(
|
|
c"kernel_qwen_route_top10",
|
|
&self.scratch.route_ids,
|
|
Some(&self.scratch.route_weights),
|
|
Some(&self.scratch.router),
|
|
None,
|
|
&[],
|
|
&route_args,
|
|
1,
|
|
1,
|
|
)?;
|
|
commands.finish()?;
|
|
|
|
let mut ids = [0_i32; EXPERTS_USED];
|
|
let mut weights = [0.0_f32; EXPERTS_USED];
|
|
self.scratch.route_ids.read_i32(&mut ids)?;
|
|
self.scratch.route_weights.read_f32(&mut weights)?;
|
|
self.scratch.moe.fill(0.0, HIDDEN.into())?;
|
|
let commands = Commands::begin()?;
|
|
for (&expert, &weight) in ids.iter().zip(&weights) {
|
|
if !(0..EXPERTS as i32).contains(&expert) || !weight.is_finite() || weight < 0.0 {
|
|
let mut router = vec![0.0; EXPERTS as usize];
|
|
self.scratch.router.read_f32(&mut router)?;
|
|
let non_finite = router.iter().filter(|value| !value.is_finite()).count();
|
|
let mut block = vec![0.0; HIDDEN as usize];
|
|
self.scratch.block.read_f32(&mut block)?;
|
|
let block_non_finite = block.iter().filter(|value| !value.is_finite()).count();
|
|
return Err(format!(
|
|
"Qwen layer {layer} router produced invalid expert {expert} with weight {weight} ({non_finite} non-finite logits, {block_non_finite} non-finite inputs)"
|
|
));
|
|
}
|
|
self.expert(&prefix, expert as u32, weight)?;
|
|
}
|
|
self.shared_expert(&prefix)?;
|
|
let mut inject_args = args();
|
|
inject_args.u[0] = HIDDEN;
|
|
self.dispatch(
|
|
c"kernel_qwen_hyper_inject",
|
|
&self.scratch.hc_norm,
|
|
Some(&self.scratch.hc),
|
|
Some(&self.scratch.moe),
|
|
Some(&self.scratch.injection),
|
|
&[],
|
|
&inject_args,
|
|
HC_WIDTH,
|
|
1,
|
|
)?;
|
|
self.scratch.hc.copy_from(
|
|
0,
|
|
&self.scratch.hc_norm,
|
|
0,
|
|
u64::from(HC_WIDTH) * 4,
|
|
"committing Qwen MoE hyper streams",
|
|
)?;
|
|
commands.finish()
|
|
}
|
|
|
|
fn hyper_read(&self, prefix: &str) -> Result<(), String> {
|
|
let norm = self.weight(&format!("{prefix}.hc_norm.weight"))?;
|
|
let down = self.weight(&format!("{prefix}.input_mix_weight_down.weight"))?;
|
|
let up = self.weight(&format!("{prefix}.input_mix_weight_up.weight"))?;
|
|
let inject = self.weight(&format!("{prefix}.block_inject_weight.weight"))?;
|
|
let mut rms = args();
|
|
rms.u[0] = HC_WIDTH;
|
|
rms.u[1] = HIDDEN;
|
|
rms.f[0] = 1.0e-6;
|
|
self.dispatch(
|
|
c"kernel_qwen_zero_rms",
|
|
&self.scratch.hc_norm,
|
|
Some(&self.scratch.hc),
|
|
None,
|
|
None,
|
|
&[self.view(norm)],
|
|
&rms,
|
|
HC,
|
|
1,
|
|
)?;
|
|
self.bf16_mv(
|
|
down,
|
|
&self.scratch.hc_norm,
|
|
&self.scratch.rank,
|
|
HC_WIDTH,
|
|
HC_RANK,
|
|
)?;
|
|
let mut unary = args();
|
|
unary.u[0] = HC_RANK;
|
|
self.dispatch(
|
|
c"kernel_qwen_silu_div4",
|
|
&self.scratch.rank,
|
|
Some(&self.scratch.rank),
|
|
None,
|
|
None,
|
|
&[],
|
|
&unary,
|
|
HC_RANK,
|
|
1,
|
|
)?;
|
|
self.bf16_mv(
|
|
up,
|
|
&self.scratch.rank,
|
|
&self.scratch.hc_mix,
|
|
HC_RANK,
|
|
HC_WIDTH,
|
|
)?;
|
|
unary.u[0] = HC_WIDTH;
|
|
self.dispatch(
|
|
c"kernel_qwen_sigmoid",
|
|
&self.scratch.hc_mix,
|
|
Some(&self.scratch.hc_mix),
|
|
None,
|
|
None,
|
|
&[],
|
|
&unary,
|
|
HC_WIDTH,
|
|
1,
|
|
)?;
|
|
let mut mix = args();
|
|
mix.u[0] = HIDDEN;
|
|
self.dispatch(
|
|
c"kernel_qwen_hyper_mix",
|
|
&self.scratch.block,
|
|
Some(&self.scratch.hc_norm),
|
|
Some(&self.scratch.hc_mix),
|
|
None,
|
|
&[],
|
|
&mix,
|
|
HIDDEN,
|
|
1,
|
|
)?;
|
|
self.bf16_mv(
|
|
inject,
|
|
&self.scratch.hc_norm,
|
|
&self.scratch.injection,
|
|
HC_WIDTH,
|
|
HC,
|
|
)?;
|
|
unary.u[0] = HC;
|
|
self.dispatch(
|
|
c"kernel_qwen_sigmoid2_div4",
|
|
&self.scratch.injection,
|
|
Some(&self.scratch.injection),
|
|
None,
|
|
None,
|
|
&[],
|
|
&unary,
|
|
HC,
|
|
1,
|
|
)
|
|
}
|
|
|
|
fn hyper_write(&self) -> Result<(), String> {
|
|
let mut values = args();
|
|
values.u[0] = HIDDEN;
|
|
self.dispatch(
|
|
c"kernel_qwen_hyper_inject",
|
|
&self.scratch.hc_norm,
|
|
Some(&self.scratch.hc),
|
|
Some(&self.scratch.hidden),
|
|
Some(&self.scratch.injection),
|
|
&[],
|
|
&values,
|
|
HC_WIDTH,
|
|
1,
|
|
)?;
|
|
self.scratch.hc.copy_from(
|
|
0,
|
|
&self.scratch.hc_norm,
|
|
0,
|
|
u64::from(HC_WIDTH) * 4,
|
|
"committing Qwen attention hyper streams",
|
|
)
|
|
}
|
|
|
|
fn gdn(&self, prefix: &str, layer: usize) -> Result<(), String> {
|
|
let LayerState::Gdn { conv, recurrent } = &self.states[layer] else {
|
|
return Err("Qwen GDN graph received attention state".into());
|
|
};
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.linear_attn.in_proj_qkv"),
|
|
HIDDEN,
|
|
GDN_QKV,
|
|
None,
|
|
)?,
|
|
&self.scratch.block,
|
|
&self.scratch.qkv,
|
|
HIDDEN,
|
|
GDN_QKV,
|
|
)?;
|
|
for (name, offset, width) in [
|
|
("in_proj_z", 0_u64, GDN_VALUE),
|
|
("in_proj_b", u64::from(GDN_VALUE), GDN_HEADS_V),
|
|
("in_proj_a", u64::from(GDN_VALUE + GDN_HEADS_V), GDN_HEADS_V),
|
|
] {
|
|
let target = self
|
|
.scratch
|
|
.controls
|
|
.view(offset * 4, u64::from(width) * 4)?;
|
|
self.affine_mv_into(
|
|
&self.affine(&format!("{prefix}.linear_attn.{name}"), HIDDEN, width, None)?,
|
|
&self.scratch.block,
|
|
&target,
|
|
HIDDEN,
|
|
width,
|
|
)?;
|
|
}
|
|
let mut conv_args = args();
|
|
conv_args.u[0] = GDN_QKV;
|
|
self.dispatch(
|
|
c"kernel_qwen_conv_silu",
|
|
&self.scratch.qkv,
|
|
Some(&self.scratch.qkv),
|
|
Some(conv),
|
|
None,
|
|
&[self.view(self.weight(&format!("{prefix}.linear_attn.conv1d.weight"))?)],
|
|
&conv_args,
|
|
GDN_QKV,
|
|
1,
|
|
)?;
|
|
let mut step = args();
|
|
step.u[0] = HEAD_DIM;
|
|
step.u[1] = GDN_HEADS_K;
|
|
step.u[2] = GDN_HEADS_V;
|
|
step.u[3] = GDN_VALUE;
|
|
step.u[4] = GDN_VALUE + GDN_HEADS_V;
|
|
step.f[0] = 1.0e-6;
|
|
self.dispatch(
|
|
c"kernel_qwen_gdn_step",
|
|
&self.scratch.gdn_raw,
|
|
Some(&self.scratch.gdn_out),
|
|
Some(&self.scratch.controls),
|
|
Some(recurrent),
|
|
&[
|
|
self.view(self.weight(&format!("{prefix}.linear_attn.A_log"))?),
|
|
self.view(self.weight(&format!("{prefix}.linear_attn.dt_bias"))?),
|
|
],
|
|
&step,
|
|
HEAD_DIM,
|
|
GDN_HEADS_V,
|
|
)?;
|
|
let mut gate = args();
|
|
gate.u[0] = HEAD_DIM;
|
|
gate.u[1] = GDN_HEADS_V;
|
|
gate.f[0] = 1.0e-6;
|
|
self.dispatch(
|
|
c"kernel_qwen_gdn_norm_gate",
|
|
&self.scratch.gdn_out,
|
|
Some(&self.scratch.gdn_raw),
|
|
Some(&self.scratch.controls),
|
|
None,
|
|
&[self.view(self.weight(&format!("{prefix}.linear_attn.norm.weight"))?)],
|
|
&gate,
|
|
GDN_HEADS_V,
|
|
1,
|
|
)?;
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.linear_attn.out_proj"),
|
|
GDN_VALUE,
|
|
HIDDEN,
|
|
None,
|
|
)?,
|
|
&self.scratch.gdn_out,
|
|
&self.scratch.hidden,
|
|
GDN_VALUE,
|
|
HIDDEN,
|
|
)
|
|
}
|
|
|
|
fn attention(&self, prefix: &str, layer: usize) -> Result<(), String> {
|
|
if self.position + 1 > DENSE_BUDGET {
|
|
return Err(
|
|
"Qwen sparse QSA selection is required beyond 2048 tokens; native QSA belongs to issue #97"
|
|
.into(),
|
|
);
|
|
}
|
|
let LayerState::Attention { kv } = &self.states[layer] else {
|
|
return Err("Qwen attention graph received GDN state".into());
|
|
};
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.self_attn.q_proj"),
|
|
HIDDEN,
|
|
ATTN_WIDTH * 2,
|
|
None,
|
|
)?,
|
|
&self.scratch.block,
|
|
&self.scratch.q_packed,
|
|
HIDDEN,
|
|
ATTN_WIDTH * 2,
|
|
)?;
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.self_attn.k_proj"),
|
|
HIDDEN,
|
|
ATTN_KV_WIDTH,
|
|
None,
|
|
)?,
|
|
&self.scratch.block,
|
|
&self.scratch.k,
|
|
HIDDEN,
|
|
ATTN_KV_WIDTH,
|
|
)?;
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.self_attn.v_proj"),
|
|
HIDDEN,
|
|
ATTN_KV_WIDTH,
|
|
None,
|
|
)?,
|
|
&self.scratch.block,
|
|
&self.scratch.v,
|
|
HIDDEN,
|
|
ATTN_KV_WIDTH,
|
|
)?;
|
|
let mut split = args();
|
|
split.u[0] = ATTN_HEADS;
|
|
split.u[1] = ATTN_DIM;
|
|
self.dispatch(
|
|
c"kernel_qwen_split_q_gate",
|
|
&self.scratch.q,
|
|
Some(&self.scratch.q_packed),
|
|
Some(&self.scratch.q_gate),
|
|
None,
|
|
&[],
|
|
&split,
|
|
ATTN_WIDTH,
|
|
1,
|
|
)?;
|
|
self.head_norm_rope(
|
|
&self.scratch.q,
|
|
&self.scratch.attention,
|
|
self.weight(&format!("{prefix}.self_attn.q_norm.weight"))?,
|
|
ATTN_HEADS,
|
|
)?;
|
|
self.head_norm_rope(
|
|
&self.scratch.k,
|
|
&self.scratch.k_rope,
|
|
self.weight(&format!("{prefix}.self_attn.k_norm.weight"))?,
|
|
ATTN_KV_HEADS,
|
|
)?;
|
|
let mut store = args();
|
|
store.u[0] = ATTN_KV_WIDTH;
|
|
store.u[1] = self.position;
|
|
self.dispatch(
|
|
c"kernel_qwen_store_kv_bf16",
|
|
kv,
|
|
Some(&self.scratch.k_rope),
|
|
Some(&self.scratch.v),
|
|
None,
|
|
&[],
|
|
&store,
|
|
ATTN_KV_WIDTH,
|
|
1,
|
|
)?;
|
|
let mut dense = args();
|
|
dense.u[0] = ATTN_HEADS;
|
|
dense.u[1] = ATTN_KV_HEADS;
|
|
dense.u[2] = ATTN_DIM;
|
|
dense.u[3] = self.position + 1;
|
|
self.dispatch(
|
|
c"kernel_qwen_dense_attention",
|
|
&self.scratch.q,
|
|
Some(&self.scratch.attention),
|
|
Some(kv),
|
|
None,
|
|
&[],
|
|
&dense,
|
|
ATTN_HEADS,
|
|
1,
|
|
)?;
|
|
let mut gate = args();
|
|
gate.u[0] = ATTN_WIDTH;
|
|
self.dispatch(
|
|
c"kernel_qwen_gate_attention",
|
|
&self.scratch.attention,
|
|
Some(&self.scratch.q),
|
|
Some(&self.scratch.q_gate),
|
|
None,
|
|
&[],
|
|
&gate,
|
|
ATTN_WIDTH,
|
|
1,
|
|
)?;
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.self_attn.o_proj"),
|
|
ATTN_WIDTH,
|
|
HIDDEN,
|
|
None,
|
|
)?,
|
|
&self.scratch.attention,
|
|
&self.scratch.hidden,
|
|
ATTN_WIDTH,
|
|
HIDDEN,
|
|
)
|
|
}
|
|
|
|
fn head_norm_rope(
|
|
&self,
|
|
input: &Buffer,
|
|
output: &Buffer,
|
|
weight: Weight<'_>,
|
|
heads: u32,
|
|
) -> Result<(), String> {
|
|
let mut values = args();
|
|
values.u[0] = ATTN_DIM;
|
|
values.u[1] = 64;
|
|
values.u[2] = heads;
|
|
values.u[3] = self.position;
|
|
values.f[0] = 1.0e-6;
|
|
values.f[1] = 10_000_000.0;
|
|
self.dispatch(
|
|
c"kernel_qwen_head_norm_rope",
|
|
output,
|
|
Some(input),
|
|
None,
|
|
None,
|
|
&[self.view(weight)],
|
|
&values,
|
|
ATTN_DIM,
|
|
heads,
|
|
)
|
|
}
|
|
|
|
fn expert(&self, prefix: &str, expert: u32, weight: f32) -> Result<(), String> {
|
|
for (name, target) in [
|
|
("gate_proj", &self.scratch.gate),
|
|
("up_proj", &self.scratch.up),
|
|
] {
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.mlp.switch_mlp.{name}"),
|
|
HIDDEN,
|
|
EXPERT_WIDTH,
|
|
Some(expert),
|
|
)?,
|
|
&self.scratch.block,
|
|
target,
|
|
HIDDEN,
|
|
EXPERT_WIDTH,
|
|
)?;
|
|
}
|
|
let mut swiglu = args();
|
|
swiglu.u[0] = EXPERT_WIDTH;
|
|
self.dispatch(
|
|
c"kernel_qwen_swiglu",
|
|
&self.scratch.mid,
|
|
Some(&self.scratch.gate),
|
|
Some(&self.scratch.up),
|
|
None,
|
|
&[],
|
|
&swiglu,
|
|
EXPERT_WIDTH,
|
|
1,
|
|
)?;
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.mlp.switch_mlp.down_proj"),
|
|
EXPERT_WIDTH,
|
|
HIDDEN,
|
|
Some(expert),
|
|
)?,
|
|
&self.scratch.mid,
|
|
&self.scratch.expert,
|
|
EXPERT_WIDTH,
|
|
HIDDEN,
|
|
)?;
|
|
let mut accumulate = args();
|
|
accumulate.u[0] = HIDDEN;
|
|
accumulate.f[0] = weight;
|
|
self.dispatch(
|
|
c"kernel_qwen_accumulate",
|
|
&self.scratch.moe,
|
|
Some(&self.scratch.expert),
|
|
None,
|
|
None,
|
|
&[],
|
|
&accumulate,
|
|
HIDDEN,
|
|
1,
|
|
)
|
|
}
|
|
|
|
fn shared_expert(&self, prefix: &str) -> Result<(), String> {
|
|
for (name, target) in [
|
|
("gate_proj", &self.scratch.gate),
|
|
("up_proj", &self.scratch.up),
|
|
] {
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.mlp.shared_expert.{name}"),
|
|
HIDDEN,
|
|
EXPERT_WIDTH,
|
|
None,
|
|
)?,
|
|
&self.scratch.block,
|
|
target,
|
|
HIDDEN,
|
|
EXPERT_WIDTH,
|
|
)?;
|
|
}
|
|
let mut swiglu = args();
|
|
swiglu.u[0] = EXPERT_WIDTH;
|
|
self.dispatch(
|
|
c"kernel_qwen_swiglu",
|
|
&self.scratch.mid,
|
|
Some(&self.scratch.gate),
|
|
Some(&self.scratch.up),
|
|
None,
|
|
&[],
|
|
&swiglu,
|
|
EXPERT_WIDTH,
|
|
1,
|
|
)?;
|
|
self.affine_mv_into(
|
|
&self.affine(
|
|
&format!("{prefix}.mlp.shared_expert.down_proj"),
|
|
EXPERT_WIDTH,
|
|
HIDDEN,
|
|
None,
|
|
)?,
|
|
&self.scratch.mid,
|
|
&self.scratch.shared,
|
|
EXPERT_WIDTH,
|
|
HIDDEN,
|
|
)?;
|
|
self.affine_mv_into(
|
|
&self.affine(&format!("{prefix}.mlp.shared_expert_gate"), HIDDEN, 1, None)?,
|
|
&self.scratch.block,
|
|
&self.scratch.gate,
|
|
HIDDEN,
|
|
1,
|
|
)?;
|
|
let mut accumulate = args();
|
|
accumulate.u[0] = HIDDEN;
|
|
self.dispatch(
|
|
c"kernel_qwen_accumulate_sigmoid_scalar",
|
|
&self.scratch.moe,
|
|
Some(&self.scratch.shared),
|
|
Some(&self.scratch.gate),
|
|
None,
|
|
&[],
|
|
&accumulate,
|
|
HIDDEN,
|
|
1,
|
|
)
|
|
}
|
|
|
|
fn final_output(&mut self) -> Result<(), String> {
|
|
let commands = Commands::begin()?;
|
|
self.final_mix()?;
|
|
self.affine_mv_into(
|
|
&self.affine("language_model.lm_head", HIDDEN, VOCAB, None)?,
|
|
&self.scratch.block,
|
|
&self.scratch.logits,
|
|
HIDDEN,
|
|
VOCAB,
|
|
)?;
|
|
commands.finish()?;
|
|
self.scratch.logits.read_f32(&mut self.logits)
|
|
}
|
|
|
|
fn final_mix(&self) -> Result<(), String> {
|
|
let prefix = "language_model.model.hyper_connection_mixer";
|
|
let norm = self.weight(&format!("{prefix}.hc_norm.weight"))?;
|
|
let down = self.weight(&format!("{prefix}.input_mix_weight_down.weight"))?;
|
|
let up = self.weight(&format!("{prefix}.input_mix_weight_up.weight"))?;
|
|
let mut rms = args();
|
|
rms.u[0] = HC_WIDTH;
|
|
rms.u[1] = HIDDEN;
|
|
rms.f[0] = 1.0e-6;
|
|
self.dispatch(
|
|
c"kernel_qwen_zero_rms",
|
|
&self.scratch.hc_norm,
|
|
Some(&self.scratch.hc),
|
|
None,
|
|
None,
|
|
&[self.view(norm)],
|
|
&rms,
|
|
HC,
|
|
1,
|
|
)?;
|
|
self.bf16_mv(
|
|
down,
|
|
&self.scratch.hc_norm,
|
|
&self.scratch.rank,
|
|
HC_WIDTH,
|
|
HC_RANK,
|
|
)?;
|
|
let mut unary = args();
|
|
unary.u[0] = HC_RANK;
|
|
self.dispatch(
|
|
c"kernel_qwen_silu_div4",
|
|
&self.scratch.rank,
|
|
Some(&self.scratch.rank),
|
|
None,
|
|
None,
|
|
&[],
|
|
&unary,
|
|
HC_RANK,
|
|
1,
|
|
)?;
|
|
self.bf16_mv(
|
|
up,
|
|
&self.scratch.rank,
|
|
&self.scratch.hc_mix,
|
|
HC_RANK,
|
|
HC_WIDTH,
|
|
)?;
|
|
unary.u[0] = HC_WIDTH;
|
|
self.dispatch(
|
|
c"kernel_qwen_sigmoid",
|
|
&self.scratch.hc_mix,
|
|
Some(&self.scratch.hc_mix),
|
|
None,
|
|
None,
|
|
&[],
|
|
&unary,
|
|
HC_WIDTH,
|
|
1,
|
|
)?;
|
|
let mut mix = args();
|
|
mix.u[0] = HIDDEN;
|
|
self.dispatch(
|
|
c"kernel_qwen_hyper_mix",
|
|
&self.scratch.block,
|
|
Some(&self.scratch.hc_norm),
|
|
Some(&self.scratch.hc_mix),
|
|
None,
|
|
&[],
|
|
&mix,
|
|
HIDDEN,
|
|
1,
|
|
)
|
|
}
|
|
|
|
fn affine(
|
|
&self,
|
|
prefix: &str,
|
|
in_dim: u32,
|
|
out_dim: u32,
|
|
expert: Option<u32>,
|
|
) -> Result<Affine<'_>, String> {
|
|
let packed = self.weight(&format!("{prefix}.weight"))?;
|
|
let scales = self.weight(&format!("{prefix}.scales"))?;
|
|
let biases = self.weight(&format!("{prefix}.biases"))?;
|
|
let bits = packed
|
|
.tensor
|
|
.quant_bits
|
|
.ok_or_else(|| format!("{prefix} is not affine quantized"))?;
|
|
let group = packed
|
|
.tensor
|
|
.group_size
|
|
.ok_or_else(|| format!("{prefix} has no affine group size"))?
|
|
as u32;
|
|
if !matches!(bits, 2 | 4 | 8) || !in_dim.is_multiple_of(group) {
|
|
return Err(format!("{prefix} has an incompatible affine layout"));
|
|
}
|
|
let expected = if expert.is_some() {
|
|
vec![
|
|
EXPERTS as u64,
|
|
out_dim as u64,
|
|
(in_dim / (32 / bits)) as u64,
|
|
]
|
|
} else {
|
|
vec![out_dim as u64, (in_dim / (32 / bits)) as u64]
|
|
};
|
|
if packed.tensor.shape != expected {
|
|
return Err(format!("{prefix} has an incompatible executor shape"));
|
|
}
|
|
let expected_parameters = if expert.is_some() {
|
|
vec![EXPERTS as u64, out_dim as u64, (in_dim / group) as u64]
|
|
} else {
|
|
vec![out_dim as u64, (in_dim / group) as u64]
|
|
};
|
|
if packed.tensor.dtype != "U32"
|
|
|| scales.tensor.dtype != "BF16"
|
|
|| biases.tensor.dtype != "BF16"
|
|
|| scales.tensor.shape != expected_parameters
|
|
|| biases.tensor.shape != expected_parameters
|
|
{
|
|
return Err(format!("{prefix} has incompatible affine parameters"));
|
|
}
|
|
Ok(Affine {
|
|
packed: Weight {
|
|
tensor: packed.tensor,
|
|
expert,
|
|
},
|
|
scales: Weight {
|
|
tensor: scales.tensor,
|
|
expert,
|
|
},
|
|
biases: Weight {
|
|
tensor: biases.tensor,
|
|
expert,
|
|
},
|
|
bits,
|
|
group,
|
|
})
|
|
}
|
|
|
|
fn weight(&self, name: &str) -> Result<Weight<'_>, String> {
|
|
self.model.tensor(name).map(|tensor| Weight {
|
|
tensor,
|
|
expert: None,
|
|
})
|
|
}
|
|
|
|
fn view(&self, weight: Weight<'_>) -> QwenWeightView {
|
|
let (map, _) = self.model.map(weight.tensor.map);
|
|
let experts = weight.expert.map_or(1, |_| EXPERTS as u64);
|
|
let bytes = (weight.tensor.range.end - weight.tensor.range.start) / experts;
|
|
let offset = weight.tensor.range.start + u64::from(weight.expert.unwrap_or(0)) * bytes;
|
|
QwenWeightView {
|
|
map: map.as_ptr().cast(),
|
|
size: map.len() as u64,
|
|
offset,
|
|
bytes,
|
|
}
|
|
}
|
|
|
|
fn affine_mv_into(
|
|
&self,
|
|
weight: &Affine<'_>,
|
|
input: &Buffer,
|
|
output: &Buffer,
|
|
in_dim: u32,
|
|
out_dim: u32,
|
|
) -> Result<(), String> {
|
|
let mut values = args();
|
|
values.u[0] = in_dim;
|
|
values.u[1] = out_dim;
|
|
values.u[2] = weight.bits;
|
|
values.u[3] = weight.group;
|
|
self.dispatch(
|
|
c"kernel_qwen_affine_mv",
|
|
output,
|
|
Some(input),
|
|
None,
|
|
None,
|
|
&[
|
|
self.view(weight.packed),
|
|
self.view(weight.scales),
|
|
self.view(weight.biases),
|
|
],
|
|
&values,
|
|
out_dim,
|
|
1,
|
|
)
|
|
}
|
|
|
|
fn bf16_mv(
|
|
&self,
|
|
weight: Weight<'_>,
|
|
input: &Buffer,
|
|
output: &Buffer,
|
|
in_dim: u32,
|
|
out_dim: u32,
|
|
) -> Result<(), String> {
|
|
if weight.tensor.dtype != "BF16" || weight.tensor.shape != [out_dim as u64, in_dim as u64] {
|
|
return Err(format!(
|
|
"{} has an incompatible BF16 matrix shape",
|
|
weight.tensor.name
|
|
));
|
|
}
|
|
let mut values = args();
|
|
values.u[0] = in_dim;
|
|
values.u[1] = out_dim;
|
|
self.dispatch(
|
|
c"kernel_qwen_bf16_mv",
|
|
output,
|
|
Some(input),
|
|
None,
|
|
None,
|
|
&[self.view(weight)],
|
|
&values,
|
|
out_dim,
|
|
1,
|
|
)
|
|
}
|
|
|
|
#[allow(clippy::too_many_arguments)]
|
|
fn dispatch(
|
|
&self,
|
|
kernel: &CStr,
|
|
out: &Buffer,
|
|
a: Option<&Buffer>,
|
|
b: Option<&Buffer>,
|
|
c: Option<&Buffer>,
|
|
weights: &[QwenWeightView],
|
|
values: &QwenKernelArgs,
|
|
grid_x: u32,
|
|
grid_y: u32,
|
|
) -> Result<(), String> {
|
|
dispatch_qwen(kernel, out, a, b, c, weights, values, grid_x, grid_y)
|
|
}
|
|
|
|
pub(super) fn prefill(
|
|
&mut self,
|
|
tokens: &[i32],
|
|
mut progress: impl FnMut(u32) -> bool,
|
|
) -> Result<usize, String> {
|
|
let mut completed = 0;
|
|
for &token in tokens {
|
|
if !progress(completed) {
|
|
break;
|
|
}
|
|
self.eval(token)?;
|
|
completed += 1;
|
|
}
|
|
progress(completed);
|
|
Ok(completed as usize)
|
|
}
|
|
|
|
pub(super) fn logits(&self) -> &[f32] {
|
|
&self.logits
|
|
}
|
|
pub(super) fn execution_stats(&self) -> ExecutionStats {
|
|
ExecutionStats::default()
|
|
}
|
|
pub(super) fn model(&self) -> &QwenModel {
|
|
&self.model
|
|
}
|
|
pub(super) fn context(&self) -> u32 {
|
|
self.context
|
|
}
|
|
pub(super) fn position(&self) -> u32 {
|
|
self.position
|
|
}
|
|
pub(super) fn tokens(&self) -> &[i32] {
|
|
&self.tokens
|
|
}
|
|
pub(super) fn checkpoint_tag(&self) -> [u8; 32] {
|
|
self.checkpoint_tag
|
|
}
|
|
pub(super) fn note_checkpoint_tag(&mut self, tag: [u8; 32]) {
|
|
self.checkpoint_tag = tag;
|
|
}
|
|
|
|
pub(super) fn reset(&mut self) -> Result<(), String> {
|
|
self.states = allocate_states(self.context)?;
|
|
self.ple_state = allocate_ple_state()?;
|
|
self.logits.fill(0.0);
|
|
self.tokens.clear();
|
|
self.position = 0;
|
|
self.checkpoint_tag = [0; 32];
|
|
Ok(())
|
|
}
|
|
|
|
pub(super) fn align_prompt(&mut self, tokens: &[i32]) -> Result<usize, String> {
|
|
if !tokens.starts_with(&self.tokens) {
|
|
self.reset()?;
|
|
}
|
|
Ok(self.tokens.len())
|
|
}
|
|
|
|
fn blank_resident(&self) -> Result<QwenResidentState, String> {
|
|
Ok(QwenResidentState {
|
|
states: allocate_states(self.context)?,
|
|
ple_state: allocate_ple_state()?,
|
|
logits: vec![0.0; VOCAB as usize],
|
|
tokens: Vec::new(),
|
|
position: 0,
|
|
checkpoint_tag: [0; 32],
|
|
})
|
|
}
|
|
|
|
pub(super) fn swap_resident_state(
|
|
&mut self,
|
|
state: &mut Option<QwenResidentState>,
|
|
) -> Result<(), String> {
|
|
let mut incoming = state.take().map_or_else(|| self.blank_resident(), Ok)?;
|
|
std::mem::swap(&mut self.states, &mut incoming.states);
|
|
std::mem::swap(&mut self.ple_state, &mut incoming.ple_state);
|
|
std::mem::swap(&mut self.logits, &mut incoming.logits);
|
|
std::mem::swap(&mut self.tokens, &mut incoming.tokens);
|
|
std::mem::swap(&mut self.position, &mut incoming.position);
|
|
std::mem::swap(&mut self.checkpoint_tag, &mut incoming.checkpoint_tag);
|
|
*state = Some(incoming);
|
|
Ok(())
|
|
}
|
|
|
|
pub(in crate::engine) fn save_checkpoint(
|
|
&mut self,
|
|
path: &Path,
|
|
tag: [u8; 32],
|
|
progress: &mut impl FnMut(u64),
|
|
) -> Result<(), String> {
|
|
if let Some(parent) = path.parent() {
|
|
fs::create_dir_all(parent).map_err(|error| error.to_string())?;
|
|
}
|
|
let temporary = path.with_extension("tmp");
|
|
let mut file = File::create(&temporary).map_err(|error| error.to_string())?;
|
|
file.write_all(CHECKPOINT_MAGIC)
|
|
.map_err(|error| error.to_string())?;
|
|
for value in [
|
|
CHECKPOINT_VERSION,
|
|
self.context,
|
|
self.position,
|
|
VOCAB,
|
|
LAYERS as u32,
|
|
] {
|
|
write_u32(&mut file, value)?;
|
|
}
|
|
file.write_all(&self.model.checkpoint_identity())
|
|
.map_err(|error| error.to_string())?;
|
|
file.write_all(&tag).map_err(|error| error.to_string())?;
|
|
for token in self.ple_state.history {
|
|
write_u32(&mut file, token as u32)?;
|
|
}
|
|
for &token in &self.tokens {
|
|
write_u32(&mut file, token as u32)?;
|
|
}
|
|
for &logit in &self.logits {
|
|
write_u32(&mut file, logit.to_bits())?;
|
|
}
|
|
let mut chunk = vec![0; CHECKPOINT_CHUNK];
|
|
write_buffer(
|
|
&mut file,
|
|
&self.ple_state.conv,
|
|
0,
|
|
u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2,
|
|
&mut chunk,
|
|
progress,
|
|
)?;
|
|
for state in &self.states {
|
|
match state {
|
|
LayerState::Gdn { conv, recurrent } => {
|
|
write_buffer(
|
|
&mut file,
|
|
conv,
|
|
0,
|
|
u64::from(GDN_QKV) * 3 * 2,
|
|
&mut chunk,
|
|
progress,
|
|
)?;
|
|
write_buffer(
|
|
&mut file,
|
|
recurrent,
|
|
0,
|
|
u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64 * 4,
|
|
&mut chunk,
|
|
progress,
|
|
)?;
|
|
}
|
|
LayerState::Attention { kv } => {
|
|
write_buffer(
|
|
&mut file,
|
|
kv,
|
|
0,
|
|
u64::from(self.position) * ATTN_KV_WIDTH as u64 * 4,
|
|
&mut chunk,
|
|
progress,
|
|
)?;
|
|
}
|
|
}
|
|
}
|
|
file.sync_all().map_err(|error| error.to_string())?;
|
|
fs::rename(temporary, path).map_err(|error| error.to_string())?;
|
|
self.checkpoint_tag = tag;
|
|
Ok(())
|
|
}
|
|
|
|
pub(in crate::engine) fn load_checkpoint(
|
|
&mut self,
|
|
path: &Path,
|
|
progress: &mut impl FnMut(u64),
|
|
) -> Result<bool, String> {
|
|
let mut file = match File::open(path) {
|
|
Ok(file) => file,
|
|
Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(false),
|
|
Err(error) => return Err(error.to_string()),
|
|
};
|
|
let mut magic = [0; 8];
|
|
file.read_exact(&mut magic)
|
|
.map_err(|error| error.to_string())?;
|
|
if &magic != CHECKPOINT_MAGIC {
|
|
return Err("Qwen checkpoint has an invalid signature".into());
|
|
}
|
|
for expected in [CHECKPOINT_VERSION, self.context] {
|
|
if read_u32(&mut file)? != expected {
|
|
return Err("Qwen checkpoint does not match the current executor".into());
|
|
}
|
|
}
|
|
let position = read_u32(&mut file)?;
|
|
if position > self.context
|
|
|| read_u32(&mut file)? != VOCAB
|
|
|| read_u32(&mut file)? != LAYERS as u32
|
|
{
|
|
return Err("Qwen checkpoint shape is invalid".into());
|
|
}
|
|
let mut identity = [0; 32];
|
|
file.read_exact(&mut identity)
|
|
.map_err(|error| error.to_string())?;
|
|
if identity != self.model.checkpoint_identity() {
|
|
return Err("Qwen checkpoint model identity changed".into());
|
|
}
|
|
let mut tag = [0; 32];
|
|
file.read_exact(&mut tag)
|
|
.map_err(|error| error.to_string())?;
|
|
let mut ple_history = [0; PLE_HISTORY];
|
|
for token in &mut ple_history {
|
|
let value = read_u32(&mut file)?;
|
|
if value >= VOCAB {
|
|
return Err("Qwen checkpoint PLE history is invalid".into());
|
|
}
|
|
*token = value as i32;
|
|
}
|
|
let mut tokens = Vec::with_capacity(position as usize);
|
|
for _ in 0..position {
|
|
let token = read_u32(&mut file)?;
|
|
if token >= VOCAB {
|
|
return Err("Qwen checkpoint token is invalid".into());
|
|
}
|
|
tokens.push(token as i32);
|
|
}
|
|
let mut logits = Vec::with_capacity(VOCAB as usize);
|
|
for _ in 0..VOCAB {
|
|
logits.push(f32::from_bits(read_u32(&mut file)?));
|
|
}
|
|
self.reset()?;
|
|
let mut chunk = vec![0; CHECKPOINT_CHUNK];
|
|
read_buffer(
|
|
&mut file,
|
|
&self.ple_state.conv,
|
|
0,
|
|
u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2,
|
|
&mut chunk,
|
|
progress,
|
|
)?;
|
|
for state in &self.states {
|
|
match state {
|
|
LayerState::Gdn { conv, recurrent } => {
|
|
read_buffer(
|
|
&mut file,
|
|
conv,
|
|
0,
|
|
u64::from(GDN_QKV) * 3 * 2,
|
|
&mut chunk,
|
|
progress,
|
|
)?;
|
|
read_buffer(
|
|
&mut file,
|
|
recurrent,
|
|
0,
|
|
u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64 * 4,
|
|
&mut chunk,
|
|
progress,
|
|
)?;
|
|
}
|
|
LayerState::Attention { kv } => read_buffer(
|
|
&mut file,
|
|
kv,
|
|
0,
|
|
u64::from(position) * ATTN_KV_WIDTH as u64 * 4,
|
|
&mut chunk,
|
|
progress,
|
|
)?,
|
|
}
|
|
}
|
|
let mut trailing = [0];
|
|
if file
|
|
.read(&mut trailing)
|
|
.map_err(|error| error.to_string())?
|
|
!= 0
|
|
{
|
|
self.reset()?;
|
|
return Err("Qwen checkpoint has trailing data".into());
|
|
}
|
|
self.position = position;
|
|
self.tokens = tokens;
|
|
self.logits = logits;
|
|
self.ple_state.history = ple_history;
|
|
self.checkpoint_tag = tag;
|
|
Ok(true)
|
|
}
|
|
}
|
|
|
|
fn allocate_states(context: u32) -> Result<Vec<LayerState>, String> {
|
|
(0..LAYERS)
|
|
.map(|layer| {
|
|
if layer % 4 == 3 {
|
|
Ok(LayerState::Attention {
|
|
kv: Buffer::bytes(u64::from(context) * ATTN_KV_WIDTH as u64 * 4)?,
|
|
})
|
|
} else {
|
|
let conv = Buffer::bytes(u64::from(GDN_QKV) * 3 * 2)?;
|
|
let recurrent =
|
|
Buffer::floats(u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64)?;
|
|
conv.fill(0.0, u64::from(GDN_QKV) * 3 / 2)?;
|
|
recurrent.fill(
|
|
0.0,
|
|
u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64,
|
|
)?;
|
|
Ok(LayerState::Gdn { conv, recurrent })
|
|
}
|
|
})
|
|
.collect()
|
|
}
|
|
|
|
fn allocate_ple_state() -> Result<PleState, String> {
|
|
let conv = Buffer::bytes(u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2)?;
|
|
conv.fill(0.0, u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 / 2)?;
|
|
Ok(PleState {
|
|
history: [EOS_TOKEN; PLE_HISTORY],
|
|
conv,
|
|
})
|
|
}
|
|
|
|
fn ple_contract(model: &QwenModel) -> Result<PleContract, String> {
|
|
let multipliers = read_i64_array::<3>(
|
|
model,
|
|
"language_model.model.layers.1.ple.ple_embedding.layer_multipliers",
|
|
)?;
|
|
let sizes = read_i64_array::<PLE_HEADS>(
|
|
model,
|
|
"language_model.model.layers.1.ple.ple_embedding.ngram_heads_vocab_sizes",
|
|
)?;
|
|
let offsets = read_i64_array::<PLE_HEADS>(
|
|
model,
|
|
"language_model.model.layers.1.ple.ple_embedding.ngram_heads_offsets",
|
|
)?;
|
|
let expected_multipliers = official_ple_multipliers();
|
|
let mut expected_sizes = [0; PLE_HEADS];
|
|
let mut expected_offsets = [0; PLE_HEADS];
|
|
let mut total = 0_i64;
|
|
let mut prime = 19_999_999_i64;
|
|
for head in 0..PLE_HEADS {
|
|
prime = next_prime(prime);
|
|
expected_sizes[head] = prime;
|
|
expected_offsets[head] = total;
|
|
total += prime;
|
|
}
|
|
if multipliers != expected_multipliers || sizes != expected_sizes || offsets != expected_offsets
|
|
{
|
|
return Err("Qwen PLE hash parameters do not match the official contract".into());
|
|
}
|
|
for (name, shape, bits, group) in [
|
|
("ngram.weight", [320_001_536, 20], Some(4), Some(32)),
|
|
("ngram.scales", [320_001_536, 5], Some(4), Some(32)),
|
|
("ngram.biases", [320_001_536, 5], Some(4), Some(32)),
|
|
] {
|
|
let tensor = model.tensor(name)?;
|
|
if tensor.shape != shape || tensor.quant_bits != bits || tensor.group_size != group {
|
|
return Err(format!("{name} does not match the Qwen PLE row layout"));
|
|
}
|
|
}
|
|
Ok(PleContract {
|
|
multipliers,
|
|
sizes,
|
|
offsets,
|
|
})
|
|
}
|
|
|
|
fn read_i64_array<const N: usize>(model: &QwenModel, name: &str) -> Result<[i64; N], String> {
|
|
let tensor = model.tensor(name)?;
|
|
if tensor.dtype != "I64" || tensor.shape != [N as u64] {
|
|
return Err(format!("{name} does not match the Qwen PLE integer layout"));
|
|
}
|
|
let bytes = model.tensor_bytes(tensor)?;
|
|
if bytes.len() != N * 8 {
|
|
return Err(format!("{name} has an invalid byte length"));
|
|
}
|
|
Ok(std::array::from_fn(|index| {
|
|
i64::from_le_bytes(bytes[index * 8..index * 8 + 8].try_into().unwrap())
|
|
}))
|
|
}
|
|
|
|
fn official_ple_multipliers() -> [i64; 3] {
|
|
const GAMMA: u64 = 0x9e37_79b9_7f4a_7c15;
|
|
const M1: u64 = 0xbf58_476d_1ce4_e5b9;
|
|
const M2: u64 = 0x94d0_49bb_1331_11eb;
|
|
let bound = (i64::MAX / VOCAB as i64 / 2) as u64;
|
|
std::array::from_fn(|index| {
|
|
let mut value = 1234_u64.wrapping_add(GAMMA.wrapping_mul(index as u64 + 1));
|
|
value = value.wrapping_add(GAMMA);
|
|
value = (value ^ (value >> 30)).wrapping_mul(M1);
|
|
value = (value ^ (value >> 27)).wrapping_mul(M2);
|
|
value ^= value >> 31;
|
|
(2 * (value % bound) + 1) as i64
|
|
})
|
|
}
|
|
|
|
fn next_prime(mut value: i64) -> i64 {
|
|
loop {
|
|
value += 1;
|
|
if value % 2 != 0
|
|
&& (3..=((value as f64).sqrt() as i64))
|
|
.step_by(2)
|
|
.all(|divisor| value % divisor != 0)
|
|
{
|
|
return value;
|
|
}
|
|
}
|
|
}
|
|
|
|
fn ple_rows(
|
|
contract: &PleContract,
|
|
history: [i32; PLE_HISTORY],
|
|
token: i32,
|
|
) -> Result<[u64; PLE_HEADS], String> {
|
|
if token < 0 || token as u32 >= VOCAB {
|
|
return Err(format!("token {token} is outside the Qwen vocabulary"));
|
|
}
|
|
let shifted = [token as i64, history[1] as i64, history[0] as i64];
|
|
let two = shifted[0].wrapping_mul(contract.multipliers[0])
|
|
^ shifted[1].wrapping_mul(contract.multipliers[1]);
|
|
let three = two ^ shifted[2].wrapping_mul(contract.multipliers[2]);
|
|
Ok(std::array::from_fn(|head| {
|
|
let mixed = if head < 8 { two } else { three };
|
|
(contract.offsets[head] + mixed.rem_euclid(contract.sizes[head])) as u64
|
|
}))
|
|
}
|
|
|
|
fn advance_ple_history(history: [i32; PLE_HISTORY], token: i32) -> [i32; PLE_HISTORY] {
|
|
if token == EOS_TOKEN {
|
|
[EOS_TOKEN; PLE_HISTORY]
|
|
} else {
|
|
[history[1], token]
|
|
}
|
|
}
|
|
|
|
fn gather_ple_row(
|
|
row: u64,
|
|
packed: &[u8],
|
|
scales: &[u8],
|
|
biases: &[u8],
|
|
) -> Result<[u8; PLE_ROW_BYTES], String> {
|
|
let row = usize::try_from(row).map_err(|_| "Qwen PLE row exceeds this platform")?;
|
|
let mut value = [0; PLE_ROW_BYTES];
|
|
value[..80].copy_from_slice(ple_row_slice(packed, row, 80)?);
|
|
value[80..90].copy_from_slice(ple_row_slice(scales, row, 10)?);
|
|
value[90..].copy_from_slice(ple_row_slice(biases, row, 10)?);
|
|
Ok(value)
|
|
}
|
|
|
|
fn ple_row_slice(data: &[u8], row: usize, width: usize) -> Result<&[u8], String> {
|
|
let start = row
|
|
.checked_mul(width)
|
|
.ok_or_else(|| "Qwen PLE row offset overflows".to_owned())?;
|
|
data.get(start..start + width)
|
|
.ok_or_else(|| format!("Qwen PLE row {row} is truncated"))
|
|
}
|
|
|
|
fn args() -> QwenKernelArgs {
|
|
QwenKernelArgs::default()
|
|
}
|
|
|
|
#[allow(clippy::too_many_arguments)]
|
|
fn dispatch_qwen(
|
|
kernel: &CStr,
|
|
out: &Buffer,
|
|
a: Option<&Buffer>,
|
|
b: Option<&Buffer>,
|
|
c: Option<&Buffer>,
|
|
weights: &[QwenWeightView],
|
|
values: &QwenKernelArgs,
|
|
grid_x: u32,
|
|
grid_y: u32,
|
|
) -> Result<(), String> {
|
|
call(
|
|
unsafe {
|
|
ds4_gpu_qwen_dispatch(
|
|
kernel.as_ptr(),
|
|
out.raw(),
|
|
a.map_or(std::ptr::null(), |buffer| buffer.raw().cast_const()),
|
|
b.map_or(std::ptr::null(), |buffer| buffer.raw().cast_const()),
|
|
c.map_or(std::ptr::null(), |buffer| buffer.raw().cast_const()),
|
|
weights.as_ptr(),
|
|
weights.len() as u32,
|
|
values,
|
|
grid_x,
|
|
grid_y,
|
|
)
|
|
},
|
|
kernel.to_str().unwrap_or("running a Qwen Metal kernel"),
|
|
)
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use memmap2::MmapOptions;
|
|
use sha2::{Digest, Sha256};
|
|
use std::fs;
|
|
use std::path::PathBuf;
|
|
|
|
fn bf16(value: f32) -> u16 {
|
|
let bits = value.to_bits();
|
|
((bits + 0x7fff + ((bits >> 16) & 1)) >> 16) as u16
|
|
}
|
|
|
|
fn close(actual: f32, expected: f32) {
|
|
assert!(
|
|
(actual - expected).abs() <= 2.0e-4 * expected.abs().max(1.0),
|
|
"{actual} != {expected}"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn qwen_ple_hash_contract_matches_golden_boundaries() {
|
|
assert_eq!(
|
|
official_ple_multipliers(),
|
|
[23_703_573_157_769, 20_109_073_645_365, 8_052_911_324_071]
|
|
);
|
|
let mut contract = PleContract {
|
|
multipliers: official_ple_multipliers(),
|
|
sizes: [0; PLE_HEADS],
|
|
offsets: [0; PLE_HEADS],
|
|
};
|
|
let mut prime = 19_999_999;
|
|
let mut offset = 0;
|
|
for head in 0..PLE_HEADS {
|
|
prime = next_prime(prime);
|
|
contract.sizes[head] = prime;
|
|
contract.offsets[head] = offset;
|
|
offset += prime;
|
|
}
|
|
let initial = ple_rows(&contract, [EOS_TOKEN; 2], 1).unwrap();
|
|
let repeated = ple_rows(&contract, [1, 1], 1).unwrap();
|
|
let boundary = ple_rows(&contract, [EOS_TOKEN, 42], 43).unwrap();
|
|
assert_eq!(
|
|
initial,
|
|
[
|
|
16_121_432,
|
|
28_938_500,
|
|
59_087_997,
|
|
73_487_090,
|
|
81_148_277,
|
|
104_500_129,
|
|
120_276_032,
|
|
149_373_875,
|
|
176_283_436,
|
|
184_305_849,
|
|
216_528_839,
|
|
231_080_079,
|
|
257_961_536,
|
|
266_068_568,
|
|
289_043_455,
|
|
305_959_965,
|
|
]
|
|
);
|
|
assert_eq!(
|
|
repeated,
|
|
[
|
|
6_868_091,
|
|
38_325_817,
|
|
54_054_700,
|
|
68_075_137,
|
|
82_949_816,
|
|
101_241_419,
|
|
138_678_867,
|
|
155_262_032,
|
|
176_541_251,
|
|
196_154_476,
|
|
215_703_237,
|
|
234_413_824,
|
|
254_220_543,
|
|
274_027_268,
|
|
293_962_951,
|
|
313_640_732,
|
|
]
|
|
);
|
|
assert_eq!(
|
|
boundary,
|
|
[
|
|
18_529_343,
|
|
23_547_650,
|
|
56_056_978,
|
|
73_570_159,
|
|
88_581_601,
|
|
113_585_506,
|
|
121_091_299,
|
|
151_099_148,
|
|
175_585_266,
|
|
184_439_538,
|
|
216_587_431,
|
|
222_082_137,
|
|
250_284_866,
|
|
278_847_169,
|
|
281_781_121,
|
|
317_050_322,
|
|
]
|
|
);
|
|
assert_eq!(contract.offsets[0], 0);
|
|
assert_eq!(contract.offsets[8], 160_000_374);
|
|
assert_eq!(contract.offsets[15] + contract.sizes[15], 320_001_446);
|
|
for (head, &row) in initial.iter().enumerate() {
|
|
assert!(
|
|
(contract.offsets[head]..contract.offsets[head] + contract.sizes[head])
|
|
.contains(&(row as i64))
|
|
);
|
|
}
|
|
assert_eq!(320_001_536 % 128, 0);
|
|
let history = advance_ple_history([EOS_TOKEN; 2], 1);
|
|
assert_eq!(advance_ple_history(history, 2), [1, 2]);
|
|
assert_eq!(advance_ple_history([1, 2], EOS_TOKEN), [EOS_TOKEN; 2]);
|
|
|
|
let hash_chunks = |chunk: usize| {
|
|
let mut history = [EOS_TOKEN; PLE_HISTORY];
|
|
[1, 2, EOS_TOKEN, 3, 4]
|
|
.chunks(chunk)
|
|
.flat_map(|tokens| {
|
|
tokens
|
|
.iter()
|
|
.map(|&token| {
|
|
let rows = ple_rows(&contract, history, token).unwrap();
|
|
history = advance_ple_history(history, token);
|
|
rows
|
|
})
|
|
.collect::<Vec<_>>()
|
|
})
|
|
.collect::<Vec<_>>()
|
|
};
|
|
assert_eq!(hash_chunks(1), hash_chunks(2));
|
|
assert_eq!(hash_chunks(1), hash_chunks(5));
|
|
}
|
|
|
|
#[test]
|
|
fn qwen_ple_gather_copies_only_the_requested_row() {
|
|
let packed = (0..240).map(|value| value as u8).collect::<Vec<_>>();
|
|
let scales = (0..30).map(|value| (value + 17) as u8).collect::<Vec<_>>();
|
|
let biases = (0..30).map(|value| (value + 47) as u8).collect::<Vec<_>>();
|
|
let direct = gather_ple_row(1, &packed, &scales, &biases).unwrap();
|
|
assert_eq!(&direct[..80], &packed[80..160]);
|
|
assert_eq!(&direct[80..90], &scales[10..20]);
|
|
assert_eq!(&direct[90..], &biases[10..20]);
|
|
assert_eq!(
|
|
gather_ple_row(1, &packed, &scales, &biases).unwrap(),
|
|
direct
|
|
);
|
|
assert!(
|
|
gather_ple_row(3, &packed, &scales, &biases)
|
|
.unwrap_err()
|
|
.contains("truncated")
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "requires Apple Metal"]
|
|
fn qwen_metal_primitives_match_reference_vectors() {
|
|
configure_sources().unwrap();
|
|
let _context = Context::open_qwen(0).unwrap();
|
|
|
|
let input = Buffer::floats(64).unwrap();
|
|
input.write_f32(&vec![1.0; 64]).unwrap();
|
|
let output = Buffer::floats(64).unwrap();
|
|
let path = std::env::temp_dir().join(format!(
|
|
"ds4-qwen96-weights-{}-{}",
|
|
std::process::id(),
|
|
std::thread::current().name().unwrap_or("test")
|
|
));
|
|
// Safetensors headers are not guaranteed to align the data region.
|
|
let mut bytes = vec![0];
|
|
for _ in 0..8 {
|
|
bytes.extend_from_slice(&0x3333_3333_u32.to_le_bytes());
|
|
}
|
|
for value in [
|
|
0.5, -1.0, 0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 1.0, 1.0, 0.0, 0.0, 0.0, 0.0,
|
|
] {
|
|
bytes.extend_from_slice(&bf16(value).to_le_bytes());
|
|
}
|
|
fs::write(&path, bytes).unwrap();
|
|
let file = File::open(&path).unwrap();
|
|
// SAFETY: this test owns the read-only file for the lifetime of the mapping.
|
|
let map = unsafe { MmapOptions::new().map(&file).unwrap() };
|
|
let view = |offset, size| QwenWeightView {
|
|
map: map.as_ptr().cast(),
|
|
size: map.len() as u64,
|
|
offset,
|
|
bytes: size,
|
|
};
|
|
|
|
let mut affine = args();
|
|
affine.u[0] = 64;
|
|
affine.u[1] = 1;
|
|
affine.u[2] = 4;
|
|
affine.u[3] = 64;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_affine_mv",
|
|
&output,
|
|
Some(&input),
|
|
None,
|
|
None,
|
|
&[view(1, 32), view(33, 2), view(35, 2)],
|
|
&affine,
|
|
1,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut scalar = [0.0];
|
|
output.read_f32(&mut scalar).unwrap();
|
|
close(scalar[0], 32.0);
|
|
|
|
let conv_input = Buffer::floats(1).unwrap();
|
|
conv_input.write_f32(&[2.0]).unwrap();
|
|
let conv_state = Buffer::bytes(6).unwrap();
|
|
conv_state.write(0, &[0; 6]).unwrap();
|
|
let conv_output = Buffer::floats(1).unwrap();
|
|
let mut conv = args();
|
|
conv.u[0] = 1;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_conv_silu",
|
|
&conv_output,
|
|
Some(&conv_input),
|
|
Some(&conv_state),
|
|
None,
|
|
&[view(41, 8)],
|
|
&conv,
|
|
1,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
conv_output.read_f32(&mut scalar).unwrap();
|
|
close(scalar[0], 8.0 / (1.0 + (-8.0_f32).exp()));
|
|
assert_eq!(
|
|
{
|
|
let mut state = [0; 6];
|
|
conv_state.read(0, &mut state).unwrap();
|
|
state
|
|
},
|
|
[0, 0, 0, 0, 0, 64]
|
|
);
|
|
|
|
let qkv = Buffer::floats(6).unwrap();
|
|
qkv.write_f32(&[3.0, 4.0, 0.0, 2.0, 5.0, 7.0]).unwrap();
|
|
let controls = Buffer::floats(4).unwrap();
|
|
controls.write_f32(&[0.0; 4]).unwrap();
|
|
let recurrent = Buffer::floats(4).unwrap();
|
|
recurrent.write_f32(&[1.0, 2.0, 3.0, 4.0]).unwrap();
|
|
let raw = Buffer::floats(2).unwrap();
|
|
let gated = Buffer::floats(2).unwrap();
|
|
let mut step = args();
|
|
step.u[0] = 2;
|
|
step.u[1] = 1;
|
|
step.u[2] = 1;
|
|
step.u[3] = 2;
|
|
step.u[4] = 3;
|
|
step.f[0] = 1.0e-6;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_gdn_step",
|
|
&raw,
|
|
Some(&qkv),
|
|
Some(&controls),
|
|
Some(&recurrent),
|
|
&[view(37, 2), view(39, 2)],
|
|
&step,
|
|
2,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut raw_values = [0.0; 2];
|
|
raw.read_f32(&mut raw_values).unwrap();
|
|
close(raw_values[0], 2.0506096);
|
|
close(raw_values[1], 2.9698484);
|
|
let mut norm = args();
|
|
norm.u[0] = 2;
|
|
norm.u[1] = 1;
|
|
norm.f[0] = 1.0e-6;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_gdn_norm_gate",
|
|
&gated,
|
|
Some(&raw),
|
|
Some(&controls),
|
|
None,
|
|
&[view(49, 4)],
|
|
&norm,
|
|
1,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut gated_values = [0.0; 2];
|
|
gated.read_f32(&mut gated_values).unwrap();
|
|
let rms = ((raw_values[0].powi(2) + raw_values[1].powi(2)) * 0.5 + 1.0e-6).sqrt();
|
|
close(gated_values[0], raw_values[0] / rms * 0.5);
|
|
close(gated_values[1], raw_values[1] / rms * 0.5);
|
|
|
|
let router = Buffer::floats(EXPERTS.into()).unwrap();
|
|
let mut logits = vec![-100.0; EXPERTS as usize];
|
|
for (index, value) in logits.iter_mut().take(EXPERTS_USED).enumerate() {
|
|
*value = index as f32;
|
|
}
|
|
router.write_f32(&logits).unwrap();
|
|
let ids = Buffer::bytes((EXPERTS_USED * 4) as u64).unwrap();
|
|
let weights = Buffer::floats(EXPERTS_USED as u64).unwrap();
|
|
let mut route = args();
|
|
route.u[0] = EXPERTS;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_route_top10",
|
|
&ids,
|
|
Some(&weights),
|
|
Some(&router),
|
|
None,
|
|
&[],
|
|
&route,
|
|
1,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut selected = [0; EXPERTS_USED];
|
|
ids.read_i32(&mut selected).unwrap();
|
|
assert_eq!(selected, [9, 8, 7, 6, 5, 4, 3, 2, 1, 0]);
|
|
let mut probabilities = [0.0; EXPERTS_USED];
|
|
weights.read_f32(&mut probabilities).unwrap();
|
|
close(probabilities.iter().sum(), 1.0);
|
|
|
|
let norm_input = Buffer::floats(4).unwrap();
|
|
norm_input.write_f32(&[1.0, 2.0, 3.0, 4.0]).unwrap();
|
|
let norm_output = Buffer::floats(4).unwrap();
|
|
let mut zero_norm = args();
|
|
zero_norm.u[0] = 4;
|
|
zero_norm.u[1] = 2;
|
|
zero_norm.f[0] = 1.0e-6;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_zero_rms",
|
|
&norm_output,
|
|
Some(&norm_input),
|
|
None,
|
|
None,
|
|
&[view(53, 8)],
|
|
&zero_norm,
|
|
2,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut normalized = [0.0; 4];
|
|
norm_output.read_f32(&mut normalized).unwrap();
|
|
for (actual, expected) in normalized.into_iter().zip([
|
|
1.0 / (2.5_f32 + 1.0e-6).sqrt(),
|
|
2.0 / (2.5_f32 + 1.0e-6).sqrt(),
|
|
3.0 / (12.5_f32 + 1.0e-6).sqrt(),
|
|
4.0 / (12.5_f32 + 1.0e-6).sqrt(),
|
|
]) {
|
|
close(actual, expected);
|
|
}
|
|
|
|
let hyper_input = Buffer::floats(8).unwrap();
|
|
hyper_input
|
|
.write_f32(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0])
|
|
.unwrap();
|
|
let hyper_weights = Buffer::floats(8).unwrap();
|
|
hyper_weights.write_f32(&[1.0; 8]).unwrap();
|
|
let mixed = Buffer::floats(2).unwrap();
|
|
let mut hyper = args();
|
|
hyper.u[0] = 2;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_hyper_mix",
|
|
&mixed,
|
|
Some(&hyper_input),
|
|
Some(&hyper_weights),
|
|
None,
|
|
&[],
|
|
&hyper,
|
|
2,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut mixed_values = [0.0; 2];
|
|
mixed.read_f32(&mut mixed_values).unwrap();
|
|
assert_eq!(mixed_values, [4.0, 5.0]);
|
|
let injection = Buffer::floats(4).unwrap();
|
|
injection.write_f32(&[1.0, 2.0, 3.0, 4.0]).unwrap();
|
|
let combined = Buffer::floats(8).unwrap();
|
|
dispatch_qwen(
|
|
c"kernel_qwen_hyper_inject",
|
|
&combined,
|
|
Some(&hyper_input),
|
|
Some(&mixed),
|
|
Some(&injection),
|
|
&[],
|
|
&hyper,
|
|
8,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut combined_values = [0.0; 8];
|
|
combined.read_f32(&mut combined_values).unwrap();
|
|
assert_eq!(
|
|
combined_values,
|
|
[5.0, 7.0, 11.0, 14.0, 17.0, 21.0, 23.0, 28.0]
|
|
);
|
|
|
|
let expert_gate = Buffer::floats(2).unwrap();
|
|
expert_gate.write_f32(&[0.0, 2.0]).unwrap();
|
|
let expert_up = Buffer::floats(2).unwrap();
|
|
expert_up.write_f32(&[4.0, 3.0]).unwrap();
|
|
let expert = Buffer::floats(2).unwrap();
|
|
let mut moe = args();
|
|
moe.u[0] = 2;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_swiglu",
|
|
&expert,
|
|
Some(&expert_gate),
|
|
Some(&expert_up),
|
|
None,
|
|
&[],
|
|
&moe,
|
|
2,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let accumulated = Buffer::floats(2).unwrap();
|
|
accumulated.fill(0.0, 2).unwrap();
|
|
moe.f[0] = 0.25;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_accumulate",
|
|
&accumulated,
|
|
Some(&expert),
|
|
None,
|
|
None,
|
|
&[],
|
|
&moe,
|
|
2,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let shared = Buffer::floats(2).unwrap();
|
|
shared.write_f32(&[2.0, 4.0]).unwrap();
|
|
let shared_gate = Buffer::floats(1).unwrap();
|
|
shared_gate.write_f32(&[0.0]).unwrap();
|
|
dispatch_qwen(
|
|
c"kernel_qwen_accumulate_sigmoid_scalar",
|
|
&accumulated,
|
|
Some(&shared),
|
|
Some(&shared_gate),
|
|
None,
|
|
&[],
|
|
&moe,
|
|
2,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut moe_values = [0.0; 2];
|
|
accumulated.read_f32(&mut moe_values).unwrap();
|
|
close(moe_values[0], 1.0);
|
|
close(
|
|
moe_values[1],
|
|
0.25 * (2.0 / (1.0 + (-2.0_f32).exp()) * 3.0) + 2.0,
|
|
);
|
|
|
|
let rope_input = Buffer::floats(4).unwrap();
|
|
rope_input.write_f32(&[1.0, 2.0, 3.0, 4.0]).unwrap();
|
|
let rope_output = Buffer::floats(4).unwrap();
|
|
let mut rope = args();
|
|
rope.u[0] = 4;
|
|
rope.u[1] = 4;
|
|
rope.u[2] = 1;
|
|
rope.u[3] = 1;
|
|
rope.f[0] = 1.0e-6;
|
|
rope.f[1] = 10_000.0;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_head_norm_rope",
|
|
&rope_output,
|
|
Some(&rope_input),
|
|
None,
|
|
None,
|
|
&[view(53, 8)],
|
|
&rope,
|
|
4,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let rms = (7.5_f32 + 1.0e-6).sqrt();
|
|
let theta = [1.0_f32, 0.01];
|
|
let normalized = [1.0 / rms, 2.0 / rms, 3.0 / rms, 4.0 / rms];
|
|
let expected_rope = [
|
|
normalized[0] * theta[0].cos() - normalized[2] * theta[0].sin(),
|
|
normalized[1] * theta[1].cos() - normalized[3] * theta[1].sin(),
|
|
normalized[2] * theta[0].cos() + normalized[0] * theta[0].sin(),
|
|
normalized[3] * theta[1].cos() + normalized[1] * theta[1].sin(),
|
|
];
|
|
let mut actual_rope = [0.0; 4];
|
|
rope_output.read_f32(&mut actual_rope).unwrap();
|
|
for (actual, expected) in actual_rope.into_iter().zip(expected_rope) {
|
|
close(actual, expected);
|
|
}
|
|
|
|
let cache = Buffer::bytes(16).unwrap();
|
|
let key = Buffer::floats(2).unwrap();
|
|
let value = Buffer::floats(2).unwrap();
|
|
let mut store = args();
|
|
store.u[0] = 2;
|
|
for (position, (keys, values)) in [([1.0, 0.0], [2.0, 3.0]), ([0.0, 1.0], [5.0, 7.0])]
|
|
.into_iter()
|
|
.enumerate()
|
|
{
|
|
key.write_f32(&keys).unwrap();
|
|
value.write_f32(&values).unwrap();
|
|
store.u[1] = position as u32;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_store_kv_bf16",
|
|
&cache,
|
|
Some(&key),
|
|
Some(&value),
|
|
None,
|
|
&[],
|
|
&store,
|
|
2,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
}
|
|
let query = Buffer::floats(2).unwrap();
|
|
query.write_f32(&[1.0, 0.0]).unwrap();
|
|
let attention = Buffer::floats(2).unwrap();
|
|
let mut dense = args();
|
|
dense.u[0] = 1;
|
|
dense.u[1] = 1;
|
|
dense.u[2] = 2;
|
|
dense.u[3] = 2;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_dense_attention",
|
|
&attention,
|
|
Some(&query),
|
|
Some(&cache),
|
|
None,
|
|
&[],
|
|
&dense,
|
|
1,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let first = (1.0_f32 / 2.0_f32.sqrt()).exp();
|
|
let probability = first / (first + 1.0);
|
|
let mut actual_attention = [0.0; 2];
|
|
attention.read_f32(&mut actual_attention).unwrap();
|
|
close(
|
|
actual_attention[0],
|
|
probability * 2.0 + (1.0 - probability) * 5.0,
|
|
);
|
|
close(
|
|
actual_attention[1],
|
|
probability * 3.0 + (1.0 - probability) * 7.0,
|
|
);
|
|
|
|
let ple_packed = Buffer::bytes(80).unwrap();
|
|
ple_packed.write(0, &[0x33; 80]).unwrap();
|
|
let ple_scales = Buffer::bytes(10).unwrap();
|
|
ple_scales
|
|
.write(0, &bf16(0.5).to_le_bytes().repeat(5))
|
|
.unwrap();
|
|
let ple_biases = Buffer::bytes(10).unwrap();
|
|
ple_biases
|
|
.write(0, &bf16(-1.0).to_le_bytes().repeat(5))
|
|
.unwrap();
|
|
let ple_embedding = Buffer::floats(160).unwrap();
|
|
let mut dequant = args();
|
|
dequant.u[0] = 160;
|
|
dequant.u[1] = 1;
|
|
dequant.u[2] = 4;
|
|
dequant.u[3] = 32;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_ple_dequant",
|
|
&ple_embedding,
|
|
Some(&ple_packed),
|
|
Some(&ple_scales),
|
|
Some(&ple_biases),
|
|
&[],
|
|
&dequant,
|
|
160,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let mut embedding = [0.0; 160];
|
|
ple_embedding.read_f32(&mut embedding).unwrap();
|
|
assert!(embedding.into_iter().all(|value| value == 0.5));
|
|
|
|
let ple_key = Buffer::floats(8).unwrap();
|
|
ple_key.write_f32(&[1.0; 8]).unwrap();
|
|
let ple_query = Buffer::floats(8).unwrap();
|
|
ple_query.write_f32(&[1.0; 8]).unwrap();
|
|
let ple_value = Buffer::floats(2).unwrap();
|
|
ple_value.write_f32(&[2.0, 3.0]).unwrap();
|
|
let ple_gated = Buffer::floats(8).unwrap();
|
|
let mut gate = args();
|
|
gate.u[0] = 2;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_ple_gate",
|
|
&ple_gated,
|
|
Some(&ple_key),
|
|
Some(&ple_query),
|
|
Some(&ple_value),
|
|
&[],
|
|
&gate,
|
|
4,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
let expected_gate = 1.0 / (1.0 + (-2.0_f32.sqrt().sqrt()).exp());
|
|
let mut gated = [0.0; 8];
|
|
ple_gated.read_f32(&mut gated).unwrap();
|
|
for stream in 0..4 {
|
|
close(gated[stream * 2], 2.0 * expected_gate);
|
|
close(gated[stream * 2 + 1], 3.0 * expected_gate);
|
|
}
|
|
|
|
let ple_normalized = Buffer::floats(1).unwrap();
|
|
ple_normalized.write_f32(&[2.0]).unwrap();
|
|
let ple_gate_value = Buffer::floats(1).unwrap();
|
|
ple_gate_value.write_f32(&[0.5]).unwrap();
|
|
let ple_state = Buffer::bytes(18).unwrap();
|
|
ple_state.write(0, &[0; 18]).unwrap();
|
|
let ple_output = Buffer::floats(1).unwrap();
|
|
let mut conv = args();
|
|
conv.u[0] = 1;
|
|
dispatch_qwen(
|
|
c"kernel_qwen_ple_conv",
|
|
&ple_output,
|
|
Some(&ple_gate_value),
|
|
Some(&ple_normalized),
|
|
Some(&ple_state),
|
|
&[view(41, 8)],
|
|
&conv,
|
|
1,
|
|
1,
|
|
)
|
|
.unwrap();
|
|
ple_output.read_f32(&mut scalar).unwrap();
|
|
close(scalar[0], 0.5 + 8.0 / (1.0 + (-8.0_f32).exp()));
|
|
let mut ple_history = [0; 18];
|
|
ple_state.read(0, &mut ple_history).unwrap();
|
|
assert_eq!(&ple_history[16..], &bf16(2.0).to_le_bytes());
|
|
|
|
drop(map);
|
|
drop(file);
|
|
fs::remove_file(path).unwrap();
|
|
}
|
|
|
|
#[test]
|
|
#[ignore = "requires the pinned 105 GB Qwen artifact set and Apple Metal"]
|
|
fn qwen_core_boundary_and_checkpoint_are_stable() {
|
|
configure_sources().unwrap();
|
|
let root = std::env::var_os("DS4SERVER_QWEN38_SOURCE")
|
|
.map(PathBuf::from)
|
|
.expect("set DS4SERVER_QWEN38_SOURCE to the pinned artifact directory");
|
|
let model = QwenModel::open(&root, 4).unwrap();
|
|
let residency_before = model.mapped_residency().unwrap();
|
|
let mut executor = QwenExecutor::open(model, 4).unwrap();
|
|
executor.eval(1).unwrap();
|
|
let residency_after = executor.model().mapped_residency().unwrap();
|
|
assert!(residency_after.0 <= executor.model().memory().resident_core);
|
|
assert!(residency_after.1 <= executor.model().memory().mapped_ple);
|
|
eprintln!(
|
|
"Qwen mapped residency core/PLE before {:?}, after {:?}",
|
|
residency_before, residency_after
|
|
);
|
|
assert!(executor.logits.iter().all(|value| value.is_finite()));
|
|
let mut digest = Sha256::new();
|
|
for value in &executor.logits {
|
|
digest.update(value.to_bits().to_le_bytes());
|
|
}
|
|
let digest: [u8; 32] = digest.finalize().into();
|
|
assert_eq!(
|
|
digest,
|
|
[
|
|
137, 244, 133, 253, 201, 214, 196, 144, 249, 130, 28, 63, 124, 75, 32, 40, 16, 148,
|
|
148, 123, 5, 50, 90, 165, 101, 44, 223, 164, 62, 54, 164, 86,
|
|
]
|
|
);
|
|
let reference = [
|
|
executor.logits[0],
|
|
executor.logits[1],
|
|
executor.logits[1000],
|
|
executor.logits[VOCAB as usize - 1],
|
|
];
|
|
let next = executor
|
|
.logits
|
|
.iter()
|
|
.enumerate()
|
|
.max_by(|a, b| a.1.total_cmp(b.1))
|
|
.unwrap()
|
|
.0 as i32;
|
|
for (actual, expected) in
|
|
reference
|
|
.into_iter()
|
|
.zip([6.406_557, 2.082_818_3, -3.058_045_6, -0.121_010_3])
|
|
{
|
|
close(actual, expected);
|
|
}
|
|
assert_eq!(next, 89_648);
|
|
|
|
let checkpoint = std::env::temp_dir().join(format!(
|
|
"ds4-qwen96-checkpoint-{}-{}",
|
|
std::process::id(),
|
|
std::thread::current().name().unwrap_or("test")
|
|
));
|
|
executor
|
|
.save_checkpoint(&checkpoint, [7; 32], &mut |_| {})
|
|
.unwrap();
|
|
|
|
executor.eval(next).unwrap();
|
|
let continued = [
|
|
executor.logits[0],
|
|
executor.logits[1],
|
|
executor.logits[1000],
|
|
executor.logits[VOCAB as usize - 1],
|
|
];
|
|
let LayerState::Gdn { conv, recurrent } = &executor.states[0] else {
|
|
unreachable!()
|
|
};
|
|
let mut continued_conv = vec![0; GDN_QKV as usize * 3 * 2];
|
|
conv.read(0, &mut continued_conv).unwrap();
|
|
let mut continued_recurrent = vec![0; 4 * 1024];
|
|
recurrent.read(0, &mut continued_recurrent).unwrap();
|
|
let mut continued_ple = vec![0; HC_WIDTH as usize * PLE_CONV_STATE as usize * 2];
|
|
executor.ple_state.conv.read(0, &mut continued_ple).unwrap();
|
|
|
|
assert!(executor.load_checkpoint(&checkpoint, &mut |_| {}).unwrap());
|
|
assert_eq!(executor.position, 1);
|
|
assert_eq!(executor.tokens, [1]);
|
|
assert_eq!(executor.ple_state.history, [EOS_TOKEN, 1]);
|
|
assert_eq!(executor.checkpoint_tag, [7; 32]);
|
|
assert_eq!(
|
|
[
|
|
executor.logits[0],
|
|
executor.logits[1],
|
|
executor.logits[1000],
|
|
executor.logits[VOCAB as usize - 1],
|
|
],
|
|
reference
|
|
);
|
|
executor.eval(next).unwrap();
|
|
for (actual, expected) in [
|
|
executor.logits[0],
|
|
executor.logits[1],
|
|
executor.logits[1000],
|
|
executor.logits[VOCAB as usize - 1],
|
|
]
|
|
.into_iter()
|
|
.zip(continued)
|
|
{
|
|
close(actual, expected);
|
|
}
|
|
let LayerState::Gdn { conv, recurrent } = &executor.states[0] else {
|
|
unreachable!()
|
|
};
|
|
let mut resumed_conv = vec![0; continued_conv.len()];
|
|
conv.read(0, &mut resumed_conv).unwrap();
|
|
assert_eq!(resumed_conv, continued_conv);
|
|
let mut resumed_recurrent = vec![0; continued_recurrent.len()];
|
|
recurrent.read(0, &mut resumed_recurrent).unwrap();
|
|
assert_eq!(resumed_recurrent, continued_recurrent);
|
|
let mut resumed_ple = vec![0; continued_ple.len()];
|
|
executor.ple_state.conv.read(0, &mut resumed_ple).unwrap();
|
|
assert_eq!(resumed_ple, continued_ple);
|
|
|
|
executor.reset().unwrap();
|
|
assert_eq!(executor.position, 0);
|
|
assert_eq!(executor.ple_state.history, [EOS_TOKEN; PLE_HISTORY]);
|
|
let mut reset_ple = vec![1; continued_ple.len()];
|
|
executor.ple_state.conv.read(0, &mut reset_ple).unwrap();
|
|
assert!(reset_ple.into_iter().all(|byte| byte == 0));
|
|
executor.eval(1).unwrap();
|
|
for (actual, expected) in [
|
|
executor.logits[0],
|
|
executor.logits[1],
|
|
executor.logits[1000],
|
|
executor.logits[VOCAB as usize - 1],
|
|
]
|
|
.into_iter()
|
|
.zip(reference)
|
|
{
|
|
close(actual, expected);
|
|
}
|
|
|
|
executor.reset().unwrap();
|
|
executor.eval(EOS_TOKEN).unwrap();
|
|
assert_eq!(executor.ple_state.history, [EOS_TOKEN; PLE_HISTORY]);
|
|
|
|
executor.position = DENSE_BUDGET;
|
|
executor.context = DENSE_BUDGET + 1;
|
|
let error = executor.eval(1).unwrap_err();
|
|
assert!(error.contains("issue #97"));
|
|
fs::remove_file(checkpoint).unwrap();
|
|
}
|
|
}
|