#![allow(dead_code)]
use anyhow::{anyhow, Context, Result};
use mlx_native::metal::MTLSize;
use mlx_native::ops::dense_mm_bf16::{dense_matmul_bf16_f32_tensor, DenseMmBf16F32Params};
use mlx_native::ops::elementwise::{cast, elementwise_add, CastDirection};
use mlx_native::ops::encode_helpers::KernelArg;
use mlx_native::ops::gather::dispatch_gather_f32;
use mlx_native::ops::l2_norm::dispatch_l2_norm;
use mlx_native::{CommandEncoder, DType, KernelRegistry, MlxBuffer, MlxDevice};
const BERT_CUSTOM_SHADERS_SOURCE: &str = r#"
#include <metal_stdlib>
using namespace metal;
struct LayerNormParams {
uint hidden;
uint batch;
float eps;
};
// Per-row LayerNorm: out[r, h] = (x[r, h] - mean(x[r, :])) /
// sqrt(var(x[r, :]) + eps) * gamma[h] + beta[h]
//
// Dispatch: one threadgroup per row, `threads_per_threadgroup.x` set to
// the chosen reduction width (≤ hidden). Threadgroup memory at index 0
// is `4 * threads_per_threadgroup.x` bytes.
//
// Two-pass: pass 1 computes the row mean via parallel reduction, pass 2
// computes the variance via the same reduction pattern using the mean
// from pass 1, then a final write applies the affine transform. F32
// throughout — BERT weights are F16 in GGUF but every dequant target
// is F32 in this loader for parity with the CPU reference.
kernel void bert_layer_norm_f32(
device const float* input [[buffer(0)]],
device const float* gamma [[buffer(1)]],
device const float* beta [[buffer(2)]],
device float* output [[buffer(3)]],
constant LayerNormParams& params [[buffer(4)]],
threadgroup float* shmem [[threadgroup(0)]],
uint tid [[thread_position_in_threadgroup]],
uint bid [[threadgroup_position_in_grid]],
uint ntg [[threads_per_threadgroup]]
) {
if (bid >= params.batch) return;
uint row_off = bid * params.hidden;
// ----- Pass 1: row sum -> mean -----
float sum = 0.0;
for (uint i = tid; i < params.hidden; i += ntg) {
sum += input[row_off + i];
}
shmem[tid] = sum;
threadgroup_barrier(mem_flags::mem_threadgroup);
// Parallel reduction; ntg is a power of two by caller construction.
for (uint stride = ntg / 2u; stride > 0u; stride >>= 1u) {
if (tid < stride) {
shmem[tid] += shmem[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float mean = shmem[0] / float(params.hidden);
threadgroup_barrier(mem_flags::mem_threadgroup);
// ----- Pass 2: row variance -> inv_std -----
float var_sum = 0.0;
for (uint i = tid; i < params.hidden; i += ntg) {
float d = input[row_off + i] - mean;
var_sum += d * d;
}
shmem[tid] = var_sum;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = ntg / 2u; stride > 0u; stride >>= 1u) {
if (tid < stride) {
shmem[tid] += shmem[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float inv_std = rsqrt(shmem[0] / float(params.hidden) + params.eps);
threadgroup_barrier(mem_flags::mem_threadgroup);
// ----- Apply: (x - mean) * inv_std * gamma + beta -----
for (uint i = tid; i < params.hidden; i += ntg) {
float x = input[row_off + i];
output[row_off + i] = ((x - mean) * inv_std) * gamma[i] + beta[i];
}
}
// Fused per-row residual_add + LayerNorm:
// tmp[r, h] = input[r, h] + residual[r, h]
// out[r, h] = (tmp[r, h] - mean(tmp[r, :])) /
// sqrt(var(tmp[r, :]) + eps) * gamma[h] + beta[h]
//
// Replaces the bert_residual_add_gpu → bert_layer_norm_gpu pair with a
// single kernel that reads the residual ONCE and avoids the
// intermediate writethrough of `tmp` to global memory. Each layer of
// nomic-bert / BERT calls this twice (post-attention residual+norm and
// post-FFN residual+norm) — at 12 layers that's 48 dispatches saved
// per forward pass (iter-86 perf optimization).
//
// Same dispatch envelope as bert_layer_norm_f32: one threadgroup per
// row, threads_per_threadgroup = power-of-two ≤ hidden, threadgroup
// memory at index 0 = 4 * threads_per_threadgroup bytes.
//
// Numerical equivalence: identical output to running residual_add then
// layer_norm sequentially. The fused kernel computes the SAME mean,
// inv_std, and affine transform — just without round-tripping `tmp`
// through device memory between the two passes.
kernel void bert_residual_layer_norm_f32(
device const float* input [[buffer(0)]],
device const float* residual [[buffer(1)]],
device const float* gamma [[buffer(2)]],
device const float* beta [[buffer(3)]],
device float* output [[buffer(4)]],
constant LayerNormParams& params [[buffer(5)]],
threadgroup float* shmem [[threadgroup(0)]],
uint tid [[thread_position_in_threadgroup]],
uint bid [[threadgroup_position_in_grid]],
uint ntg [[threads_per_threadgroup]]
) {
if (bid >= params.batch) return;
uint row_off = bid * params.hidden;
// Pass 1: row sum of (input + residual) -> mean.
// Reads BOTH input and residual once per element across the row.
float sum = 0.0;
for (uint i = tid; i < params.hidden; i += ntg) {
sum += input[row_off + i] + residual[row_off + i];
}
shmem[tid] = sum;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = ntg / 2u; stride > 0u; stride >>= 1u) {
if (tid < stride) {
shmem[tid] += shmem[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float mean = shmem[0] / float(params.hidden);
threadgroup_barrier(mem_flags::mem_threadgroup);
// Pass 2: variance over (input + residual).
// Re-reads input + residual; cheaper than writing intermediate sum
// to global memory between passes.
float var_sum = 0.0;
for (uint i = tid; i < params.hidden; i += ntg) {
float d = (input[row_off + i] + residual[row_off + i]) - mean;
var_sum += d * d;
}
shmem[tid] = var_sum;
threadgroup_barrier(mem_flags::mem_threadgroup);
for (uint stride = ntg / 2u; stride > 0u; stride >>= 1u) {
if (tid < stride) {
shmem[tid] += shmem[tid + stride];
}
threadgroup_barrier(mem_flags::mem_threadgroup);
}
float inv_std = rsqrt(shmem[0] / float(params.hidden) + params.eps);
threadgroup_barrier(mem_flags::mem_threadgroup);
// Apply: ((tmp - mean) * inv_std) * gamma + beta.
// Final pass re-reads input + residual one more time. The total
// global reads per element across the kernel are 3 reads of input
// and 3 of residual — same as the unfused pair (residual_add reads
// 2 + writes 1 = 3 ops; layer_norm reads 3 + writes 1 = 4 ops; the
// fused kernel saves the residual_add's write + the layer_norm's
// first read of tmp — net 2 fewer global memory ops per element).
for (uint i = tid; i < params.hidden; i += ntg) {
float t = input[row_off + i] + residual[row_off + i];
output[row_off + i] = ((t - mean) * inv_std) * gamma[i] + beta[i];
}
}
struct BiasAddParams {
uint rows;
uint cols;
};
// out[r, c] = input[r, c] + bias[c]. The matmul produces `[rows, cols]`
// row-major; this kernel broadcasts the per-column bias along rows.
//
// In-place is supported (input == output): each thread reads `input` and
// writes `output` at the same offset, so no aliasing hazard.
kernel void bert_bias_add_f32(
device const float* input [[buffer(0)]],
device const float* bias [[buffer(1)]],
device float* output [[buffer(2)]],
constant BiasAddParams& params [[buffer(3)]],
uint2 gid [[thread_position_in_grid]]
) {
uint c = gid.x;
uint r = gid.y;
if (c >= params.cols || r >= params.rows) return;
uint idx = r * params.cols + c;
output[idx] = input[idx] + bias[c];
}
struct PoolMeanParams {
uint seq_len;
uint hidden;
};
struct AttnMaskAddParams {
uint num_heads;
uint seq_q;
uint seq_k;
};
// Broadcast-add a `[seq_q, seq_k]` mask to a `[num_heads, seq_q, seq_k]`
// scores tensor. Mask[r, c] is 0.0 at valid positions and -INF
// (encoded as a very large negative float) at padded positions; after
// softmax the padded contributions vanish.
//
// In-place safe (input == output): each thread reads from `scores`
// at one offset and writes to `output` at the same offset. The kernel
// is bandwidth-bound; no reduction.
kernel void bert_attention_mask_add_f32(
device const float* scores [[buffer(0)]],
device const float* mask [[buffer(1)]],
device float* output [[buffer(2)]],
constant AttnMaskAddParams& params [[buffer(3)]],
uint3 gid [[thread_position_in_grid]]
) {
uint c = gid.x;
uint r = gid.y;
uint h = gid.z;
if (c >= params.seq_k || r >= params.seq_q || h >= params.num_heads) return;
uint scores_idx = (h * params.seq_q + r) * params.seq_k + c;
uint mask_idx = r * params.seq_k + c;
output[scores_idx] = scores[scores_idx] + mask[mask_idx];
}
// Mean-pool across the sequence dim:
// out[h] = (1/seq_len) * sum_{s=0}^{seq_len-1} input[s * hidden + h]
//
// Input shape `[seq_len, hidden]` row-major; output shape `[hidden]`.
// One thread per column of `hidden`; the loop iterates over seq_len.
// For seq_len up to a few thousand, this is fully bandwidth-bound and a
// fancier reduction (parallel sum) would not help on M5 Max. Keeping
// the kernel scalar-per-thread also makes correctness obvious.
kernel void bert_pool_mean_f32(
device const float* input [[buffer(0)]],
device float* output [[buffer(1)]],
constant PoolMeanParams& params [[buffer(2)]],
uint gid [[thread_position_in_grid]]
) {
if (gid >= params.hidden) return;
float acc = 0.0;
for (uint s = 0; s < params.seq_len; ++s) {
acc += input[s * params.hidden + gid];
}
output[gid] = acc / float(params.seq_len);
}
"#;
#[repr(C)]
#[derive(Clone, Copy)]
struct LayerNormGpuParams {
hidden: u32,
batch: u32,
eps: f32,
}
#[repr(C)]
#[derive(Clone, Copy)]
struct BiasAddGpuParams {
rows: u32,
cols: u32,
}
#[repr(C)]
#[derive(Clone, Copy)]
struct PoolMeanGpuParams {
seq_len: u32,
hidden: u32,
}
#[repr(C)]
#[derive(Clone, Copy)]
struct AttnMaskAddGpuParams {
num_heads: u32,
seq_q: u32,
seq_k: u32,
}
fn pod_as_bytes<T: Copy>(p: &T) -> &[u8] {
unsafe { std::slice::from_raw_parts(p as *const T as *const u8, std::mem::size_of::<T>()) }
}
pub fn register_bert_custom_shaders(registry: &mut KernelRegistry) {
registry.register_source("bert_layer_norm_f32", BERT_CUSTOM_SHADERS_SOURCE);
registry.register_source("bert_residual_layer_norm_f32", BERT_CUSTOM_SHADERS_SOURCE);
registry.register_source("bert_bias_add_f32", BERT_CUSTOM_SHADERS_SOURCE);
registry.register_source("bert_pool_mean_f32", BERT_CUSTOM_SHADERS_SOURCE);
registry.register_source("bert_attention_mask_add_f32", BERT_CUSTOM_SHADERS_SOURCE);
mlx_native::ops::gelu::register(registry);
mlx_native::ops::softmax::register(registry);
mlx_native::ops::sigmoid_mul::register(registry);
mlx_native::ops::gather::register(registry);
mlx_native::ops::l2_norm::register(registry);
}
pub fn bert_layer_norm_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
gamma: &MlxBuffer,
beta: &MlxBuffer,
eps: f32,
batch: u32,
hidden: u32,
) -> Result<MlxBuffer> {
if batch == 0 || hidden == 0 {
return Err(anyhow!(
"bert_layer_norm_gpu: batch ({}) and hidden ({}) must be > 0",
batch,
hidden
));
}
let total = (batch as usize) * (hidden as usize);
let output = device
.alloc_buffer(total * 4, DType::F32, vec![batch as usize, hidden as usize])
.map_err(|e| anyhow!("alloc bert_layer_norm output: {e}"))?;
let pipeline = registry
.get_pipeline("bert_layer_norm_f32", device.metal_device())
.map_err(|e| anyhow!("bert_layer_norm_gpu: get_pipeline: {e}"))?;
let params = LayerNormGpuParams { hidden, batch, eps };
let bytes = pod_as_bytes(¶ms);
let cap = hidden.min(256);
let ntg = prev_pow2(cap.max(1));
let threadgroups = MTLSize::new(batch as u64, 1, 1);
let threadgroup_size = MTLSize::new(ntg as u64, 1, 1);
let shmem_bytes = (ntg as u64) * 4;
encoder.encode_threadgroups_with_args_and_shared(
pipeline,
&[
(0, KernelArg::Buffer(input)),
(1, KernelArg::Buffer(gamma)),
(2, KernelArg::Buffer(beta)),
(3, KernelArg::Buffer(&output)),
(4, KernelArg::Bytes(bytes)),
],
&[(0, shmem_bytes)],
threadgroups,
threadgroup_size,
);
Ok(output)
}
fn prev_pow2(n: u32) -> u32 {
debug_assert!(n >= 1);
1u32 << (31 - n.leading_zeros())
}
#[allow(clippy::too_many_arguments)]
pub fn bert_residual_layer_norm_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
residual: &MlxBuffer,
gamma: &MlxBuffer,
beta: &MlxBuffer,
eps: f32,
batch: u32,
hidden: u32,
) -> Result<MlxBuffer> {
if batch == 0 || hidden == 0 {
return Err(anyhow!(
"bert_residual_layer_norm_gpu: batch ({}) and hidden ({}) must be > 0",
batch,
hidden
));
}
let total = (batch as usize) * (hidden as usize);
let output = device
.alloc_buffer(total * 4, DType::F32, vec![batch as usize, hidden as usize])
.map_err(|e| anyhow!("alloc bert_residual_layer_norm output: {e}"))?;
let pipeline = registry
.get_pipeline("bert_residual_layer_norm_f32", device.metal_device())
.map_err(|e| anyhow!("bert_residual_layer_norm_gpu: get_pipeline: {e}"))?;
let params = LayerNormGpuParams { hidden, batch, eps };
let bytes = pod_as_bytes(¶ms);
let cap = hidden.min(256);
let ntg = prev_pow2(cap.max(1));
let threadgroups = MTLSize::new(batch as u64, 1, 1);
let threadgroup_size = MTLSize::new(ntg as u64, 1, 1);
let shmem_bytes = (ntg as u64) * 4;
encoder.encode_threadgroups_with_args_and_shared(
pipeline,
&[
(0, KernelArg::Buffer(input)),
(1, KernelArg::Buffer(residual)),
(2, KernelArg::Buffer(gamma)),
(3, KernelArg::Buffer(beta)),
(4, KernelArg::Buffer(&output)),
(5, KernelArg::Bytes(bytes)),
],
&[(0, shmem_bytes)],
threadgroups,
threadgroup_size,
);
Ok(output)
}
pub fn bert_bias_add_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
bias: &MlxBuffer,
rows: u32,
cols: u32,
) -> Result<MlxBuffer> {
if rows == 0 || cols == 0 {
return Err(anyhow!(
"bert_bias_add_gpu: rows ({}) and cols ({}) must be > 0",
rows,
cols
));
}
let total = (rows as usize) * (cols as usize);
let output = device
.alloc_buffer(total * 4, DType::F32, vec![rows as usize, cols as usize])
.map_err(|e| anyhow!("alloc bert_bias_add output: {e}"))?;
let pipeline = registry
.get_pipeline("bert_bias_add_f32", device.metal_device())
.map_err(|e| anyhow!("bert_bias_add_gpu: get_pipeline: {e}"))?;
let params = BiasAddGpuParams { rows, cols };
let bytes = pod_as_bytes(¶ms);
let grid = MTLSize::new(cols as u64, rows as u64, 1);
let tg_x = std::cmp::min(64, cols as u64);
let tg = MTLSize::new(tg_x, 1, 1);
encoder.encode_with_args(
pipeline,
&[
(0, KernelArg::Buffer(input)),
(1, KernelArg::Buffer(bias)),
(2, KernelArg::Buffer(&output)),
(3, KernelArg::Bytes(bytes)),
],
grid,
tg,
);
Ok(output)
}
#[allow(clippy::too_many_arguments)]
pub fn bert_linear_bf16_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight_bf16: &MlxBuffer,
bias_opt: Option<&MlxBuffer>,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<MlxBuffer> {
if in_features < 32 {
return Err(anyhow!(
"bert_linear_bf16_gpu: in_features ({}) must be >= 32",
in_features
));
}
if seq_len == 0 || out_features == 0 {
return Err(anyhow!(
"bert_linear_bf16_gpu: seq_len ({}) and out_features ({}) must be > 0",
seq_len,
out_features
));
}
if let Some(b) = bias_opt {
if b.element_count() != out_features as usize {
return Err(anyhow!(
"bert_linear_bf16_gpu: bias element_count ({}) != out_features ({})",
b.element_count(),
out_features
));
}
}
let out_bytes = (seq_len as usize) * (out_features as usize) * 4;
let mut matmul_out = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![seq_len as usize, out_features as usize],
)
.map_err(|e| anyhow!("alloc matmul output: {e}"))?;
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
weight_bf16,
input,
&mut matmul_out,
¶ms,
)
.context("bert_linear_bf16_gpu: dense_matmul_bf16_f32_tensor")?;
if let Some(bias) = bias_opt {
encoder.memory_barrier();
bert_bias_add_gpu(
encoder,
registry,
device,
&matmul_out,
bias,
seq_len,
out_features,
)
.context("bert_linear_bf16_gpu: bias_add")
} else {
Ok(matmul_out)
}
}
pub fn bert_linear_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
weight_f32: &MlxBuffer,
bias_opt: Option<&MlxBuffer>,
seq_len: u32,
in_features: u32,
out_features: u32,
) -> Result<MlxBuffer> {
if in_features < 32 {
return Err(anyhow!(
"bert_linear_gpu: in_features ({}) must be >= 32",
in_features
));
}
if seq_len == 0 || out_features == 0 {
return Err(anyhow!(
"bert_linear_gpu: seq_len ({}) and out_features ({}) must be > 0",
seq_len,
out_features
));
}
if let Some(b) = bias_opt {
if b.element_count() != out_features as usize {
return Err(anyhow!(
"bert_linear_gpu: bias element_count ({}) != out_features ({})",
b.element_count(),
out_features
));
}
}
let metal_dev = device.metal_device();
let n_w = (out_features as usize) * (in_features as usize);
let weight_bf16 = device
.alloc_buffer(
n_w * 2,
DType::BF16,
vec![out_features as usize, in_features as usize],
)
.map_err(|e| anyhow!("alloc weight_bf16: {e}"))?;
cast(
encoder,
registry,
metal_dev,
weight_f32,
&weight_bf16,
n_w,
CastDirection::F32ToBF16,
)
.context("bert_linear_gpu: F32→BF16 cast")?;
encoder.memory_barrier();
let out_bytes = (seq_len as usize) * (out_features as usize) * 4;
let mut matmul_out = device
.alloc_buffer(
out_bytes,
DType::F32,
vec![seq_len as usize, out_features as usize],
)
.map_err(|e| anyhow!("alloc matmul output: {e}"))?;
let params = DenseMmBf16F32Params {
m: seq_len,
n: out_features,
k: in_features,
src0_batch: 1,
src1_batch: 1,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&weight_bf16,
input,
&mut matmul_out,
¶ms,
)
.context("bert_linear_gpu: dense_matmul_bf16_f32_tensor")?;
if let Some(bias) = bias_opt {
encoder.memory_barrier();
bert_bias_add_gpu(
encoder,
registry,
device,
&matmul_out,
bias,
seq_len,
out_features,
)
.context("bert_linear_gpu: bias_add")
} else {
Ok(matmul_out)
}
}
pub fn bert_gelu_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
) -> Result<MlxBuffer> {
let n = input.element_count();
if n == 0 {
return Err(anyhow!(
"bert_gelu_gpu: input must have at least one element"
));
}
let bytes = match input.dtype() {
DType::F32 => n * 4,
DType::F16 | DType::BF16 => n * 2,
other => {
return Err(anyhow!(
"bert_gelu_gpu: unsupported dtype {:?} (expected F32/F16/BF16)",
other
));
}
};
let output = device
.alloc_buffer(bytes, input.dtype(), input.shape().to_vec())
.map_err(|e| anyhow!("alloc bert_gelu output: {e}"))?;
mlx_native::ops::gelu::dispatch_gelu(encoder, registry, device.metal_device(), input, &output)
.map_err(|e| anyhow!("bert_gelu_gpu: dispatch_gelu: {e}"))?;
Ok(output)
}
pub fn bert_attention_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
seq_len: u32,
num_heads: u32,
head_dim: u32,
scale: f32,
) -> Result<MlxBuffer> {
crate::inference::vision::vit_gpu::vit_attention_gpu(
encoder,
registry,
device,
q_seq_major,
k_seq_major,
v_seq_major,
seq_len,
num_heads,
head_dim,
scale,
)
.context("bert_attention_gpu: delegate to vit_attention_gpu")
}
pub fn bert_attention_mask_add_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
scores: &MlxBuffer,
mask: &MlxBuffer,
num_heads: u32,
seq_q: u32,
seq_k: u32,
) -> Result<MlxBuffer> {
if num_heads == 0 || seq_q == 0 || seq_k == 0 {
return Err(anyhow!(
"bert_attention_mask_add_gpu: num_heads/seq_q/seq_k must be > 0"
));
}
let total = (num_heads as usize) * (seq_q as usize) * (seq_k as usize);
let output = device
.alloc_buffer(
total * 4,
DType::F32,
vec![num_heads as usize, seq_q as usize, seq_k as usize],
)
.map_err(|e| anyhow!("alloc mask_add output: {e}"))?;
let pipeline = registry
.get_pipeline("bert_attention_mask_add_f32", device.metal_device())
.map_err(|e| anyhow!("bert_attention_mask_add_gpu: get_pipeline: {e}"))?;
let params = AttnMaskAddGpuParams {
num_heads,
seq_q,
seq_k,
};
let bytes = pod_as_bytes(¶ms);
let grid = MTLSize::new(seq_k as u64, seq_q as u64, num_heads as u64);
let tg = MTLSize::new(std::cmp::min(64, seq_k as u64), 1, 1);
encoder.encode_with_args(
pipeline,
&[
(0, KernelArg::Buffer(scores)),
(1, KernelArg::Buffer(mask)),
(2, KernelArg::Buffer(&output)),
(3, KernelArg::Bytes(bytes)),
],
grid,
tg,
);
Ok(output)
}
pub fn alloc_bert_attention_mask(
device: &MlxDevice,
seq_len: u32,
valid_len: u32,
) -> Result<MlxBuffer> {
if seq_len == 0 {
return Err(anyhow!("alloc_bert_attention_mask: seq_len must be > 0"));
}
let n = (seq_len as usize) * (seq_len as usize);
let buf = device
.alloc_buffer(n * 4, DType::F32, vec![seq_len as usize, seq_len as usize])
.map_err(|e| anyhow!("alloc attention mask: {e}"))?;
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, n) };
let valid = valid_len.min(seq_len) as usize;
let seq = seq_len as usize;
for r in 0..seq {
for c in 0..seq {
s[r * seq + c] = if c < valid { 0.0 } else { -1e30 };
}
}
Ok(buf)
}
#[allow(clippy::too_many_arguments)]
pub fn bert_attention_with_mask_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
q_seq_major: &MlxBuffer,
k_seq_major: &MlxBuffer,
v_seq_major: &MlxBuffer,
mask: &MlxBuffer,
seq_len: u32,
num_heads: u32,
head_dim: u32,
scale: f32,
) -> Result<MlxBuffer> {
use crate::inference::vision::vit_gpu::{vit_attention_scores_gpu, vit_softmax_last_dim_gpu};
use mlx_native::ops::transpose::{permute_021_f32, transpose_last2_bf16};
if head_dim < 32 {
return Err(anyhow!(
"bert_attention_with_mask_gpu: head_dim ({}) must be >= 32",
head_dim
));
}
if seq_len == 0 || num_heads == 0 {
return Err(anyhow!(
"bert_attention_with_mask_gpu: seq_len/num_heads must be > 0"
));
}
let metal_dev = device.metal_device();
let scores = vit_attention_scores_gpu(
encoder,
registry,
device,
q_seq_major,
k_seq_major,
seq_len,
num_heads,
head_dim,
scale,
)?;
encoder.memory_barrier();
let masked_scores = bert_attention_mask_add_gpu(
encoder, registry, device, &scores, mask, num_heads, seq_len, seq_len,
)?;
encoder.memory_barrier();
let n_rows = (num_heads as u64) * (seq_len as u64);
let softmaxed = vit_softmax_last_dim_gpu(
encoder,
registry,
device,
&masked_scores,
n_rows as u32,
seq_len,
)?;
encoder.memory_barrier();
let n_v = (seq_len as usize) * (num_heads as usize) * (head_dim as usize);
let v_perm = device
.alloc_buffer(
n_v * 4,
DType::F32,
vec![num_heads as usize, seq_len as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc v_perm: {e}"))?;
permute_021_f32(
encoder,
registry,
metal_dev,
v_seq_major,
&v_perm,
seq_len as usize,
num_heads as usize,
head_dim as usize,
)
.context("permute V seq→head major")?;
encoder.memory_barrier();
let v_bf16 = device
.alloc_buffer(
n_v * 2,
DType::BF16,
vec![num_heads as usize, seq_len as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc v_bf16: {e}"))?;
cast(
encoder,
registry,
metal_dev,
&v_perm,
&v_bf16,
n_v,
CastDirection::F32ToBF16,
)
.context("cast V f32→bf16")?;
encoder.memory_barrier();
let v_t_bf16 = device
.alloc_buffer(
n_v * 2,
DType::BF16,
vec![num_heads as usize, head_dim as usize, seq_len as usize],
)
.map_err(|e| anyhow!("alloc v_t_bf16: {e}"))?;
transpose_last2_bf16(
encoder,
registry,
metal_dev,
&v_bf16,
&v_t_bf16,
num_heads as usize,
seq_len as usize,
head_dim as usize,
)
.context("transpose V last 2")?;
encoder.memory_barrier();
let n_attn = (num_heads as usize) * (seq_len as usize) * (head_dim as usize);
let mut attn_head_major = device
.alloc_buffer(
n_attn * 4,
DType::F32,
vec![num_heads as usize, seq_len as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc attn_head_major: {e}"))?;
let params = DenseMmBf16F32Params {
m: seq_len,
n: head_dim,
k: seq_len,
src0_batch: num_heads,
src1_batch: num_heads,
};
dense_matmul_bf16_f32_tensor(
encoder,
registry,
device,
&v_t_bf16,
&softmaxed,
&mut attn_head_major,
¶ms,
)
.context("attention scores @ V matmul")?;
encoder.memory_barrier();
let attn_seq_major = device
.alloc_buffer(
n_attn * 4,
DType::F32,
vec![seq_len as usize, num_heads as usize, head_dim as usize],
)
.map_err(|e| anyhow!("alloc attn_seq_major: {e}"))?;
permute_021_f32(
encoder,
registry,
metal_dev,
&attn_head_major,
&attn_seq_major,
num_heads as usize,
seq_len as usize,
head_dim as usize,
)
.context("permute attn head→seq major")?;
Ok(attn_seq_major)
}
pub fn bert_residual_add_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
a: &MlxBuffer,
b: &MlxBuffer,
n_elements: u32,
) -> Result<MlxBuffer> {
if n_elements == 0 {
return Err(anyhow!("bert_residual_add_gpu: n_elements must be > 0"));
}
let out = device
.alloc_buffer(
(n_elements as usize) * 4,
DType::F32,
vec![n_elements as usize],
)
.map_err(|e| anyhow!("alloc residual add output: {e}"))?;
elementwise_add(
encoder,
registry,
device.metal_device(),
a,
b,
&out,
n_elements as usize,
DType::F32,
)
.context("bert_residual_add_gpu: elementwise_add")?;
Ok(out)
}
pub struct BertEncoderBlockTensors<'a> {
pub q_w: &'a MlxBuffer,
pub q_b: Option<&'a MlxBuffer>,
pub k_w: &'a MlxBuffer,
pub k_b: Option<&'a MlxBuffer>,
pub v_w: &'a MlxBuffer,
pub v_b: Option<&'a MlxBuffer>,
pub o_w: &'a MlxBuffer,
pub o_b: 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 down_w: &'a MlxBuffer,
pub down_b: Option<&'a MlxBuffer>,
pub ffn_norm_gamma: &'a MlxBuffer,
pub ffn_norm_beta: &'a MlxBuffer,
}
#[allow(clippy::too_many_arguments)]
pub fn apply_bert_encoder_block_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
tensors: &BertEncoderBlockTensors<'_>,
attention_mask: Option<&MlxBuffer>,
seq_len: u32,
hidden: u32,
num_heads: u32,
intermediate: u32,
eps: f32,
) -> Result<MlxBuffer> {
if hidden == 0 || num_heads == 0 || hidden % num_heads != 0 {
return Err(anyhow!(
"apply_bert_encoder_block_gpu: hidden ({}) must be > 0 and divisible by num_heads ({})",
hidden,
num_heads
));
}
let head_dim = hidden / num_heads;
let n_hidden = (seq_len as usize) * (hidden as usize);
let n_inter = (seq_len as usize) * (intermediate as usize);
let scale = 1.0_f32 / (head_dim as f32).sqrt();
let q = bert_linear_gpu(
encoder,
registry,
device,
input,
tensors.q_w,
tensors.q_b,
seq_len,
hidden,
hidden,
)
.context("encoder block: Q projection")?;
encoder.memory_barrier();
let k = bert_linear_gpu(
encoder,
registry,
device,
input,
tensors.k_w,
tensors.k_b,
seq_len,
hidden,
hidden,
)
.context("encoder block: K projection")?;
encoder.memory_barrier();
let v = bert_linear_gpu(
encoder,
registry,
device,
input,
tensors.v_w,
tensors.v_b,
seq_len,
hidden,
hidden,
)
.context("encoder block: V projection")?;
encoder.memory_barrier();
let attn_out = match attention_mask {
Some(mask) => bert_attention_with_mask_gpu(
encoder, registry, device, &q, &k, &v, mask, seq_len, num_heads, head_dim, scale,
)
.context("encoder block: masked bidirectional attention")?,
None => bert_attention_gpu(
encoder, registry, device, &q, &k, &v, seq_len, num_heads, head_dim, scale,
)
.context("encoder block: bidirectional attention")?,
};
encoder.memory_barrier();
let o_proj = bert_linear_gpu(
encoder,
registry,
device,
&attn_out,
tensors.o_w,
tensors.o_b,
seq_len,
hidden,
hidden,
)
.context("encoder block: attention output projection")?;
encoder.memory_barrier();
let _ = n_hidden; let post_attn = bert_residual_layer_norm_gpu(
encoder,
registry,
device,
input,
&o_proj,
tensors.attn_norm_gamma,
tensors.attn_norm_beta,
eps,
seq_len,
hidden,
)
.context("encoder block: post-attn residual+LayerNorm (fused)")?;
encoder.memory_barrier();
let ffn_up = bert_linear_gpu(
encoder,
registry,
device,
&post_attn,
tensors.up_w,
tensors.up_b,
seq_len,
hidden,
intermediate,
)
.context("encoder block: FFN up projection")?;
encoder.memory_barrier();
let ffn_act =
bert_gelu_gpu(encoder, registry, device, &ffn_up).context("encoder block: GeLU")?;
encoder.memory_barrier();
let ffn_down = bert_linear_gpu(
encoder,
registry,
device,
&ffn_act,
tensors.down_w,
tensors.down_b,
seq_len,
intermediate,
hidden,
)
.context("encoder block: FFN down projection")?;
let _ = n_inter; encoder.memory_barrier();
bert_residual_layer_norm_gpu(
encoder,
registry,
device,
&post_attn,
&ffn_down,
tensors.ffn_norm_gamma,
tensors.ffn_norm_beta,
eps,
seq_len,
hidden,
)
.context("encoder block: post-FFN residual+LayerNorm (fused)")
}
pub fn bert_embed_gather_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
table: &MlxBuffer,
ids: &MlxBuffer,
vocab: u32,
hidden: u32,
n_ids: u32,
) -> Result<MlxBuffer> {
if vocab == 0 || hidden == 0 || n_ids == 0 {
return Err(anyhow!(
"bert_embed_gather_gpu: vocab/hidden/n_ids must all be > 0 (got {} / {} / {})",
vocab,
hidden,
n_ids
));
}
let total = (n_ids as usize) * (hidden as usize);
let output = device
.alloc_buffer(total * 4, DType::F32, vec![n_ids as usize, hidden as usize])
.map_err(|e| anyhow!("alloc embed_gather output: {e}"))?;
dispatch_gather_f32(
encoder,
registry,
device.metal_device(),
table,
ids,
&output,
vocab,
hidden,
n_ids,
)
.map_err(|e| anyhow!("bert_embed_gather_gpu: dispatch_gather_f32: {e}"))?;
Ok(output)
}
fn alloc_position_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 position_ids: {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)
}
#[allow(clippy::too_many_arguments)]
pub fn bert_embeddings_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input_ids: &MlxBuffer,
type_ids_opt: Option<&MlxBuffer>,
token_embd: &MlxBuffer,
position_embd: &MlxBuffer,
token_types_opt: Option<&MlxBuffer>,
embed_norm_gamma: &MlxBuffer,
embed_norm_beta: &MlxBuffer,
eps: f32,
seq_len: u32,
hidden: u32,
vocab: u32,
max_pos: u32,
type_vocab: u32,
) -> Result<MlxBuffer> {
if seq_len == 0 || hidden == 0 {
return Err(anyhow!(
"bert_embeddings_gpu: seq_len ({}) and hidden ({}) must be > 0",
seq_len,
hidden
));
}
if seq_len > max_pos {
return Err(anyhow!(
"bert_embeddings_gpu: seq_len ({}) exceeds max_pos ({})",
seq_len,
max_pos
));
}
match (type_ids_opt.is_some(), token_types_opt.is_some()) {
(true, true) | (false, false) => {}
(a, b) => {
return Err(anyhow!(
"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("embeddings: token gather")?;
encoder.memory_barrier();
let pos_ids = alloc_position_ids(device, seq_len)?;
let pos = bert_embed_gather_gpu(
encoder,
registry,
device,
position_embd,
&pos_ids,
max_pos,
hidden,
seq_len,
)
.context("embeddings: position gather")?;
encoder.memory_barrier();
let tok_pos = bert_residual_add_gpu(encoder, registry, device, &tok, &pos, n_hidden as u32)
.context("embeddings: token + position add")?;
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("embeddings: type gather")?;
encoder.memory_barrier();
let s = bert_residual_add_gpu(encoder, registry, device, &tok_pos, &typ, n_hidden as u32)
.context("embeddings: + token_types add")?;
encoder.memory_barrier();
s
} else {
tok_pos
};
bert_layer_norm_gpu(
encoder,
registry,
device,
&summed,
embed_norm_gamma,
embed_norm_beta,
eps,
seq_len,
hidden,
)
.context("embeddings: post-sum LayerNorm")
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BertPoolKind {
Mean,
Cls,
Last,
}
pub fn bert_pool_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
kind: BertPoolKind,
seq_len: u32,
hidden: u32,
) -> Result<MlxBuffer> {
if seq_len == 0 || hidden == 0 {
return Err(anyhow!(
"bert_pool_gpu: seq_len ({}) and hidden ({}) must be > 0",
seq_len,
hidden
));
}
match kind {
BertPoolKind::Mean => {
let output = device
.alloc_buffer((hidden as usize) * 4, DType::F32, vec![hidden as usize])
.map_err(|e| anyhow!("alloc pool_mean output: {e}"))?;
let pipeline = registry
.get_pipeline("bert_pool_mean_f32", device.metal_device())
.map_err(|e| anyhow!("bert_pool_gpu: get_pipeline: {e}"))?;
let params = PoolMeanGpuParams { seq_len, hidden };
let bytes = pod_as_bytes(¶ms);
let grid = MTLSize::new(hidden as u64, 1, 1);
let tg_x = std::cmp::min(64, hidden as u64);
let tg = MTLSize::new(tg_x, 1, 1);
encoder.encode_with_args(
pipeline,
&[
(0, KernelArg::Buffer(input)),
(1, KernelArg::Buffer(&output)),
(2, KernelArg::Bytes(bytes)),
],
grid,
tg,
);
Ok(output)
}
BertPoolKind::Cls | BertPoolKind::Last => {
let idx_val: u32 = match kind {
BertPoolKind::Cls => 0,
BertPoolKind::Last => seq_len - 1,
BertPoolKind::Mean => unreachable!(),
};
let idx_buf = device
.alloc_buffer(4, DType::U32, vec![1])
.map_err(|e| anyhow!("alloc pool index buffer: {e}"))?;
let s: &mut [u32] =
unsafe { std::slice::from_raw_parts_mut(idx_buf.contents_ptr() as *mut u32, 1) };
s[0] = idx_val;
bert_embed_gather_gpu(
encoder, registry, device, input, &idx_buf, seq_len, hidden, 1,
)
}
}
}
pub fn bert_l2_normalize_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input: &MlxBuffer,
eps: f32,
rows: u32,
dim: u32,
) -> Result<MlxBuffer> {
if rows == 0 || dim == 0 {
return Err(anyhow!(
"bert_l2_normalize_gpu: rows ({}) and dim ({}) must be > 0",
rows,
dim
));
}
let total = (rows as usize) * (dim as usize);
let output = device
.alloc_buffer(total * 4, DType::F32, vec![rows as usize, dim as usize])
.map_err(|e| anyhow!("alloc l2_norm output: {e}"))?;
let params_buf = device
.alloc_buffer(8, DType::F32, vec![2])
.map_err(|e| anyhow!("alloc l2_norm params: {e}"))?;
{
let s: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(params_buf.contents_ptr() as *mut f32, 2) };
s[0] = eps;
s[1] = dim as f32;
}
dispatch_l2_norm(
encoder,
registry,
device.metal_device(),
input,
&output,
¶ms_buf,
rows,
dim,
)
.map_err(|e| anyhow!("bert_l2_normalize_gpu: dispatch_l2_norm: {e}"))?;
Ok(output)
}
pub fn apply_bert_full_forward_gpu(
encoder: &mut CommandEncoder,
registry: &mut KernelRegistry,
device: &MlxDevice,
input_ids: &MlxBuffer,
type_ids_opt: Option<&MlxBuffer>,
weights: &super::weights::LoadedBertWeights,
cfg: &super::config::BertConfig,
seq_len: u32,
valid_token_count: u32,
) -> Result<MlxBuffer> {
use super::config::PoolingType;
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_bert_full_forward_gpu: pooling_type=None is not a single-vector embedding output"
));
}
PoolingType::Rank => {
return Err(anyhow!(
"apply_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 eps = cfg.layer_norm_eps;
if seq_len < 32 {
return Err(anyhow!(
"apply_bert_full_forward_gpu: seq_len ({}) must be >= 32 (post-softmax matmul K floor)",
seq_len
));
}
if seq_len > max_pos {
return Err(anyhow!(
"apply_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 = bert_embeddings_gpu(
encoder,
registry,
device,
input_ids,
effective_type_ids,
weights.token_embd_weight()?,
weights.position_embd_weight()?,
token_types_for_call,
weights.embed_norm_weight()?,
weights.embed_norm_bias()?,
eps,
seq_len,
hidden,
vocab,
max_pos,
type_vocab,
)
.context("full forward: embeddings")?;
encoder.memory_barrier();
let mask_opt: Option<MlxBuffer> = Some(alloc_bert_attention_mask(
device,
seq_len,
valid_token_count,
)?);
let mask_ref = mask_opt.as_ref();
for layer_idx in 0..cfg.num_hidden_layers {
let tensors = BertEncoderBlockTensors {
q_w: weights.block_required(layer_idx, "attn_q.weight")?,
q_b: weights.block_optional(layer_idx, "attn_q.bias"),
k_w: weights.block_required(layer_idx, "attn_k.weight")?,
k_b: weights.block_optional(layer_idx, "attn_k.bias"),
v_w: weights.block_required(layer_idx, "attn_v.weight")?,
v_b: weights.block_optional(layer_idx, "attn_v.bias"),
o_w: weights.block_required(layer_idx, "attn_output.weight")?,
o_b: weights.block_optional(layer_idx, "attn_output.bias"),
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"),
down_w: weights.block_required(layer_idx, "ffn_down.weight")?,
down_b: weights.block_optional(layer_idx, "ffn_down.bias"),
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_bert_encoder_block_gpu(
encoder,
registry,
device,
&hidden_states,
&tensors,
mask_ref,
seq_len,
hidden,
num_heads,
intermediate,
eps,
)
.with_context(|| format!("full forward: encoder block {}", layer_idx))?;
encoder.memory_barrier();
}
let pooled = bert_pool_gpu(
encoder,
registry,
device,
&hidden_states,
pool_kind,
valid_token_count,
hidden,
)
.context("full forward: pool")?;
encoder.memory_barrier();
bert_l2_normalize_gpu(encoder, registry, device, &pooled, eps, 1, hidden)
.context("full forward: l2 normalize")
}
#[cfg(test)]
fn bert_linear_cpu_ref(
input: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
seq: usize,
in_features: usize,
out_features: usize,
) -> Vec<f32> {
assert_eq!(input.len(), seq * in_features);
assert_eq!(weight.len(), out_features * in_features);
if let Some(b) = bias {
assert_eq!(b.len(), out_features);
}
let mut out = vec![0.0f32; seq * out_features];
for m in 0..seq {
for n in 0..out_features {
let mut acc = 0.0f64;
for k in 0..in_features {
acc += (input[m * in_features + k] as f64) * (weight[n * in_features + k] as f64);
}
let mut v = acc as f32;
if let Some(b) = bias {
v += b[n];
}
out[m * out_features + n] = v;
}
}
out
}
#[cfg(test)]
fn bert_gelu_cpu_ref(input: &[f32]) -> Vec<f32> {
const GELU_COEF_A: f32 = 0.044715;
const SQRT_2_OVER_PI: f32 = 0.7978845608028654;
input
.iter()
.map(|&x| {
let inner = SQRT_2_OVER_PI * (x + GELU_COEF_A * x * x * x);
0.5 * x * (1.0 + inner.tanh())
})
.collect()
}
#[cfg(test)]
fn bert_attention_cpu_ref(
q: &[f32],
k: &[f32],
v: &[f32],
seq_len: usize,
num_heads: usize,
head_dim: usize,
scale: f32,
) -> Vec<f32> {
let n = seq_len * num_heads * head_dim;
assert_eq!(q.len(), n);
assert_eq!(k.len(), n);
assert_eq!(v.len(), n);
fn row(buf: &[f32], s: usize, h: usize, num_heads: usize, head_dim: usize) -> &[f32] {
let off = (s * num_heads + h) * head_dim;
&buf[off..off + head_dim]
}
let mut out = vec![0.0f32; n];
for h in 0..num_heads {
let mut scores = vec![0.0f32; seq_len * seq_len];
for sq in 0..seq_len {
for sk in 0..seq_len {
let mut acc = 0.0f64;
let qi = row(q, sq, h, num_heads, head_dim);
let ki = row(k, sk, h, num_heads, head_dim);
for d in 0..head_dim {
acc += qi[d] as f64 * ki[d] as f64;
}
scores[sq * seq_len + sk] = (acc as f32) * scale;
}
}
for sq in 0..seq_len {
let row_off = sq * seq_len;
let mut m = scores[row_off];
for sk in 1..seq_len {
m = m.max(scores[row_off + sk]);
}
let mut sum = 0.0f64;
for sk in 0..seq_len {
let e = ((scores[row_off + sk] - m) as f64).exp();
scores[row_off + sk] = e as f32;
sum += e;
}
let inv = (1.0 / sum) as f32;
for sk in 0..seq_len {
scores[row_off + sk] *= inv;
}
}
for sq in 0..seq_len {
let out_off = (sq * num_heads + h) * head_dim;
for d in 0..head_dim {
let mut acc = 0.0f64;
for sk in 0..seq_len {
let vi = row(v, sk, h, num_heads, head_dim);
acc += (scores[sq * seq_len + sk] as f64) * (vi[d] as f64);
}
out[out_off + d] = acc as f32;
}
}
}
out
}
#[cfg(test)]
#[allow(clippy::too_many_arguments)]
fn apply_bert_encoder_block_cpu_ref(
input: &[f32],
q_w: &[f32],
q_b: Option<&[f32]>,
k_w: &[f32],
k_b: Option<&[f32]>,
v_w: &[f32],
v_b: Option<&[f32]>,
o_w: &[f32],
o_b: Option<&[f32]>,
attn_gamma: &[f32],
attn_beta: &[f32],
up_w: &[f32],
up_b: Option<&[f32]>,
down_w: &[f32],
down_b: Option<&[f32]>,
ffn_gamma: &[f32],
ffn_beta: &[f32],
seq_len: usize,
hidden: usize,
num_heads: usize,
intermediate: usize,
eps: f32,
) -> Vec<f32> {
let head_dim = hidden / num_heads;
let scale = 1.0_f32 / (head_dim as f32).sqrt();
let q = bert_linear_cpu_ref(input, q_w, q_b, seq_len, hidden, hidden);
let k = bert_linear_cpu_ref(input, k_w, k_b, seq_len, hidden, hidden);
let v = bert_linear_cpu_ref(input, v_w, v_b, seq_len, hidden, hidden);
let attn = bert_attention_cpu_ref(&q, &k, &v, seq_len, num_heads, head_dim, scale);
let o = bert_linear_cpu_ref(&attn, o_w, o_b, seq_len, hidden, hidden);
let n_hidden = seq_len * hidden;
let mut residual = vec![0.0f32; n_hidden];
for i in 0..n_hidden {
residual[i] = input[i] + o[i];
}
let post_attn = bert_layer_norm_cpu_ref(&residual, attn_gamma, attn_beta, eps, seq_len, hidden);
let ffn_up = bert_linear_cpu_ref(&post_attn, up_w, up_b, seq_len, hidden, intermediate);
let ffn_act = bert_gelu_cpu_ref(&ffn_up);
let ffn_down = bert_linear_cpu_ref(&ffn_act, down_w, down_b, seq_len, intermediate, hidden);
let mut residual2 = vec![0.0f32; n_hidden];
for i in 0..n_hidden {
residual2[i] = post_attn[i] + ffn_down[i];
}
bert_layer_norm_cpu_ref(&residual2, ffn_gamma, ffn_beta, eps, seq_len, hidden)
}
#[cfg(test)]
fn bert_layer_norm_cpu_ref(
input: &[f32],
gamma: &[f32],
beta: &[f32],
eps: f32,
batch: usize,
hidden: usize,
) -> Vec<f32> {
assert_eq!(input.len(), batch * hidden);
assert_eq!(gamma.len(), hidden);
assert_eq!(beta.len(), hidden);
let mut out = vec![0.0f32; batch * hidden];
for r in 0..batch {
let row = &input[r * hidden..(r + 1) * hidden];
let mean: f32 = row.iter().sum::<f32>() / hidden as f32;
let var: f32 = row.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / hidden as f32;
let inv_std = 1.0 / (var + eps).sqrt();
for h in 0..hidden {
out[r * hidden + h] = (row[h] - mean) * inv_std * gamma[h] + beta[h];
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use mlx_native::GraphExecutor;
fn upload_f32(device: &MlxDevice, data: &[f32], shape: Vec<usize>) -> MlxBuffer {
let bytes = data.len() * 4;
let buf = device.alloc_buffer(bytes, DType::F32, shape).unwrap();
let slice: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut f32, data.len()) };
slice.copy_from_slice(data);
buf
}
fn readback_f32(buf: &MlxBuffer, expected_len: usize) -> Vec<f32> {
let slice: &[f32] = buf.as_slice::<f32>().expect("readback as_slice");
assert_eq!(slice.len(), expected_len, "readback length mismatch");
slice.to_vec()
}
fn run_layer_norm(
input_data: &[f32],
gamma_data: &[f32],
beta_data: &[f32],
eps: f32,
batch: usize,
hidden: usize,
) -> Vec<f32> {
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let input = upload_f32(device, input_data, vec![batch, hidden]);
let gamma = upload_f32(device, gamma_data, vec![hidden]);
let beta = upload_f32(device, beta_data, vec![hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let output = bert_layer_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&input,
&gamma,
&beta,
eps,
batch as u32,
hidden as u32,
)
.expect("gpu dispatch");
session.finish().expect("finish");
readback_f32(&output, batch * hidden)
}
#[test]
fn prev_pow2_table() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
assert_eq!(prev_pow2(1), 1);
assert_eq!(prev_pow2(2), 2);
assert_eq!(prev_pow2(3), 2);
assert_eq!(prev_pow2(255), 128);
assert_eq!(prev_pow2(256), 256);
assert_eq!(prev_pow2(257), 256);
assert_eq!(prev_pow2(384), 256);
assert_eq!(prev_pow2(1024), 1024);
}
#[test]
fn cpu_ref_matches_known_value() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let out = bert_layer_norm_cpu_ref(
&[1.0, 2.0, 3.0, 4.0],
&[1.0, 1.0, 1.0, 1.0],
&[0.0, 0.0, 0.0, 0.0],
0.0,
1,
4,
);
let inv_std = 1.0 / 1.25f32.sqrt();
for (got, expected) in out.iter().zip([
(1.0 - 2.5) * inv_std,
(2.0 - 2.5) * inv_std,
(3.0 - 2.5) * inv_std,
(4.0 - 2.5) * inv_std,
]) {
assert!(
(got - expected).abs() < 1e-6,
"got {got}, expected {expected}"
);
}
}
#[test]
fn gpu_matches_cpu_on_synthetic_small_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let batch = 4usize;
let hidden = 8usize;
let input: Vec<f32> = (0..batch * hidden)
.map(|i| 0.1 * (i as f32) - 0.4)
.collect();
let gamma: Vec<f32> = (0..hidden).map(|i| 1.0 + 0.05 * i as f32).collect();
let beta: Vec<f32> = (0..hidden).map(|i| 0.01 * i as f32).collect();
let eps = 1e-12;
let cpu = bert_layer_norm_cpu_ref(&input, &gamma, &beta, eps, batch, hidden);
let gpu = run_layer_norm(&input, &gamma, &beta, eps, batch, hidden);
for (i, (g, c)) in gpu.iter().zip(cpu.iter()).enumerate() {
assert!(
(g - c).abs() < 1e-5,
"row {} col {}: gpu={} cpu={} diff={}",
i / hidden,
i % hidden,
g,
c,
(g - c).abs()
);
}
}
#[test]
fn gpu_constant_input_yields_bias_only_output() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let batch = 2usize;
let hidden = 16usize;
let input = vec![3.5f32; batch * hidden];
let gamma = vec![2.0f32; hidden]; let beta: Vec<f32> = (0..hidden).map(|i| 0.1 * i as f32 - 0.7).collect();
let eps = 1e-5;
let gpu = run_layer_norm(&input, &gamma, &beta, eps, batch, hidden);
for r in 0..batch {
for h in 0..hidden {
let got = gpu[r * hidden + h];
let want = beta[h];
assert!(
(got - want).abs() < 1e-6,
"row {} col {}: got {} want {}",
r,
h,
got,
want
);
}
}
}
#[test]
fn gpu_matches_cpu_at_bge_small_hidden_384() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let batch = 32usize; let hidden = 384usize; let input: Vec<f32> = (0..batch * hidden)
.map(|i| ((i.wrapping_mul(2654435761) % 1000) as f32) * 0.001 - 0.5)
.collect();
let gamma: Vec<f32> = (0..hidden).map(|i| 1.0 - 0.001 * i as f32).collect();
let beta: Vec<f32> = (0..hidden).map(|i| 0.0001 * i as f32).collect();
let eps = 1e-12;
let cpu = bert_layer_norm_cpu_ref(&input, &gamma, &beta, eps, batch, hidden);
let gpu = run_layer_norm(&input, &gamma, &beta, eps, batch, hidden);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(max_diff < 1e-4, "max_diff at bge-small shape: {max_diff}");
}
#[test]
fn gpu_matches_cpu_at_mxbai_large_hidden_1024() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let batch = 8usize;
let hidden = 1024usize; let input: Vec<f32> = (0..batch * hidden)
.map(|i| ((i.wrapping_mul(2246822519) % 700) as f32) * 0.001 - 0.35)
.collect();
let gamma: Vec<f32> = (0..hidden).map(|i| 0.5 + 0.001 * i as f32).collect();
let beta = vec![0.0f32; hidden];
let eps = 1e-12;
let cpu = bert_layer_norm_cpu_ref(&input, &gamma, &beta, eps, batch, hidden);
let gpu = run_layer_norm(&input, &gamma, &beta, eps, batch, hidden);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(max_diff < 2e-4, "max_diff at mxbai shape: {max_diff}");
}
#[test]
fn gpu_rejects_zero_dimensions() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let device = match MlxDevice::new() {
Ok(d) => d,
Err(e) => {
eprintln!("skipping: no Metal device available: {e}");
return;
}
};
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let buf = upload_f32(device, &[0.0; 4], vec![1, 4]);
let err = bert_layer_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&buf,
&buf,
&buf,
1e-12,
0,
4,
);
assert!(err.is_err(), "batch=0 must error");
let err = bert_layer_norm_gpu(
session.encoder_mut(),
&mut registry,
device,
&buf,
&buf,
&buf,
1e-12,
4,
0,
);
assert!(err.is_err(), "hidden=0 must error");
drop(session);
}
fn run_linear(
input_data: &[f32],
weight_data: &[f32],
bias_data: Option<&[f32]>,
seq_len: usize,
in_features: usize,
out_features: usize,
) -> Vec<f32> {
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let input = upload_f32(device, input_data, vec![seq_len, in_features]);
let weight = upload_f32(device, weight_data, vec![out_features, in_features]);
let bias_buf = bias_data.map(|b| upload_f32(device, b, vec![out_features]));
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let output = bert_linear_gpu(
session.encoder_mut(),
&mut registry,
device,
&input,
&weight,
bias_buf.as_ref(),
seq_len as u32,
in_features as u32,
out_features as u32,
)
.expect("gpu dispatch");
session.finish().expect("finish");
readback_f32(&output, seq_len * out_features)
}
#[test]
fn linear_no_bias_matches_cpu_on_small_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 4usize;
let in_f = 64usize; let out_f = 32usize;
let input: Vec<f32> = (0..seq * in_f).map(|i| 0.01 * (i as f32) - 0.3).collect();
let weight: Vec<f32> = (0..out_f * in_f)
.map(|i| 0.005 * (i as f32) - 0.4)
.collect();
let cpu = bert_linear_cpu_ref(&input, &weight, None, seq, in_f, out_f);
let gpu = run_linear(&input, &weight, None, seq, in_f, out_f);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(max_diff < 0.20, "max_diff {} > 0.20", max_diff);
}
#[test]
fn linear_with_bias_matches_cpu_on_small_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 4usize;
let in_f = 64usize;
let out_f = 32usize;
let input: Vec<f32> = (0..seq * in_f).map(|i| 0.01 * (i as f32) - 0.3).collect();
let weight: Vec<f32> = (0..out_f * in_f)
.map(|i| 0.005 * (i as f32) - 0.4)
.collect();
let bias: Vec<f32> = (0..out_f).map(|i| 0.1 * i as f32 - 1.0).collect();
let cpu = bert_linear_cpu_ref(&input, &weight, Some(&bias), seq, in_f, out_f);
let gpu = run_linear(&input, &weight, Some(&bias), seq, in_f, out_f);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(max_diff < 0.20, "max_diff {} > 0.20 (bias path)", max_diff);
}
#[test]
fn linear_at_bge_small_qkv_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 32usize;
let hidden = 384usize;
let input: Vec<f32> = (0..seq * hidden)
.map(|i| ((i.wrapping_mul(2654435761) % 1000) as f32) * 0.001 - 0.5)
.collect();
let weight: Vec<f32> = (0..hidden * hidden)
.map(|i| ((i.wrapping_mul(40503) % 600) as f32) * 0.001 - 0.3)
.collect();
let bias: Vec<f32> = (0..hidden).map(|i| 0.0001 * i as f32).collect();
let cpu = bert_linear_cpu_ref(&input, &weight, Some(&bias), seq, hidden, hidden);
let gpu = run_linear(&input, &weight, Some(&bias), seq, hidden, hidden);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(
max_diff < 5e-2,
"max_diff {} > 0.05 at bge-small QKV shape",
max_diff
);
}
#[test]
fn linear_rejects_in_features_below_32() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let input = upload_f32(device, &[0.0; 16], vec![1, 16]);
let weight = upload_f32(device, &[0.0; 16 * 8], vec![8, 16]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let err = bert_linear_gpu(
session.encoder_mut(),
&mut registry,
device,
&input,
&weight,
None,
1,
16, 8,
);
assert!(err.is_err(), "in_features=16 must error");
drop(session);
}
#[test]
fn linear_rejects_bias_size_mismatch() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let input = upload_f32(device, &[0.0; 64], vec![1, 64]);
let weight = upload_f32(device, &[0.0; 64 * 8], vec![8, 64]);
let bias_wrong = upload_f32(device, &[0.0; 4], vec![4]); let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let err = bert_linear_gpu(
session.encoder_mut(),
&mut registry,
device,
&input,
&weight,
Some(&bias_wrong),
1,
64,
8,
);
assert!(err.is_err(), "bias size mismatch must error");
drop(session);
}
fn run_gelu(input_data: &[f32]) -> Vec<f32> {
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let input = upload_f32(device, input_data, vec![input_data.len()]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let output =
bert_gelu_gpu(session.encoder_mut(), &mut registry, device, &input).expect("gpu");
session.finish().expect("finish");
readback_f32(&output, input_data.len())
}
#[test]
fn gelu_cpu_ref_known_values() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let out = bert_gelu_cpu_ref(&[0.0, 1.0, -1.0]);
assert!(out[0].abs() < 1e-6);
assert!((out[1] - 0.8411920071).abs() < 1e-4);
assert!((out[2] - -0.1588079929).abs() < 1e-4);
}
#[test]
fn gelu_gpu_matches_cpu_small_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let input: Vec<f32> = (-32..32).map(|i| 0.1 * i as f32).collect();
let cpu = bert_gelu_cpu_ref(&input);
let gpu = run_gelu(&input);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(max_diff < 1e-5, "gelu max_diff: {max_diff}");
}
#[test]
fn gelu_gpu_matches_cpu_at_bge_small_ffn_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let n: usize = 32 * 1536;
let input: Vec<f32> = (0..n)
.map(|i| ((i.wrapping_mul(2654435761usize) % 1000) as f32) * 0.005 - 2.5)
.collect();
let cpu = bert_gelu_cpu_ref(&input);
let gpu = run_gelu(&input);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(
max_diff < 1e-5,
"gelu max_diff at bge FFN shape: {max_diff}"
);
}
#[test]
fn bias_add_gpu_matches_cpu_small() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let rows = 5usize;
let cols = 7usize;
let input: Vec<f32> = (0..rows * cols).map(|i| 0.01 * i as f32).collect();
let bias: Vec<f32> = (0..cols).map(|i| 0.1 * i as f32 - 0.3).collect();
let inp_buf = upload_f32(device, &input, vec![rows, cols]);
let bias_buf = upload_f32(device, &bias, vec![cols]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out_buf = bert_bias_add_gpu(
session.encoder_mut(),
&mut registry,
device,
&inp_buf,
&bias_buf,
rows as u32,
cols as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out_buf, rows * cols);
for r in 0..rows {
for c in 0..cols {
let expected = input[r * cols + c] + bias[c];
let actual = got[r * cols + c];
assert!(
(expected - actual).abs() < 1e-6,
"row {} col {}: got {} expected {}",
r,
c,
actual,
expected
);
}
}
}
fn run_attention(
q_data: &[f32],
k_data: &[f32],
v_data: &[f32],
seq_len: usize,
num_heads: usize,
head_dim: usize,
scale: f32,
) -> Vec<f32> {
let device = MlxDevice::new().expect("device");
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let shape = vec![seq_len, num_heads, head_dim];
let q = upload_f32(device, q_data, shape.clone());
let k = upload_f32(device, k_data, shape.clone());
let v = upload_f32(device, v_data, shape);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_attention_gpu(
session.encoder_mut(),
&mut registry,
device,
&q,
&k,
&v,
seq_len as u32,
num_heads as u32,
head_dim as u32,
scale,
)
.expect("dispatch");
session.finish().expect("finish");
readback_f32(&out, seq_len * num_heads * head_dim)
}
#[test]
fn cpu_ref_attention_simple_softmax() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
let q = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0];
let k = q.clone();
let v = vec![10.0, 20.0, 30.0, 40.0, 50.0, 60.0, 70.0, 80.0];
let scale = 100.0; let out = bert_attention_cpu_ref(&q, &k, &v, 2, 1, 4, scale);
for d in 0..4 {
assert!(
(out[d] - v[d]).abs() < 1e-3,
"row 0 d={d}: got {}, want {}",
out[d],
v[d]
);
assert!(
(out[4 + d] - v[4 + d]).abs() < 1e-3,
"row 1 d={d}: got {}, want {}",
out[4 + d],
v[4 + d]
);
}
}
#[test]
fn attention_gpu_matches_cpu_at_synthetic_small_input() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 32usize;
let num_heads = 1usize;
let head_dim = 32usize; let n = seq * num_heads * head_dim;
let q: Vec<f32> = (0..n).map(|i| 0.05 * (i as f32) - 0.5).collect();
let k: Vec<f32> = (0..n).map(|i| 0.04 * (i as f32) - 0.4).collect();
let v: Vec<f32> = (0..n).map(|i| 0.03 * (i as f32) - 0.3).collect();
let scale = 1.0 / (head_dim as f32).sqrt();
let cpu = bert_attention_cpu_ref(&q, &k, &v, seq, num_heads, head_dim, scale);
let gpu = run_attention(&q, &k, &v, seq, num_heads, head_dim, scale);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(max_diff < 0.20, "max_diff at synthetic small: {max_diff}");
}
#[test]
fn attention_gpu_matches_cpu_at_bge_small_attention_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 32usize;
let num_heads = 12usize;
let head_dim = 32usize;
let n = seq * num_heads * head_dim;
let q: Vec<f32> = (0..n)
.map(|i| ((i.wrapping_mul(2654435761usize) % 1000) as f32) * 0.001 - 0.5)
.collect();
let k: Vec<f32> = (0..n)
.map(|i| ((i.wrapping_mul(40503usize) % 700) as f32) * 0.001 - 0.35)
.collect();
let v: Vec<f32> = (0..n)
.map(|i| ((i.wrapping_mul(2246822519usize) % 800) as f32) * 0.001 - 0.4)
.collect();
let scale = 1.0 / (head_dim as f32).sqrt();
let cpu = bert_attention_cpu_ref(&q, &k, &v, seq, num_heads, head_dim, scale);
let gpu = run_attention(&q, &k, &v, seq, num_heads, head_dim, scale);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(max_diff < 0.10, "max_diff at bge-small attn: {max_diff}");
}
#[test]
fn attention_gpu_rejects_small_head_dim() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let buf = upload_f32(device, &[0.0; 16], vec![1, 1, 16]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let err = bert_attention_gpu(
session.encoder_mut(),
&mut registry,
device,
&buf,
&buf,
&buf,
1,
1,
16, 1.0,
);
assert!(err.is_err(), "head_dim=16 must error");
drop(session);
}
#[test]
fn encoder_block_gpu_matches_cpu_ref_at_minimal_shape() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 32usize;
let hidden = 64usize;
let num_heads = 2usize;
let intermediate = 128usize;
let eps = 1e-12f32;
let prand = |seed: usize, n: usize, scale: f32, offset: f32| -> Vec<f32> {
(0..n)
.map(|i| {
((i.wrapping_mul(2654435761usize).wrapping_add(seed) % 1000) as f32) * scale
+ offset
})
.collect()
};
let input = prand(1, seq * hidden, 0.001, -0.5);
let q_w = prand(2, hidden * hidden, 0.0005, -0.15);
let q_b = prand(3, hidden, 0.001, -0.05);
let k_w = prand(4, hidden * hidden, 0.0005, -0.15);
let k_b = prand(5, hidden, 0.001, -0.05);
let v_w = prand(6, hidden * hidden, 0.0005, -0.15);
let v_b = prand(7, hidden, 0.001, -0.05);
let o_w = prand(8, hidden * hidden, 0.0005, -0.15);
let o_b = prand(9, hidden, 0.001, -0.05);
let attn_gamma = prand(10, hidden, 0.0001, 1.0); let attn_beta = prand(11, hidden, 0.001, -0.5);
let up_w = prand(12, intermediate * hidden, 0.0005, -0.15);
let up_b = prand(13, intermediate, 0.001, -0.5);
let down_w = prand(14, hidden * intermediate, 0.0005, -0.15);
let down_b = prand(15, hidden, 0.001, -0.5);
let ffn_gamma = prand(16, hidden, 0.0001, 1.0);
let ffn_beta = prand(17, hidden, 0.001, -0.5);
let cpu_out = apply_bert_encoder_block_cpu_ref(
&input,
&q_w,
Some(&q_b),
&k_w,
Some(&k_b),
&v_w,
Some(&v_b),
&o_w,
Some(&o_b),
&attn_gamma,
&attn_beta,
&up_w,
Some(&up_b),
&down_w,
Some(&down_b),
&ffn_gamma,
&ffn_beta,
seq,
hidden,
num_heads,
intermediate,
eps,
);
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let in_buf = upload_f32(device, &input, vec![seq, hidden]);
let q_w_b = upload_f32(device, &q_w, vec![hidden, hidden]);
let q_b_b = upload_f32(device, &q_b, vec![hidden]);
let k_w_b = upload_f32(device, &k_w, vec![hidden, hidden]);
let k_b_b = upload_f32(device, &k_b, vec![hidden]);
let v_w_b = upload_f32(device, &v_w, vec![hidden, hidden]);
let v_b_b = upload_f32(device, &v_b, vec![hidden]);
let o_w_b = upload_f32(device, &o_w, vec![hidden, hidden]);
let o_b_b = upload_f32(device, &o_b, vec![hidden]);
let attn_gamma_b = upload_f32(device, &attn_gamma, vec![hidden]);
let attn_beta_b = upload_f32(device, &attn_beta, vec![hidden]);
let up_w_b = upload_f32(device, &up_w, vec![intermediate, hidden]);
let up_b_b = upload_f32(device, &up_b, vec![intermediate]);
let down_w_b = upload_f32(device, &down_w, vec![hidden, intermediate]);
let down_b_b = upload_f32(device, &down_b, vec![hidden]);
let ffn_gamma_b = upload_f32(device, &ffn_gamma, vec![hidden]);
let ffn_beta_b = upload_f32(device, &ffn_beta, vec![hidden]);
let tensors = BertEncoderBlockTensors {
q_w: &q_w_b,
q_b: Some(&q_b_b),
k_w: &k_w_b,
k_b: Some(&k_b_b),
v_w: &v_w_b,
v_b: Some(&v_b_b),
o_w: &o_w_b,
o_b: Some(&o_b_b),
attn_norm_gamma: &attn_gamma_b,
attn_norm_beta: &attn_beta_b,
up_w: &up_w_b,
up_b: Some(&up_b_b),
down_w: &down_w_b,
down_b: Some(&down_b_b),
ffn_norm_gamma: &ffn_gamma_b,
ffn_norm_beta: &ffn_beta_b,
};
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out_buf = apply_bert_encoder_block_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
&tensors,
None, seq as u32,
hidden as u32,
num_heads as u32,
intermediate as u32,
eps,
)
.expect("encoder block dispatch");
session.finish().expect("finish");
let gpu_out = readback_f32(&out_buf, seq * hidden);
let mut max_diff = 0.0f32;
let mut argmax = 0usize;
for (i, (g, c)) in gpu_out.iter().zip(cpu_out.iter()).enumerate() {
let d = (g - c).abs();
if d > max_diff {
max_diff = d;
argmax = i;
}
}
assert!(
max_diff < 0.50,
"encoder-block max_diff {} > 0.50 at i={} (gpu={}, cpu={})",
max_diff,
argmax,
gpu_out[argmax],
cpu_out[argmax]
);
}
fn upload_u32(device: &MlxDevice, data: &[u32], shape: Vec<usize>) -> MlxBuffer {
let bytes = data.len() * 4;
let buf = device.alloc_buffer(bytes, DType::U32, shape).unwrap();
let slice: &mut [u32] =
unsafe { std::slice::from_raw_parts_mut(buf.contents_ptr() as *mut u32, data.len()) };
slice.copy_from_slice(data);
buf
}
#[test]
fn embed_gather_gpu_picks_correct_rows() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let vocab = 5usize;
let hidden = 8usize;
let table: Vec<f32> = (0..vocab * hidden)
.map(|i| ((i / hidden) * 100 + (i % hidden)) as f32)
.collect();
let ids: Vec<u32> = vec![3, 0, 4, 1, 2];
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let table_buf = upload_f32(device, &table, vec![vocab, hidden]);
let ids_buf = upload_u32(device, &ids, vec![ids.len()]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_embed_gather_gpu(
session.encoder_mut(),
&mut registry,
device,
&table_buf,
&ids_buf,
vocab as u32,
hidden as u32,
ids.len() as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, ids.len() * hidden);
for (i, &id) in ids.iter().enumerate() {
for h in 0..hidden {
let expected = (id as usize * 100 + h) as f32;
let actual = got[i * hidden + h];
assert!(
(expected - actual).abs() < 1e-6,
"row {} h {}: got {}, want {}",
i,
h,
actual,
expected
);
}
}
}
fn embeddings_cpu_ref(
input_ids: &[u32],
type_ids: Option<&[u32]>,
token_embd: &[f32],
position_embd: &[f32],
token_types: Option<&[f32]>,
embed_gamma: &[f32],
embed_beta: &[f32],
eps: f32,
seq_len: usize,
hidden: usize,
) -> Vec<f32> {
let mut summed = vec![0.0f32; seq_len * hidden];
for s in 0..seq_len {
let tid = input_ids[s] as usize;
for h in 0..hidden {
let v = token_embd[tid * hidden + h] + position_embd[s * hidden + h];
let v = if let (Some(tids), Some(tt)) = (type_ids, token_types) {
let typ = tids[s] as usize;
v + tt[typ * hidden + h]
} else {
v
};
summed[s * hidden + h] = v;
}
}
bert_layer_norm_cpu_ref(&summed, embed_gamma, embed_beta, eps, seq_len, hidden)
}
#[test]
fn embeddings_gpu_matches_cpu_with_token_types() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 32usize; let hidden = 64usize;
let vocab = 100usize;
let max_pos = 128usize;
let type_vocab = 2usize;
let eps = 1e-12f32;
let prand_f32 = |seed: usize, n: usize| -> Vec<f32> {
(0..n)
.map(|i| {
((i.wrapping_mul(2654435761usize).wrapping_add(seed) % 1000) as f32) * 0.001
- 0.5
})
.collect()
};
let token_embd = prand_f32(1, vocab * hidden);
let position_embd = prand_f32(2, max_pos * hidden);
let token_types = prand_f32(3, type_vocab * hidden);
let embed_gamma = prand_f32(4, hidden);
let embed_beta = prand_f32(5, hidden);
let input_ids: Vec<u32> = (0..seq)
.map(|i| (i.wrapping_mul(7) % vocab) as u32)
.collect();
let type_ids: Vec<u32> = (0..seq).map(|i| (i & 1) as u32).collect();
let cpu = embeddings_cpu_ref(
&input_ids,
Some(&type_ids),
&token_embd,
&position_embd,
Some(&token_types),
&embed_gamma,
&embed_beta,
eps,
seq,
hidden,
);
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let token_embd_b = upload_f32(device, &token_embd, vec![vocab, hidden]);
let position_embd_b = upload_f32(device, &position_embd, vec![max_pos, hidden]);
let token_types_b = upload_f32(device, &token_types, vec![type_vocab, hidden]);
let embed_gamma_b = upload_f32(device, &embed_gamma, vec![hidden]);
let embed_beta_b = upload_f32(device, &embed_beta, vec![hidden]);
let input_ids_b = upload_u32(device, &input_ids, vec![seq]);
let type_ids_b = upload_u32(device, &type_ids, vec![seq]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_embeddings_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_ids_b,
Some(&type_ids_b),
&token_embd_b,
&position_embd_b,
Some(&token_types_b),
&embed_gamma_b,
&embed_beta_b,
eps,
seq as u32,
hidden as u32,
vocab as u32,
max_pos as u32,
type_vocab as u32,
)
.expect("embeddings dispatch");
session.finish().expect("finish");
let gpu = readback_f32(&out, seq * hidden);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(max_diff < 1e-4, "embeddings max_diff: {max_diff}");
}
#[test]
fn embeddings_gpu_without_token_types_path_works() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 32usize;
let hidden = 64usize;
let vocab = 50usize;
let max_pos = 64usize;
let eps = 1e-12f32;
let token_embd: Vec<f32> = (0..vocab * hidden).map(|i| (i as f32) * 0.001).collect();
let position_embd: Vec<f32> = (0..max_pos * hidden).map(|i| (i as f32) * 0.0005).collect();
let embed_gamma = vec![1.0f32; hidden];
let embed_beta = vec![0.0f32; hidden];
let input_ids: Vec<u32> = (0..seq).map(|i| (i % vocab) as u32).collect();
let cpu = embeddings_cpu_ref(
&input_ids,
None,
&token_embd,
&position_embd,
None,
&embed_gamma,
&embed_beta,
eps,
seq,
hidden,
);
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let token_embd_b = upload_f32(device, &token_embd, vec![vocab, hidden]);
let position_embd_b = upload_f32(device, &position_embd, vec![max_pos, hidden]);
let embed_gamma_b = upload_f32(device, &embed_gamma, vec![hidden]);
let embed_beta_b = upload_f32(device, &embed_beta, vec![hidden]);
let input_ids_b = upload_u32(device, &input_ids, vec![seq]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_embeddings_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_ids_b,
None,
&token_embd_b,
&position_embd_b,
None,
&embed_gamma_b,
&embed_beta_b,
eps,
seq as u32,
hidden as u32,
vocab as u32,
max_pos as u32,
0,
)
.expect("embeddings dispatch (no token_types)");
session.finish().expect("finish");
let gpu = readback_f32(&out, seq * hidden);
let mut max_diff = 0.0f32;
for (g, c) in gpu.iter().zip(cpu.iter()) {
max_diff = max_diff.max((g - c).abs());
}
assert!(
max_diff < 1e-4,
"embeddings (no types) max_diff: {max_diff}"
);
}
#[test]
fn embeddings_rejects_inconsistent_token_types_args() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let buf = upload_f32(device, &[0.0; 64], vec![64]);
let id_buf = upload_u32(device, &[0; 32], vec![32]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let err = bert_embeddings_gpu(
session.encoder_mut(),
&mut registry,
device,
&id_buf,
Some(&id_buf), &buf,
&buf,
None, &buf,
&buf,
1e-12,
32,
1,
10,
64,
2,
);
assert!(err.is_err(), "type_ids without token_types must error");
drop(session);
}
#[test]
fn pool_mean_gpu_matches_cpu_average() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 5usize;
let hidden = 8usize;
let input: Vec<f32> = (0..seq * hidden)
.map(|i| (i / hidden) as f32 + (i % hidden) as f32 * 0.1)
.collect();
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let in_buf = upload_f32(device, &input, vec![seq, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_pool_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
BertPoolKind::Mean,
seq as u32,
hidden as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, hidden);
for h in 0..hidden {
let mut expected = 0.0f32;
for s in 0..seq {
expected += s as f32 + h as f32 * 0.1;
}
expected /= seq as f32;
assert!(
(got[h] - expected).abs() < 1e-5,
"h={}: got {}, want {}",
h,
got[h],
expected
);
}
}
#[test]
fn pool_cls_gpu_returns_first_row() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 5usize;
let hidden = 8usize;
let input: Vec<f32> = (0..seq * hidden).map(|i| i as f32).collect();
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let in_buf = upload_f32(device, &input, vec![seq, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_pool_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
BertPoolKind::Cls,
seq as u32,
hidden as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, hidden);
for h in 0..hidden {
assert!(
(got[h] - h as f32).abs() < 1e-6,
"CLS h={}: got {}, want {}",
h,
got[h],
h
);
}
}
#[test]
fn pool_last_gpu_returns_last_row() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 5usize;
let hidden = 8usize;
let input: Vec<f32> = (0..seq * hidden).map(|i| i as f32).collect();
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let in_buf = upload_f32(device, &input, vec![seq, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_pool_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
BertPoolKind::Last,
seq as u32,
hidden as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, hidden);
for h in 0..hidden {
let expected = ((seq - 1) * hidden + h) as f32;
assert!(
(got[h] - expected).abs() < 1e-6,
"Last h={}: got {}, want {}",
h,
got[h],
expected
);
}
}
#[test]
fn pool_last_padded_returns_valid_last_row_not_seq_len_minus_one() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq_len: usize = 8;
let valid: usize = 5;
let hidden: usize = 4;
let input: Vec<f32> = (0..seq_len * hidden).map(|i| i as f32).collect();
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let in_buf = upload_f32(device, &input, vec![seq_len, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_pool_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
BertPoolKind::Last,
valid as u32,
hidden as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, hidden);
for h in 0..hidden {
let expected = ((valid - 1) * hidden + h) as f32;
assert!(
(got[h] - expected).abs() < 1e-6,
"Last padded h={h}: got {}, want {} (row {}); \
if got {} that is the buggy seq_len-1 row",
got[h],
expected,
valid - 1,
((seq_len - 1) * hidden + h) as f32,
);
}
}
#[test]
fn pool_mean_padded_divides_by_valid_not_seq_len() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq_len: usize = 8;
let valid: usize = 5;
let hidden: usize = 4;
let mut input = vec![0.0f32; seq_len * hidden];
for s in 0..valid {
for h in 0..hidden {
input[s * hidden + h] = (s * hidden + h) as f32;
}
}
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let in_buf = upload_f32(device, &input, vec![seq_len, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_pool_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
BertPoolKind::Mean,
valid as u32,
hidden as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, hidden);
for h in 0..hidden {
let sum: f32 = (0..valid).map(|s| (s * hidden + h) as f32).sum();
let expected = sum / valid as f32;
let buggy = sum / seq_len as f32;
assert!(
(got[h] - expected).abs() < 1e-5,
"Mean padded h={h}: got {:.6}, want {:.6} (buggy would be {:.6})",
got[h],
expected,
buggy,
);
}
}
#[test]
fn pool_cls_padded_still_returns_row_zero() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq_len: usize = 8;
let valid: usize = 5;
let hidden: usize = 4;
let input: Vec<f32> = (0..seq_len * hidden).map(|i| i as f32).collect();
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let in_buf = upload_f32(device, &input, vec![seq_len, hidden]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_pool_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
BertPoolKind::Cls,
valid as u32, hidden as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, hidden);
for h in 0..hidden {
let expected = h as f32; assert!(
(got[h] - expected).abs() < 1e-6,
"Cls padded h={h}: got {}, want {}",
got[h],
expected,
);
}
}
#[test]
fn l2_normalize_gpu_produces_unit_norm() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let rows = 1usize;
let dim = 384usize; let input: Vec<f32> = (0..rows * dim).map(|i| (i as f32) * 0.01 - 1.0).collect();
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let in_buf = upload_f32(device, &input, vec![rows, dim]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = bert_l2_normalize_gpu(
session.encoder_mut(),
&mut registry,
device,
&in_buf,
1e-12,
rows as u32,
dim as u32,
)
.expect("dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, rows * dim);
for r in 0..rows {
let mut sum_sq = 0.0f64;
for d in 0..dim {
let v = got[r * dim + d] as f64;
sum_sq += v * v;
}
let norm = sum_sq.sqrt() as f32;
assert!((norm - 1.0).abs() < 1e-4, "row {} norm {} != 1.0", r, norm);
}
}
#[test]
fn full_forward_gpu_produces_unit_norm_output_at_minimal_config() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let seq = 32usize;
let hidden = 64usize;
let num_heads = 2usize;
let intermediate = 128usize;
let vocab = 100usize;
let max_pos = 64usize;
let type_vocab = 2usize;
let num_layers = 2usize;
let cfg = super::super::config::BertConfig {
hidden_size: hidden,
num_attention_heads: num_heads,
num_hidden_layers: num_layers,
intermediate_size: intermediate,
max_position_embeddings: max_pos,
vocab_size: vocab,
type_vocab_size: type_vocab,
layer_norm_eps: 1e-12,
hidden_act: "gelu".into(),
pooling_type: super::super::config::PoolingType::Mean,
causal_attention: false,
};
let prand = |seed: usize, n: usize, scale: f32, offset: f32| -> Vec<f32> {
(0..n)
.map(|i| {
((i.wrapping_mul(2654435761usize).wrapping_add(seed) % 1000) as f32) * scale
+ offset
})
.collect()
};
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let mut tensors: std::collections::HashMap<String, MlxBuffer> =
std::collections::HashMap::new();
tensors.insert(
super::super::TENSOR_TOKEN_EMBD.into(),
upload_f32(
device,
&prand(1, vocab * hidden, 0.001, -0.5),
vec![vocab, hidden],
),
);
tensors.insert(
super::super::TENSOR_POS_EMBD.into(),
upload_f32(
device,
&prand(2, max_pos * hidden, 0.001, -0.5),
vec![max_pos, hidden],
),
);
tensors.insert(
super::super::TENSOR_TOKEN_TYPES.into(),
upload_f32(
device,
&prand(3, type_vocab * hidden, 0.001, -0.5),
vec![type_vocab, hidden],
),
);
tensors.insert(
super::super::TENSOR_EMBED_NORM_WEIGHT.into(),
upload_f32(device, &prand(4, hidden, 0.0001, 1.0), vec![hidden]),
);
tensors.insert(
super::super::TENSOR_EMBED_NORM_BIAS.into(),
upload_f32(device, &prand(5, hidden, 0.001, -0.5), vec![hidden]),
);
for layer in 0..num_layers {
let key = |s: &str| -> String { super::super::config::bert_layer_tensor(layer, s) };
let qkvo_seed_base = 100 + layer * 100;
for (i, name) in [
"attn_q.weight",
"attn_k.weight",
"attn_v.weight",
"attn_output.weight",
]
.iter()
.enumerate()
{
tensors.insert(
key(name),
upload_f32(
device,
&prand(qkvo_seed_base + i * 10, hidden * hidden, 0.0005, -0.15),
vec![hidden, hidden],
),
);
}
for (i, name) in [
"attn_q.bias",
"attn_k.bias",
"attn_v.bias",
"attn_output.bias",
]
.iter()
.enumerate()
{
tensors.insert(
key(name),
upload_f32(
device,
&prand(qkvo_seed_base + 50 + i * 10, hidden, 0.001, -0.05),
vec![hidden],
),
);
}
tensors.insert(
key("attn_output_norm.weight"),
upload_f32(
device,
&prand(qkvo_seed_base + 60, hidden, 0.0001, 1.0),
vec![hidden],
),
);
tensors.insert(
key("attn_output_norm.bias"),
upload_f32(
device,
&prand(qkvo_seed_base + 61, hidden, 0.001, -0.5),
vec![hidden],
),
);
tensors.insert(
key("ffn_up.weight"),
upload_f32(
device,
&prand(qkvo_seed_base + 70, intermediate * hidden, 0.0005, -0.15),
vec![intermediate, hidden],
),
);
tensors.insert(
key("ffn_up.bias"),
upload_f32(
device,
&prand(qkvo_seed_base + 71, intermediate, 0.001, -0.5),
vec![intermediate],
),
);
tensors.insert(
key("ffn_down.weight"),
upload_f32(
device,
&prand(qkvo_seed_base + 72, hidden * intermediate, 0.0005, -0.15),
vec![hidden, intermediate],
),
);
tensors.insert(
key("ffn_down.bias"),
upload_f32(
device,
&prand(qkvo_seed_base + 73, hidden, 0.001, -0.5),
vec![hidden],
),
);
tensors.insert(
key("layer_output_norm.weight"),
upload_f32(
device,
&prand(qkvo_seed_base + 80, hidden, 0.0001, 1.0),
vec![hidden],
),
);
tensors.insert(
key("layer_output_norm.bias"),
upload_f32(
device,
&prand(qkvo_seed_base + 81, hidden, 0.001, -0.5),
vec![hidden],
),
);
}
let weights = super::super::weights::LoadedBertWeights::from_tensors_for_test(
tensors,
MlxDevice::new().unwrap(),
);
let input_ids: Vec<u32> = (0..seq)
.map(|i| (i.wrapping_mul(7) % vocab) as u32)
.collect();
let type_ids: Vec<u32> = (0..seq).map(|i| (i & 1) as u32).collect();
let input_ids_b = upload_u32(device, &input_ids, vec![seq]);
let type_ids_b = upload_u32(device, &type_ids, vec![seq]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let out = apply_bert_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_ids_b,
Some(&type_ids_b),
&weights,
&cfg,
seq as u32,
seq as u32, )
.expect("full forward dispatch");
session.finish().expect("finish");
let got = readback_f32(&out, hidden);
for (i, &v) in got.iter().enumerate() {
assert!(v.is_finite(), "output[{i}] not finite: {v}");
}
let mut sum_sq = 0.0f64;
for &v in &got {
sum_sq += (v as f64) * (v as f64);
}
let norm = sum_sq.sqrt() as f32;
assert!(
(norm - 1.0).abs() < 1e-4,
"full-forward output not unit L2 norm: {} (expected 1.0)",
norm
);
}
#[test]
fn full_forward_rejects_pooling_type_none() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let cfg = super::super::config::BertConfig {
hidden_size: 64,
num_attention_heads: 2,
num_hidden_layers: 1,
intermediate_size: 128,
max_position_embeddings: 64,
vocab_size: 100,
type_vocab_size: 2,
layer_norm_eps: 1e-12,
hidden_act: "gelu".into(),
pooling_type: super::super::config::PoolingType::None,
causal_attention: false,
};
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let weights = super::super::weights::LoadedBertWeights::empty(MlxDevice::new().unwrap());
let input_ids = upload_u32(device, &[0u32; 32], vec![32]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let err = apply_bert_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_ids,
None,
&weights,
&cfg,
32,
32,
);
assert!(err.is_err(), "pooling_type=None must error");
let msg = format!("{}", err.unwrap_err());
assert!(
msg.contains("pooling_type=None"),
"error must name pooling_type=None: {msg}"
);
drop(session);
}
#[test]
fn full_forward_rejects_seq_len_below_floor() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let cfg = super::super::config::BertConfig {
hidden_size: 64,
num_attention_heads: 2,
num_hidden_layers: 1,
intermediate_size: 128,
max_position_embeddings: 64,
vocab_size: 100,
type_vocab_size: 2,
layer_norm_eps: 1e-12,
hidden_act: "gelu".into(),
pooling_type: super::super::config::PoolingType::Mean,
causal_attention: false,
};
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let weights = super::super::weights::LoadedBertWeights::empty(MlxDevice::new().unwrap());
let input_ids = upload_u32(device, &[0u32; 16], vec![16]);
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let err = apply_bert_full_forward_gpu(
session.encoder_mut(),
&mut registry,
device,
&input_ids,
None,
&weights,
&cfg,
16, 16,
);
assert!(err.is_err(), "seq_len < 32 must error");
drop(session);
}
#[test]
fn encoder_block_rejects_hidden_not_divisible_by_num_heads() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
if MlxDevice::new().is_err() {
eprintln!("skipping: no Metal device available");
return;
}
let device = MlxDevice::new().unwrap();
let executor = GraphExecutor::new(device);
let device_ref: *const MlxDevice = executor.device() as *const _;
let device: &MlxDevice = unsafe { &*device_ref };
let buf = upload_f32(device, &[0.0; 32 * 65], vec![32, 65]);
let small = upload_f32(device, &[0.0; 65], vec![65]);
let small_w = upload_f32(device, &[0.0; 65 * 65], vec![65, 65]);
let tensors = BertEncoderBlockTensors {
q_w: &small_w,
q_b: None,
k_w: &small_w,
k_b: None,
v_w: &small_w,
v_b: None,
o_w: &small_w,
o_b: None,
attn_norm_gamma: &small,
attn_norm_beta: &small,
up_w: &small_w,
up_b: None,
down_w: &small_w,
down_b: None,
ffn_norm_gamma: &small,
ffn_norm_beta: &small,
};
let mut session = executor.begin().expect("begin");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let err = apply_bert_encoder_block_gpu(
session.encoder_mut(),
&mut registry,
device,
&buf,
&tensors,
None,
32,
65, 2, 128,
1e-12,
);
assert!(err.is_err(), "hidden % num_heads != 0 must error");
drop(session);
}
#[rustfmt::skip]
const BGE_GROUND_TRUTH_HELLO_WORLD: [f32; 384] = [
1.5245200e-02f32, -2.2694700e-02f32, 8.5904000e-03f32, -7.4250800e-02f32,
3.8820000e-03f32, 2.7334000e-03f32, -3.1268600e-02f32, 4.4632100e-02f32,
4.4025500e-02f32, -7.8715000e-03f32, -2.5214700e-02f32, -3.3416000e-02f32,
1.4413500e-02f32, 4.6410600e-02f32, 8.6145000e-03f32, -1.6080500e-02f32,
7.4887000e-03f32, -1.8985500e-02f32, -1.1456510e-01f32, -1.8153900e-02f32,
1.2628550e-01f32, 2.9765000e-02f32, 2.5295900e-02f32, -3.4199300e-02f32,
-4.1072000e-02f32, 6.6354000e-03f32, 1.0337600e-02f32, 2.2416600e-02f32,
4.4387000e-03f32, -1.2735740e-01f32, -1.6072700e-02f32, -2.0363500e-02f32,
4.7341200e-02f32, 1.1577300e-02f32, 6.8168700e-02f32, 7.3822000e-03f32,
-1.7995200e-02f32, 4.0913100e-02f32, -1.0225900e-02f32, 2.3681300e-02f32,
1.0515600e-02f32, -2.8547400e-02f32, 8.1570000e-03f32, -1.5209600e-02f32,
3.0877500e-02f32, -6.5933600e-02f32, -2.2227600e-02f32, 5.3976200e-02f32,
2.6597000e-03f32, 2.2453900e-02f32, -9.1692300e-02f32, -4.5221700e-02f32,
-4.2084000e-03f32, -5.5980000e-03f32, -5.4037000e-03f32, 9.8472100e-02f32,
6.0502400e-02f32, 7.4229000e-03f32, 1.3882200e-02f32, 2.6891000e-03f32,
4.7620100e-02f32, 2.8700200e-02f32, -1.5526920e-01f32, 6.8964400e-02f32,
3.0246500e-02f32, -1.7961500e-02f32, 2.0926600e-02f32, 2.1297000e-02f32,
1.4070900e-02f32, 2.0100000e-03f32, 2.6726000e-03f32, 3.9320000e-03f32,
4.1020500e-02f32, 6.5808500e-02f32, -6.1782000e-03f32, -1.6395200e-02f32,
8.2713000e-03f32, -4.9068900e-02f32, -2.1074200e-02f32, -3.0809100e-02f32,
-4.0618500e-02f32, 5.9307100e-02f32, 1.8110800e-02f32, -4.4221500e-02f32,
7.0220000e-04f32, -2.7995600e-02f32, -4.0511900e-02f32, -1.1288300e-02f32,
-2.4989100e-02f32, 9.6083000e-03f32, -1.7424800e-02f32, -2.7018600e-02f32,
-1.5503500e-02f32, -5.5615000e-03f32, -4.1466900e-02f32, 7.1377000e-03f32,
7.0550000e-03f32, 9.7625000e-03f32, 6.9740000e-04f32, 3.4366410e-01f32,
-9.5348500e-02f32, -2.0391000e-03f32, 2.8154300e-02f32, -9.1379400e-02f32,
5.9625200e-02f32, 2.4833900e-02f32, -1.6368100e-02f32, -2.9081900e-02f32,
-8.2563000e-03f32, 1.5810900e-02f32, 1.2835300e-02f32, -6.4291800e-02f32,
1.4440600e-02f32, -1.3786200e-02f32, 1.0548000e-03f32, -1.9413000e-02f32,
5.0069200e-02f32, -2.7956000e-03f32, 9.3218700e-02f32, -2.9476600e-02f32,
-8.1660000e-03f32, 3.0733500e-02f32, -4.4089600e-02f32, -4.1517000e-03f32,
5.2847200e-02f32, -6.4513800e-02f32, 5.8293600e-02f32, 7.7707000e-02f32,
1.1471900e-02f32, 6.9841600e-02f32, -5.4241000e-03f32, 5.9807100e-02f32,
-2.6379400e-02f32, -8.6844000e-03f32, 2.7463500e-02f32, -1.4279000e-02f32,
-1.8396800e-02f32, -1.3933100e-02f32, 3.5396400e-02f32, -5.6796000e-02f32,
8.2031000e-03f32, -7.6783700e-02f32, -2.2572900e-02f32, -1.1289320e-01f32,
3.3520000e-04f32, 3.0491700e-02f32, -7.3323700e-02f32, 2.4644000e-02f32,
-1.9628000e-02f32, -2.4151000e-02f32, -3.8864300e-02f32, 7.8660000e-02f32,
4.9657000e-03f32, -1.6394300e-02f32, 7.9038000e-03f32, 5.4956400e-02f32,
-1.2885100e-02f32, 6.8558700e-02f32, 7.6886000e-03f32, 8.7973000e-03f32,
-1.8497000e-03f32, -1.2487900e-02f32, -1.3336700e-02f32, 6.6981000e-03f32,
-1.7702300e-02f32, -1.2828940e-01f32, 9.9775000e-03f32, 1.9581100e-02f32,
-7.2369000e-03f32, 8.7490000e-04f32, 3.3446000e-03f32, 1.6448400e-02f32,
-3.9572100e-02f32, 2.8781000e-02f32, 1.0959640e-01f32, 7.5352000e-03f32,
-4.0438000e-03f32, 4.4518800e-02f32, -4.7362600e-02f32, 2.4964000e-02f32,
6.0060000e-02f32, -5.0817200e-02f32, -4.1762900e-02f32, 1.9128400e-02f32,
2.8170000e-02f32, -2.5297200e-02f32, -2.0765900e-02f32, -3.0412200e-02f32,
6.2295800e-02f32, 6.7060100e-02f32, -2.3187700e-02f32, 1.0734500e-02f32,
-3.1954700e-02f32, -3.4281000e-02f32, -8.4273100e-02f32, 3.2644000e-03f32,
3.3854400e-02f32, -8.1071600e-02f32, 1.3342300e-02f32, -2.1539200e-02f32,
1.4645080e-01f32, 5.3104200e-02f32, 3.9373000e-03f32, 2.8793300e-02f32,
5.2210000e-04f32, 4.2053000e-03f32, 4.0630400e-02f32, 6.3329000e-03f32,
4.4674100e-02f32, 1.3358900e-02f32, -2.4087900e-02f32, -1.5090700e-02f32,
7.3155600e-02f32, -6.5968000e-03f32, 2.1938400e-02f32, -4.2962600e-02f32,
-1.0111400e-02f32, 7.4725400e-02f32, 2.3953400e-02f32, 4.7187100e-02f32,
-3.9791600e-02f32, 1.0862900e-02f32, -2.2080500e-02f32, -2.6233610e-01f32,
1.8241100e-02f32, 8.3272000e-03f32, -3.3388000e-03f32, -3.4754400e-02f32,
2.3121200e-02f32, 3.8076700e-02f32, -5.1632600e-02f32, 1.0183660e-01f32,
-9.0856000e-03f32, 8.7150100e-02f32, -5.9687000e-02f32, -8.4066000e-03f32,
-3.6316700e-02f32, 1.7526600e-02f32, 2.3206900e-02f32, -1.4118300e-02f32,
1.6098100e-02f32, -1.0101000e-02f32, -2.2810100e-02f32, 2.8593400e-02f32,
2.2963600e-02f32, 4.3393100e-02f32, -4.7477900e-02f32, 4.4480400e-02f32,
-5.9617800e-02f32, 1.4664930e-01f32, 8.3677600e-02f32, -2.0353100e-02f32,
2.4272500e-02f32, 3.6327000e-02f32, -2.8053600e-02f32, -9.2974000e-03f32,
-1.1975370e-01f32, -2.5626800e-02f32, 7.3649500e-02f32, -3.4565200e-02f32,
-6.7101600e-02f32, -9.6649400e-02f32, -2.2320300e-02f32, -1.2381200e-02f32,
1.3812800e-02f32, -4.1045900e-02f32, -4.1482000e-03f32, -2.4175200e-02f32,
-7.4918800e-02f32, -5.2708000e-02f32, 9.8539000e-03f32, -5.2184100e-02f32,
-1.2405700e-02f32, -1.1600300e-02f32, 2.2357500e-02f32, 5.7255500e-02f32,
5.9955200e-02f32, 1.9182300e-02f32, -4.6006900e-02f32, 1.5220000e-03f32,
-5.6150000e-04f32, -1.1537300e-02f32, 3.2410700e-02f32, -1.4726500e-02f32,
-2.2231900e-02f32, 1.5891200e-02f32, -3.6683100e-02f32, 1.1616600e-02f32,
3.5017200e-02f32, -6.1055400e-02f32, -2.4773600e-02f32, 4.9781600e-02f32,
-1.7490500e-02f32, -1.8016700e-02f32, -3.5791000e-02f32, 2.1147000e-02f32,
-1.6487700e-02f32, 3.6262300e-02f32, 1.4208100e-02f32, -4.6354000e-03f32,
-2.3338500e-02f32, -3.9657500e-02f32, -2.8176300e-02f32, -5.5650000e-03f32,
1.1481100e-02f32, 5.8364500e-02f32, 1.4236800e-02f32, 3.2630400e-02f32,
5.4060100e-02f32, 6.4761900e-02f32, 7.6997000e-03f32, 3.5554700e-02f32,
-1.6068700e-02f32, -1.2872200e-02f32, 4.1160900e-02f32, -5.3329000e-03f32,
-6.9779300e-02f32, 1.1296000e-02f32, 1.6109300e-02f32, -2.9520580e-01f32,
2.7797900e-02f32, -3.0543000e-03f32, 2.1346000e-02f32, 4.0104000e-03f32,
2.1070500e-02f32, 4.1065700e-02f32, -2.6940000e-04f32, -5.7391000e-02f32,
2.2403100e-02f32, -7.7541100e-02f32, 2.0342000e-02f32, 1.6253200e-02f32,
-6.6938100e-02f32, 8.0250000e-04f32, 2.0220900e-02f32, -2.4065000e-03f32,
-1.1099900e-02f32, 1.7103700e-02f32, -1.9624700e-02f32, 2.0817000e-03f32,
2.2252900e-02f32, 2.2973500e-01f32, -2.3067100e-02f32, 5.6796500e-02f32,
3.9116400e-02f32, -9.2091000e-03f32, 4.5408000e-03f32, 5.4877500e-02f32,
1.9218000e-02f32, -9.8172600e-02f32, -1.4690000e-04f32, 3.1628200e-02f32,
-1.5647900e-02f32, 3.5432100e-02f32, 1.1029000e-02f32, -6.8190800e-02f32,
-2.8878400e-02f32, 2.3841100e-02f32, -5.3042700e-02f32, -2.5050200e-02f32,
2.2473200e-02f32, -4.5994900e-02f32, 7.0341900e-02f32, 3.4573300e-02f32,
-7.7199700e-02f32, -1.3507200e-02f32, -4.9011200e-02f32, -3.9882000e-03f32,
3.7250900e-02f32, -2.8179500e-02f32, -7.9685900e-02f32, 5.6049000e-03f32,
3.2035800e-02f32, -3.0435800e-02f32, 1.5001600e-02f32, 1.4733000e-02f32,
-8.9100000e-03f32, 1.6160500e-02f32, -6.3414600e-02f32, 2.1236900e-02f32,
-6.2079000e-03f32, 4.9397100e-02f32, 2.2779600e-02f32, 2.6059400e-02f32,
];
#[test]
fn bge_full_forward_matches_llama_embedding_on_hello_world() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::config::{BertConfig, PoolingType};
use super::super::tokenizer::{BertVocab, BertWpmTokenizer};
use super::super::weights::LoadedBertWeights;
use mlx_native::gguf::GgufFile;
use std::path::Path;
let model_path = Path::new("/opt/hf2q/models/bert-test/bge-small-en-v1.5-f16.gguf");
if !model_path.exists() {
eprintln!("skipping: bge GGUF fixture not at {}", model_path.display());
return;
}
let gguf = GgufFile::open(model_path).expect("open bge GGUF");
let cfg = BertConfig::from_gguf(&gguf).expect("parse bge config");
assert_eq!(cfg.hidden_size, 384, "expected bge hidden=384");
assert_eq!(cfg.num_hidden_layers, 12);
assert_eq!(cfg.num_attention_heads, 12);
assert_eq!(
cfg.pooling_type,
PoolingType::Cls,
"bge GGUF must declare CLS pool"
);
let vocab = BertVocab::from_gguf(&gguf).expect("parse bge vocab");
let tok = BertWpmTokenizer::new(&vocab);
let real_ids = tok.encode("hello world", true);
let valid_token_count: u32 = real_ids.len() as u32;
assert_eq!(
real_ids.len(),
4,
"bge tokenizer must produce 4 tokens for hello world, 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 device = MlxDevice::new().expect("create device");
let mut registry = KernelRegistry::new();
register_bert_custom_shaders(&mut registry);
let weights =
LoadedBertWeights::load_from_path(model_path, &cfg).expect("load bge 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_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
valid_token_count,
)
.expect("bge full forward");
encoder.commit_and_wait().expect("commit_and_wait");
let view: &[f32] = pooled.as_slice::<f32>().expect("read pooled f32");
assert_eq!(view.len(), 384);
let truth: &[f32] = &BGE_GROUND_TRUTH_HELLO_WORLD;
let dot: f32 = view.iter().zip(truth.iter()).map(|(a, b)| a * b).sum();
let na: f32 = 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 = view
.iter()
.zip(truth.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0_f32, f32::max);
eprintln!(
"[bge parity] cosine={:.6}, ||hf2q||_2={:.6}, ||truth||_2={:.6}, max_abs_diff={:.4e}",
cosine, na, nb, max_abs_diff
);
eprintln!(
" hf2q first4 = {:?}
truth first4 = {:?}",
&view[..4],
&truth[..4]
);
assert!(
cosine >= 0.999,
"bge cosine {cosine:.6} below 0.999 — shared primitive bug"
);
}
#[test]
fn bge_full_forward_padding_invariance_at_seq_lens_32_64_128_256_512() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::config::BertConfig;
use super::super::tokenizer::{BertVocab, BertWpmTokenizer};
use super::super::weights::LoadedBertWeights;
use mlx_native::gguf::GgufFile;
use std::path::Path;
let model_path = Path::new("/opt/hf2q/models/bert-test/bge-small-en-v1.5-f16.gguf");
if !model_path.exists() {
eprintln!("skipping: bge GGUF fixture not at {}", model_path.display());
return;
}
let gguf = GgufFile::open(model_path).expect("open bge GGUF");
let cfg = BertConfig::from_gguf(&gguf).expect("parse bge config");
let vocab = BertVocab::from_gguf(&gguf).expect("parse bge vocab");
let tok = BertWpmTokenizer::new(&vocab);
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_bert_custom_shaders(&mut registry);
let weights =
LoadedBertWeights::load_from_path(model_path, &cfg).expect("load bge 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_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
valid_token_count,
)
.unwrap_or_else(|e| panic!("bge 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(), 384, "seq_len={seq_len}: hidden_size mismatch");
for &v in view {
assert!(v.is_finite(), "seq_len={seq_len}: non-finite component");
}
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!(
"[bge pad-invariance] seq_len 32 vs {sl}: cosine={cosine:.7}, drift={drift:.2e}"
);
assert!(
cosine >= 0.99999,
"bge: seq_len 32 vs {sl}: cosine {cosine:.7} below 0.99999 padding-invariance \
gate (drift {drift:.2e}). Padding mask leak in bert_attention_with_mask_gpu."
);
}
eprintln!(
"[bge pad-invariance] PASS — max drift across {} seq_lens = {:.2e}",
seq_lens.len(),
max_drift
);
}
#[rustfmt::skip]
const MXBAI_GROUND_TRUTH_HELLO_WORLD: [f32; 1024] = [
2.2862000e-02f32, 3.2229400e-02f32, 1.6500900e-02f32, -4.0425400e-02f32,
-2.1881100e-02f32, -7.9486000e-03f32, 2.9895900e-02f32, 3.7125100e-02f32,
4.6133100e-02f32, 2.4655500e-02f32, 6.9248000e-03f32, 1.9999800e-02f32,
-1.2427500e-02f32, 7.4361000e-03f32, -4.6717900e-02f32, 2.0224000e-02f32,
-2.5140500e-02f32, -3.0057000e-02f32, -6.0605700e-02f32, 3.2786500e-02f32,
-2.3757700e-02f32, 1.2677700e-02f32, -8.3747500e-02f32, -3.3079300e-02f32,
-2.5537300e-02f32, 4.1600900e-02f32, 2.7574400e-02f32, -3.2820000e-04f32,
4.1534600e-02f32, 4.3373900e-02f32, -1.1165400e-02f32, 3.2116000e-03f32,
1.2129000e-03f32, -5.0809700e-02f32, -1.1083400e-02f32, -1.5287500e-02f32,
4.4852200e-02f32, -2.4152300e-02f32, -2.4532900e-02f32, -4.1632600e-02f32,
2.9234100e-02f32, 3.1365000e-03f32, 3.0496200e-02f32, -3.9785200e-02f32,
-6.9592800e-02f32, -1.3217600e-02f32, 2.6155400e-02f32, 4.8777000e-03f32,
4.6322700e-02f32, -4.0573000e-03f32, 1.8286000e-03f32, 4.0606800e-02f32,
9.0854000e-03f32, -1.2301000e-03f32, -1.1784400e-02f32, 1.3525800e-02f32,
-6.0879800e-02f32, 1.0414100e-02f32, -3.1798900e-02f32, 1.0662600e-02f32,
2.4546300e-02f32, 5.2817700e-02f32, 2.0446700e-02f32, -7.2727700e-02f32,
1.2185700e-02f32, 4.2775700e-02f32, 5.5470000e-03f32, -5.2566000e-03f32,
3.6138200e-02f32, 1.5169800e-02f32, -3.7171000e-02f32, 5.3496600e-02f32,
-4.1447000e-03f32, -3.9818200e-02f32, -5.4784500e-02f32, 3.5411800e-02f32,
1.6109200e-02f32, 2.3299900e-02f32, -1.8856800e-02f32, 3.2218300e-02f32,
-1.1987400e-02f32, 3.3749400e-02f32, 4.9290000e-04f32, 3.0040000e-03f32,
-5.4527400e-02f32, -3.4486900e-02f32, 2.0522100e-02f32, 1.6838200e-02f32,
-1.7051900e-02f32, -2.0391800e-02f32, 4.1376400e-02f32, 3.1528700e-02f32,
-1.0009300e-02f32, -3.6536000e-03f32, 6.0232300e-02f32, 3.3995600e-02f32,
-2.9666300e-02f32, 1.6734300e-02f32, -3.0788400e-02f32, -4.2687000e-03f32,
3.6202200e-02f32, 4.0201700e-02f32, -4.2336900e-02f32, -1.5622000e-03f32,
-4.6224600e-02f32, -3.4100500e-02f32, -1.0826500e-02f32, 8.9550000e-03f32,
-3.7400800e-02f32, -3.1266400e-02f32, -4.2501000e-03f32, 3.7303000e-02f32,
3.9644000e-03f32, 1.5097500e-02f32, 3.0286000e-03f32, 4.5656400e-02f32,
8.6289000e-03f32, 3.2128600e-02f32, -1.0707300e-02f32, -1.0184800e-02f32,
2.1930800e-02f32, 2.8468000e-02f32, -1.4027300e-02f32, -2.8737800e-02f32,
-2.0836200e-02f32, -2.0604900e-02f32, -2.4155600e-02f32, 5.7173500e-02f32,
5.1529000e-03f32, 1.0762000e-02f32, 7.3470000e-04f32, -1.9931700e-02f32,
3.3788400e-02f32, 2.8987700e-02f32, -1.8344000e-03f32, 3.0667000e-03f32,
-1.3084600e-02f32, 3.5200500e-02f32, 3.8240900e-02f32, -4.5845000e-02f32,
-4.3980000e-04f32, 3.7770000e-04f32, -2.3862600e-02f32, 8.0980900e-02f32,
1.0245500e-02f32, 3.1511700e-02f32, -1.4256700e-02f32, 2.4348600e-02f32,
-6.9770000e-04f32, 4.3189900e-02f32, -2.5497100e-02f32, 2.2618500e-02f32,
8.9268000e-03f32, 4.2109800e-02f32, -1.5989200e-02f32, -2.7807000e-03f32,
-8.9620000e-03f32, 1.4604200e-02f32, 1.8095000e-02f32, -8.7595000e-03f32,
7.6450000e-04f32, 2.9396000e-02f32, -2.7347500e-02f32, 3.6878000e-03f32,
-4.0618800e-02f32, 7.6515000e-03f32, -4.9296600e-02f32, -5.4022000e-03f32,
9.1581000e-03f32, -1.3371800e-02f32, 6.6258000e-03f32, -3.6590400e-02f32,
-2.5426100e-02f32, 7.9233000e-03f32, 2.0431000e-02f32, 1.8890300e-02f32,
1.1466300e-02f32, 2.8679000e-03f32, 3.4525400e-02f32, 1.3935100e-02f32,
-3.1265700e-02f32, 1.2692000e-02f32, 1.5685000e-03f32, 4.9095600e-02f32,
-1.7892000e-02f32, -1.7392000e-03f32, 8.6984000e-03f32, 1.0384900e-02f32,
-2.2696900e-02f32, -2.6924800e-02f32, 2.4078400e-02f32, 2.2032200e-02f32,
-6.1268300e-02f32, 1.6198200e-02f32, 5.1380000e-03f32, 2.5814000e-03f32,
-2.9176000e-02f32, 1.8927600e-02f32, -3.3073200e-02f32, -8.3106400e-02f32,
-3.7652300e-02f32, 5.7429000e-03f32, 1.0024900e-02f32, 2.0627700e-02f32,
1.0066600e-02f32, -6.5196000e-03f32, 8.0183000e-03f32, 2.1712700e-02f32,
-1.6242600e-02f32, -1.9620300e-02f32, 8.3928000e-03f32, -1.4324000e-03f32,
-1.9611100e-02f32, 3.4865000e-03f32, 1.3618400e-02f32, -2.7884200e-02f32,
-2.3681700e-02f32, 1.6979900e-02f32, -1.5290000e-03f32, -1.6604000e-03f32,
-2.2413800e-02f32, 1.0234000e-02f32, 4.0514800e-02f32, 6.1869000e-02f32,
2.0578700e-02f32, 4.3876000e-03f32, -5.5726200e-02f32, 1.8319400e-02f32,
-3.0017300e-02f32, 3.1964000e-03f32, -1.8740100e-02f32, 3.0228900e-02f32,
-1.0270000e-03f32, 4.5232400e-02f32, 1.4132900e-02f32, 8.6312000e-03f32,
5.0155000e-02f32, 5.2359800e-02f32, 1.0172700e-02f32, -1.9043400e-02f32,
-4.0117600e-02f32, 1.7401000e-03f32, 2.2032500e-02f32, 1.1584200e-02f32,
-3.6885000e-03f32, -2.9159900e-02f32, 4.8161700e-02f32, -1.9127000e-02f32,
2.6407000e-02f32, 2.4215900e-02f32, -3.3261500e-02f32, 4.9436300e-02f32,
3.0378100e-02f32, -3.6395000e-03f32, -4.3533100e-02f32, -2.8479000e-02f32,
2.8138000e-03f32, 6.6386700e-02f32, -4.7156500e-02f32, -2.5913400e-02f32,
3.7953000e-03f32, 2.9152400e-02f32, 1.6483000e-02f32, 3.8055500e-02f32,
-1.9844900e-02f32, 3.5440000e-03f32, 2.4980000e-03f32, -2.2053000e-03f32,
-1.7802000e-02f32, -6.8874000e-02f32, -1.2386800e-02f32, -1.8708000e-03f32,
-3.3156200e-02f32, 2.5248600e-02f32, -2.8189300e-02f32, -1.2455000e-03f32,
1.4453700e-02f32, 1.1176400e-02f32, 2.8258600e-02f32, -1.8777800e-02f32,
2.2680400e-02f32, -1.0154300e-02f32, 3.9323500e-02f32, 3.0262300e-02f32,
1.4747500e-02f32, -9.7326000e-03f32, -2.3373700e-02f32, -4.7098600e-02f32,
-1.7440600e-02f32, 3.2133600e-02f32, -5.1420000e-02f32, 2.1719800e-02f32,
2.7737600e-02f32, 1.5388800e-02f32, -2.8993300e-02f32, 1.4752100e-02f32,
-3.0294100e-02f32, 2.2744200e-02f32, -2.6081300e-02f32, -4.3239700e-02f32,
-2.2400500e-02f32, 9.4560000e-04f32, 2.7710000e-03f32, -4.1011900e-02f32,
-2.6324400e-02f32, 2.3616400e-02f32, -1.0622200e-02f32, -8.1068000e-03f32,
5.1571700e-02f32, 1.1321800e-02f32, -3.0840500e-02f32, 6.0576500e-02f32,
-9.3060000e-04f32, -1.6983700e-02f32, -5.1242400e-02f32, 3.6244500e-02f32,
2.1792400e-02f32, 3.8190600e-02f32, -3.1378900e-02f32, -3.0171200e-02f32,
-2.1146700e-02f32, -2.8735700e-02f32, -5.2895000e-03f32, -6.3761200e-02f32,
-2.3351000e-03f32, 3.1823100e-02f32, 1.8874400e-02f32, -4.8763500e-02f32,
2.1885000e-02f32, -4.0564000e-02f32, -7.2895800e-02f32, -2.6741200e-02f32,
-2.0427900e-02f32, 2.6241500e-02f32, 2.6279000e-03f32, -2.7142000e-03f32,
-1.5128100e-02f32, 6.4017000e-03f32, 1.6682100e-02f32, -3.7042000e-03f32,
3.4899600e-02f32, -4.8226600e-02f32, 2.2115500e-02f32, 4.8111500e-02f32,
-1.8555800e-02f32, -9.3284000e-03f32, -5.8455000e-03f32, -3.6137200e-02f32,
-6.6335000e-03f32, 3.7441900e-02f32, -3.2261000e-03f32, -1.0439500e-02f32,
7.1557000e-03f32, 3.2592600e-02f32, -2.1904000e-03f32, 4.8681400e-02f32,
-6.7723200e-02f32, -4.5323000e-03f32, 5.8523500e-02f32, -2.8170300e-02f32,
3.9120400e-02f32, 3.0864400e-02f32, 7.6855000e-03f32, -6.3554500e-02f32,
4.5573000e-03f32, -5.1020600e-02f32, 1.3798600e-02f32, 1.1104200e-02f32,
2.6145700e-02f32, -5.3097200e-02f32, 6.2102600e-02f32, -1.6956400e-02f32,
-2.7684600e-02f32, -1.3683900e-02f32, -3.1626100e-02f32, -2.6886800e-02f32,
2.4274300e-02f32, -7.4259000e-03f32, 2.3702000e-03f32, -4.5291800e-02f32,
2.4134100e-02f32, -1.8067900e-02f32, 4.9146000e-03f32, 8.8902000e-02f32,
5.6667900e-02f32, 1.2717500e-02f32, 7.0911000e-03f32, -6.1002000e-03f32,
-3.1570000e-03f32, -1.1723200e-02f32, -6.8555200e-02f32, -1.9305900e-02f32,
6.0335900e-02f32, -4.7615000e-02f32, -7.4321000e-03f32, -7.0127000e-02f32,
4.8954700e-02f32, 1.3645300e-02f32, 6.0050000e-02f32, 1.5846000e-02f32,
-3.0289000e-02f32, -3.3252000e-03f32, -1.7062600e-02f32, 2.5397000e-03f32,
-3.0811500e-02f32, -3.3390000e-04f32, -3.8029400e-02f32, 3.2091100e-02f32,
7.3723000e-03f32, -2.2332200e-02f32, -4.8648200e-02f32, -8.4179000e-03f32,
-1.8427000e-02f32, 4.6247600e-02f32, 2.2795900e-02f32, 2.3600600e-02f32,
-1.1928700e-02f32, 3.4484400e-02f32, 2.9937000e-03f32, -2.0899100e-02f32,
-4.2119200e-02f32, -1.8230600e-02f32, -6.3242300e-02f32, -2.9463000e-03f32,
3.0099500e-02f32, -2.5254500e-02f32, -3.5661100e-02f32, -2.7630700e-02f32,
1.2882500e-02f32, 3.1902200e-02f32, 7.4414000e-03f32, -3.6057700e-02f32,
6.2990000e-04f32, -2.4982300e-02f32, -3.6968000e-02f32, 2.1778200e-02f32,
1.4077000e-02f32, -2.1269900e-02f32, 4.7625000e-02f32, -6.2774900e-02f32,
2.7494900e-02f32, 2.8338000e-02f32, -5.4001000e-03f32, -2.3043000e-03f32,
-8.6676000e-03f32, -1.1847200e-02f32, 6.5945000e-03f32, 2.5272900e-02f32,
-7.4520000e-03f32, -2.3510500e-02f32, 3.6342800e-02f32, -3.8237200e-02f32,
-3.1844000e-03f32, -4.4990000e-04f32, 2.1312500e-02f32, -1.4886600e-02f32,
-1.1578300e-02f32, 9.7378000e-03f32, 5.1158300e-02f32, -3.0004000e-03f32,
6.7589000e-03f32, -1.6115200e-02f32, 3.6568000e-02f32, -4.9268000e-03f32,
-3.6748700e-02f32, 5.9463000e-02f32, -1.9746600e-02f32, -5.1079000e-03f32,
3.1014800e-02f32, -8.1931000e-03f32, -2.8500900e-02f32, 2.0540300e-02f32,
3.0086400e-02f32, -2.9973000e-03f32, -9.9483000e-03f32, -6.8647000e-03f32,
6.9795000e-03f32, 4.0454800e-02f32, -7.6346000e-03f32, 1.9463300e-02f32,
-3.3533900e-02f32, 1.4442000e-03f32, -1.6911600e-02f32, -4.0769900e-02f32,
8.8688000e-03f32, -9.2398500e-02f32, 1.8769300e-02f32, 1.0338500e-02f32,
1.4680100e-02f32, 1.1304100e-02f32, -1.0749500e-02f32, 2.9636500e-02f32,
-1.4964900e-02f32, -6.1511400e-02f32, -1.0479500e-02f32, -3.2020500e-02f32,
-4.4911300e-02f32, -1.7740400e-02f32, -2.2508000e-03f32, -2.0031800e-02f32,
2.3014400e-02f32, -2.2459900e-02f32, -5.2564300e-02f32, -1.1192300e-02f32,
-1.1211800e-02f32, 6.1159000e-03f32, -1.2372700e-02f32, 3.1708000e-02f32,
-4.3072900e-02f32, -2.3828000e-03f32, -1.3262300e-02f32, 2.4795900e-02f32,
-2.4822000e-02f32, -2.5076600e-02f32, 8.4510000e-04f32, 4.8709100e-02f32,
9.8562000e-03f32, 2.9879900e-02f32, -5.1196200e-02f32, 2.7311800e-02f32,
9.7970000e-04f32, -5.0924000e-02f32, -2.6263700e-02f32, 2.6038500e-02f32,
-1.3252000e-03f32, 4.1169100e-02f32, 2.6987800e-02f32, 1.2051100e-02f32,
-5.8767000e-03f32, -4.0797400e-02f32, -2.8732900e-02f32, -2.3355400e-02f32,
-3.6151200e-02f32, -5.1868400e-02f32, -1.8994800e-02f32, 2.4010000e-02f32,
7.7268000e-03f32, -8.2449000e-03f32, -6.4542200e-02f32, -1.8828200e-02f32,
-1.3966900e-02f32, 3.4285600e-02f32, -5.2552100e-02f32, -2.5276400e-02f32,
-2.9173300e-02f32, 3.6389000e-02f32, -3.1427800e-02f32, 7.4602800e-02f32,
-4.7566800e-02f32, 1.8356500e-02f32, -1.4352100e-02f32, 1.6278200e-02f32,
-1.0875500e-02f32, 6.2132000e-02f32, -5.2906800e-02f32, 4.9864000e-03f32,
2.7260000e-02f32, 1.0520200e-02f32, 1.1069100e-02f32, 2.6004400e-02f32,
5.5627000e-03f32, 1.3730400e-02f32, -1.6293100e-02f32, 9.4372000e-03f32,
-5.0403700e-02f32, -1.4747900e-02f32, -2.1884600e-02f32, 1.4559000e-02f32,
1.4931100e-02f32, 1.4597300e-02f32, -1.6728800e-02f32, -2.2293300e-02f32,
6.8085800e-02f32, 3.7395000e-03f32, -3.6064500e-02f32, -4.3834500e-02f32,
-4.2494000e-02f32, -7.1476000e-02f32, -2.3417100e-02f32, 9.2438000e-03f32,
-1.3847600e-02f32, -1.7655800e-02f32, -5.2257300e-02f32, 2.6359900e-02f32,
2.0828700e-02f32, -2.1759700e-02f32, 3.3311000e-02f32, 1.0801800e-01f32,
5.6134100e-02f32, -8.8320000e-04f32, -8.9354100e-02f32, 4.4191400e-02f32,
6.3227000e-03f32, -1.6213600e-02f32, -1.2138600e-02f32, 2.2647500e-02f32,
-1.8122200e-02f32, -6.7824000e-03f32, -2.8051700e-02f32, -4.9404100e-02f32,
-5.9882000e-02f32, -1.6048600e-02f32, 4.9397000e-02f32, 3.7778000e-03f32,
1.2607000e-02f32, -5.3785000e-03f32, -2.3814700e-02f32, -6.9269300e-02f32,
2.6107100e-02f32, 6.6044000e-02f32, -2.7619100e-02f32, 6.2448900e-02f32,
5.9607400e-02f32, 1.8743100e-02f32, 5.8603500e-02f32, -2.8879700e-02f32,
1.7019100e-02f32, 1.1876000e-03f32, 7.3308600e-02f32, 4.1337800e-02f32,
1.4845400e-02f32, 4.4626000e-02f32, 4.8008800e-02f32, -8.0565800e-02f32,
-5.3585100e-02f32, -3.8113000e-03f32, -1.0737600e-02f32, -1.5536000e-03f32,
-3.4258300e-02f32, 1.0620000e-03f32, -3.5405000e-02f32, -2.8307000e-03f32,
5.3051900e-02f32, 2.1197000e-02f32, -5.3689000e-03f32, 1.7494000e-02f32,
3.6831000e-03f32, 4.2539100e-02f32, -5.9184000e-03f32, 1.0358500e-02f32,
6.3987600e-02f32, -3.7929100e-02f32, 1.8180500e-02f32, -2.6493900e-02f32,
1.3841000e-03f32, -2.7468400e-02f32, 9.7875000e-03f32, 1.4954200e-02f32,
1.6891500e-02f32, 1.8076200e-02f32, 5.9886300e-02f32, 3.7067000e-03f32,
2.3840700e-02f32, -3.8954300e-02f32, -8.9651000e-03f32, 6.7597000e-03f32,
-5.9056100e-02f32, -4.8183200e-02f32, -2.1049700e-02f32, -7.1333000e-03f32,
-1.8852600e-02f32, 2.4794800e-02f32, -1.9173500e-02f32, -1.0800500e-02f32,
-1.5372000e-03f32, 1.4693000e-02f32, -7.6761000e-03f32, -3.8464400e-02f32,
2.8865000e-03f32, -8.4996900e-02f32, -5.7256500e-02f32, -1.3738900e-02f32,
1.0133300e-02f32, -1.3737400e-02f32, 1.6240800e-02f32, 1.5023500e-02f32,
2.1985300e-02f32, -1.3351400e-02f32, 1.1235500e-02f32, -9.8606000e-03f32,
2.6557900e-02f32, -8.9650000e-03f32, 5.1829700e-02f32, -6.4708600e-02f32,
-7.3609000e-02f32, 7.8220000e-04f32, 1.7445700e-02f32, -1.0260900e-02f32,
-2.1395700e-02f32, -2.2310400e-02f32, -1.0673300e-02f32, -1.0567400e-02f32,
6.4216200e-02f32, -1.1993000e-03f32, 1.6724400e-02f32, -2.2510200e-02f32,
-1.9426000e-02f32, 3.0697900e-02f32, 7.9458000e-03f32, -7.3166000e-03f32,
8.2411000e-03f32, -5.6915000e-03f32, -7.0201700e-02f32, 2.5255000e-03f32,
-1.6971600e-02f32, 5.8360000e-04f32, -2.4168200e-02f32, -3.5294000e-03f32,
5.3846000e-03f32, 1.9681000e-03f32, -3.3645600e-02f32, -2.8378800e-02f32,
-1.4100000e-02f32, -3.3058100e-02f32, -6.2496000e-03f32, -9.5420000e-04f32,
1.2760400e-02f32, -1.4827200e-02f32, 1.9506900e-02f32, -3.1506500e-02f32,
2.7964700e-02f32, -8.4771000e-03f32, -1.4876000e-03f32, -6.2906000e-03f32,
-1.6269000e-03f32, -2.8008000e-03f32, -1.3981900e-02f32, 1.0373800e-02f32,
-5.9039000e-03f32, 1.6205700e-02f32, 2.2784000e-02f32, 4.9244400e-02f32,
2.8444200e-02f32, -8.1399000e-03f32, 2.3860600e-02f32, -3.0506600e-02f32,
-1.6948200e-02f32, -2.3872300e-02f32, -9.7571000e-03f32, 7.2810000e-03f32,
4.4397700e-02f32, -1.3552000e-03f32, -4.5923700e-02f32, 4.0605000e-03f32,
-5.4009000e-03f32, -2.7175900e-02f32, 4.4955400e-02f32, -6.9098700e-02f32,
-3.5846400e-02f32, 2.6923600e-02f32, -5.4106100e-02f32, 2.9440500e-02f32,
-9.0675000e-03f32, -3.1219600e-02f32, 4.2961700e-02f32, 1.0904700e-02f32,
-1.5930900e-02f32, -2.5379300e-02f32, 1.9928300e-02f32, 1.6682200e-02f32,
-1.5719000e-03f32, -7.4102400e-02f32, -2.2706800e-02f32, 1.7079900e-02f32,
1.3165000e-02f32, 4.2992600e-02f32, -4.8893900e-02f32, -3.3629000e-03f32,
1.8718400e-02f32, 8.8008000e-03f32, -1.7933500e-02f32, 9.5796000e-03f32,
5.3268000e-03f32, -1.0987700e-02f32, 3.5844200e-02f32, 1.6112800e-02f32,
1.9731500e-02f32, -6.2491100e-02f32, 2.9250300e-02f32, 3.6052000e-03f32,
2.4888400e-02f32, -6.4298200e-02f32, 7.2592000e-03f32, 3.4341000e-03f32,
2.1791900e-02f32, -1.8989000e-03f32, 5.8870000e-02f32, 2.3985600e-02f32,
2.1390100e-02f32, 3.2579800e-02f32, 3.1067000e-03f32, 9.5311000e-03f32,
3.3476700e-02f32, -3.2642000e-02f32, -3.3315000e-03f32, 3.6116900e-02f32,
3.7278500e-02f32, 3.7887300e-02f32, 7.3874000e-03f32, 8.9610000e-03f32,
1.7817700e-02f32, -3.0485000e-03f32, 1.2165900e-02f32, 6.8150000e-03f32,
1.8155600e-02f32, 3.7112000e-03f32, 1.0940300e-02f32, 2.0073700e-02f32,
-2.2687000e-02f32, 1.0770200e-02f32, -2.5991100e-02f32, 1.3413000e-02f32,
-6.1886000e-03f32, -2.9126000e-02f32, -5.1982500e-02f32, -5.6037000e-03f32,
5.0463700e-02f32, 1.4155500e-02f32, 8.6375000e-03f32, -2.0995000e-03f32,
-1.5770400e-02f32, 1.0537100e-02f32, 4.8329100e-02f32, -7.8195000e-03f32,
3.8701900e-02f32, 1.2228900e-02f32, 2.4737000e-02f32, 2.7029600e-02f32,
6.8640000e-04f32, 2.8675600e-02f32, 3.1105900e-02f32, -5.6329900e-02f32,
-3.3821600e-02f32, 6.3955000e-03f32, 1.6503100e-02f32, -1.3584000e-03f32,
4.6793000e-03f32, -3.4737900e-02f32, 4.2422700e-02f32, 8.8776000e-03f32,
-9.4292000e-03f32, 1.1102500e-02f32, -2.4456300e-02f32, -4.3121600e-02f32,
-5.6746100e-02f32, 3.2851000e-02f32, -8.1746000e-03f32, 8.3245000e-03f32,
5.4157300e-02f32, -1.0709200e-02f32, -1.5350000e-03f32, 1.0397000e-02f32,
1.4033000e-02f32, 2.6737000e-02f32, 7.5933000e-03f32, 2.1814800e-02f32,
-2.4173700e-02f32, 3.6758100e-02f32, 4.9254500e-02f32, -4.1397900e-02f32,
1.1131600e-02f32, 1.3043300e-02f32, -1.5823600e-02f32, 1.2876600e-02f32,
1.9750300e-02f32, -5.1226900e-02f32, -8.9014000e-03f32, -1.9618700e-02f32,
-5.2303400e-02f32, 2.8494500e-02f32, 3.6790600e-02f32, -3.1983400e-02f32,
-1.3976300e-02f32, 5.8967000e-03f32, 2.0737600e-02f32, -2.0991200e-02f32,
3.8876500e-02f32, 1.6713600e-02f32, 2.1521800e-02f32, -3.6749300e-02f32,
-1.6150000e-04f32, -2.4888700e-02f32, 9.3091000e-03f32, 1.7954000e-03f32,
-6.6524500e-02f32, 4.3178000e-03f32, -4.4832400e-02f32, 7.5906000e-03f32,
-1.8633900e-02f32, -4.0724200e-02f32, -5.5274000e-03f32, 2.1824400e-02f32,
3.4358000e-03f32, -2.9230400e-02f32, 6.5207000e-03f32, 3.3028000e-02f32,
1.1731000e-03f32, 1.0430900e-02f32, -2.0097300e-02f32, 2.4722600e-02f32,
7.8836000e-03f32, 8.2255100e-02f32, -6.2831000e-03f32, 5.9372000e-02f32,
1.9079600e-02f32, -8.1807000e-03f32, 3.1516400e-02f32, -2.2834000e-03f32,
3.6261700e-02f32, 4.2615000e-03f32, -1.0707800e-02f32, -2.9460600e-02f32,
1.5319400e-02f32, -1.4550700e-02f32, -6.9470000e-03f32, 1.1964500e-02f32,
6.7788300e-02f32, -8.3592000e-03f32, -2.3782000e-03f32, -4.0738900e-02f32,
-2.1628900e-02f32, -4.8629900e-02f32, -3.4740200e-02f32, -4.6098200e-02f32,
6.5833300e-02f32, -2.4617000e-03f32, -2.7654000e-03f32, -1.1778400e-02f32,
-4.6229600e-02f32, 2.3633460e-01f32, 4.6884300e-02f32, 4.0941100e-02f32,
5.3436900e-02f32, 7.3879000e-03f32, -8.4473000e-03f32, 1.9252200e-02f32,
-1.9778100e-02f32, 3.7206000e-03f32, -3.7005400e-02f32, 1.1873900e-02f32,
-1.0385000e-03f32, 2.8184400e-02f32, 6.3196400e-02f32, 3.0511000e-03f32,
5.1730300e-02f32, -4.5365700e-02f32, 9.5138000e-03f32, -1.1105300e-02f32,
-2.5750400e-02f32, -6.7887400e-02f32, 1.8028100e-02f32, 2.4626400e-02f32,
5.5506000e-03f32, 1.5091600e-02f32, 1.9322400e-02f32, 1.6044300e-02f32,
-4.6013300e-02f32, -7.8463000e-03f32, -1.2125400e-02f32, 5.4351600e-02f32,
-2.8350000e-04f32, 1.1774300e-02f32, 7.3179000e-03f32, -3.3139200e-02f32,
4.4295800e-02f32, -1.8367200e-02f32, -9.3501000e-03f32, 4.7690000e-03f32,
4.3934100e-02f32, -1.8728600e-02f32, -3.2801800e-02f32, 1.4522500e-02f32,
3.4296000e-03f32, 6.6564000e-03f32, 5.4273300e-02f32, -4.1352100e-02f32,
6.3561000e-03f32, 6.8160000e-03f32, -2.4460900e-02f32, 3.6838400e-02f32,
-2.9931100e-02f32, 1.9422200e-02f32, -4.1083400e-02f32, -4.8003400e-02f32,
-1.9515000e-02f32, 3.4448000e-03f32, -1.1905000e-03f32, -3.6424400e-02f32,
7.1766000e-03f32, -1.1291100e-02f32, 1.2760800e-02f32, -2.5732300e-02f32,
-8.8609000e-03f32, -4.7951300e-02f32, -6.5900000e-04f32, 4.3316800e-02f32,
6.3517700e-02f32, -3.8809000e-03f32, -4.6437300e-02f32, -3.7860400e-02f32,
-1.2300700e-02f32, 2.3294000e-03f32, -5.7220700e-02f32, 1.2738000e-03f32,
5.8221900e-02f32, -4.3387700e-02f32, -9.4290000e-04f32, 1.6419400e-02f32,
-3.8486000e-03f32, -3.3620000e-04f32, 1.8537200e-02f32, 6.1911000e-03f32,
-1.1911800e-02f32, -2.9750000e-04f32, 3.2422600e-02f32, -8.7210000e-03f32,
-9.4019000e-03f32, -1.1974200e-02f32, 1.1888000e-02f32, 5.3618100e-02f32,
5.8220500e-02f32, -7.5546000e-03f32, -5.6425000e-03f32, 1.9206100e-02f32,
];
#[test]
fn mxbai_full_forward_matches_llama_embedding_on_hello_world() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::config::{BertConfig, PoolingType};
use super::super::tokenizer::{BertVocab, BertWpmTokenizer};
use super::super::weights::LoadedBertWeights;
use mlx_native::gguf::GgufFile;
use std::path::Path;
let model_path = Path::new("/opt/hf2q/models/bert-test/mxbai-embed-large-v1-f16.gguf");
if !model_path.exists() {
eprintln!(
"skipping: mxbai GGUF fixture not at {}",
model_path.display()
);
return;
}
let gguf = GgufFile::open(model_path).expect("open mxbai GGUF");
let cfg = BertConfig::from_gguf(&gguf).expect("parse mxbai config");
assert_eq!(cfg.hidden_size, 1024, "expected mxbai hidden=1024");
assert_eq!(cfg.num_hidden_layers, 24);
assert_eq!(cfg.num_attention_heads, 16);
assert_eq!(cfg.pooling_type, PoolingType::Cls);
let vocab = BertVocab::from_gguf(&gguf).expect("parse mxbai vocab");
let tok = BertWpmTokenizer::new(&vocab);
let real_ids = tok.encode("hello world", true);
let valid_token_count: u32 = real_ids.len() as u32;
assert_eq!(real_ids.len(), 4);
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_bert_custom_shaders(&mut registry);
let weights =
LoadedBertWeights::load_from_path(model_path, &cfg).expect("load mxbai 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_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
valid_token_count,
)
.expect("mxbai full forward");
encoder.commit_and_wait().expect("commit_and_wait");
let view: &[f32] = pooled.as_slice::<f32>().expect("read pooled f32");
assert_eq!(view.len(), 1024);
let truth: &[f32] = &MXBAI_GROUND_TRUTH_HELLO_WORLD;
let dot: f32 = view.iter().zip(truth.iter()).map(|(a, b)| a * b).sum();
let na: f32 = 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 = view
.iter()
.zip(truth.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0_f32, f32::max);
eprintln!(
"[mxbai parity] cosine={:.6}, ||hf2q||_2={:.6}, ||truth||_2={:.6}, max_abs_diff={:.4e}",
cosine, na, nb, max_abs_diff
);
eprintln!(" hf2q first4 = {:?}", &view[..4]);
eprintln!(" truth first4 = {:?}", &truth[..4]);
assert!(cosine >= 0.999, "mxbai cosine {cosine:.6} below 0.999");
}
#[test]
fn mxbai_full_forward_padding_invariance_at_seq_lens_32_64_128_256_512() {
let _gpu = crate::inference::hf2q_gpu_test_lock();
use super::super::config::BertConfig;
use super::super::tokenizer::{BertVocab, BertWpmTokenizer};
use super::super::weights::LoadedBertWeights;
use mlx_native::gguf::GgufFile;
use std::path::Path;
let model_path = Path::new("/opt/hf2q/models/bert-test/mxbai-embed-large-v1-f16.gguf");
if !model_path.exists() {
eprintln!(
"skipping: mxbai GGUF fixture not at {}",
model_path.display()
);
return;
}
let gguf = GgufFile::open(model_path).expect("open mxbai GGUF");
let cfg = BertConfig::from_gguf(&gguf).expect("parse mxbai config");
let vocab = BertVocab::from_gguf(&gguf).expect("parse mxbai vocab");
let tok = BertWpmTokenizer::new(&vocab);
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_bert_custom_shaders(&mut registry);
let weights =
LoadedBertWeights::load_from_path(model_path, &cfg).expect("load mxbai 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_bert_full_forward_gpu(
&mut encoder,
&mut registry,
&device,
&input_ids,
None,
&weights,
&cfg,
seq_len,
valid_token_count,
)
.unwrap_or_else(|e| panic!("mxbai 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(), 1024, "seq_len={seq_len}: hidden_size mismatch");
for &v in view {
assert!(v.is_finite(), "seq_len={seq_len}: non-finite component");
}
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!(
"[mxbai pad-invariance] seq_len 32 vs {sl}: cosine={cosine:.7}, drift={drift:.2e}"
);
assert!(
cosine >= 0.99999,
"mxbai: seq_len 32 vs {sl}: cosine {cosine:.7} below 0.99999 padding-invariance \
gate (drift {drift:.2e})."
);
}
eprintln!(
"[mxbai pad-invariance] PASS — max drift across {} seq_lens = {:.2e}",
seq_lens.len(),
max_drift
);
}
}