Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion crates/metaltile-core/src/protocol.rs
Original file line number Diff line number Diff line change
Expand Up @@ -323,7 +323,7 @@ mod tests {
assert_eq!(op_group("mt_softmax_f32"), "softmax");
assert_eq!(op_group("mt_softmax_bf16"), "softmax");
assert_eq!(op_group("softmax"), "softmax");
assert_eq!(op_group("ffai_vector_add_f16"), "ffai_vector_add");
assert_eq!(op_group("mt_vector_add_f16"), "vector_add");
}

#[test]
Expand Down
174 changes: 0 additions & 174 deletions crates/metaltile-std/src/ffai/arg_reduce.rs

This file was deleted.

5 changes: 0 additions & 5 deletions crates/metaltile-std/src/ffai/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@
//! counterpart lands in mainline at a future pin, the file moves to
//! `mlx/`.

pub mod arg_reduce;
pub mod attn_head_gate;
pub mod aura_dequant_rotated;
pub mod aura_encode;
Expand All @@ -26,7 +25,6 @@ pub mod aura_flash_sdpa;
pub mod aura_score;
pub mod aura_value;
pub mod avg_pool2d_nhwc;
pub mod axpy_scalar_inplace;
pub mod batched_4_block_scaled_qgemv;
pub mod batched_4_block_scaled_qmm;
pub mod batched_4_qgemv;
Expand All @@ -35,7 +33,6 @@ pub mod batched_qkv_block_scaled_qgemv;
pub mod batched_qkv_block_scaled_qmm;
pub mod batched_qkv_qgemv;
pub mod batched_qkv_qmm;
pub mod clamp_scalar;
pub mod dequant_gather;
pub mod dequant_gather_block_scaled;
pub mod dequant_gemv;
Expand All @@ -61,7 +58,6 @@ pub mod gated_delta_prep;
pub mod gated_delta_prep_chunk;
pub mod gated_delta_replay;
pub mod gated_delta_wy;
pub mod gather;
pub mod gelu_erf;
pub mod gemm;
pub mod gemm_q4_mpp;
Expand Down Expand Up @@ -112,7 +108,6 @@ pub mod moe_mpp_int8;
pub mod moe_mpp_shared;
pub mod moe_router_sigmoid_bias;
pub mod moe_router_sqrtsoftplus;
pub mod mt_vector_add;
pub mod patch_embed;
pub mod patch_embed_block_scaled;
pub mod patch_embed_mma;
Expand Down
1 change: 1 addition & 0 deletions crates/metaltile-std/src/kernels/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,5 +9,6 @@

pub mod convolution;
pub mod norm;
pub mod ops;
pub mod rope;
pub mod sampling;
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
use metaltile::kernel;

#[kernel]
pub fn ffai_axpy_scalar_inplace<T>(src: Tensor<T>, mut accum: Tensor<T>, #[constexpr] scalar: f32) {
pub fn mt_axpy_scalar_inplace<T>(src: Tensor<T>, mut accum: Tensor<T>, #[constexpr] scalar: f32) {
let i = tid;
let s = load(src[i]).cast::<f32>();
let a = load(accum[i]).cast::<f32>();
Expand All @@ -21,7 +21,7 @@ pub fn ffai_axpy_scalar_inplace<T>(src: Tensor<T>, mut accum: Tensor<T>, #[const
pub mod kernel_tests {
use metaltile::{test::*, test_kernel};

use super::ffai_axpy_scalar_inplace;
use super::mt_axpy_scalar_inplace;
use crate::utils::{pack_f32, unpack_f32};

fn setup(n: usize, dt: DType) -> TestSetup {
Expand All @@ -32,7 +32,7 @@ pub mod kernel_tests {
let accum_dt = unpack_f32(&pack_f32(&accum_in, dt), dt);
let expected: Vec<f32> =
accum_dt.iter().zip(&src_dt).map(|(a, s)| a + scalar * s).collect();
TestSetup::new(ffai_axpy_scalar_inplace::kernel_ir_for(dt))
TestSetup::new(mt_axpy_scalar_inplace::kernel_ir_for(dt))
.input(TestBuffer::from_vec("src", pack_f32(&src, dt), dt))
.input(TestBuffer::from_vec("accum", pack_f32(&accum_in, dt), dt))
.constexpr("scalar", scalar)
Expand All @@ -47,12 +47,12 @@ pub mod kernel_tests {
pub mod kernel_benches {
use metaltile::{bench, test::*};

use super::ffai_axpy_scalar_inplace;
use super::mt_axpy_scalar_inplace;

#[bench(dtypes = [f32, f16, bf16])]
fn bench_axpy_scalar(dt: DType) -> BenchSetup {
let n = 4096usize;
BenchSetup::new(ffai_axpy_scalar_inplace::kernel_ir_for(dt))
BenchSetup::new(mt_axpy_scalar_inplace::kernel_ir_for(dt))
.buffer(BenchBuffer::random("src", n, dt))
.buffer(BenchBuffer::random("accum", n, dt).output())
.constexpr("scalar", 0.5f32)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@
use metaltile::kernel;

#[kernel]
pub fn ffai_clamp_scalar<T>(input: Tensor<T>, out: Tensor<T>, lo: Tensor<f32>, hi: Tensor<f32>) {
pub fn mt_clamp_scalar<T>(input: Tensor<T>, out: Tensor<T>, lo: Tensor<f32>, hi: Tensor<f32>) {
let i = program_id::<0>();
let lo_v = load(lo[0]);
let hi_v = load(hi[0]);
Expand All @@ -39,12 +39,12 @@ pub fn ffai_clamp_scalar<T>(input: Tensor<T>, out: Tensor<T>, lo: Tensor<f32>, h
store(out[i], clamped.cast::<T>());
}

/// New-syntax correctness for `ffai_clamp_scalar`. Grid3D, grid `[n,1,1]`,
/// New-syntax correctness for `mt_clamp_scalar`. Grid3D, grid `[n,1,1]`,
/// tpg `[1,1,1]`. Oracle clamps each element to `[lo, hi]`.
pub mod kernel_tests {
use metaltile::{test::*, test_kernel};

use super::ffai_clamp_scalar;
use super::mt_clamp_scalar;
use crate::utils::{pack_f32, unpack_f32};

fn f32_bytes(v: f32) -> Vec<u8> { v.to_le_bytes().to_vec() }
Expand All @@ -58,7 +58,7 @@ pub mod kernel_tests {
let input_f: Vec<f32> = (0..n).map(|i| (i as f32 - 128.0) * 0.1).collect();
let input = unpack_f32(&pack_f32(&input_f, dt), dt);
let exp: Vec<f32> = input.iter().map(|&x| x.max(lo).min(hi)).collect();
TestSetup::new(ffai_clamp_scalar::kernel_ir_for(dt))
TestSetup::new(mt_clamp_scalar::kernel_ir_for(dt))
.mode(KernelMode::Grid3D)
.input(TestBuffer::from_vec("input", pack_f32(&input_f, dt), dt))
.input(TestBuffer::from_vec("lo", f32_bytes(lo), DType::F32))
Expand All @@ -69,16 +69,16 @@ pub mod kernel_tests {
}
}

/// New-syntax benchmark for `ffai_clamp_scalar`.
/// New-syntax benchmark for `mt_clamp_scalar`.
pub mod kernel_benches {
use metaltile::{bench, test::*};

use super::ffai_clamp_scalar;
use super::mt_clamp_scalar;

#[bench(dtypes = [f32, f16, bf16])]
fn bench_clamp_scalar(dt: DType) -> BenchSetup {
let n = 576 * 768usize; // SigLIP patch-grid activation size
BenchSetup::new(ffai_clamp_scalar::kernel_ir_for(dt))
BenchSetup::new(mt_clamp_scalar::kernel_ir_for(dt))
.mode(KernelMode::Grid3D)
.buffer(BenchBuffer::random("input", n, dt))
.buffer(BenchBuffer::random("lo", 1, DType::F32))
Expand Down
Loading
Loading