use super::config::DFlashConfig;
use super::weights::{DFlashWeights, WeightsError};
use mlx_native::{DType, MlxBuffer, MlxDevice, MlxError};
use safetensors::tensor::TensorView;
#[derive(Debug, thiserror::Error)]
pub enum TensorsError {
#[error("dflash tensors mlx: {0}")]
Mlx(#[from] MlxError),
#[error("dflash tensors weights: {0}")]
Weights(#[from] WeightsError),
#[error("dflash tensors: missing manifest entry `{0}`")]
MissingEntry(String),
}
pub struct DFlashLayerTensors {
pub input_layernorm: MlxBuffer,
pub post_attention_layernorm: MlxBuffer,
pub q_proj: MlxBuffer,
pub k_proj: MlxBuffer,
pub v_proj: MlxBuffer,
pub o_proj: MlxBuffer,
pub q_norm: MlxBuffer,
pub k_norm: MlxBuffer,
pub mlp_gate: MlxBuffer,
pub mlp_up: MlxBuffer,
pub mlp_down: MlxBuffer,
}
pub struct DFlashModelTensors {
pub fc: MlxBuffer,
pub hidden_norm: MlxBuffer,
pub final_norm: MlxBuffer,
pub layers: Vec<DFlashLayerTensors>,
}
fn upload_bf16(device: &MlxDevice, view: &TensorView<'_>) -> Result<MlxBuffer, TensorsError> {
let shape: Vec<usize> = view.shape().to_vec();
let byte_len = view.data().len();
let mut buf = device.alloc_buffer(byte_len, DType::BF16, shape)?;
let dst: &mut [u8] = buf
.as_mut_slice::<u8>()
.map_err(|e| TensorsError::Mlx(MlxError::InvalidArgument(format!("buffer slice: {e}"))))?;
debug_assert_eq!(dst.len(), byte_len);
dst.copy_from_slice(view.data());
Ok(buf)
}
pub(super) fn decode_bf16_bytes_to_f32(bytes: &[u8]) -> Result<Vec<f32>, TensorsError> {
let n_elem = bytes.len() / 2;
if bytes.len() != n_elem * 2 {
return Err(TensorsError::Mlx(MlxError::InvalidArgument(format!(
"decode_bf16_bytes_to_f32: data len {} not even (not BF16-aligned)",
bytes.len()
))));
}
let mut out = Vec::with_capacity(n_elem);
for i in 0..n_elem {
let lo = bytes[i * 2] as u32;
let hi = bytes[i * 2 + 1] as u32;
let bf16_bits = lo | (hi << 8);
let f32_bits = bf16_bits << 16;
out.push(f32::from_bits(f32_bits));
}
Ok(out)
}
fn upload_bf16_as_f32(
device: &MlxDevice,
view: &TensorView<'_>,
) -> Result<MlxBuffer, TensorsError> {
let f32_values = decode_bf16_bytes_to_f32(view.data())?;
let n_elem = f32_values.len();
let shape: Vec<usize> = view.shape().to_vec();
let mut buf = device.alloc_buffer(n_elem * 4, DType::F32, shape)?;
let dst: &mut [f32] = buf
.as_mut_slice::<f32>()
.map_err(|e| TensorsError::Mlx(MlxError::InvalidArgument(format!("f32 slice: {e}"))))?;
debug_assert_eq!(dst.len(), n_elem);
dst.copy_from_slice(&f32_values);
Ok(buf)
}
impl DFlashModelTensors {
pub fn upload(
device: &MlxDevice,
cfg: &DFlashConfig,
weights: &DFlashWeights<'_>,
) -> Result<Self, TensorsError> {
let fc = upload_bf16(device, fetch(weights, "fc.weight")?)?;
let hidden_norm = upload_bf16_as_f32(device, fetch(weights, "hidden_norm.weight")?)?;
let final_norm = upload_bf16_as_f32(device, fetch(weights, "norm.weight")?)?;
let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
for i in 0..cfg.num_hidden_layers {
let p = format!("layers.{i}");
let lw = DFlashLayerTensors {
input_layernorm: upload_bf16_as_f32(
device,
fetch(weights, &format!("{p}.input_layernorm.weight"))?,
)?,
post_attention_layernorm: upload_bf16_as_f32(
device,
fetch(weights, &format!("{p}.post_attention_layernorm.weight"))?,
)?,
q_proj: upload_bf16(
device,
fetch(weights, &format!("{p}.self_attn.q_proj.weight"))?,
)?,
k_proj: upload_bf16(
device,
fetch(weights, &format!("{p}.self_attn.k_proj.weight"))?,
)?,
v_proj: upload_bf16(
device,
fetch(weights, &format!("{p}.self_attn.v_proj.weight"))?,
)?,
o_proj: upload_bf16(
device,
fetch(weights, &format!("{p}.self_attn.o_proj.weight"))?,
)?,
q_norm: upload_bf16_as_f32(
device,
fetch(weights, &format!("{p}.self_attn.q_norm.weight"))?,
)?,
k_norm: upload_bf16_as_f32(
device,
fetch(weights, &format!("{p}.self_attn.k_norm.weight"))?,
)?,
mlp_gate: upload_bf16(
device,
fetch(weights, &format!("{p}.mlp.gate_proj.weight"))?,
)?,
mlp_up: upload_bf16(device, fetch(weights, &format!("{p}.mlp.up_proj.weight"))?)?,
mlp_down: upload_bf16(
device,
fetch(weights, &format!("{p}.mlp.down_proj.weight"))?,
)?,
};
layers.push(lw);
}
debug_assert_eq!(
fc.dtype(),
DType::BF16,
"fc weight must stay BF16 for dense_matmul"
);
debug_assert_eq!(
hidden_norm.dtype(),
DType::F32,
"hidden_norm weight must be F32 for rms_norm_f32 kernel"
);
debug_assert_eq!(
final_norm.dtype(),
DType::F32,
"final_norm weight must be F32 for rms_norm_f32 kernel"
);
for (idx, l) in layers.iter().enumerate() {
debug_assert_eq!(
l.input_layernorm.dtype(),
DType::F32,
"layer {idx}: input_layernorm weight must be F32 for rms_norm_f32 kernel"
);
debug_assert_eq!(
l.post_attention_layernorm.dtype(),
DType::F32,
"layer {idx}: post_attention_layernorm weight must be F32 for rms_norm_f32 kernel"
);
debug_assert_eq!(
l.q_norm.dtype(),
DType::F32,
"layer {idx}: q_norm weight must be F32 for rms_norm_f32 head_norm kernel"
);
debug_assert_eq!(
l.k_norm.dtype(),
DType::F32,
"layer {idx}: k_norm weight must be F32 for rms_norm_f32 head_norm kernel"
);
debug_assert_eq!(
l.q_proj.dtype(),
DType::BF16,
"layer {idx}: q_proj must stay BF16 for dense_matmul"
);
debug_assert_eq!(
l.k_proj.dtype(),
DType::BF16,
"layer {idx}: k_proj must stay BF16 for dense_matmul"
);
debug_assert_eq!(
l.v_proj.dtype(),
DType::BF16,
"layer {idx}: v_proj must stay BF16 for dense_matmul"
);
debug_assert_eq!(
l.o_proj.dtype(),
DType::BF16,
"layer {idx}: o_proj must stay BF16 for dense_matmul"
);
debug_assert_eq!(
l.mlp_gate.dtype(),
DType::BF16,
"layer {idx}: mlp_gate must stay BF16 for dense_matmul"
);
debug_assert_eq!(
l.mlp_up.dtype(),
DType::BF16,
"layer {idx}: mlp_up must stay BF16 for dense_matmul"
);
debug_assert_eq!(
l.mlp_down.dtype(),
DType::BF16,
"layer {idx}: mlp_down must stay BF16 for dense_matmul"
);
}
Ok(DFlashModelTensors {
fc,
hidden_norm,
final_norm,
layers,
})
}
pub fn gpu_resident_bytes(&self) -> usize {
let layer_bytes: usize = self
.layers
.iter()
.map(|l| {
l.input_layernorm.byte_len()
+ l.post_attention_layernorm.byte_len()
+ l.q_proj.byte_len()
+ l.k_proj.byte_len()
+ l.v_proj.byte_len()
+ l.o_proj.byte_len()
+ l.q_norm.byte_len()
+ l.k_norm.byte_len()
+ l.mlp_gate.byte_len()
+ l.mlp_up.byte_len()
+ l.mlp_down.byte_len()
})
.sum();
self.fc.byte_len() + self.hidden_norm.byte_len() + self.final_norm.byte_len() + layer_bytes
}
}
fn fetch<'a, 'b>(
weights: &'a DFlashWeights<'b>,
name: &str,
) -> Result<&'a TensorView<'b>, TensorsError> {
weights
.tensor(name)
.ok_or_else(|| TensorsError::MissingEntry(name.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::inference::spec_decode::dflash::config::DFlashConfig;
use crate::inference::spec_decode::dflash::weights::{DFlashWeights, DFlashWeightsFile};
fn gemma4_26b_a4b_dflash_config() -> DFlashConfig {
DFlashConfig::from_json_str(super::super::config::tests::GEMMA4_26B_A4B_DFLASH_CONFIG)
.expect("test fixture must parse")
}
fn bf16_bytes_from_f32(values: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(values.len() * 2);
for v in values {
let bits = v.to_bits();
let bf16_bits = (bits >> 16) as u16;
out.push((bf16_bits & 0xff) as u8);
out.push(((bf16_bits >> 8) & 0xff) as u8);
}
out
}
#[test]
fn decode_bf16_bytes_round_trips_canonical_values() {
let canonical = [0.0f32, 1.0, -1.0, 2.0, 0.5, -0.5, 256.0, -256.0];
let bytes = bf16_bytes_from_f32(&canonical);
let decoded = decode_bf16_bytes_to_f32(&bytes).expect("decode");
assert_eq!(decoded.len(), canonical.len());
for (i, (got, want)) in decoded.iter().zip(canonical.iter()).enumerate() {
assert_eq!(
got.to_bits(),
want.to_bits(),
"canonical[{i}] = {want} round-trip got {got}"
);
}
}
#[test]
fn bf16_over_f32_misinterpret_signature() {
let canonical = [1.0f32, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0];
let bf16_bytes = bf16_bytes_from_f32(&canonical);
assert_eq!(bf16_bytes.len(), canonical.len() * 2);
assert_eq!(&bf16_bytes[0..2], &[0x80, 0x3f]);
let mut misread = Vec::with_capacity(canonical.len() / 2);
for chunk in bf16_bytes.chunks_exact(4) {
let bits = u32::from_le_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]);
misread.push(f32::from_bits(bits));
}
assert_eq!(misread.len(), 4);
let bug_value = f32::from_bits(0x3f803f80);
assert_eq!(
misread[0].to_bits(),
bug_value.to_bits(),
"pre-iter-106 bug: reads BF16 1.0 + 1.0 as F32 ≈ {} (NOT 1.0)",
bug_value
);
assert_ne!(
misread[0], 1.0,
"this is the silent-corruption signature iter-106 fixed"
);
}
#[test]
fn decode_bf16_bytes_rejects_odd_length() {
let bytes = vec![0x80, 0x3f, 0x00]; let err = decode_bf16_bytes_to_f32(&bytes).unwrap_err();
assert!(format!("{err}").contains("not even"));
}
#[test]
#[ignore = "requires Metal device + drafter HF cache"]
fn uploads_real_drafter_to_gpu() {
let cfg = gemma4_26b_a4b_dflash_config();
let device = MlxDevice::new().expect("Metal device available on M5 Max");
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 safetensors_data_bytes = weights.total_data_bytes();
let tensors = DFlashModelTensors::upload(&device, &cfg, &weights).expect("GPU upload");
let gpu_bytes = tensors.gpu_resident_bytes();
let per_layer_cast_elems = 2 * cfg.head_dim + 2 * cfg.hidden_size;
let model_cast_elems = 2 * cfg.hidden_size;
let cast_expansion_bytes =
(cfg.num_hidden_layers * per_layer_cast_elems + model_cast_elems) * 2;
assert_eq!(
gpu_bytes,
safetensors_data_bytes + cast_expansion_bytes,
"GPU resident bytes must equal safetensors data + F32 cast expansion \
(cfg.num_hidden_layers={}, head_dim={}, hidden_size={}, expansion={})",
cfg.num_hidden_layers,
cfg.head_dim,
cfg.hidden_size,
cast_expansion_bytes,
);
assert_eq!(tensors.layers.len(), cfg.num_hidden_layers);
}
}