use core::marker::PhantomData;
use teeny_core::dtype::Float;
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison, Tensor},
*,
};
#[kernel]
pub fn flash_attention2_forward<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>,
l_ptr: T::Pointer<D>,
n_ctx_q: i32,
n_ctx_k: i32,
softmax_scale: f32, neg_inf: 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_m = T::program_id(Axis::X); let pid_bh = T::program_id(Axis::Y);
let kv_bh_base = pid_bh * n_ctx_k * HEAD_DIM;
let q_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
let o_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
let l_row_base = pid_bh * n_ctx_q + pid_m;
let d = T::arange(0, HEAD_DIM);
let q_vec = T::load(
q_ptr.add_offsets(d + q_row_base),
None, None, &[], None, None, None, false,
);
let mut acc = T::zeros::<D>(&[HEAD_DIM]);
let mut m_i = T::full(&[HEAD_DIM], D::from_f64(neg_inf as f64));
let mut l_i = T::zeros::<D>(&[HEAD_DIM]);
let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
for k_row in 0..n_ctx_k {
let kv_row_base = kv_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), true) * scale_t;
let m_new = T::maximum(m_i, qk); let exp_diff = T::exp(m_i - m_new); let p = T::exp(qk - m_new);
l_i = exp_diff * l_i + p;
acc = exp_diff * acc + p * v_vec;
m_i = m_new;
}
let o_row = acc / l_i; let l_save_sum = T::sum(m_i + T::log(l_i), Some(0), false); let l_save = l_save_sum / T::full(&[1], D::from_f64(HEAD_DIM as f64));
T::store(o_ptr.add_offsets(d + o_row_base), o_row, None, &[], None, None);
T::store(l_ptr.add_offsets(T::arange(0, 1) + l_row_base), l_save, None, &[], None, None);
}
#[kernel]
pub fn flash_attention2_backward_dq<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>,
n_ctx_q: i32,
n_ctx_k: 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_m = T::program_id(Axis::X);
let pid_bh = T::program_id(Axis::Y);
let q_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
let kv_bh_base = pid_bh * n_ctx_k * HEAD_DIM;
let l_row_base = pid_bh * n_ctx_q + pid_m;
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 + q_row_base), None, None, &[], None, None, None, false);
let o_vec = T::load(o_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
let do_vec = T::load(do_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
let d_q = T::sum(o_vec * do_vec, Some(0), false);
let l_q_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_row_base), None, None, &[], None, None, None, false);
let l_q = T::sum(l_q_raw, Some(0), false);
let mut dq_acc = T::zeros::<D>(&[HEAD_DIM]);
for k_row in 0..n_ctx_k {
let kv_row_base = kv_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_q);
let do_dot_v = T::sum(do_vec * v_vec, Some(0), false);
let ds = p * (do_dot_v - d_q);
dq_acc = dq_acc + ds * k_vec * scale_t;
}
T::store(dq_ptr.add_offsets(d + q_row_base), dq_acc, None, &[], None, None);
}
#[kernel]
pub fn flash_attention2_backward_dkv<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>,
dk_ptr: T::Pointer<D>,
dv_ptr: T::Pointer<D>,
n_ctx_q: i32,
n_ctx_k: 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 q_bh_base = pid_bh * n_ctx_q * HEAD_DIM;
let kv_row_base = pid_bh * n_ctx_k * HEAD_DIM + pid_n * HEAD_DIM;
let l_bh_base = pid_bh * n_ctx_q;
let d = T::arange(0, HEAD_DIM);
let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
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 mut dk_acc = T::zeros::<D>(&[HEAD_DIM]);
let mut dv_acc = T::zeros::<D>(&[HEAD_DIM]);
for q_row in 0..n_ctx_q {
let q_row_base = q_bh_base + q_row * HEAD_DIM;
let l_row_base = l_bh_base + q_row;
let q_vec_m = T::load(q_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
let o_vec_m = T::load(o_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
let do_vec_m = T::load(do_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
let l_m_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_row_base), 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, Some(0), false) * scale_t;
let p = T::exp(qk - l_m);
dv_acc = dv_acc + p * do_vec_m;
let do_dot_v = T::sum(do_vec_m * v_vec, Some(0), false);
let ds = p * (do_dot_v - d_m);
dk_acc = dk_acc + ds * q_vec_m * scale_t;
}
T::store(dk_ptr.add_offsets(d + kv_row_base), dk_acc, None, &[], None, None);
T::store(dv_ptr.add_offsets(d + kv_row_base), dv_acc, None, &[], None, None);
}
pub struct FlashAttention2Op<'a, D: Float + Send + Sync + 'static> {
pub forward: FlashAttention2Forward<D>,
pub backward_dq: FlashAttention2BackwardDq<D>,
pub backward_dkv: FlashAttention2BackwardDkv<D>,
_marker: PhantomData<&'a ()>,
}
impl<'a, D: Float + Send + Sync + 'static> FlashAttention2Op<'a, D> {
pub fn new(head_dim: i32) -> Self {
Self {
forward: FlashAttention2Forward::<D>::new(head_dim),
backward_dq: FlashAttention2BackwardDq::<D>::new(head_dim),
backward_dkv: FlashAttention2BackwardDkv::<D>::new(head_dim),
_marker: PhantomData,
}
}
}