#![allow(dead_code)]
use anyhow::{anyhow, Context, Result};
use mlx_native::ops::elementwise::{cast, CastDirection};
use mlx_native::ops::flash_attn_prefill::{
self as flash_attn_prefill, dispatch_flash_attn_prefill_bf16_d64, FlashAttnPrefillLayout,
FlashAttnPrefillParams,
};
use mlx_native::ops::rope::dispatch_rope_neox_f32;
use mlx_native::ops::silu_mul::dispatch_silu_mul;
use mlx_native::{CommandEncoder, DType, KernelRegistry, MlxBuffer, MlxDevice};
use super::super::bert::bert_gpu::{
bert_embed_gather_gpu, bert_l2_normalize_gpu, bert_layer_norm_gpu, bert_linear_bf16_gpu,
bert_linear_gpu, bert_pool_gpu, bert_residual_add_gpu, bert_residual_layer_norm_gpu,
register_bert_custom_shaders, BertPoolKind,
};
use super::super::bert::config::PoolingType;
use super::config::NomicBertConfig;
use super::weights::LoadedNomicBertWeights;
pub fn register_nomic_bert_kernels(registry: &mut KernelRegistry) {
register_bert_custom_shaders(registry);
mlx_native::ops::rope::register(registry);
mlx_native::ops::silu_mul::register(registry);
flash_attn_prefill::register(registry);
}
#[allow(clippy::too_many_arguments)]
pub fn nomic_bert_embeddings_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input_ids: &MlxBuffer,
type_ids_opt: Option<&MlxBuffer>,
token_embd: &MlxBuffer,
token_types_opt: Option<&MlxBuffer>,
embed_norm_gamma: &MlxBuffer,
embed_norm_beta: &MlxBuffer,
eps: f32,
seq_len: u32,
hidden: u32,
vocab: u32,
type_vocab: u32,
) -> Result<MlxBuffer> {
if seq_len == 0 || hidden == 0 {
return Err(anyhow!(
"nomic_bert_embeddings_gpu: seq_len ({}) and hidden ({}) must be > 0",
seq_len,
hidden
));
}
match (type_ids_opt.is_some(), token_types_opt.is_some()) {
(true, true) | (false, false) => {}
(a, b) => {
return Err(anyhow!(
"nomic_bert_embeddings_gpu: type_ids and token_types must both be Some or both None (got {} / {})",
a, b
));
}
}
let n_hidden = (seq_len as usize) * (hidden as usize);
let tok = bert_embed_gather_gpu(
encoder, registry, device, token_embd, input_ids, vocab, hidden, seq_len,
)
.context("nomic embeddings: token gather")?;
encoder.memory_barrier();
let summed = if let (Some(type_ids), Some(token_types)) = (type_ids_opt, token_types_opt) {
let typ = bert_embed_gather_gpu(
encoder,
registry,
device,
token_types,
type_ids,
type_vocab,
hidden,
seq_len,
)
.context("nomic embeddings: type gather")?;
encoder.memory_barrier();
let s = bert_residual_add_gpu(encoder, registry, device, &tok, &typ, n_hidden as u32)
.context("nomic embeddings: token + type add")?;
encoder.memory_barrier();
s
} else {
tok
};
bert_layer_norm_gpu(
encoder,
registry,
device,
&summed,
embed_norm_gamma,
embed_norm_beta,
eps,
seq_len,
hidden,
)
.context("nomic embeddings: post-sum LayerNorm")
}
fn alloc_rope_positions(device: &MlxDevice, seq_len: u32) -> Result<MlxBuffer> {
let n = seq_len as usize;
let buf = device
.alloc_buffer(n * 4, DType::U32, vec![n])
.map_err(|e| anyhow!("alloc rope positions: {e}"))?;
let slice: &mut [u32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut u32, n) };
for (i, slot) in slice.iter_mut().enumerate() {
*slot = i as u32;
}
Ok(buf)
}
fn alloc_zero_type_ids(device: &MlxDevice, seq_len: u32) -> Result<MlxBuffer> {
let n = seq_len as usize;
let buf = device
.alloc_buffer(n * 4, DType::U32, vec![n])
.map_err(|e| anyhow!("alloc zero type_ids: {e}"))?;
let slice: &mut [u32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut u32, n) };
for slot in slice.iter_mut() {
*slot = 0;
}
Ok(buf)
}
fn alloc_rope_params(
device: &MlxDevice,
theta: f32,
head_dim: u32,
rope_dim: u32,
) -> Result<MlxBuffer> {
let buf = device
.alloc_buffer(16, DType::F32, vec![4])
.map_err(|e| anyhow!("alloc rope params: {e}"))?;
let slice: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, 4) };
slice[0] = theta;
slice[1] = head_dim as f32;
slice[2] = rope_dim as f32;
slice[3] = 0.0;
Ok(buf)
}
fn alloc_nomic_attn_mask_bf16(
device: &MlxDevice,
seq_len: u32,
valid_len: u32,
) -> Result<MlxBuffer> {
if seq_len == 0 {
return Err(anyhow!("alloc_nomic_attn_mask_bf16: seq_len must be > 0"));
}
let n = (seq_len as usize) * (seq_len as usize);
let buf = device
.alloc_buffer(n * 2, DType::BF16, vec![seq_len as usize, seq_len as usize])
.map_err(|e| anyhow!("alloc bf16 attention mask: {e}"))?;
let s: &mut [half::bf16] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut half::bf16, n) };
let valid = valid_len.min(seq_len) as usize;
let seq = seq_len as usize;
let zero = half::bf16::from_f32(0.0);
let neg_inf = half::bf16::from_f32(f32::NEG_INFINITY);
for r in 0..seq {
for c in 0..seq {
s[r * seq + c] = if c < valid { zero } else { neg_inf };
}
}
Ok(buf)
}
fn alloc_silu_mul_params(device: &MlxDevice, n: u32) -> Result<MlxBuffer> {
let buf = device
.alloc_buffer(4, DType::U32, vec![1])
.map_err(|e| anyhow!("alloc silu_mul params: {e}"))?;
let slice: &mut [u32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut u32, 1) };
slice[0] = n;
Ok(buf)
}
pub struct NomicBertEncoderBlockTensors<'a> {
pub qkv_w: &'a MlxBuffer,
pub qkv_b: Option<&'a MlxBuffer>,
pub qkv_w_bf16: Option<&'a MlxBuffer>,
pub o_w: &'a MlxBuffer,
pub o_b: Option<&'a MlxBuffer>,
pub o_w_bf16: Option<&'a MlxBuffer>,
pub attn_norm_gamma: &'a MlxBuffer,
pub attn_norm_beta: &'a MlxBuffer,
pub up_w: &'a MlxBuffer,
pub up_b: Option<&'a MlxBuffer>,
pub up_w_bf16: Option<&'a MlxBuffer>,
pub gate_w: &'a MlxBuffer,
pub gate_b: Option<&'a MlxBuffer>,
pub gate_w_bf16: Option<&'a MlxBuffer>,
pub down_w: &'a MlxBuffer,
pub down_b: Option<&'a MlxBuffer>,
pub down_w_bf16: Option<&'a MlxBuffer>,
pub ffn_norm_gamma: &'a MlxBuffer,
pub ffn_norm_beta: &'a MlxBuffer,
}
#[allow(clippy::too_many_arguments)]
pub fn apply_nomic_bert_encoder_block_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
tensors: &NomicBertEncoderBlockTensors<'_>,
mask_bf16: &MlxBuffer,
rope_positions: &MlxBuffer,
rope_params: &MlxBuffer,
seq_len: u32,
hidden: u32,
num_heads: u32,
intermediate: u32,
eps: f32,
) -> Result<MlxBuffer> {
if hidden % num_heads != 0 {
return Err(anyhow!(
"apply_nomic_bert_encoder_block_gpu: hidden ({}) not divisible by num_heads ({})",
hidden,
num_heads
));
}
let head_dim = hidden / num_heads;
let n_hidden_elems = (seq_len as usize) * (hidden as usize);
let weight_elems_per_block = (hidden as usize) * (hidden as usize);
let weight_bytes_per_block_f32 = weight_elems_per_block * 4;
let weight_bytes_per_block_bf16 = weight_elems_per_block * 2;
let q_w = tensors.qkv_w.slice_view(0, weight_elems_per_block);
let k_w = tensors
.qkv_w
.slice_view(weight_bytes_per_block_f32 as u64, weight_elems_per_block);
let v_w = tensors.qkv_w.slice_view(
(2 * weight_bytes_per_block_f32) as u64,
weight_elems_per_block,
);
let (q_b, k_b, v_b) = match tensors.qkv_b {
None => (None, None, None),
Some(qkvb) => {
let q = qkvb.slice_view(0, hidden as usize);
let k = qkvb.slice_view((hidden as usize * 4) as u64, hidden as usize);
let v = qkvb.slice_view((2 * hidden as usize * 4) as u64, hidden as usize);
(Some(q), Some(k), Some(v))
}
};
let q_b_ref = q_b.as_ref();
let k_b_ref = k_b.as_ref();
let v_b_ref = v_b.as_ref();
let (q_proj, k_proj, v_proj) = if let Some(qkv_bf16) = tensors.qkv_w_bf16 {
let q_w_bf16 = qkv_bf16.slice_view(0, weight_elems_per_block);
let k_w_bf16 =
qkv_bf16.slice_view(weight_bytes_per_block_bf16 as u64, weight_elems_per_block);
let v_w_bf16 = qkv_bf16.slice_view(
(2 * weight_bytes_per_block_bf16) as u64,
weight_elems_per_block,
);
let q = bert_linear_bf16_gpu(
encoder, registry, device, input, &q_w_bf16, q_b_ref, seq_len, hidden, hidden,
)
.context("nomic block: q_proj bf16 linear")?;
let k = bert_linear_bf16_gpu(
encoder, registry, device, input, &k_w_bf16, k_b_ref, seq_len, hidden, hidden,
)
.context("nomic block: k_proj bf16 linear")?;
let v = bert_linear_bf16_gpu(
encoder, registry, device, input, &v_w_bf16, v_b_ref, seq_len, hidden, hidden,
)
.context("nomic block: v_proj bf16 linear")?;
encoder.memory_barrier();
(q, k, v)
} else {
let q = bert_linear_gpu(
encoder, registry, device, input, &q_w, q_b_ref, seq_len, hidden, hidden,
)
.context("nomic block: q_proj linear (f32 fallback)")?;
let k = bert_linear_gpu(
encoder, registry, device, input, &k_w, k_b_ref, seq_len, hidden, hidden,
)
.context("nomic block: k_proj linear (f32 fallback)")?;
let v = bert_linear_gpu(
encoder, registry, device, input, &v_w, v_b_ref, seq_len, hidden, hidden,
)
.context("nomic block: v_proj linear (f32 fallback)")?;
encoder.memory_barrier();
(q, k, v)
};
let q_rotated = device
.alloc_buffer(
n_hidden_elems * 4,
DType::F32,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc q_rotated: {e}"))?;
let k_rotated = device
.alloc_buffer(
n_hidden_elems * 4,
DType::F32,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc k_rotated: {e}"))?;
dispatch_rope_neox_f32(
encoder,
registry,
device.metal_device(),
&q_proj,
&q_rotated,
rope_params,
rope_positions,
None,
seq_len,
num_heads,
head_dim,
head_dim,
)
.map_err(|e| anyhow!("nomic block: rope on Q: {e}"))?;
dispatch_rope_neox_f32(
encoder,
registry,
device.metal_device(),
&k_proj,
&k_rotated,
rope_params,
rope_positions,
None,
seq_len,
num_heads,
head_dim,
head_dim,
)
.map_err(|e| anyhow!("nomic block: rope on K: {e}"))?;
encoder.memory_barrier();
let n_attn_elems = (seq_len as usize) * (hidden as usize);
let q_bf16 = device
.alloc_buffer(
n_attn_elems * 2,
DType::BF16,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc q_bf16: {e}"))?;
let k_bf16 = device
.alloc_buffer(
n_attn_elems * 2,
DType::BF16,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc k_bf16: {e}"))?;
let v_bf16 = device
.alloc_buffer(
n_attn_elems * 2,
DType::BF16,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc v_bf16: {e}"))?;
cast(
encoder,
registry,
device.metal_device(),
&q_rotated,
&q_bf16,
n_attn_elems,
CastDirection::F32ToBF16,
)
.context("nomic block: cast Q F32→BF16")?;
cast(
encoder,
registry,
device.metal_device(),
&k_rotated,
&k_bf16,
n_attn_elems,
CastDirection::F32ToBF16,
)
.context("nomic block: cast K F32→BF16")?;
cast(
encoder,
registry,
device.metal_device(),
&v_proj,
&v_bf16,
n_attn_elems,
CastDirection::F32ToBF16,
)
.context("nomic block: cast V F32→BF16")?;
encoder.memory_barrier();
let mut attn_out_bf16 = device
.alloc_buffer(
n_attn_elems * 2,
DType::BF16,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc attn_out_bf16: {e}"))?;
let scale = 1.0_f32 / (head_dim as f32).sqrt();
let fa_params = FlashAttnPrefillParams {
n_heads: num_heads,
n_kv_heads: num_heads, head_dim,
seq_len_q: seq_len,
seq_len_k: seq_len,
batch: 1,
scale,
do_causal: false, };
dispatch_flash_attn_prefill_bf16_d64(
encoder,
device,
registry,
&q_bf16,
&k_bf16,
&v_bf16,
Some(mask_bf16),
&mut attn_out_bf16,
&fa_params,
FlashAttnPrefillLayout::SeqMajor,
)
.map_err(|e| anyhow!("nomic block: flash_attn_prefill_bf16_d64: {e}"))?;
encoder.memory_barrier();
let attn_out = device
.alloc_buffer(
n_attn_elems * 4,
DType::F32,
vec![seq_len as usize, hidden as usize],
)
.map_err(|e| anyhow!("alloc attn_out F32: {e}"))?;
cast(
encoder,
registry,
device.metal_device(),
&attn_out_bf16,
&attn_out,
n_attn_elems,
CastDirection::BF16ToF32,
)
.context("nomic block: cast attn_out BF16→F32")?;
encoder.memory_barrier();
let attn_proj = if let Some(o_w_bf16) = tensors.o_w_bf16 {
bert_linear_bf16_gpu(
encoder,
registry,
device,
&attn_out,
o_w_bf16,
tensors.o_b,
seq_len,
hidden,
hidden,
)
.context("nomic block: attention output bf16 projection")?
} else {
bert_linear_gpu(
encoder,
registry,
device,
&attn_out,
tensors.o_w,
tensors.o_b,
seq_len,
hidden,
hidden,
)
.context("nomic block: attention output projection (f32 fallback)")?
};
encoder.memory_barrier();
let _ = n_hidden_elems; let after_attn_norm = bert_residual_layer_norm_gpu(
encoder,
registry,
device,
&attn_proj,
input,
tensors.attn_norm_gamma,
tensors.attn_norm_beta,
eps,
seq_len,
hidden,
)
.context("nomic block: post-attention residual+LayerNorm (fused)")?;
encoder.memory_barrier();
let n_ffn = (seq_len as usize) * (intermediate as usize);
let up_proj = if let Some(up_w_bf16) = tensors.up_w_bf16 {
bert_linear_bf16_gpu(
encoder,
registry,
device,
&after_attn_norm,
up_w_bf16,
tensors.up_b,
seq_len,
hidden,
intermediate,
)
.context("nomic block: ffn_up bf16 linear")?
} else {
bert_linear_gpu(
encoder,
registry,
device,
&after_attn_norm,
tensors.up_w,
tensors.up_b,
seq_len,
hidden,
intermediate,
)
.context("nomic block: ffn_up linear (f32 fallback)")?
};
let gate_proj = if let Some(gate_w_bf16) = tensors.gate_w_bf16 {
bert_linear_bf16_gpu(
encoder,
registry,
device,
&after_attn_norm,
gate_w_bf16,
tensors.gate_b,
seq_len,
hidden,
intermediate,
)
.context("nomic block: ffn_gate bf16 linear")?
} else {
bert_linear_gpu(
encoder,
registry,
device,
&after_attn_norm,
tensors.gate_w,
tensors.gate_b,
seq_len,
hidden,
intermediate,
)
.context("nomic block: ffn_gate linear (f32 fallback)")?
};
encoder.memory_barrier();
let silu_gated = device
.alloc_buffer(
n_ffn * 4,
DType::F32,
vec![seq_len as usize, intermediate as usize],
)
.map_err(|e| anyhow!("alloc silu_gated: {e}"))?;
let _silu_params = alloc_silu_mul_params(device, n_ffn as u32)?;
dispatch_silu_mul(
encoder,
registry,
device.metal_device(),
&gate_proj,
&up_proj,
&silu_gated,
&_silu_params,
n_ffn as u32,
)
.map_err(|e| anyhow!("nomic block: silu_mul: {e}"))?;
encoder.memory_barrier();
let down_proj = if let Some(down_w_bf16) = tensors.down_w_bf16 {
bert_linear_bf16_gpu(
encoder,
registry,
device,
&silu_gated,
down_w_bf16,
tensors.down_b,
seq_len,
intermediate,
hidden,
)
.context("nomic block: ffn_down bf16 linear")?
} else {
bert_linear_gpu(
encoder,
registry,
device,
&silu_gated,
tensors.down_w,
tensors.down_b,
seq_len,
intermediate,
hidden,
)
.context("nomic block: ffn_down linear (f32 fallback)")?
};
encoder.memory_barrier();
let block_out = bert_residual_layer_norm_gpu(
encoder,
registry,
device,
&down_proj,
&after_attn_norm,
tensors.ffn_norm_gamma,
tensors.ffn_norm_beta,
eps,
seq_len,
hidden,
)
.context("nomic block: post-FFN residual+LayerNorm (fused)")?;
encoder.memory_barrier();
drop(_silu_params);
Ok(block_out)
}
pub fn apply_nomic_bert_full_forward_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input_ids: &MlxBuffer,
type_ids_opt: Option<&MlxBuffer>,
weights: &LoadedNomicBertWeights,
cfg: &NomicBertConfig,
seq_len: u32,
valid_token_count: u32,
) -> Result<MlxBuffer> {
let pool_kind = match cfg.pooling_type {
PoolingType::Mean => BertPoolKind::Mean,
PoolingType::Cls => BertPoolKind::Cls,
PoolingType::Last => BertPoolKind::Last,
PoolingType::None => {
return Err(anyhow!(
"apply_nomic_bert_full_forward_gpu: pooling_type=None is not a single-vector embedding"
));
}
PoolingType::Rank => {
return Err(anyhow!(
"apply_nomic_bert_full_forward_gpu: pooling_type=Rank is reranker-only (out of scope for /v1/embeddings)"
));
}
};
let hidden = cfg.hidden_size as u32;
let num_heads = cfg.num_attention_heads as u32;
let intermediate = cfg.intermediate_size as u32;
let vocab = cfg.vocab_size as u32;
let max_pos = cfg.max_position_embeddings as u32;
let type_vocab = cfg.type_vocab_size as u32;
let head_dim = (hidden / num_heads) as u32;
let eps = cfg.layer_norm_eps;
if seq_len == 0 {
return Err(anyhow!(
"apply_nomic_bert_full_forward_gpu: seq_len must be > 0"
));
}
if head_dim != 64 {
return Err(anyhow!(
"apply_nomic_bert_full_forward_gpu: head_dim must be 64 \
(only `flash_attn_prefill_bf16_d64` instantiation registered); got {}",
head_dim
));
}
if seq_len > max_pos {
return Err(anyhow!(
"apply_nomic_bert_full_forward_gpu: seq_len ({}) > max_position_embeddings ({})",
seq_len,
max_pos
));
}
let synthesized_type_ids: Option<MlxBuffer> = match (type_ids_opt, weights.token_types_weight())
{
(None, Some(_)) => Some(alloc_zero_type_ids(device, seq_len)?),
_ => None,
};
let effective_type_ids: Option<&MlxBuffer> = match (type_ids_opt, synthesized_type_ids.as_ref())
{
(Some(b), _) => Some(b),
(None, Some(b)) => Some(b),
(None, None) => None,
};
let token_types_for_call = if effective_type_ids.is_some() {
weights.token_types_weight()
} else {
None
};
let mut hidden_states = nomic_bert_embeddings_gpu(
encoder,
registry,
device,
input_ids,
effective_type_ids,
weights.token_embd_weight()?,
token_types_for_call,
weights.embed_norm_weight()?,
weights.embed_norm_bias()?,
eps,
seq_len,
hidden,
vocab,
type_vocab,
)
.context("nomic full-forward: embeddings")?;
encoder.memory_barrier();
let mask_bf16 = alloc_nomic_attn_mask_bf16(device, seq_len, valid_token_count)?;
let rope_positions = alloc_rope_positions(device, seq_len)?;
let rope_params = alloc_rope_params(device, cfg.rope_freq_base, head_dim, head_dim)?;
for layer_idx in 0..cfg.num_hidden_layers {
let tensors = NomicBertEncoderBlockTensors {
qkv_w: weights.block_required(layer_idx, "attn_qkv.weight")?,
qkv_b: weights.block_optional(layer_idx, "attn_qkv.bias"),
qkv_w_bf16: weights.block_weight_bf16(layer_idx, "attn_qkv.weight"),
o_w: weights.block_required(layer_idx, "attn_output.weight")?,
o_b: weights.block_optional(layer_idx, "attn_output.bias"),
o_w_bf16: weights.block_weight_bf16(layer_idx, "attn_output.weight"),
attn_norm_gamma: weights.block_required(layer_idx, "attn_output_norm.weight")?,
attn_norm_beta: weights.block_required(layer_idx, "attn_output_norm.bias")?,
up_w: weights.block_required(layer_idx, "ffn_up.weight")?,
up_b: weights.block_optional(layer_idx, "ffn_up.bias"),
up_w_bf16: weights.block_weight_bf16(layer_idx, "ffn_up.weight"),
gate_w: weights.block_required(layer_idx, "ffn_gate.weight")?,
gate_b: weights.block_optional(layer_idx, "ffn_gate.bias"),
gate_w_bf16: weights.block_weight_bf16(layer_idx, "ffn_gate.weight"),
down_w: weights.block_required(layer_idx, "ffn_down.weight")?,
down_b: weights.block_optional(layer_idx, "ffn_down.bias"),
down_w_bf16: weights.block_weight_bf16(layer_idx, "ffn_down.weight"),
ffn_norm_gamma: weights.block_required(layer_idx, "layer_output_norm.weight")?,
ffn_norm_beta: weights.block_required(layer_idx, "layer_output_norm.bias")?,
};
hidden_states = apply_nomic_bert_encoder_block_gpu(
encoder,
registry,
device,
&hidden_states,
&tensors,
&mask_bf16,
&rope_positions,
&rope_params,
seq_len,
hidden,
num_heads,
intermediate,
eps,
)
.with_context(|| format!("nomic full-forward: block {}", layer_idx))?;
encoder.memory_barrier();
}
let pooled = bert_pool_gpu(
encoder,
registry,
device,
&hidden_states,
pool_kind,
valid_token_count,
hidden,
)
.context("nomic full-forward: pool")?;
encoder.memory_barrier();
bert_l2_normalize_gpu(encoder, registry, device, &pooled, 1e-12, 1, hidden)
.context("nomic full-forward: l2 normalize")
}
#[cfg(test)]
pub(crate) mod tests {
use super::*;
use std::collections::HashMap;
fn synthetic_min_cfg(num_layers: usize) -> NomicBertConfig {
NomicBertConfig {
hidden_size: 128, num_attention_heads: 2,
num_hidden_layers: num_layers,
intermediate_size: 256, max_position_embeddings: 128,
vocab_size: 100,
type_vocab_size: 2,
layer_norm_eps: 1e-12,
pooling_type: PoolingType::Mean,
rope_freq_base: 1000.0,
causal_attention: false,
}
}
fn synthetic_weights(
device: &MlxDevice,
cfg: &NomicBertConfig,
) -> Result<HashMap<String, MlxBuffer>> {
let make_buf = |device: &MlxDevice, n: usize, key: &str| -> Result<MlxBuffer> {
let buf = device
.alloc_buffer(n * 4, DType::F32, vec![n])
.map_err(|e| anyhow!("alloc {key}: {e}"))?;
let slice: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
let key_hash: u32 = key.bytes().fold(2166136261u32, |acc, b| {
acc.wrapping_mul(16777619).wrapping_add(b as u32)
});
for (i, slot) in slice.iter_mut().enumerate() {
let h = key_hash
.wrapping_mul(2654435761)
.wrapping_add((i as u32).wrapping_mul(2246822519));
*slot = ((h as i32) as f32 / i32::MAX as f32) * 0.05;
}
Ok(buf)
};
let make_ones = |device: &MlxDevice, n: usize, key: &str| -> Result<MlxBuffer> {
let buf = device
.alloc_buffer(n * 4, DType::F32, vec![n])
.map_err(|e| anyhow!("alloc {key}: {e}"))?;
let slice: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
slice.fill(1.0);
Ok(buf)
};
let make_zeros = |device: &MlxDevice, n: usize, key: &str| -> Result<MlxBuffer> {
let buf = device
.alloc_buffer(n * 4, DType::F32, vec![n])
.map_err(|e| anyhow!("alloc {key}: {e}"))?;
let slice: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
slice.fill(0.0);
Ok(buf)
};
let hidden = cfg.hidden_size;
let intermediate = cfg.intermediate_size;
let vocab = cfg.vocab_size;
let type_vocab = cfg.type_vocab_size;
let mut tensors: HashMap<String, MlxBuffer> = HashMap::new();
tensors.insert(
"token_embd.weight".into(),
make_buf(device, vocab * hidden, "token_embd")?,
);
tensors.insert(
"token_types.weight".into(),
make_buf(device, type_vocab * hidden, "token_types")?,
);
tensors.insert(
"token_embd_norm.weight".into(),
make_ones(device, hidden, "embd_norm_w")?,
);
tensors.insert(
"token_embd_norm.bias".into(),
make_zeros(device, hidden, "embd_norm_b")?,
);
for il in 0..cfg.num_hidden_layers {
tensors.insert(
format!("blk.{}.attn_qkv.weight", il),
make_buf(device, 3 * hidden * hidden, &format!("qkv_w_{il}"))?,
);
tensors.insert(
format!("blk.{}.attn_output.weight", il),
make_buf(device, hidden * hidden, &format!("o_w_{il}"))?,
);
tensors.insert(
format!("blk.{}.attn_output_norm.weight", il),
make_ones(device, hidden, &format!("attn_norm_w_{il}"))?,
);
tensors.insert(
format!("blk.{}.attn_output_norm.bias", il),
make_zeros(device, hidden, &format!("attn_norm_b_{il}"))?,
);
tensors.insert(
format!("blk.{}.ffn_up.weight", il),
make_buf(device, intermediate * hidden, &format!("up_w_{il}"))?,
);
tensors.insert(
format!("blk.{}.ffn_gate.weight", il),
make_buf(device, intermediate * hidden, &format!("gate_w_{il}"))?,
);
tensors.insert(
format!("blk.{}.ffn_down.weight", il),
make_buf(device, hidden * intermediate, &format!("down_w_{il}"))?,
);
tensors.insert(
format!("blk.{}.layer_output_norm.weight", il),
make_ones(device, hidden, &format!("ffn_norm_w_{il}"))?,
);
tensors.insert(
format!("blk.{}.layer_output_norm.bias", il),
make_zeros(device, hidden, &format!("ffn_norm_b_{il}"))?,
);
}
Ok(tensors)
}
pub(crate) fn run_synthetic_min_forward_for_cross_family_test() {
let cfg = synthetic_min_cfg(2);
let device = MlxDevice::new().expect("create device");
let mut registry = KernelRegistry::new();
register_nomic_bert_kernels(&mut registry);
let tensors = synthetic_weights(&device, &cfg).expect("build synthetic weights");
let weights_device = device.clone();
let weights = LoadedNomicBertWeights::from_tensors_for_test(tensors, weights_device);
let seq_len: u32 = 32;
let input_ids = device
.alloc_buffer((seq_len as usize) * 4, DType::U32, vec![seq_len as usize])
.expect("alloc input_ids");
{
let slice: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(
input_ids.contents_ptr() as *mut u32,
seq_len as usize,
)
};
for (i, slot) in slice.iter_mut().enumerate() {
*slot = (i as u32 * 7 + 3) % cfg.vocab_size as u32;
}
}
let mut encoder = device.command_encoder().expect("command_encoder");
let pooled = apply_nomic_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
seq_len,
)
.expect("nomic full forward");
encoder.commit_and_wait().expect("commit_and_wait");
assert_eq!(pooled.element_count(), cfg.hidden_size);
let view: &[f32] = pooled.as_slice::<f32>().expect("read pooled f32");
let norm: f32 = view.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-3,
"expected ||y||₂ ≈ 1.0 (post-l2-normalize), got {norm}"
);
let n_nonzero = view.iter().filter(|v| v.abs() > 1e-12).count();
assert!(
n_nonzero >= cfg.hidden_size / 2,
"expected most components non-zero, got {n_nonzero} of {}",
cfg.hidden_size
);
for &v in view {
assert!(v.is_finite(), "non-finite output element: {v}");
}
}
#[test]
fn full_forward_at_synthetic_min_config_produces_unit_norm_output() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
run_synthetic_min_forward_for_cross_family_test();
}
#[test]
fn full_forward_rejects_non_64_head_dim() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let cfg = NomicBertConfig {
hidden_size: 64, num_attention_heads: 2,
num_hidden_layers: 1,
intermediate_size: 128,
max_position_embeddings: 128,
vocab_size: 100,
type_vocab_size: 2,
layer_norm_eps: 1e-12,
pooling_type: PoolingType::Mean,
rope_freq_base: 1000.0,
causal_attention: false,
};
let device = MlxDevice::new().expect("create device");
let mut registry = KernelRegistry::new();
register_nomic_bert_kernels(&mut registry);
let tensors = synthetic_weights(&device, &cfg).expect("build synthetic weights");
let weights_device = device.clone();
let weights = LoadedNomicBertWeights::from_tensors_for_test(tensors, weights_device);
let seq_len: u32 = 32;
let input_ids = device
.alloc_buffer((seq_len as usize) * 4, DType::U32, vec![seq_len as usize])
.expect("alloc input_ids");
let mut encoder = device.command_encoder().expect("command_encoder");
let err = apply_nomic_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
seq_len,
)
.expect_err("head_dim != 64 must reject");
let msg = format!("{err}");
assert!(
msg.contains("head_dim") && msg.contains("64"),
"error message must reference head_dim and the required 64 constraint: {msg}"
);
}
#[rustfmt::skip]
const LLAMA_EMBEDDING_GROUND_TRUTH_HELLO_WORLD: [f32; 768] = [
-6.6696000e-03f32, -1.3524000e-03f32, -1.7149610e-01f32, 8.4113000e-03f32,
5.8636000e-03f32, 6.9821200e-02f32, -2.0240000e-04f32, -4.3022800e-02f32,
-1.4626900e-02f32, -5.4056500e-02f32, 5.4160000e-04f32, 3.9272200e-02f32,
2.7769300e-02f32, 8.0812800e-02f32, 4.5334100e-02f32, -6.2951900e-02f32,
1.0281800e-02f32, -2.9656100e-02f32, -4.2753000e-02f32, 2.9597000e-02f32,
-3.7053000e-03f32, -9.4301000e-02f32, -7.5451000e-03f32, 3.8064000e-02f32,
9.2231700e-02f32, -1.4276000e-02f32, -1.4984500e-02f32, 6.1637500e-02f32,
6.4217000e-03f32, -2.1997000e-02f32, -1.1787000e-03f32, -1.0889600e-02f32,
-2.0770000e-04f32, 1.5721300e-02f32, 3.9444200e-02f32, 2.7844000e-03f32,
3.2542000e-02f32, 1.7387900e-02f32, 1.6315700e-02f32, 5.8692000e-03f32,
-4.7176000e-03f32, -1.4858700e-02f32, 1.1955900e-02f32, 1.0195000e-02f32,
6.5921400e-02f32, -1.5323000e-03f32, -4.1892000e-03f32, 2.5850000e-04f32,
8.6810500e-02f32, -6.0505200e-02f32, -1.8267700e-02f32, 5.3402000e-03f32,
-9.7460000e-04f32, 6.0159100e-02f32, 6.7261100e-02f32, 3.5314900e-02f32,
4.9696500e-02f32, -6.1601000e-02f32, 2.4186600e-02f32, 3.4579300e-02f32,
2.1759600e-02f32, 4.3670700e-02f32, 3.2912600e-02f32, 6.5303800e-02f32,
-1.7461700e-02f32, -3.3584400e-02f32, -2.5229100e-02f32, 3.5515100e-02f32,
-2.7050000e-03f32, 1.8090300e-02f32, 7.3137600e-02f32, 4.2705000e-03f32,
1.0861100e-02f32, 1.4041100e-02f32, 2.4605300e-02f32, 2.8004000e-02f32,
1.8594200e-02f32, 7.9048000e-03f32, -9.1900000e-04f32, -1.1905300e-02f32,
3.5421100e-02f32, -4.1623100e-02f32, 5.5484100e-02f32, -4.4268500e-02f32,
-3.2420500e-02f32, -7.3050200e-02f32, -4.0210600e-02f32, 1.5107500e-02f32,
-7.4528400e-02f32, -2.3277800e-02f32, 7.2323200e-02f32, 2.4769200e-02f32,
-5.1810000e-04f32, -1.8537900e-02f32, -4.0699700e-02f32, 2.0469300e-02f32,
-4.6032500e-02f32, -1.4164000e-03f32, -2.0057200e-02f32, -1.0966100e-02f32,
1.5139200e-02f32, -2.6177000e-02f32, -2.1159000e-03f32, 3.9117600e-02f32,
7.4003400e-02f32, 3.3914800e-02f32, -4.2793600e-02f32, -1.1042500e-02f32,
-5.4241700e-02f32, -3.8041700e-02f32, -3.3099100e-02f32, 1.0952800e-02f32,
-7.7578000e-03f32, 1.6444700e-02f32, 1.1090300e-02f32, -1.8598600e-02f32,
3.7661700e-02f32, -8.0060300e-02f32, 1.3430200e-02f32, 3.8800400e-02f32,
-3.6834600e-02f32, -3.0368000e-03f32, -3.9599900e-02f32, 1.6384100e-02f32,
5.1860500e-02f32, 5.1746800e-02f32, -7.9243300e-02f32, -1.9026100e-02f32,
2.3724100e-02f32, 5.8361000e-03f32, 1.3606100e-02f32, -3.2963600e-02f32,
-2.5896500e-02f32, -2.4711400e-02f32, -2.5973700e-02f32, 1.4827800e-02f32,
-1.5112600e-02f32, -1.5942600e-02f32, 5.6933500e-02f32, 3.3677600e-02f32,
8.9373000e-03f32, 1.3058600e-02f32, 2.4422200e-02f32, -4.9721500e-02f32,
-4.0742700e-02f32, -3.7903000e-02f32, 6.7944700e-02f32, -5.4161300e-02f32,
2.5137400e-02f32, -2.9522100e-02f32, 1.1179000e-03f32, 4.3306200e-02f32,
4.2004700e-02f32, 3.6765400e-02f32, 1.7319600e-02f32, -9.2276000e-03f32,
-1.9733600e-02f32, 2.4374200e-02f32, 3.7715700e-02f32, -4.6261000e-02f32,
1.3235600e-02f32, -8.5385000e-03f32, -7.7219000e-03f32, 1.0578100e-02f32,
3.6785100e-02f32, -5.5686800e-02f32, 3.2216600e-02f32, 6.3634400e-02f32,
4.7028000e-03f32, 3.2951100e-02f32, -3.6591600e-02f32, -3.8097500e-02f32,
3.5200100e-02f32, -3.2147600e-02f32, -2.0294400e-02f32, -8.4025000e-03f32,
-8.9155000e-03f32, -2.8797300e-02f32, 2.7853000e-02f32, 6.2297000e-03f32,
1.7037400e-02f32, -4.1399900e-02f32, 5.3571000e-03f32, 2.3880700e-02f32,
-1.3950300e-02f32, -2.4504100e-02f32, -2.2816600e-02f32, 2.1248000e-03f32,
2.1516600e-02f32, -3.5409700e-02f32, -8.3450000e-03f32, 1.7045600e-02f32,
-6.2701600e-02f32, -3.6372200e-02f32, 2.0275800e-02f32, -6.0924000e-03f32,
1.9449000e-02f32, 1.3768700e-02f32, 2.3274300e-02f32, -8.7041200e-02f32,
-4.0821800e-02f32, -1.5951000e-03f32, -2.7441100e-02f32, -2.3361100e-02f32,
-1.2320000e-04f32, 8.0487700e-02f32, 4.9427200e-02f32, 3.5657900e-02f32,
3.8404300e-02f32, 1.3656900e-02f32, 6.7240500e-02f32, -6.1013000e-02f32,
-5.1268500e-02f32, -1.9204600e-02f32, -1.8456100e-02f32, -9.4280000e-03f32,
-1.3013600e-02f32, -3.9915000e-02f32, -6.8722000e-03f32, -5.2381000e-03f32,
4.1953500e-02f32, -3.7336400e-02f32, 4.6914900e-02f32, 1.5287100e-02f32,
6.0664000e-02f32, -4.3220000e-04f32, -7.6455200e-02f32, 3.9420000e-04f32,
-6.3508300e-02f32, 2.0475900e-02f32, -3.1551200e-02f32, -6.4586800e-02f32,
1.8433200e-02f32, 2.4937000e-02f32, -1.7175900e-02f32, 3.8528800e-02f32,
4.5732900e-02f32, 6.8505000e-02f32, -1.3920300e-02f32, -2.9826000e-02f32,
-2.9718900e-02f32, 3.1809800e-02f32, 1.2525400e-02f32, -3.7864700e-02f32,
-4.5741600e-02f32, 2.0945800e-02f32, 1.3378600e-02f32, 3.9765000e-03f32,
-1.6016100e-02f32, 7.9120000e-04f32, -5.0325400e-02f32, -2.6335300e-02f32,
-2.0792100e-02f32, 1.3855000e-03f32, 2.7455600e-02f32, -1.1552200e-02f32,
-1.8247500e-02f32, 1.2089200e-02f32, 1.1794200e-02f32, -1.8306400e-02f32,
2.4806100e-02f32, -1.1313790e-01f32, 3.8121800e-02f32, -2.6435400e-02f32,
-4.8520900e-02f32, 3.4552000e-02f32, -4.8187700e-02f32, 3.4003100e-02f32,
3.5869800e-02f32, -4.3102900e-02f32, 1.2138600e-02f32, 9.6684000e-03f32,
9.5008000e-03f32, 2.7620800e-02f32, -4.3571100e-02f32, -6.1012000e-03f32,
2.4550600e-02f32, 1.4137200e-02f32, -2.1751800e-02f32, 1.8641100e-02f32,
-2.5030500e-02f32, -3.0357800e-02f32, -1.2105300e-02f32, -3.4376800e-02f32,
8.0324000e-03f32, 1.1767000e-02f32, -9.7419000e-03f32, 1.2958400e-02f32,
-3.3330700e-02f32, -1.3954200e-02f32, 1.3599500e-02f32, 4.6106700e-02f32,
3.0477600e-02f32, 6.9333100e-02f32, 1.2361600e-02f32, 2.9699600e-02f32,
2.6872200e-02f32, 2.8872900e-02f32, -1.1028400e-02f32, -9.7210000e-03f32,
-1.2504600e-02f32, -3.0737000e-03f32, 5.5915000e-02f32, -5.5938000e-03f32,
-1.8363800e-02f32, 2.6282000e-03f32, 7.4948300e-02f32, -1.8678300e-02f32,
4.8992800e-02f32, -4.4942000e-03f32, -4.6219300e-02f32, 7.8223400e-02f32,
-9.1623600e-02f32, -2.8647000e-03f32, -5.4467100e-02f32, 4.3308000e-02f32,
-2.0742300e-02f32, 3.2399100e-02f32, 6.9515000e-02f32, 1.8378400e-02f32,
-3.8602500e-02f32, -5.0964700e-02f32, 2.5752500e-02f32, -3.2207600e-02f32,
-2.0015000e-03f32, 2.2829800e-02f32, 3.4840500e-02f32, 2.6006800e-02f32,
2.3608000e-02f32, -4.8204200e-02f32, -1.4688700e-02f32, 1.2953900e-02f32,
-4.0061300e-02f32, -4.3547100e-02f32, -1.5404500e-02f32, 4.3021400e-02f32,
1.4217300e-02f32, -1.9293600e-02f32, -4.3386300e-02f32, 5.0931100e-02f32,
-5.2732000e-03f32, 2.3002600e-02f32, 4.0056100e-02f32, -1.2371900e-02f32,
-2.2102400e-02f32, -1.9255200e-02f32, -2.2667500e-02f32, 1.6477400e-02f32,
-1.5024500e-02f32, 2.3886200e-02f32, 6.4518000e-03f32, 2.5857200e-02f32,
-3.7417800e-02f32, 2.5491300e-02f32, 6.3017000e-03f32, 2.4324100e-02f32,
1.1121200e-02f32, 3.2617300e-02f32, 9.7421000e-03f32, -6.4688500e-02f32,
-2.6930400e-02f32, -1.0810000e-03f32, 1.4429800e-02f32, -1.4482400e-02f32,
-2.9594800e-02f32, 3.6383500e-02f32, 2.7113400e-02f32, 6.3338000e-03f32,
2.5931000e-02f32, -1.8233900e-02f32, -9.2634000e-03f32, -2.1566000e-02f32,
1.7372000e-03f32, 3.1390600e-02f32, 2.2766800e-02f32, 6.7679000e-03f32,
-4.8544900e-02f32, -3.4086800e-02f32, 8.0462000e-03f32, 4.0696200e-02f32,
-1.6917700e-02f32, -2.5898600e-02f32, 3.7125400e-02f32, -2.8145100e-02f32,
1.1707400e-02f32, 3.4267900e-02f32, 1.5698300e-02f32, -2.7624200e-02f32,
6.9315000e-03f32, 3.6654900e-02f32, 1.4375900e-02f32, -2.4399900e-02f32,
3.5763000e-03f32, 2.1591000e-03f32, -4.3111000e-03f32, -2.9810300e-02f32,
6.8243000e-03f32, -2.2369500e-02f32, -2.1174000e-02f32, 6.8368000e-03f32,
-6.7607000e-03f32, 1.2870100e-02f32, 2.8253000e-02f32, -7.1764500e-02f32,
6.3303000e-03f32, -9.4630000e-04f32, -5.1895500e-02f32, -2.5373100e-02f32,
8.5210000e-03f32, -1.5810700e-02f32, 6.4544500e-02f32, 6.3795000e-02f32,
4.5600200e-02f32, -5.5528200e-02f32, -4.1763200e-02f32, 1.1410400e-02f32,
3.6577900e-02f32, -6.8033600e-02f32, -1.2944800e-02f32, -1.1250000e-03f32,
1.7747800e-02f32, 7.7825100e-02f32, 9.6088000e-03f32, -1.3749000e-02f32,
-3.6817300e-02f32, 6.7867900e-02f32, 2.8122900e-02f32, 2.5646100e-02f32,
1.0362000e-03f32, -3.9197900e-02f32, -9.8872000e-03f32, 1.4315400e-02f32,
1.8575000e-02f32, 1.2935500e-02f32, -2.2592300e-02f32, -2.4799100e-02f32,
4.4715800e-02f32, -1.6318900e-02f32, 3.3302100e-02f32, 3.5868800e-02f32,
6.6783100e-02f32, -3.4930700e-02f32, -5.7269400e-02f32, -5.9972000e-03f32,
-7.8350000e-03f32, 1.2211830e-01f32, 8.8283900e-02f32, -9.9352000e-03f32,
-5.0083900e-02f32, -5.3087000e-03f32, -2.0628500e-02f32, 1.7979200e-02f32,
4.4053000e-02f32, 9.6263000e-03f32, 8.6672300e-02f32, -5.0582700e-02f32,
4.5380800e-02f32, -1.9539800e-02f32, -4.5610000e-04f32, 4.2559600e-02f32,
-1.9513300e-02f32, 6.2631000e-03f32, -3.1664400e-02f32, 8.3836000e-03f32,
1.1540100e-02f32, -5.4867100e-02f32, 5.8530000e-04f32, 1.3838000e-03f32,
-7.2131000e-03f32, 9.4528000e-03f32, -4.6331800e-02f32, -5.2411900e-02f32,
-1.9209000e-02f32, -1.3257400e-02f32, -6.0218300e-02f32, 2.2961200e-02f32,
-1.6933200e-02f32, -1.3100100e-02f32, -1.4583500e-02f32, 2.6643300e-02f32,
2.7661400e-02f32, 2.8252300e-02f32, -6.9593600e-02f32, 1.6248700e-02f32,
5.0440000e-02f32, 6.9895800e-02f32, -1.1571000e-03f32, 1.3644400e-02f32,
-1.7439800e-02f32, 7.1650000e-03f32, -2.8896000e-03f32, -2.4393600e-02f32,
3.6425400e-02f32, -2.9890000e-04f32, -2.7375400e-02f32, -4.7094000e-03f32,
-4.5289300e-02f32, 3.5725500e-02f32, -4.5007200e-02f32, 1.1070000e-03f32,
3.8081600e-02f32, -1.1230100e-02f32, 7.1999000e-03f32, 3.3610400e-02f32,
6.8639000e-03f32, 2.3139700e-02f32, 2.6155600e-02f32, -5.7708000e-02f32,
9.9852000e-03f32, 2.1354700e-02f32, 3.1218800e-02f32, -1.1297100e-02f32,
-2.3065700e-02f32, 4.2179300e-02f32, 9.3753200e-02f32, -4.2773900e-02f32,
2.5180000e-02f32, -1.1069100e-02f32, -3.0348900e-02f32, 5.0103700e-02f32,
1.1354600e-02f32, -1.8240100e-02f32, -3.2781300e-02f32, 1.2266100e-02f32,
-7.6380600e-02f32, 6.8787400e-02f32, -1.2997500e-02f32, -5.4016200e-02f32,
2.7228400e-02f32, 2.6439900e-02f32, 1.4635100e-02f32, 1.0713200e-02f32,
-2.3172800e-02f32, -4.4703300e-02f32, 2.8427700e-02f32, -6.5769000e-03f32,
4.3078000e-03f32, 1.5217600e-02f32, 2.5405900e-02f32, 1.5774300e-02f32,
-2.9051000e-02f32, -4.0942000e-03f32, -8.7075000e-03f32, -3.5142900e-02f32,
4.2014600e-02f32, 3.4278600e-02f32, -5.1490600e-02f32, 1.6976500e-02f32,
2.2717700e-02f32, 3.1094700e-02f32, -6.9700300e-02f32, -4.6632600e-02f32,
-2.8557600e-02f32, -1.0152700e-02f32, -4.7277100e-02f32, -5.7679200e-02f32,
-5.0320000e-04f32, 2.1446700e-02f32, -3.2161400e-02f32, -7.4619700e-02f32,
4.6693000e-02f32, 4.7707500e-02f32, -2.4216800e-02f32, 1.5537500e-02f32,
3.1047200e-02f32, -5.2569000e-03f32, -1.5880500e-02f32, -1.2518600e-02f32,
1.3172300e-02f32, -1.7127400e-02f32, -2.9639000e-02f32, -4.1154700e-02f32,
2.2904300e-02f32, -2.9324900e-02f32, 1.6750600e-02f32, -4.9756000e-03f32,
4.0822200e-02f32, -4.2819000e-03f32, -6.4213000e-02f32, -1.8390500e-02f32,
-7.2800000e-05f32, -5.5501300e-02f32, -5.8984000e-03f32, 5.1533200e-02f32,
-1.3947800e-02f32, 1.2336100e-02f32, 5.5780000e-04f32, -7.4789400e-02f32,
3.9419300e-02f32, -3.4653400e-02f32, -2.4352800e-02f32, 2.6578200e-02f32,
5.3825000e-02f32, -2.4245500e-02f32, -3.2574900e-02f32, 4.9105500e-02f32,
-3.9906800e-02f32, -4.3803200e-02f32, -1.7663800e-02f32, -4.4167900e-02f32,
-2.9276000e-02f32, 6.4075000e-03f32, 6.0689900e-02f32, -6.9809000e-02f32,
4.9774200e-02f32, 7.8172000e-02f32, 8.2345000e-03f32, 4.1580200e-02f32,
1.8442300e-02f32, 1.5560500e-02f32, 7.5401900e-02f32, 2.9353300e-02f32,
-2.2047500e-02f32, 8.5528000e-03f32, 2.7844900e-02f32, -1.5099400e-02f32,
4.2347800e-02f32, -2.0616000e-03f32, -1.7948600e-02f32, -6.9906400e-02f32,
-3.4608900e-02f32, -1.5580200e-02f32, 4.9552700e-02f32, 2.4922100e-02f32,
2.7784400e-02f32, -6.3345000e-03f32, -4.4251600e-02f32, -5.0236200e-02f32,
-5.7502200e-02f32, 6.2764400e-02f32, 4.0139600e-02f32, -6.8978000e-03f32,
-6.4271800e-02f32, 2.3647000e-03f32, 1.6232200e-02f32, 2.9681700e-02f32,
1.9716900e-02f32, -2.7960000e-03f32, -3.1999600e-02f32, 1.7260500e-02f32,
6.1012600e-02f32, 1.3232500e-02f32, 1.8163400e-02f32, 1.7620000e-04f32,
1.4968500e-02f32, -4.0804300e-02f32, 4.3764700e-02f32, 2.4680400e-02f32,
5.5778100e-02f32, 4.4632200e-02f32, 7.5896700e-02f32, 6.1313200e-02f32,
4.8259900e-02f32, -1.3964600e-02f32, -2.7013200e-02f32, -1.1387000e-02f32,
1.2016400e-02f32, -2.7300600e-02f32, -8.4480900e-02f32, 2.0433600e-02f32,
-1.0788300e-02f32, 2.6292000e-03f32, -6.6455100e-02f32, -2.4444200e-02f32,
3.3388000e-02f32, -2.1442300e-02f32, -3.2666300e-02f32, 1.9507800e-02f32,
-9.2234800e-02f32, 1.3595900e-02f32, -1.5368600e-02f32, -2.0472100e-02f32,
-2.8691200e-02f32, -4.4806300e-02f32, -2.7665200e-02f32, 3.8195300e-02f32,
2.7114000e-02f32, 2.2422500e-02f32, 2.9953900e-02f32, 2.4472000e-03f32,
1.1154500e-02f32, -1.4125200e-02f32, -4.3632000e-02f32, 3.4539100e-02f32,
4.5745500e-02f32, -4.3739300e-02f32, 6.4070700e-02f32, -1.9190600e-02f32,
-7.7880300e-02f32, -6.0991400e-02f32, -1.2944600e-02f32, -1.5316700e-02f32,
-5.9819000e-03f32, -3.1322900e-02f32, -2.6103200e-02f32, 2.9772100e-02f32,
-1.2600200e-02f32, 1.2044100e-02f32, -3.9712600e-02f32, 3.5522000e-02f32,
-4.1178100e-02f32, -1.2571100e-02f32, 2.2523000e-02f32, -7.8828000e-03f32,
4.6103000e-03f32, -3.9207100e-02f32, 1.3137100e-02f32, 4.1068200e-02f32,
-9.2080000e-03f32, -1.5108000e-03f32, -1.3505600e-02f32, 6.4108600e-02f32,
1.5352400e-02f32, 3.4981400e-02f32, -9.6561000e-03f32, 4.0101400e-02f32,
-2.7272300e-02f32, -1.0268500e-02f32, -6.3159000e-03f32, 5.9788300e-02f32,
7.2369200e-02f32, 4.2342700e-02f32, -4.1509400e-02f32, -2.3098600e-02f32,
-2.6804800e-02f32, 2.0771000e-03f32, 1.8563900e-02f32, -3.3814200e-02f32,
1.5673800e-02f32, -3.7488400e-02f32, -2.7946300e-02f32, -3.7747300e-02f32,
-3.2442600e-02f32, 2.8004300e-02f32, -2.6214700e-02f32, 2.7615200e-02f32,
-7.3416000e-03f32, -5.4686100e-02f32, 5.0802000e-03f32, -3.3259000e-02f32,
-2.3903400e-02f32, -7.0778800e-02f32, 1.7292100e-02f32, 6.2792200e-02f32,
-4.9236000e-03f32, -2.3950700e-02f32, 3.4221000e-02f32, 7.2967300e-02f32,
-9.6511000e-03f32, -2.0971600e-02f32, 2.4748800e-02f32, 5.5330000e-04f32,
-1.0364100e-02f32, -7.1326500e-02f32, -4.3360000e-04f32, 3.6105500e-02f32,
9.9634000e-03f32, 2.2385100e-02f32, 6.2977100e-02f32, -4.1682900e-02f32,
4.3001200e-02f32, -1.4988600e-02f32, -2.2700000e-04f32, 9.6763000e-03f32,
2.5719400e-02f32, -2.6735100e-02f32, -5.0922400e-02f32, -4.4518000e-03f32,
];
#[test]
fn full_forward_at_production_scale_on_real_nomic_gguf_produces_unit_norm_output() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::gguf::GgufFile;
use std::path::Path;
use super::super::tokenizer::build_nomic_wordpiece_tokenizer;
let model_path = Path::new("/opt/hf2q/models/bert-test/nomic-embed-text-v1.5-f16.gguf");
if !model_path.exists() {
eprintln!(
"skipping: nomic GGUF fixture not at {}",
model_path.display()
);
return;
}
let gguf = GgufFile::open(model_path).expect("open nomic GGUF");
let cfg = NomicBertConfig::from_gguf(&gguf).expect("parse nomic config");
assert_eq!(
cfg.hidden_size, 768,
"expected nomic-embed-text-v1.5 hidden=768"
);
assert_eq!(cfg.num_hidden_layers, 12, "expected 12 blocks");
assert_eq!(cfg.num_attention_heads, 12, "expected 12 heads");
assert_eq!(cfg.intermediate_size, 3072, "expected n_ff=3072");
let tok = build_nomic_wordpiece_tokenizer(model_path).expect("build tokenizer");
let real_ids = tok.encode("hello world", true);
assert_eq!(
real_ids.len(),
4,
"expected [CLS] hello world [SEP], got {real_ids:?}"
);
let seq_len: u32 = 32;
let pad_id = tok.specials().pad;
let mut padded_ids: Vec<u32> = real_ids.clone();
while padded_ids.len() < seq_len as usize {
padded_ids.push(pad_id);
}
let valid_token_count: u32 = real_ids.len() as u32;
let device = MlxDevice::new().expect("create device");
let mut registry = KernelRegistry::new();
register_nomic_bert_kernels(&mut registry);
let weights = LoadedNomicBertWeights::load_from_path(model_path, &cfg)
.expect("load real nomic weights");
assert!(weights.len() > 0, "loader returned zero tensors");
assert_eq!(
weights.len(),
112,
"expected 112 tensors loaded from nomic GGUF, got {}",
weights.len()
);
let input_ids = device
.alloc_buffer((seq_len as usize) * 4, DType::U32, vec![seq_len as usize])
.expect("alloc input_ids");
{
let slice: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(
input_ids.contents_ptr() as *mut u32,
seq_len as usize,
)
};
slice.copy_from_slice(&padded_ids);
}
let mut encoder = device.command_encoder().expect("command_encoder");
let pooled = apply_nomic_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None, &weights,
&cfg,
seq_len,
valid_token_count,
)
.expect("nomic full forward at production scale");
encoder.commit_and_wait().expect("commit_and_wait");
assert_eq!(
pooled.element_count(),
cfg.hidden_size,
"expected output dim = hidden_size = 768"
);
let view: &[f32] = pooled.as_slice::<f32>().expect("read pooled f32");
let norm: f32 = view.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-3,
"expected ||y||₂ ≈ 1.0 (post-l2-normalize), got {norm}"
);
let n_nontrivial = view.iter().filter(|v| v.abs() > 1e-6).count();
assert!(
n_nontrivial >= 700,
"expected most of 768 components non-trivial, got {n_nontrivial}"
);
for &v in view {
assert!(v.is_finite(), "non-finite output element: {v}");
}
let max_abs = view.iter().fold(0.0_f32, |acc, &v| acc.max(v.abs()));
assert!(
max_abs < 1.0,
"max |y_i| = {max_abs} unexpectedly large (post-l2 should be ≤ 1.0)"
);
eprintln!(
"[nomic real-gguf smoke] hidden={}, ||y||₂={:.6}, max|y|={:.4}, first4={:?}",
view.len(),
norm,
max_abs,
&view[..4]
);
}
#[test]
#[ignore = "perf timing test; run with --ignored --nocapture"]
fn forward_timing_10x_warm() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::gguf::GgufFile;
use std::path::Path;
use std::time::Instant;
use super::super::tokenizer::build_nomic_wordpiece_tokenizer;
let model_path = Path::new("/opt/hf2q/models/bert-test/nomic-embed-text-v1.5-f16.gguf");
if !model_path.exists() {
eprintln!(
"skipping: nomic GGUF fixture not at {}",
model_path.display()
);
return;
}
let gguf = GgufFile::open(model_path).expect("open nomic GGUF");
let cfg = NomicBertConfig::from_gguf(&gguf).expect("parse cfg");
let tok = build_nomic_wordpiece_tokenizer(model_path).expect("tok");
let real_ids = tok.encode("hello world", true);
let valid_token_count = real_ids.len() as u32;
let seq_len: u32 = 32;
let pad_id = tok.specials().pad;
let mut padded_ids: Vec<u32> = real_ids.clone();
while padded_ids.len() < seq_len as usize {
padded_ids.push(pad_id);
}
let device = MlxDevice::new().expect("device");
let mut registry = KernelRegistry::new();
register_nomic_bert_kernels(&mut registry);
let weights = LoadedNomicBertWeights::load_from_path(model_path, &cfg).expect("load");
let input_ids = device
.alloc_buffer((seq_len as usize) * 4, DType::U32, vec![seq_len as usize])
.expect("alloc input_ids");
unsafe {
let s: &mut [u32] = std::slice::from_raw_parts_mut(
input_ids.contents_ptr() as *mut u32,
seq_len as usize,
);
s.copy_from_slice(&padded_ids);
}
eprintln!("--- 10 sequential forwards (single process, no HTTP) ---");
let mut timings: Vec<f64> = Vec::with_capacity(10);
for i in 0..10 {
let t0 = Instant::now();
let mut encoder = device.command_encoder().expect("encoder");
let pooled = apply_nomic_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
valid_token_count,
)
.expect("forward");
encoder.commit_and_wait().expect("commit");
let elapsed_ms = t0.elapsed().as_secs_f64() * 1000.0;
timings.push(elapsed_ms);
let _ = pooled.as_slice::<f32>().expect("read")[0];
eprintln!(" forward {}: {:.2} ms", i + 1, elapsed_ms);
}
let mean = timings.iter().sum::<f64>() / timings.len() as f64;
let min = timings.iter().cloned().fold(f64::INFINITY, f64::min);
eprintln!(
" --> mean {:.2} ms, min {:.2} ms (llama-embedding reference: ~4.54 ms)",
mean, min
);
}
#[test]
fn full_forward_matches_llama_embedding_on_hello_world() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::gguf::GgufFile;
use std::path::Path;
use super::super::tokenizer::build_nomic_wordpiece_tokenizer;
let model_path = Path::new("/opt/hf2q/models/bert-test/nomic-embed-text-v1.5-f16.gguf");
if !model_path.exists() {
eprintln!(
"skipping: nomic GGUF fixture not at {}",
model_path.display()
);
return;
}
let gguf = GgufFile::open(model_path).expect("open nomic GGUF");
let cfg = NomicBertConfig::from_gguf(&gguf).expect("parse nomic config");
let tok = build_nomic_wordpiece_tokenizer(model_path).expect("build tokenizer");
let real_ids = tok.encode("hello world", true);
let valid_token_count: u32 = real_ids.len() as u32;
let seq_len: u32 = 32;
let pad_id = tok.specials().pad;
let mut padded_ids: Vec<u32> = real_ids.clone();
while padded_ids.len() < seq_len as usize {
padded_ids.push(pad_id);
}
let device = MlxDevice::new().expect("create device");
let mut registry = KernelRegistry::new();
register_nomic_bert_kernels(&mut registry);
let weights = LoadedNomicBertWeights::load_from_path(model_path, &cfg)
.expect("load real nomic weights");
let input_ids = device
.alloc_buffer((seq_len as usize) * 4, DType::U32, vec![seq_len as usize])
.expect("alloc input_ids");
{
let slice: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(
input_ids.contents_ptr() as *mut u32,
seq_len as usize,
)
};
slice.copy_from_slice(&padded_ids);
}
let mut encoder = device.command_encoder().expect("command_encoder");
let pooled = apply_nomic_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
valid_token_count,
)
.expect("nomic full forward");
encoder.commit_and_wait().expect("commit_and_wait");
let hf2q_view: &[f32] = pooled.as_slice::<f32>().expect("read hf2q pooled f32");
assert_eq!(hf2q_view.len(), 768);
let truth: &[f32] = &LLAMA_EMBEDDING_GROUND_TRUTH_HELLO_WORLD;
let dot: f32 = hf2q_view.iter().zip(truth.iter()).map(|(a, b)| a * b).sum();
let na: f32 = hf2q_view.iter().map(|v| v * v).sum::<f32>().sqrt();
let nb: f32 = truth.iter().map(|v| v * v).sum::<f32>().sqrt();
let cosine = dot / (na * nb);
let max_abs_diff = hf2q_view
.iter()
.zip(truth.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0_f32, f32::max);
eprintln!(
"[nomic parity] cosine={:.6}, ||hf2q||₂={:.6}, ||truth||₂={:.6}, max_abs_diff={:.4e}",
cosine, na, nb, max_abs_diff
);
eprintln!(
" hf2q first4 = {:?}\n truth first4 = {:?}",
&hf2q_view[..4],
&truth[..4]
);
assert!(
cosine >= 0.999,
"cosine {cosine:.6} below 0.999 gate; hf2q diverges from llama-embedding",
);
}
#[test]
fn full_forward_padding_invariance_at_seq_lens_32_64_128_256_512() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use mlx_native::gguf::GgufFile;
use std::path::Path;
use super::super::tokenizer::build_nomic_wordpiece_tokenizer;
let model_path = Path::new("/opt/hf2q/models/bert-test/nomic-embed-text-v1.5-f16.gguf");
if !model_path.exists() {
eprintln!(
"skipping: nomic GGUF fixture not at {}",
model_path.display()
);
return;
}
let gguf = GgufFile::open(model_path).expect("open nomic GGUF");
let cfg = NomicBertConfig::from_gguf(&gguf).expect("parse nomic config");
let tok = build_nomic_wordpiece_tokenizer(model_path).expect("build tokenizer");
let real_ids = tok.encode("hello world", true);
let valid_token_count: u32 = real_ids.len() as u32;
let pad_id = tok.specials().pad;
let device = MlxDevice::new().expect("create device");
let mut registry = KernelRegistry::new();
register_nomic_bert_kernels(&mut registry);
let weights = LoadedNomicBertWeights::load_from_path(model_path, &cfg)
.expect("load real nomic weights");
let seq_lens: &[u32] = &[32, 64, 128, 256, 512];
let mut outputs: Vec<(u32, Vec<f32>)> = Vec::with_capacity(seq_lens.len());
for &seq_len in seq_lens {
let mut padded_ids: Vec<u32> = real_ids.clone();
while padded_ids.len() < seq_len as usize {
padded_ids.push(pad_id);
}
let input_ids = device
.alloc_buffer((seq_len as usize) * 4, DType::U32, vec![seq_len as usize])
.expect("alloc input_ids");
{
let slice: &mut [u32] = unsafe {
std::slice::from_raw_parts_mut(
input_ids.contents_ptr() as *mut u32,
seq_len as usize,
)
};
slice.copy_from_slice(&padded_ids);
}
let mut encoder = device.command_encoder().expect("command_encoder");
let pooled = apply_nomic_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
valid_token_count,
)
.unwrap_or_else(|e| panic!("forward at seq_len={seq_len}: {e}"));
encoder.commit_and_wait().expect("commit_and_wait");
let view: &[f32] = pooled.as_slice::<f32>().expect("read pooled f32");
assert_eq!(view.len(), 768, "seq_len={seq_len}: hidden_size mismatch");
let n: f32 = view.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!(
(n - 1.0).abs() < 1e-3,
"seq_len={seq_len}: ||y||₂ = {n}, expected ~1.0"
);
for &v in view {
assert!(
v.is_finite(),
"seq_len={seq_len}: non-finite component encountered"
);
}
outputs.push((seq_len, view.to_vec()));
}
let baseline = &outputs[0]; let mut max_drift: f32 = 0.0;
for (sl, vec) in outputs.iter().skip(1) {
let dot: f32 = baseline.1.iter().zip(vec.iter()).map(|(a, b)| a * b).sum();
let na: f32 = baseline.1.iter().map(|v| v * v).sum::<f32>().sqrt();
let nb: f32 = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
let cosine = dot / (na * nb);
let drift = (1.0 - cosine).abs();
if drift > max_drift {
max_drift = drift;
}
eprintln!("[pad-invariance] seq_len 32 vs {sl}: cosine={cosine:.7}, drift={drift:.2e}");
assert!(
cosine >= 0.99999,
"seq_len 32 vs {sl}: cosine {cosine:.7} below 0.99999 padding-invariance gate \
(drift {drift:.2e}). Either the BF16 padding mask is leaking, flash-attn d=64 \
is unstable at long seq_len, or mean pooling is dividing by seq_len instead \
of valid_token_count."
);
}
eprintln!(
"[pad-invariance] PASS — max drift across {} seq_lens = {:.2e} (gate: 1e-5)",
seq_lens.len(),
max_drift
);
}
}