use crate::dev_trace::dev_trace_enabled;
use crate::error::{RealizarError, Result};
use crate::gguf::OwnedQuantizedModel;
use crate::quantize::{dequantize_q4_k, dequantize_q5_k, dequantize_q6_k};
#[must_use]
fn dequant_start_message(
num_layers: usize,
hidden: usize,
num_heads: usize,
num_kv_heads: usize,
intermediate: usize,
dev_trace: bool,
) -> String {
if dev_trace {
format!(
"Preparing GPU weights: dequantizing {num_layers} layers to F32 \
(hidden={hidden}, heads={num_heads}/{num_kv_heads}, intermediate={intermediate})"
)
} else {
format!("Preparing GPU weights: dequantizing {num_layers} layers to F32")
}
}
#[must_use]
fn dequant_done_message(weight_count: usize, total_bytes: usize) -> String {
format!(
"GPU weights ready: {weight_count} tensors, {:.1} MB F32",
total_bytes as f64 / 1e6
)
}
#[provable_contracts_macros::contract("wgpu-forward-pass-v1", equation = "dequant_correctness")]
pub fn dequant_model_weights(
model: &OwnedQuantizedModel,
) -> Result<Vec<(String, Vec<f32>, usize, usize)>> {
dequant_model_weights_except(model, &std::collections::HashSet::new())
}
pub fn dequant_model_weights_except<S: std::hash::BuildHasher>(
model: &OwnedQuantizedModel,
skip: &std::collections::HashSet<String, S>,
) -> Result<Vec<(String, Vec<f32>, usize, usize)>> {
let config = &model.config;
let hidden = config.hidden_dim;
let num_heads = config.num_heads;
let num_kv_heads = config.num_kv_heads;
let head_dim = config.head_dim();
let intermediate = config.intermediate_dim;
let num_layers = model.layers().len();
let mut weights = Vec::new();
macro_rules! push_w {
($name:expr, $data:expr, $rows:expr, $cols:expr $(,)?) => {{
let n: String = $name;
if !skip.contains(&n) {
weights.push((n, $data, $rows, $cols));
}
}};
}
eprintln!(
"{}",
dequant_start_message(
num_layers,
hidden,
num_heads,
num_kv_heads,
intermediate,
dev_trace_enabled(),
)
);
for (i, layer) in model.layers().iter().enumerate() {
let prefix = format!("layer.{i}");
push_w!(
format!("{prefix}.attn_norm"),
layer.attn_norm_weight.clone(),
1,
hidden,
);
if let Some(ref ffn_norm) = layer.ffn_norm_weight {
push_w!(format!("{prefix}.ffn_norm"), ffn_norm.clone(), 1, hidden);
}
let q_dim = num_heads * head_dim;
let kv_dim = num_kv_heads * head_dim;
match &layer.qkv_weight {
crate::gguf::OwnedQKVWeights::Fused(tensor) => {
let f32_data = dequant_tensor_public(tensor)?;
let total_out = q_dim + 2 * kv_dim;
let q_data = f32_data[..q_dim * hidden].to_vec();
let k_data = f32_data[q_dim * hidden..(q_dim + kv_dim) * hidden].to_vec();
let v_data = f32_data[(q_dim + kv_dim) * hidden..total_out * hidden].to_vec();
push_w!(format!("{prefix}.q_proj"), q_data, q_dim, hidden);
push_w!(format!("{prefix}.k_proj"), k_data, kv_dim, hidden);
push_w!(format!("{prefix}.v_proj"), v_data, kv_dim, hidden);
},
crate::gguf::OwnedQKVWeights::Separate { q, k, v } => {
push_w!(
format!("{prefix}.q_proj"),
dequant_tensor_public(q)?,
q_dim,
hidden,
);
push_w!(
format!("{prefix}.k_proj"),
dequant_tensor_public(k)?,
kv_dim,
hidden,
);
push_w!(
format!("{prefix}.v_proj"),
dequant_tensor_public(v)?,
kv_dim,
hidden,
);
},
}
if let Some(ref bias) = layer.qkv_bias {
if bias.len() >= q_dim + 2 * kv_dim {
push_w!(format!("{prefix}.q_bias"), bias[..q_dim].to_vec(), 1, q_dim);
push_w!(
format!("{prefix}.k_bias"),
bias[q_dim..q_dim + kv_dim].to_vec(),
1,
kv_dim,
);
push_w!(
format!("{prefix}.v_bias"),
bias[q_dim + kv_dim..q_dim + 2 * kv_dim].to_vec(),
1,
kv_dim,
);
}
}
push_w!(
format!("{prefix}.o_proj"),
dequant_tensor_public(&layer.attn_output_weight)?,
hidden,
q_dim,
);
if let Some(ref gate) = layer.ffn_gate_weight {
push_w!(
format!("{prefix}.gate_proj"),
dequant_tensor_public(gate)?,
intermediate,
hidden,
);
}
push_w!(
format!("{prefix}.up_proj"),
dequant_tensor_public(&layer.ffn_up_weight)?,
intermediate,
hidden,
);
push_w!(
format!("{prefix}.down_proj"),
dequant_tensor_public(&layer.ffn_down_weight)?,
hidden,
intermediate,
);
if (i + 1) % 7 == 0 || i == num_layers - 1 {
eprintln!(" Dequantized layer {}/{}", i + 1, num_layers);
}
}
push_w!(
"lm_head".to_string(),
dequant_tensor_public(model.lm_head_weight())?,
config.vocab_size,
hidden,
);
let total_bytes: usize = weights.iter().map(|(_, d, _, _)| d.len() * 4).sum();
eprintln!("{}", dequant_done_message(weights.len(), total_bytes));
Ok(weights)
}
pub fn raw_q4k_weights(model: &OwnedQuantizedModel) -> Vec<(String, Vec<u8>, usize, usize)> {
const GGUF_TYPE_Q4_K: u32 = 12;
let config = &model.config;
let hidden = config.hidden_dim;
let num_heads = config.num_heads;
let num_kv_heads = config.num_kv_heads;
let head_dim = config.head_dim();
let intermediate = config.intermediate_dim;
let q_dim = num_heads * head_dim;
let kv_dim = num_kv_heads * head_dim;
let mut raw = Vec::new();
for (i, layer) in model.layers().iter().enumerate() {
let prefix = format!("layer.{i}");
let projections: Vec<(&str, &crate::gguf::OwnedQuantizedTensor, usize, usize)> = vec![
("o_proj", &layer.attn_output_weight, hidden, q_dim),
("up_proj", &layer.ffn_up_weight, intermediate, hidden),
("down_proj", &layer.ffn_down_weight, hidden, intermediate),
];
if let Some(ref gate) = layer.ffn_gate_weight {
raw.push((
format!("{prefix}.gate_proj"),
gate.data.clone(),
intermediate,
hidden,
));
}
for (name, tensor, rows, cols) in projections {
if tensor.qtype == GGUF_TYPE_Q4_K {
raw.push((format!("{prefix}.{name}"), tensor.data.clone(), rows, cols));
}
}
if let crate::gguf::OwnedQKVWeights::Separate { q, k, v } = &layer.qkv_weight {
if q.qtype == GGUF_TYPE_Q4_K {
raw.push((format!("{prefix}.q_proj"), q.data.clone(), q_dim, hidden));
}
if k.qtype == GGUF_TYPE_Q4_K {
raw.push((format!("{prefix}.k_proj"), k.data.clone(), kv_dim, hidden));
}
if v.qtype == GGUF_TYPE_Q4_K {
raw.push((format!("{prefix}.v_proj"), v.data.clone(), kv_dim, hidden));
}
}
}
raw
}
pub fn dequant_tensor_public(tensor: &crate::gguf::OwnedQuantizedTensor) -> Result<Vec<f32>> {
const GGUF_TYPE_Q4_K: u32 = 12;
const GGUF_TYPE_Q6_K: u32 = 14;
const GGUF_TYPE_Q5_K: u32 = 13;
const GGUF_TYPE_F32: u32 = 0;
const GGUF_TYPE_F16: u32 = 1;
use crate::gguf::{APR_TYPE_Q4, APR_TYPE_Q8};
match tensor.qtype {
GGUF_TYPE_Q4_K => dequantize_q4_k(&tensor.data),
GGUF_TYPE_Q6_K => dequantize_q6_k(&tensor.data),
GGUF_TYPE_Q5_K => dequantize_q5_k(&tensor.data),
GGUF_TYPE_F32 => Ok(tensor
.data
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()),
GGUF_TYPE_F16 => Ok(tensor
.data
.chunks_exact(2)
.map(|c| {
let bits = u16::from_le_bytes([c[0], c[1]]);
half::f16::from_bits(bits).to_f32()
})
.collect()),
APR_TYPE_Q4 => Ok(crate::apr::dequant::dequantize_apr_q4(
&tensor.data,
tensor.in_dim * tensor.out_dim,
)),
APR_TYPE_Q8 => Ok(crate::apr::dequant::dequantize_apr_q8(
&tensor.data,
tensor.in_dim * tensor.out_dim,
)),
other => Err(RealizarError::FormatError {
reason: format!("Unsupported quantization type {} for WGPU dequant", other),
}),
}
}
#[cfg(test)]
mod ticket_free_output_tests {
use super::{dequant_done_message, dequant_start_message};
fn has_ticket_tag(line: &str) -> bool {
line.contains("PMAT-")
|| line.contains("GH-")
|| line.contains("PAR-")
|| line.contains('[')
}
#[test]
fn dequant_start_line_is_english_and_ticket_free() {
let line = dequant_start_message(28, 1536, 12, 2, 8960, false);
assert!(
!has_ticket_tag(&line),
"start line still addresses the user in ticket numbers: {line}"
);
assert_eq!(line, "Preparing GPU weights: dequantizing 28 layers to F32");
for geometry in ["hidden=", "heads=", "intermediate="] {
assert!(
!line.contains(geometry),
"developer geometry {geometry:?} leaked into default output: {line}"
);
}
}
#[test]
fn dequant_start_line_keeps_geometry_under_dev_trace() {
let line = dequant_start_message(28, 1536, 12, 2, 8960, true);
assert!(line.contains("hidden=1536"), "{line}");
assert!(line.contains("heads=12/2"), "{line}");
assert!(line.contains("intermediate=8960"), "{line}");
assert!(!has_ticket_tag(&line.replace(['(', ')'], "")), "{line}");
}
#[test]
fn dequant_done_line_is_english_and_ticket_free() {
let line = dequant_done_message(337, 6_174_900_000);
assert!(
!has_ticket_tag(&line),
"done line still addresses the user in ticket numbers: {line}"
);
assert_eq!(line, "GPU weights ready: 337 tensors, 6174.9 MB F32");
}
}
#[cfg(test)]
mod dequant_skip_2378 {
use super::*;
use crate::gguf::test_helpers::create_test_model_with_config;
use crate::gguf::GGUFConfig;
fn q4k_config() -> GGUFConfig {
GGUFConfig {
architecture: "test".to_string(),
constraints: crate::gguf::ArchConstraints::from_architecture("test"),
hidden_dim: 256,
intermediate_dim: 512,
num_layers: 1,
num_heads: 4,
num_kv_heads: 4,
vocab_size: 100,
context_length: 1024,
rope_theta: 10000.0,
eps: 1e-5,
rope_type: 0,
explicit_head_dim: None,
query_pre_attn_scalar: None,
bos_token_id: None,
eos_token_id: None,
}
}
fn f32_elements(w: &[(String, Vec<f32>, usize, usize)]) -> usize {
w.iter().map(|(_, d, _, _)| d.len()).sum()
}
#[test]
fn a_skipped_tensor_is_never_dequantized() {
let mut model = create_test_model_with_config(&q4k_config());
let raw: std::collections::HashSet<String> = raw_q4k_weights(&model)
.into_iter()
.map(|(n, _, _, _)| n)
.collect();
assert!(
!raw.is_empty(),
"the fixture has no Q4_K weights, so the skip is untested"
);
model.layers[0].ffn_up_weight.data.truncate(3);
let corrupted: Vec<&String> = raw.iter().filter(|n| n.contains("up_proj")).collect();
assert!(
!corrupted.is_empty(),
"the corrupted tensor is not among the raw-Q4K uploads, so this test \
is aimed at the wrong tensor: {raw:?}"
);
assert!(
dequant_model_weights(&model).is_err(),
"the corrupted tensor dequantized fine, so this test cannot \
distinguish lazy from eager"
);
assert!(
dequant_model_weights_except(&model, &raw).is_ok(),
"a tensor uploaded as raw Q4_K was dequantized anyway -- the skip \
happens after the work rather than instead of it"
);
}
#[test]
fn an_empty_skip_set_is_the_old_behaviour() {
let model = create_test_model_with_config(&q4k_config());
let a = dequant_model_weights(&model).expect("a");
let b = dequant_model_weights_except(&model, &std::collections::HashSet::new()).expect("b");
assert_eq!(a.len(), b.len());
assert_eq!(f32_elements(&a), f32_elements(&b));
}
}