#![allow(non_snake_case)]
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison},
*,
};
#[kernel]
pub fn muon_frob_norm_sq<T: Triton, const BLOCK_SIZE: i32>(
x_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
n_elements: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let offsets = T::arange(0, BLOCK_SIZE) + pid * BLOCK_SIZE;
let mask = offsets.lt(n_elements);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<f32>(&[BLOCK_SIZE])),
&[],
None,
None,
None,
false,
);
let partial = T::expand_dims(T::sum(x * x, Some(0), false), 0);
let out_off: T::I32Tensor = T::arange(0, 1);
let _ = T::atomic_add(out_ptr.add_offsets(out_off), partial, None, None, None);
}
#[kernel]
pub fn muon_ns_xtx<
T: Triton,
const TRANSPOSE: bool,
const BLOCK_R: i32,
const BLOCK_K: i32,
const GROUP_R: i32,
>(
x_ptr: T::Pointer<f32>,
t_ptr: T::Pointer<f32>,
M: i32,
N: i32,
stride_xm: i32,
) {
let R = if TRANSPOSE { N } else { M };
let K = if TRANSPOSE { M } else { N };
let a_stride_row = if TRANSPOSE { 1 } else { stride_xm };
let a_stride_col = if TRANSPOSE { stride_xm } else { 1 };
let pid = T::program_id(Axis::X);
let num_pid_r = T::cdiv(R, BLOCK_R);
let num_pid_in_group = GROUP_R * num_pid_r;
let group_id = pid / num_pid_in_group;
let first_pid_r = group_id * GROUP_R;
let remaining = num_pid_r - first_pid_r;
let group_size = if remaining < GROUP_R {
remaining
} else {
GROUP_R
};
let pid_in_group = pid % num_pid_in_group;
let pid_rm = first_pid_r + (pid_in_group % group_size);
let pid_rn = pid_in_group / group_size;
let a_desc = T::make_tensor_descriptor(
x_ptr,
&[R, K],
&[a_stride_row, a_stride_col],
&[BLOCK_R, BLOCK_K],
Some(PaddingOption::Zero),
);
let b_desc = T::make_tensor_descriptor(
x_ptr,
&[R, K],
&[a_stride_row, a_stride_col],
&[BLOCK_R, BLOCK_K],
Some(PaddingOption::Zero),
);
let mut acc = T::zeros::<f32>(&[BLOCK_R, BLOCK_R]);
let k_tiles = T::cdiv(K, BLOCK_K);
for k in 0..k_tiles {
let a = T::load_tensor_descriptor(a_desc, &[pid_rm * BLOCK_R, k * BLOCK_K]);
let b = T::load_tensor_descriptor(b_desc, &[pid_rn * BLOCK_R, k * BLOCK_K]);
let b_t = T::trans(b, &[1, 0]);
acc = T::dot::<f32, f32>(a, b_t, Some(acc), InputPrecision::IEEE, None);
}
let t_desc = T::make_tensor_descriptor(
t_ptr,
&[R, R],
&[R, 1],
&[BLOCK_R, BLOCK_R],
Some(PaddingOption::Zero),
);
T::store_tensor_descriptor(t_desc, &[pid_rm * BLOCK_R, pid_rn * BLOCK_R], acc);
}
#[kernel]
pub fn muon_ns_step<
T: Triton,
const TRANSPOSE: bool,
const BLOCK_M: i32,
const BLOCK_N: i32,
const BLOCK_K: i32,
const GROUP_M: i32,
>(
t_ptr: T::Pointer<f32>,
x_ptr: T::Pointer<f32>,
M: i32,
N: i32,
stride_tm: i32,
stride_xm: i32,
a: f32,
b: f32,
) {
let K = if TRANSPOSE { N } else { M };
let pid = T::program_id(Axis::X);
let num_pid_m = T::cdiv(M, BLOCK_M);
let num_pid_n = T::cdiv(N, BLOCK_N);
let num_pid_in_group = GROUP_M * num_pid_n;
let group_id = pid / num_pid_in_group;
let first_pid_m = group_id * GROUP_M;
let remaining = num_pid_m - first_pid_m;
let group_size = if remaining < GROUP_M {
remaining
} else {
GROUP_M
};
let pid_in_group = pid % num_pid_in_group;
let pid_m = first_pid_m + (pid_in_group % group_size);
let pid_n = pid_in_group / group_size;
let (a_desc, b_desc) = if TRANSPOSE {
let ad = T::make_tensor_descriptor(
x_ptr,
&[M, K],
&[stride_xm, 1],
&[BLOCK_M, BLOCK_K],
Some(PaddingOption::Zero),
);
let bd = T::make_tensor_descriptor(
t_ptr,
&[K, N],
&[stride_tm, 1],
&[BLOCK_K, BLOCK_N],
Some(PaddingOption::Zero),
);
(ad, bd)
} else {
let ad = T::make_tensor_descriptor(
t_ptr,
&[M, K],
&[stride_tm, 1],
&[BLOCK_M, BLOCK_K],
Some(PaddingOption::Zero),
);
let bd = T::make_tensor_descriptor(
x_ptr,
&[K, N],
&[stride_xm, 1],
&[BLOCK_K, BLOCK_N],
Some(PaddingOption::Zero),
);
(ad, bd)
};
let mut acc = T::zeros::<f32>(&[BLOCK_M, BLOCK_N]);
let k_tiles = T::cdiv(K, BLOCK_K);
for k in 0..k_tiles {
let av = T::load_tensor_descriptor(a_desc, &[pid_m * BLOCK_M, k * BLOCK_K]);
let bv = T::load_tensor_descriptor(b_desc, &[k * BLOCK_K, pid_n * BLOCK_N]);
acc = T::dot::<f32, f32>(av, bv, Some(acc), InputPrecision::IEEE, None);
}
let x_desc = T::make_tensor_descriptor(
x_ptr,
&[M, N],
&[stride_xm, 1],
&[BLOCK_M, BLOCK_N],
Some(PaddingOption::Zero),
);
let x_tile = T::load_tensor_descriptor(x_desc, &[pid_m * BLOCK_M, pid_n * BLOCK_N]);
let a_t = T::full::<f32>(&[BLOCK_M, BLOCK_N], a);
let b_t = T::full::<f32>(&[BLOCK_M, BLOCK_N], b);
let result = a_t * x_tile + b_t * acc;
T::store_tensor_descriptor(x_desc, &[pid_m * BLOCK_M, pid_n * BLOCK_N], result);
}
#[kernel]
pub fn muon_update<T: Triton, const BLOCK_SIZE: i32>(
params_ptr: T::Pointer<f32>,
grad_ptr: T::Pointer<f32>,
n_elements: i32,
lr: f32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let offsets = T::arange(0, BLOCK_SIZE) + pid * BLOCK_SIZE;
let mask = offsets.lt(n_elements);
let p = T::load(
params_ptr.add_offsets(offsets),
Some(mask),
None,
&[],
None,
None,
None,
false,
);
let g = T::load(
grad_ptr.add_offsets(offsets),
Some(mask),
None,
&[],
None,
None,
None,
false,
);
let lr_t = T::full::<f32>(&[BLOCK_SIZE], lr);
T::store(
params_ptr.add_offsets(offsets),
p - lr_t * g,
Some(mask),
&[],
None,
None,
);
}