#![allow(non_snake_case)]
use core::ffi::c_void;
use teeny_core::dtype::Float;
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison, Tensor},
*,
};
use super::flash_attn2::FlashAttention2Forward;
#[kernel]
pub fn psa_pack_qkv<T: Triton, D: Float, const KEY_DIM: i32>(
qkv_ptr: T::Pointer<D>,
out_ptr: T::Pointer<D>,
qkv_h: i32, H: i32,
W: i32,
B: i32,
num_heads: i32,
) where
T::I32Tensor: Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let pid = T::program_id(Axis::X); let BH: i32 = B * num_heads;
let N: i32 = H * W;
let section: i32 = pid / (BH * N);
let bh: i32 = (pid / N) % BH;
let n: i32 = pid % N;
let b: i32 = bh / num_heads;
let h: i32 = bh % num_heads;
let d = T::arange(0, KEY_DIM);
let chan_base: i32 = h * 4 * KEY_DIM + section * KEY_DIM;
let src_off = (d + chan_base) * (H * W) + (b * qkv_h * H * W + n);
let x = T::load(qkv_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
let dst_base: i32 = section * BH * N * KEY_DIM + bh * N * KEY_DIM + n * KEY_DIM;
let dst_off = d + dst_base;
T::store(out_ptr.add_offsets(dst_off), x, None, &[], None, None);
}
#[kernel]
pub fn psa_extract_v_nchw<T: Triton, D: Float, const KEY_DIM: i32>(
qkv_ptr: T::Pointer<D>,
v_ptr: T::Pointer<D>,
qkv_h: i32,
c: i32, H: i32,
W: i32,
num_heads: i32,
) where
T::I32Tensor: Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let pid = T::program_id(Axis::X); let N: i32 = H * W;
let bh: i32 = pid / N;
let n: i32 = pid % N;
let b: i32 = bh / num_heads;
let h: i32 = bh % num_heads;
let d = T::arange(0, KEY_DIM);
let src_lo_base: i32 = h * 4 * KEY_DIM + 2 * KEY_DIM;
let src_hi_base: i32 = h * 4 * KEY_DIM + 3 * KEY_DIM;
let src_off_lo = (d + src_lo_base) * (H * W) + (b * qkv_h * H * W + n);
let src_off_hi = (d + src_hi_base) * (H * W) + (b * qkv_h * H * W + n);
let x_lo = T::load(qkv_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
let x_hi = T::load(qkv_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
let dst_lo_base: i32 = h * 2 * KEY_DIM;
let dst_hi_base: i32 = h * 2 * KEY_DIM + KEY_DIM;
let dst_off_lo = (d + dst_lo_base) * (H * W) + (b * c * H * W + n);
let dst_off_hi = (d + dst_hi_base) * (H * W) + (b * c * H * W + n);
T::store(v_ptr.add_offsets(dst_off_lo), x_lo, None, &[], None, None);
T::store(v_ptr.add_offsets(dst_off_hi), x_hi, None, &[], None, None);
}
#[kernel]
pub fn psa_merge_attn_nchw<T: Triton, D: Float, const KEY_DIM: i32>(
lo_ptr: T::Pointer<D>,
hi_ptr: T::Pointer<D>,
out_ptr: T::Pointer<D>,
c: i32,
H: i32,
W: i32,
num_heads: i32,
) where
T::I32Tensor: Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let pid = T::program_id(Axis::X); let N: i32 = H * W;
let bh: i32 = pid / N;
let n: i32 = pid % N;
let b: i32 = bh / num_heads;
let h: i32 = bh % num_heads;
let d = T::arange(0, KEY_DIM);
let src_base: i32 = bh * N * KEY_DIM + n * KEY_DIM;
let src_off = d + src_base;
let x_lo = T::load(lo_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
let x_hi = T::load(hi_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
let dst_lo_base: i32 = h * 2 * KEY_DIM;
let dst_hi_base: i32 = h * 2 * KEY_DIM + KEY_DIM;
let dst_off_lo = (d + dst_lo_base) * (H * W) + (b * c * H * W + n);
let dst_off_hi = (d + dst_hi_base) * (H * W) + (b * c * H * W + n);
T::store(out_ptr.add_offsets(dst_off_lo), x_lo, None, &[], None, None);
T::store(out_ptr.add_offsets(dst_off_hi), x_hi, None, &[], None, None);
}
#[kernel]
pub fn psa_pack_qkv_backward<T: Triton, D: Float, const KEY_DIM: i32>(
d_packed_ptr: T::Pointer<D>, d_qkv_ptr: T::Pointer<D>, qkv_h: i32,
H: i32,
W: i32,
B: i32,
num_heads: i32,
) where
T::I32Tensor: Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let pid = T::program_id(Axis::X); let BH = B * num_heads;
let N = H * W;
let section = pid / (BH * N);
let bh = (pid / N) % BH;
let n = pid % N;
let b = bh / num_heads;
let h = bh % num_heads;
let d = T::arange(0, KEY_DIM);
let src_base = section * BH * N * KEY_DIM + bh * N * KEY_DIM + n * KEY_DIM;
let dx = T::load(d_packed_ptr.add_offsets(d + src_base), None, None, &[], None, None, None, false);
let chan_base = h * 4 * KEY_DIM + section * KEY_DIM;
let dst_off = (d + chan_base) * (H * W) + (b * qkv_h * H * W + n);
T::atomic_add(d_qkv_ptr.add_offsets(dst_off), dx, None, None, None);
}
#[kernel]
pub fn psa_extract_v_backward<T: Triton, D: Float, const KEY_DIM: i32>(
d_v_ptr: T::Pointer<D>, d_qkv_ptr: T::Pointer<D>, qkv_h: i32,
c: i32,
H: i32,
W: i32,
num_heads: i32,
) where
T::I32Tensor: Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let pid = T::program_id(Axis::X); let N = H * W;
let bh = pid / N;
let n = pid % N;
let b = bh / num_heads;
let h = bh % num_heads;
let d = T::arange(0, KEY_DIM);
let v_lo_src_base = h * 2 * KEY_DIM;
let v_hi_src_base = h * 2 * KEY_DIM + KEY_DIM;
let src_off_lo = (d + v_lo_src_base) * (H * W) + (b * c * H * W + n);
let src_off_hi = (d + v_hi_src_base) * (H * W) + (b * c * H * W + n);
let dx_lo = T::load(d_v_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
let dx_hi = T::load(d_v_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
let qkv_lo_base = h * 4 * KEY_DIM + 2 * KEY_DIM;
let qkv_hi_base = h * 4 * KEY_DIM + 3 * KEY_DIM;
let dst_off_lo = (d + qkv_lo_base) * (H * W) + (b * qkv_h * H * W + n);
let dst_off_hi = (d + qkv_hi_base) * (H * W) + (b * qkv_h * H * W + n);
T::atomic_add(d_qkv_ptr.add_offsets(dst_off_lo), dx_lo, None, None, None);
T::atomic_add(d_qkv_ptr.add_offsets(dst_off_hi), dx_hi, None, None, None);
}
#[kernel]
pub fn psa_merge_attn_backward<T: Triton, D: Float, const KEY_DIM: i32>(
d_merged_ptr: T::Pointer<D>, d_lo_ptr: T::Pointer<D>, d_hi_ptr: T::Pointer<D>, c: i32,
H: i32,
W: i32,
num_heads: i32,
) where
T::I32Tensor: Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let pid = T::program_id(Axis::X); let N = H * W;
let bh = pid / N;
let n = pid % N;
let b = bh / num_heads;
let h = bh % num_heads;
let d = T::arange(0, KEY_DIM);
let src_lo_base = h * 2 * KEY_DIM;
let src_hi_base = h * 2 * KEY_DIM + KEY_DIM;
let src_off_lo = (d + src_lo_base) * (H * W) + (b * c * H * W + n);
let src_off_hi = (d + src_hi_base) * (H * W) + (b * c * H * W + n);
let dx_lo = T::load(d_merged_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
let dx_hi = T::load(d_merged_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
let dst_base = bh * N * KEY_DIM + n * KEY_DIM;
let dst_off = d + dst_base;
T::store(d_lo_ptr.add_offsets(dst_off), dx_lo, None, &[], None, None);
T::store(d_hi_ptr.add_offsets(dst_off), dx_hi, None, &[], None, None);
}
pub struct PsaPackQkvRuntimeOp<D: Float + Send + Sync + 'static> {
fwd: PsaPackQkv<D>,
bwd: PsaPackQkvBackward<D>,
num_heads: usize,
}
impl<D: Float + Send + Sync + 'static> PsaPackQkvRuntimeOp<D> {
pub fn new(key_dim: i32, num_heads: usize) -> Self {
Self { fwd: PsaPackQkv::<D>::new(key_dim), bwd: PsaPackQkvBackward::<D>::new(key_dim), num_heads }
}
pub fn kernel_name(&self) -> &str { self.fwd.name }
pub fn forward_source(&self) -> &str { &self.fwd.source }
pub fn backward_source(&self) -> &str { &self.bwd.source }
}
impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaPackQkvRuntimeOp<D> {
fn n_activation_inputs(&self) -> usize { 1 }
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_params: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
_output_shape: &[usize],
_output_row_stride: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let b = inputs[0].1[0] as i32;
let qkv_h = inputs[0].1[1] as i32;
let h = inputs[0].1[2] as i32;
let w = inputs[0].1[3] as i32;
visitor.visit_ptr(inputs[0].0);
visitor.visit_ptr(output);
visitor.visit_i32(qkv_h);
visitor.visit_i32(h);
visitor.visit_i32(w);
visitor.visit_i32(b);
visitor.visit_i32(self.num_heads as i32);
}
fn block(&self) -> [u32; 3] { [128, 1, 1] }
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
[(output_shape[0] * output_shape[1] * output_shape[2] * output_shape[3]) as u32, 1, 1]
}
#[cfg(feature = "training")]
fn has_backward(&self) -> bool { true }
#[cfg(feature = "training")]
fn pack_backward_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_params: &[teeny_core::model::RawPtr],
_output: teeny_core::model::RawPtr,
output_shape: &[usize],
grad_output: teeny_core::model::RawPtr,
_grad_output_row_stride: i32,
grad_inputs: &[teeny_core::model::RawPtr],
_grad_params: &[teeny_core::model::RawPtr],
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let b = inputs[0].1[0] as i32;
let qkv_h = inputs[0].1[1] as i32;
let h = inputs[0].1[2] as i32;
let w = inputs[0].1[3] as i32;
let _ = output_shape;
visitor.visit_ptr(grad_output); visitor.visit_ptr(grad_inputs[0]); visitor.visit_i32(qkv_h);
visitor.visit_i32(h);
visitor.visit_i32(w);
visitor.visit_i32(b);
visitor.visit_i32(self.num_heads as i32);
}
#[cfg(feature = "training")]
fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
#[cfg(feature = "training")]
fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
[(output_shape[0] * output_shape[1] * output_shape[2] * output_shape[3]) as u32, 1, 1]
}
}
pub struct PsaExtractVRuntimeOp<D: Float + Send + Sync + 'static> {
fwd: PsaExtractVNchw<D>,
bwd: PsaExtractVBackward<D>,
num_heads: usize,
}
impl<D: Float + Send + Sync + 'static> PsaExtractVRuntimeOp<D> {
pub fn new(key_dim: i32, num_heads: usize) -> Self {
Self { fwd: PsaExtractVNchw::<D>::new(key_dim), bwd: PsaExtractVBackward::<D>::new(key_dim), num_heads }
}
pub fn kernel_name(&self) -> &str { self.fwd.name }
pub fn forward_source(&self) -> &str { &self.fwd.source }
pub fn backward_source(&self) -> &str { &self.bwd.source }
}
impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaExtractVRuntimeOp<D> {
fn n_activation_inputs(&self) -> usize { 1 }
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_params: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
output_shape: &[usize],
_output_row_stride: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let qkv_h = inputs[0].1[1] as i32;
let h = inputs[0].1[2] as i32;
let w = inputs[0].1[3] as i32;
let c = output_shape[1] as i32;
visitor.visit_ptr(inputs[0].0);
visitor.visit_ptr(output);
visitor.visit_i32(qkv_h);
visitor.visit_i32(c);
visitor.visit_i32(h);
visitor.visit_i32(w);
visitor.visit_i32(self.num_heads as i32);
}
fn block(&self) -> [u32; 3] { [128, 1, 1] }
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let bh = output_shape[0] * self.num_heads;
let n = output_shape[2] * output_shape[3];
[(bh * n) as u32, 1, 1]
}
#[cfg(feature = "training")]
fn has_backward(&self) -> bool { true }
#[cfg(feature = "training")]
fn pack_backward_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_params: &[teeny_core::model::RawPtr],
_output: teeny_core::model::RawPtr,
output_shape: &[usize],
grad_output: teeny_core::model::RawPtr,
_grad_output_row_stride: i32,
grad_inputs: &[teeny_core::model::RawPtr],
_grad_params: &[teeny_core::model::RawPtr],
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let qkv_h = inputs[0].1[1] as i32;
let h = inputs[0].1[2] as i32;
let w = inputs[0].1[3] as i32;
let c = output_shape[1] as i32;
visitor.visit_ptr(grad_output); visitor.visit_ptr(grad_inputs[0]); visitor.visit_i32(qkv_h);
visitor.visit_i32(c);
visitor.visit_i32(h);
visitor.visit_i32(w);
visitor.visit_i32(self.num_heads as i32);
}
#[cfg(feature = "training")]
fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
#[cfg(feature = "training")]
fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
let bh = output_shape[0] * self.num_heads;
let n = output_shape[2] * output_shape[3];
[(bh * n) as u32, 1, 1]
}
}
pub struct PsaMergeAttnRuntimeOp<D: Float + Send + Sync + 'static> {
fwd: PsaMergeAttnNchw<D>,
bwd: PsaMergeAttnBackward<D>,
num_heads: usize,
}
impl<D: Float + Send + Sync + 'static> PsaMergeAttnRuntimeOp<D> {
pub fn new(key_dim: i32, num_heads: usize) -> Self {
Self { fwd: PsaMergeAttnNchw::<D>::new(key_dim), bwd: PsaMergeAttnBackward::<D>::new(key_dim), num_heads }
}
pub fn kernel_name(&self) -> &str { self.fwd.name }
pub fn forward_source(&self) -> &str { &self.fwd.source }
pub fn backward_source(&self) -> &str { &self.bwd.source }
}
impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaMergeAttnRuntimeOp<D> {
fn n_activation_inputs(&self) -> usize { 2 }
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_params: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
output_shape: &[usize],
_output_row_stride: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let c = output_shape[1] as i32;
let h = output_shape[2] as i32;
let w = output_shape[3] as i32;
visitor.visit_ptr(inputs[0].0);
visitor.visit_ptr(inputs[1].0);
visitor.visit_ptr(output);
visitor.visit_i32(c);
visitor.visit_i32(h);
visitor.visit_i32(w);
visitor.visit_i32(self.num_heads as i32);
}
fn block(&self) -> [u32; 3] { [128, 1, 1] }
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let bh = output_shape[0] * self.num_heads;
let n = output_shape[2] * output_shape[3];
[(bh * n) as u32, 1, 1]
}
#[cfg(feature = "training")]
fn has_backward(&self) -> bool { true }
#[cfg(feature = "training")]
fn pack_backward_args(
&self,
_inputs: &[(teeny_core::model::RawPtr, &[usize])],
_params: &[teeny_core::model::RawPtr],
_output: teeny_core::model::RawPtr,
output_shape: &[usize],
grad_output: teeny_core::model::RawPtr,
_grad_output_row_stride: i32,
grad_inputs: &[teeny_core::model::RawPtr],
_grad_params: &[teeny_core::model::RawPtr],
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let c = output_shape[1] as i32;
let h = output_shape[2] as i32;
let w = output_shape[3] as i32;
visitor.visit_ptr(grad_output); visitor.visit_ptr(grad_inputs[0]); visitor.visit_ptr(grad_inputs[1]); visitor.visit_i32(c);
visitor.visit_i32(h);
visitor.visit_i32(w);
visitor.visit_i32(self.num_heads as i32);
}
#[cfg(feature = "training")]
fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
#[cfg(feature = "training")]
fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
let bh = output_shape[0] * self.num_heads;
let n = output_shape[2] * output_shape[3];
[(bh * n) as u32, 1, 1]
}
}
#[kernel]
pub fn psa_fa2_backward<T: Triton, D: Float, const HEAD_DIM: i32>(
q_ptr: T::Pointer<D>, k_ptr: T::Pointer<D>, v_ptr: T::Pointer<D>, o_ptr: T::Pointer<D>, do_ptr: T::Pointer<D>, l_ptr: T::Pointer<D>, dq_ptr: T::Pointer<D>, dk_ptr: T::Pointer<D>, dv_ptr: T::Pointer<D>, N: i32, softmax_scale: f32,
) where
T::I32Tensor: Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let pid_n = T::program_id(Axis::X); let pid_bh = T::program_id(Axis::Y);
let row_base = pid_bh * N * HEAD_DIM + pid_n * HEAD_DIM;
let l_base = pid_bh * N + pid_n;
let bh_base = pid_bh * N * HEAD_DIM;
let l_bh = pid_bh * N;
let d = T::arange(0, HEAD_DIM);
let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
let q_vec = T::load(q_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
let o_vec = T::load(o_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
let do_vec = T::load(do_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
let d_n = T::sum(o_vec * do_vec, Some(0), false);
let l_n_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_base), None, None, &[], None, None, None, false);
let l_n = T::sum(l_n_raw, Some(0), false);
let mut dq_acc = T::zeros::<D>(&[HEAD_DIM]);
for k_row in 0..N {
let kv_row_base = bh_base + k_row * HEAD_DIM;
let k_vec = T::load(k_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
let v_vec = T::load(v_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
let qk = T::sum(q_vec * k_vec, Some(0), false) * scale_t;
let p = T::exp(qk - l_n);
let do_dot_v = T::sum(do_vec * v_vec, Some(0), false);
let ds = p * (do_dot_v - d_n);
dq_acc = dq_acc + ds * k_vec * scale_t;
}
T::atomic_add(dq_ptr.add_offsets(d + row_base), dq_acc, None, None, None);
let k_vec_n = T::load(k_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
let v_vec_n = T::load(v_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
let mut dk_acc = T::zeros::<D>(&[HEAD_DIM]);
let mut dv_acc = T::zeros::<D>(&[HEAD_DIM]);
for q_row in 0..N {
let q_row_base_m = bh_base + q_row * HEAD_DIM;
let l_row_base_m = l_bh + q_row;
let q_vec_m = T::load(q_ptr.add_offsets(d + q_row_base_m), None, None, &[], None, None, None, false);
let o_vec_m = T::load(o_ptr.add_offsets(d + q_row_base_m), None, None, &[], None, None, None, false);
let do_vec_m = T::load(do_ptr.add_offsets(d + q_row_base_m), None, None, &[], None, None, None, false);
let l_m_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_row_base_m), None, None, &[], None, None, None, false);
let l_m = T::sum(l_m_raw, Some(0), false);
let d_m = T::sum(o_vec_m * do_vec_m, Some(0), false);
let qk = T::sum(q_vec_m * k_vec_n, Some(0), false) * scale_t;
let p = T::exp(qk - l_m);
dv_acc = dv_acc + p * do_vec_m;
let do_dot_v_m = T::sum(do_vec_m * v_vec_n, Some(0), false);
let ds_m = p * (do_dot_v_m - d_m);
dk_acc = dk_acc + ds_m * q_vec_m * scale_t;
}
T::atomic_add(dk_ptr.add_offsets(d + row_base), dk_acc, None, None, None);
T::store(dv_ptr.add_offsets(d + row_base), dv_acc, None, &[], None, None);
}
use std::any::Any;
use std::sync::Arc;
use teeny_core::{
graph::{CustomOp, Shape},
model::RuntimeOp,
};
pub struct PsaPackQkvOp<D: Float + Send + Sync + 'static> {
inner: Arc<PsaPackQkvRuntimeOp<D>>,
num_heads: usize,
}
impl<D: Float + Send + Sync + 'static> PsaPackQkvOp<D> {
pub fn new(key_dim: i32, num_heads: usize) -> Self {
Self { inner: Arc::new(PsaPackQkvRuntimeOp::<D>::new(key_dim, num_heads)), num_heads }
}
}
impl<D: Float + Send + Sync + 'static> CustomOp for PsaPackQkvOp<D> {
fn name(&self) -> &str { "psa_pack_qkv" }
fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
let s = input_shapes[0]; let nh = self.num_heads;
vec![
s[0], Some(4), Some(nh), s[2].and_then(|h| s[3].map(|w| h * w)), s[1].map(|qkv_h| qkv_h / (nh * 4)), ]
}
fn as_any(&self) -> &dyn Any { self }
fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
Some((
self.inner.kernel_name().to_string(),
self.inner.forward_source().to_string(),
"entry_point".to_string(),
Arc::clone(&self.inner) as Arc<dyn RuntimeOp>,
))
}
fn lower_backward_source(&self) -> String {
self.inner.backward_source().to_string()
}
}
pub struct PsaExtractVOp<D: Float + Send + Sync + 'static>(Arc<PsaExtractVRuntimeOp<D>>);
impl<D: Float + Send + Sync + 'static> PsaExtractVOp<D> {
pub fn new(key_dim: i32, num_heads: usize) -> Self {
Self(Arc::new(PsaExtractVRuntimeOp::<D>::new(key_dim, num_heads)))
}
}
impl<D: Float + Send + Sync + 'static> CustomOp for PsaExtractVOp<D> {
fn name(&self) -> &str { "psa_extract_v_nchw" }
fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
let s = input_shapes[0];
vec![s[0], s[1].map(|qkv_h| qkv_h / 2), s[2], s[3]]
}
fn as_any(&self) -> &dyn Any { self }
fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
Some((
self.0.kernel_name().to_string(),
self.0.forward_source().to_string(),
"entry_point".to_string(),
Arc::clone(&self.0) as Arc<dyn RuntimeOp>,
))
}
fn lower_backward_source(&self) -> String {
self.0.backward_source().to_string()
}
}
pub struct PsaMergeAttnOp<D: Float + Send + Sync + 'static> {
inner: Arc<PsaMergeAttnRuntimeOp<D>>,
num_heads: usize,
h: usize,
w: usize,
}
impl<D: Float + Send + Sync + 'static> PsaMergeAttnOp<D> {
pub fn new(key_dim: i32, num_heads: usize, h: usize, w: usize) -> Self {
Self { inner: Arc::new(PsaMergeAttnRuntimeOp::<D>::new(key_dim, num_heads)), num_heads, h, w }
}
}
impl<D: Float + Send + Sync + 'static> CustomOp for PsaMergeAttnOp<D> {
fn name(&self) -> &str { "psa_merge_attn_nchw" }
fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
let lo = input_shapes[0];
let nh = self.num_heads;
vec![
lo[0], lo[3].map(|kd| nh * 2 * kd), Some(self.h),
Some(self.w),
]
}
fn as_any(&self) -> &dyn Any { self }
fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
Some((
self.inner.kernel_name().to_string(),
self.inner.forward_source().to_string(),
"entry_point".to_string(),
Arc::clone(&self.inner) as Arc<dyn RuntimeOp>,
))
}
fn lower_backward_source(&self) -> String {
self.inner.backward_source().to_string()
}
}
pub struct FlashAttn2PsaOp<D: Float + Send + Sync + 'static>(Arc<FlashAttn2PsaRuntimeOp<D>>);
impl<D: Float + Send + Sync + 'static> FlashAttn2PsaOp<D> {
pub fn new_lo(key_dim: i32) -> Self {
Self(Arc::new(FlashAttn2PsaRuntimeOp::<D>::new_lo(key_dim)))
}
pub fn new_hi(key_dim: i32) -> Self {
Self(Arc::new(FlashAttn2PsaRuntimeOp::<D>::new_hi(key_dim)))
}
}
impl<D: Float + Send + Sync + 'static> CustomOp for FlashAttn2PsaOp<D> {
fn name(&self) -> &str { "flash_attention2_forward" }
fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
let s = input_shapes[0];
vec![s[0], s[2], s[3], s[4]]
}
fn as_any(&self) -> &dyn Any { self }
fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
Some((
self.0.kernel_name().to_string(),
self.0.forward_source().to_string(),
"entry_point".to_string(),
Arc::clone(&self.0) as Arc<dyn RuntimeOp>,
))
}
fn lower_backward_source(&self) -> String {
self.0.backward_source().to_string()
}
}
pub struct FlashAttn2PsaRuntimeOp<D: Float + Send + Sync + 'static> {
fwd: FlashAttention2Forward<D>,
bwd: PsaFa2Backward<D>,
v_section: usize,
}
impl<D: Float + Send + Sync + 'static> FlashAttn2PsaRuntimeOp<D> {
pub fn new_lo(key_dim: i32) -> Self {
Self { fwd: FlashAttention2Forward::<D>::new(key_dim), bwd: PsaFa2Backward::<D>::new(key_dim), v_section: 2 }
}
pub fn new_hi(key_dim: i32) -> Self {
Self { fwd: FlashAttention2Forward::<D>::new(key_dim), bwd: PsaFa2Backward::<D>::new(key_dim), v_section: 3 }
}
pub fn kernel_name(&self) -> &str { self.fwd.name }
pub fn forward_source(&self) -> &str { &self.fwd.source }
pub fn backward_source(&self) -> &str { &self.bwd.source }
}
impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for FlashAttn2PsaRuntimeOp<D> {
fn n_activation_inputs(&self) -> usize { 1 }
fn param_shapes(&self, input_shapes: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
let bh = input_shapes[0][0] * input_shapes[0][2]; let n = input_shapes[0][3];
vec![vec![bh * n]] }
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
params: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
_output_shape: &[usize],
_output_row_stride: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let b = inputs[0].1[0];
let nh = inputs[0].1[2];
let n = inputs[0].1[3];
let kd = inputs[0].1[4];
let bh = b * nh; let section_elems = bh * n * kd;
let base = inputs[0].0 as *mut D;
let q_ptr = base as *mut c_void;
let k_ptr = unsafe { base.add(section_elems) } as *mut c_void;
let v_ptr = unsafe { base.add(self.v_section * section_elems) } as *mut c_void;
let softmax_scale = 1.0_f32 / (kd as f32).sqrt();
visitor.visit_ptr(q_ptr);
visitor.visit_ptr(k_ptr);
visitor.visit_ptr(v_ptr);
visitor.visit_ptr(output);
visitor.visit_ptr(params[0]);
visitor.visit_i32(n as i32);
visitor.visit_i32(n as i32);
visitor.visit_f32(softmax_scale);
visitor.visit_f32(f32::NEG_INFINITY);
}
fn block(&self) -> [u32; 3] { [1, 1, 1] }
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
[output_shape[2] as u32, (output_shape[0] * output_shape[1]) as u32, 1]
}
#[cfg(feature = "training")]
fn has_backward(&self) -> bool { true }
#[cfg(feature = "training")]
fn pack_backward_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
params: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
_output_shape: &[usize],
grad_output: teeny_core::model::RawPtr,
_grad_output_row_stride: i32,
grad_inputs: &[teeny_core::model::RawPtr],
_grad_params: &[teeny_core::model::RawPtr],
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let b = inputs[0].1[0];
let nh = inputs[0].1[2];
let n = inputs[0].1[3];
let kd = inputs[0].1[4];
let bh = b * nh;
let section_elems = bh * n * kd;
let softmax_scale = 1.0_f32 / (kd as f32).sqrt();
let fwd_base = inputs[0].0 as *mut D;
let q_ptr = fwd_base as *mut c_void;
let k_ptr = unsafe { fwd_base.add(section_elems) } as *mut c_void;
let v_ptr = unsafe { fwd_base.add(self.v_section * section_elems) } as *mut c_void;
let d_base = grad_inputs[0] as *mut D;
let dq_ptr = d_base as *mut c_void;
let dk_ptr = unsafe { d_base.add(section_elems) } as *mut c_void;
let dv_ptr = unsafe { d_base.add(self.v_section * section_elems) } as *mut c_void;
visitor.visit_ptr(q_ptr);
visitor.visit_ptr(k_ptr);
visitor.visit_ptr(v_ptr);
visitor.visit_ptr(output); visitor.visit_ptr(grad_output); visitor.visit_ptr(params[0]); visitor.visit_ptr(dq_ptr);
visitor.visit_ptr(dk_ptr);
visitor.visit_ptr(dv_ptr);
visitor.visit_i32(n as i32); visitor.visit_f32(softmax_scale);
}
#[cfg(feature = "training")]
fn backward_block(&self) -> [u32; 3] { [1, 1, 1] }
#[cfg(feature = "training")]
fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
[output_shape[2] as u32, (output_shape[0] * output_shape[1]) as u32, 1]
}
}