diff --git a/include/infinicore/ops.hpp b/include/infinicore/ops.hpp index b5c4ff18f..8ada707aa 100644 --- a/include/infinicore/ops.hpp +++ b/include/infinicore/ops.hpp @@ -35,6 +35,7 @@ #include "ops/hardswish.hpp" #include "ops/hardtanh.hpp" #include "ops/kv_caching.hpp" +#include "ops/kimi_delta_attention.hpp" #include "ops/layer_norm.hpp" #include "ops/linear.hpp" #include "ops/mamba_selective_scan.hpp" diff --git a/include/infinicore/ops/kimi_delta_attention.hpp b/include/infinicore/ops/kimi_delta_attention.hpp new file mode 100644 index 000000000..812cb9eeb --- /dev/null +++ b/include/infinicore/ops/kimi_delta_attention.hpp @@ -0,0 +1,63 @@ +#pragma once + +#include "infinicore.h" + +#include "../device.hpp" +#include "../graph/graph.hpp" +#include "common/op.hpp" + +#include + +namespace infinicore::op { + +INFINICORE_GRAPH_OP_CLASS(KimiDeltaAttention, + Tensor, + Tensor, + std::optional, + const Tensor &, + const Tensor &, + const Tensor &, + const Tensor &, + const Tensor &, + const Tensor &, + const Tensor &, + std::optional, + std::optional, + std::optional, + float, + float, + bool); + +__export Tensor kimi_delta_attention(const Tensor &q, + const Tensor &k, + const Tensor &v, + const Tensor &g, + const Tensor &beta, + const Tensor &A_log, + const Tensor &dt_bias, + Tensor initial_state, + std::optional cu_seqlens = std::nullopt, + std::optional initial_state_indices = std::nullopt, + std::optional final_state_indices = std::nullopt, + float scale = 1.0f, + float lower_bound = -5.0f, + bool use_qk_l2norm = true); + +__export void kimi_delta_attention_(Tensor out, + Tensor initial_state, + std::optional final_state, + const Tensor &q, + const Tensor &k, + const Tensor &v, + const Tensor &g, + const Tensor &beta, + const Tensor &A_log, + const Tensor &dt_bias, + std::optional cu_seqlens, + std::optional initial_state_indices, + std::optional final_state_indices, + float scale = 1.0f, + float lower_bound = -5.0f, + bool use_qk_l2norm = true); + +} // namespace infinicore::op diff --git a/include/infiniop.h b/include/infiniop.h index 95bf75a0d..93891aead 100644 --- a/include/infiniop.h +++ b/include/infiniop.h @@ -75,6 +75,7 @@ #include "infiniop/ops/kron.h" #include "infiniop/ops/kthvalue.h" #include "infiniop/ops/kv_caching.h" +#include "infiniop/ops/kimi_delta_attention.h" #include "infiniop/ops/layer_norm.h" #include "infiniop/ops/ldexp.h" #include "infiniop/ops/lerp.h" diff --git a/include/infiniop/ops/kimi_delta_attention.h b/include/infiniop/ops/kimi_delta_attention.h new file mode 100644 index 000000000..62a58f2ce --- /dev/null +++ b/include/infiniop/ops/kimi_delta_attention.h @@ -0,0 +1,54 @@ +#ifndef __INFINIOP_KIMI_DELTA_ATTENTION_API_H__ +#define __INFINIOP_KIMI_DELTA_ATTENTION_API_H__ + +#include "../operator_descriptor.h" + +typedef struct InfiniopDescriptor *infiniopKimiDeltaAttentionDescriptor_t; + +__INFINI_C __export infiniStatus_t infiniopCreateKimiDeltaAttentionDescriptor( + infiniopHandle_t handle, + infiniopKimiDeltaAttentionDescriptor_t *desc_ptr, + infiniopTensorDescriptor_t out_desc, // [B,T,H,D] or varlen [1,total_tokens,H,D] + infiniopTensorDescriptor_t initial_state_desc, // [B,H,D,D] or indexed pool [pool_size,H,D,D] + infiniopTensorDescriptor_t final_state_desc, // null when final_state_indices_desc is provided + infiniopTensorDescriptor_t q_desc, // [B,T,H,D] or varlen [1,total_tokens,H,D] + infiniopTensorDescriptor_t k_desc, // same shape as q + infiniopTensorDescriptor_t v_desc, // same shape as q + infiniopTensorDescriptor_t g_desc, // raw KDA gate [B,T,H,D] or varlen [1,total_tokens,H,D] + infiniopTensorDescriptor_t beta_desc, // raw beta logits [B,T,H] or varlen [1,total_tokens,H] + infiniopTensorDescriptor_t A_log_desc, // [H], fp32 + infiniopTensorDescriptor_t dt_bias_desc, // [H,D], fp32 + infiniopTensorDescriptor_t cu_seqlens_desc, // nullable; [B + 1], int32/int64 + infiniopTensorDescriptor_t initial_state_indices_desc, // nullable; [B], int32/int64 + infiniopTensorDescriptor_t final_state_indices_desc, // nullable; [B], int32/int64; writes final state in-place to initial_state + float scale, + float lower_bound, + bool use_qk_l2norm); + +__INFINI_C __export infiniStatus_t infiniopGetKimiDeltaAttentionWorkspaceSize( + infiniopKimiDeltaAttentionDescriptor_t desc, + size_t *size); + +__INFINI_C __export infiniStatus_t infiniopKimiDeltaAttention( + infiniopKimiDeltaAttentionDescriptor_t desc, + void *workspace, + size_t workspace_size, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + void *stream); + +__INFINI_C __export infiniStatus_t infiniopDestroyKimiDeltaAttentionDescriptor( + infiniopKimiDeltaAttentionDescriptor_t desc); + +#endif diff --git a/python/infinicore/nn/functional/__init__.py b/python/infinicore/nn/functional/__init__.py index 0395c5626..0d87db29a 100644 --- a/python/infinicore/nn/functional/__init__.py +++ b/python/infinicore/nn/functional/__init__.py @@ -18,6 +18,7 @@ from .hinge_embedding_loss import hinge_embedding_loss from .huber_loss import huber_loss from .interpolate import interpolate +from .kimi_delta_attention import kimi_delta_attention from .layer_norm import layer_norm from .linear import linear from .linear_w8a8i8 import linear_w8a8i8 @@ -56,6 +57,7 @@ "fused_gated_delta_net_gating", "gaussian_nll_loss", "interpolate", + "kimi_delta_attention", "linear", "binary_cross_entropy_with_logits", "random_sample", diff --git a/python/infinicore/nn/functional/kimi_delta_attention.py b/python/infinicore/nn/functional/kimi_delta_attention.py new file mode 100644 index 000000000..482a3a84b --- /dev/null +++ b/python/infinicore/nn/functional/kimi_delta_attention.py @@ -0,0 +1,66 @@ +from infinicore.lib import _infinicore +from infinicore.tensor import Tensor + + +def kimi_delta_attention( + q: Tensor, + k: Tensor, + v: Tensor, + g: Tensor, + beta: Tensor, + A_log: Tensor, + dt_bias: Tensor, + initial_state: Tensor, + *, + cu_seqlens: Tensor | None = None, + initial_state_indices: Tensor | None = None, + final_state_indices: Tensor | None = None, + scale: float = 1.0, + lower_bound: float = -5.0, + use_qk_l2norm: bool = True, +) -> Tensor: + """Run Kimi Delta Attention and return only ``out``. + + Padded mode: + q/k/v/g/out: ``[B, T, H, D]`` + beta: ``[B, T, H]`` + initial_state: ``[B, H, D, D]`` + + Continuous-batch mode: + Pass ``cu_seqlens`` with shape ``[num_requests + 1]``. + q/k/v/g/out: ``[1, total_tokens, H, D]`` + beta: ``[1, total_tokens, H]`` + + Indexed state-pool mode: + initial_state is ``[pool_size, H, D, D]``. + ``initial_state_indices`` and ``final_state_indices`` are both + ``[num_requests]`` int32/int64 tensors. The final state is written + in-place to ``initial_state[final_state_indices]``. + """ + if (initial_state_indices is None) != (final_state_indices is None): + raise ValueError( + "initial_state_indices and final_state_indices must be provided together" + ) + + return Tensor( + _infinicore.kimi_delta_attention( + q._underlying, + k._underlying, + v._underlying, + g._underlying, + beta._underlying, + A_log._underlying, + dt_bias._underlying, + initial_state._underlying, + None if cu_seqlens is None else cu_seqlens._underlying, + ( + None + if initial_state_indices is None + else initial_state_indices._underlying + ), + None if final_state_indices is None else final_state_indices._underlying, + scale, + lower_bound, + use_qk_l2norm, + ) + ) diff --git a/src/infinicore/ops/kimi_delta_attention/kimi_delta_attention.cc b/src/infinicore/ops/kimi_delta_attention/kimi_delta_attention.cc new file mode 100644 index 000000000..d997ccc1a --- /dev/null +++ b/src/infinicore/ops/kimi_delta_attention/kimi_delta_attention.cc @@ -0,0 +1,194 @@ +#include "infinicore/ops/kimi_delta_attention.hpp" + +#include "../../utils.hpp" + +#include + +namespace infinicore::op { + +INFINICORE_GRAPH_OP_DISPATCHERS_IMPL(KimiDeltaAttention); + +KimiDeltaAttention::KimiDeltaAttention(Tensor out, + Tensor initial_state, + std::optional final_state, + const Tensor &q, + const Tensor &k, + const Tensor &v, + const Tensor &g, + const Tensor &beta, + const Tensor &A_log, + const Tensor &dt_bias, + std::optional cu_seqlens, + std::optional initial_state_indices, + std::optional final_state_indices, + float scale, + float lower_bound, + bool use_qk_l2norm) { + INFINICORE_ASSERT_TENSORS_SAME_DEVICE(out, initial_state, q, k, v, g, beta, A_log, dt_bias); + if (final_state.has_value()) { + INFINICORE_ASSERT_TENSORS_SAME_DEVICE(out, final_state.value()); + } + if (cu_seqlens.has_value()) { + INFINICORE_ASSERT_TENSORS_SAME_DEVICE(out, cu_seqlens.value()); + } + if (initial_state_indices.has_value()) { + INFINICORE_ASSERT_TENSORS_SAME_DEVICE(out, initial_state_indices.value()); + } + if (final_state_indices.has_value()) { + INFINICORE_ASSERT_TENSORS_SAME_DEVICE(out, final_state_indices.value()); + } + INFINICORE_GRAPH_OP_DISPATCH(out->device().getType(), + out, + initial_state, + final_state, + q, + k, + v, + g, + beta, + A_log, + dt_bias, + cu_seqlens, + initial_state_indices, + final_state_indices, + scale, + lower_bound, + use_qk_l2norm); +} + +void KimiDeltaAttention::execute(Tensor out, + Tensor initial_state, + std::optional final_state, + const Tensor &q, + const Tensor &k, + const Tensor &v, + const Tensor &g, + const Tensor &beta, + const Tensor &A_log, + const Tensor &dt_bias, + std::optional cu_seqlens, + std::optional initial_state_indices, + std::optional final_state_indices, + float scale, + float lower_bound, + bool use_qk_l2norm) { + INFINICORE_GRAPH_OP_RECORD_OR_RUN(KimiDeltaAttention, + out, + initial_state, + final_state, + q, + k, + v, + g, + beta, + A_log, + dt_bias, + cu_seqlens, + initial_state_indices, + final_state_indices, + scale, + lower_bound, + use_qk_l2norm); +} + +static void check_inputs(const Tensor &q, + const Tensor &k, + const Tensor &v, + const Tensor &g, + const Tensor &beta, + const Tensor &A_log, + const Tensor &dt_bias) { + if (q->shape().size() != 4 || k->shape() != q->shape() || v->shape() != q->shape() || g->shape() != q->shape()) { + throw std::runtime_error("kimi_delta_attention expects q/k/v/g with shape [B, T, H, D] or [1, total_tokens, H, D]"); + } + if (beta->shape().size() != 3 || beta->shape()[0] != q->shape()[0] || beta->shape()[1] != q->shape()[1] || beta->shape()[2] != q->shape()[2]) { + throw std::runtime_error("kimi_delta_attention expects beta with shape [B, T, H] or [1, total_tokens, H]"); + } + if (A_log->shape().size() != 1 || A_log->shape()[0] != q->shape()[2] || dt_bias->shape().size() != 2 || dt_bias->shape()[0] != q->shape()[2] || dt_bias->shape()[1] != q->shape()[3]) { + throw std::runtime_error("kimi_delta_attention expects A_log [H] and dt_bias [H, D]"); + } +} + +static Shape final_state_shape(const Tensor &q, std::optional cu_seqlens) { + size_t B = cu_seqlens.has_value() ? cu_seqlens.value()->shape()[0] - 1 : q->shape()[0]; + return {B, q->shape()[2], q->shape()[3], q->shape()[3]}; +} + +Tensor kimi_delta_attention(const Tensor &q, + const Tensor &k, + const Tensor &v, + const Tensor &g, + const Tensor &beta, + const Tensor &A_log, + const Tensor &dt_bias, + Tensor initial_state, + std::optional cu_seqlens, + std::optional initial_state_indices, + std::optional final_state_indices, + float scale, + float lower_bound, + bool use_qk_l2norm) { + check_inputs(q, k, v, g, beta, A_log, dt_bias); + Tensor out = Tensor::empty(v->shape(), v->dtype(), v->device()); + std::optional final_state = std::nullopt; + if (!final_state_indices.has_value()) { + final_state = Tensor::empty(final_state_shape(q, cu_seqlens), initial_state->dtype(), initial_state->device()); + } + kimi_delta_attention_(out, + initial_state, + final_state, + q, + k, + v, + g, + beta, + A_log, + dt_bias, + cu_seqlens, + initial_state_indices, + final_state_indices, + scale, + lower_bound, + use_qk_l2norm); + return out; +} + +void kimi_delta_attention_(Tensor out, + Tensor initial_state, + std::optional final_state, + const Tensor &q, + const Tensor &k, + const Tensor &v, + const Tensor &g, + const Tensor &beta, + const Tensor &A_log, + const Tensor &dt_bias, + std::optional cu_seqlens, + std::optional initial_state_indices, + std::optional final_state_indices, + float scale, + float lower_bound, + bool use_qk_l2norm) { + check_inputs(q, k, v, g, beta, A_log, dt_bias); + if (out->shape() != v->shape()) { + throw std::runtime_error("kimi_delta_attention_ output shape must match v"); + } + KimiDeltaAttention::execute(out, + initial_state, + final_state, + q, + k, + v, + g, + beta, + A_log, + dt_bias, + cu_seqlens, + initial_state_indices, + final_state_indices, + scale, + lower_bound, + use_qk_l2norm); +} + +} // namespace infinicore::op diff --git a/src/infinicore/ops/kimi_delta_attention/kimi_delta_attention_infiniop.cc b/src/infinicore/ops/kimi_delta_attention/kimi_delta_attention_infiniop.cc new file mode 100644 index 000000000..a1a332eb9 --- /dev/null +++ b/src/infinicore/ops/kimi_delta_attention/kimi_delta_attention_infiniop.cc @@ -0,0 +1,120 @@ +#include "infinicore/ops/kimi_delta_attention.hpp" + +#include "../infiniop_impl.hpp" + +namespace infinicore::op::kimi_delta_attention_impl::infiniop { + +INFINIOP_CACHABLE_DESCRIPTOR(Descriptor, KimiDeltaAttention, 100); + +struct PlannedMeta { + std::shared_ptr descriptor; + graph::GraphTensor workspace, out, initial_state, q, k, v, g, beta, A_log, dt_bias; + std::optional final_state; + std::optional cu_seqlens; + std::optional initial_state_indices; + std::optional final_state_indices; +}; + +void *plan(Tensor out, + Tensor initial_state, + std::optional final_state, + const Tensor &q, + const Tensor &k, + const Tensor &v, + const Tensor &g, + const Tensor &beta, + const Tensor &A_log, + const Tensor &dt_bias, + std::optional cu_seqlens, + std::optional initial_state_indices, + std::optional final_state_indices, + float scale, + float lower_bound, + bool use_qk_l2norm) { + size_t seed = hash_combine(out, + initial_state, + final_state, + q, + k, + v, + g, + beta, + A_log, + dt_bias, + cu_seqlens, + initial_state_indices, + final_state_indices, + scale, + lower_bound, + use_qk_l2norm); + + INFINIOP_CACHABLE_DESCRIPTOR_GET_OR_CREATE( + Descriptor, descriptor, KimiDeltaAttention, seed, + out->desc(), + initial_state->desc(), + final_state.has_value() ? final_state.value()->desc() : nullptr, + q->desc(), + k->desc(), + v->desc(), + g->desc(), + beta->desc(), + A_log->desc(), + dt_bias->desc(), + cu_seqlens.has_value() ? cu_seqlens.value()->desc() : nullptr, + initial_state_indices.has_value() ? initial_state_indices.value()->desc() : nullptr, + final_state_indices.has_value() ? final_state_indices.value()->desc() : nullptr, + scale, + lower_bound, + use_qk_l2norm); + + INFINIOP_WORKSPACE_TENSOR(workspace, KimiDeltaAttention, descriptor); + + return new PlannedMeta{ + descriptor, + graph::GraphTensor(workspace), + graph::GraphTensor(out), + graph::GraphTensor(initial_state), + graph::GraphTensor(q), + graph::GraphTensor(k), + graph::GraphTensor(v), + graph::GraphTensor(g), + graph::GraphTensor(beta), + graph::GraphTensor(A_log), + graph::GraphTensor(dt_bias), + final_state.has_value() ? std::optional(graph::GraphTensor(final_state.value())) : std::nullopt, + cu_seqlens.has_value() ? std::optional(graph::GraphTensor(cu_seqlens.value())) : std::nullopt, + initial_state_indices.has_value() ? std::optional(graph::GraphTensor(initial_state_indices.value())) : std::nullopt, + final_state_indices.has_value() ? std::optional(graph::GraphTensor(final_state_indices.value())) : std::nullopt}; +} + +void run(void *planned_meta) { + auto planned = reinterpret_cast(planned_meta); + + INFINICORE_CHECK_ERROR(infiniopKimiDeltaAttention( + planned->descriptor->desc, + planned->workspace->data(), + planned->workspace->numel(), + planned->out->data(), + planned->initial_state->data(), + planned->final_state.has_value() ? planned->final_state.value()->data() : nullptr, + planned->q->data(), + planned->k->data(), + planned->v->data(), + planned->g->data(), + planned->beta->data(), + planned->A_log->data(), + planned->dt_bias->data(), + planned->cu_seqlens.has_value() ? planned->cu_seqlens.value()->data() : nullptr, + planned->initial_state_indices.has_value() ? planned->initial_state_indices.value()->data() : nullptr, + planned->final_state_indices.has_value() ? planned->final_state_indices.value()->data() : nullptr, + context::getStream())); +} + +void cleanup(void **planned_meta_ptr) { + delete *reinterpret_cast(planned_meta_ptr); + *planned_meta_ptr = nullptr; +} + +INFINICORE_GRAPH_OP_REGISTER_ALLDEVICE(KimiDeltaAttention, &plan, &run, &cleanup); + +} // namespace infinicore::op::kimi_delta_attention_impl::infiniop diff --git a/src/infinicore/pybind11/ops.hpp b/src/infinicore/pybind11/ops.hpp index 087261382..18bbe1e23 100644 --- a/src/infinicore/pybind11/ops.hpp +++ b/src/infinicore/pybind11/ops.hpp @@ -61,6 +61,7 @@ #include "ops/index_copy.hpp" #include "ops/inner.hpp" #include "ops/interpolate.hpp" +#include "ops/kimi_delta_attention.hpp" #include "ops/kron.hpp" #include "ops/kthvalue.hpp" #include "ops/kv_caching.hpp" @@ -170,6 +171,7 @@ inline void bind(py::module &m) { bind_flash_attention(m); bind_hinge_embedding_loss(m); bind_kv_caching(m); + bind_kimi_delta_attention(m); bind_fmod(m); bind_fused_gated_delta_net_gating(m); bind_fmin(m); diff --git a/src/infinicore/pybind11/ops/kimi_delta_attention.hpp b/src/infinicore/pybind11/ops/kimi_delta_attention.hpp new file mode 100644 index 000000000..78f56c214 --- /dev/null +++ b/src/infinicore/pybind11/ops/kimi_delta_attention.hpp @@ -0,0 +1,52 @@ +#pragma once + +#include +#include + +#include "infinicore/ops/kimi_delta_attention.hpp" + +namespace py = pybind11; + +namespace infinicore::ops { + +inline void bind_kimi_delta_attention(py::module &m) { + m.def("kimi_delta_attention", + &op::kimi_delta_attention, + py::arg("q"), + py::arg("k"), + py::arg("v"), + py::arg("g"), + py::arg("beta"), + py::arg("A_log"), + py::arg("dt_bias"), + py::arg("initial_state"), + py::arg("cu_seqlens") = std::nullopt, + py::arg("initial_state_indices") = std::nullopt, + py::arg("final_state_indices") = std::nullopt, + py::arg("scale") = 1.0f, + py::arg("lower_bound") = -5.0f, + py::arg("use_qk_l2norm") = true, + R"doc(Kimi Delta Attention out-of-place.)doc"); + + m.def("kimi_delta_attention_", + &op::kimi_delta_attention_, + py::arg("out"), + py::arg("initial_state"), + py::arg("final_state"), + py::arg("q"), + py::arg("k"), + py::arg("v"), + py::arg("g"), + py::arg("beta"), + py::arg("A_log"), + py::arg("dt_bias"), + py::arg("cu_seqlens") = std::nullopt, + py::arg("initial_state_indices") = std::nullopt, + py::arg("final_state_indices") = std::nullopt, + py::arg("scale") = 1.0f, + py::arg("lower_bound") = -5.0f, + py::arg("use_qk_l2norm") = true, + R"doc(Kimi Delta Attention writing to provided output/state.)doc"); +} + +} // namespace infinicore::ops diff --git a/src/infiniop/ops/kimi_delta_attention/cuda/kernel.cuh b/src/infiniop/ops/kimi_delta_attention/cuda/kernel.cuh new file mode 100644 index 000000000..9ab83c1a3 --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/cuda/kernel.cuh @@ -0,0 +1,382 @@ +#ifndef __KIMI_DELTA_ATTENTION_CUDA_KERNEL_CUH__ +#define __KIMI_DELTA_ATTENTION_CUDA_KERNEL_CUH__ + +#include +#include + +template +__device__ inline float kdaLoadAsFloat(const T *ptr, ptrdiff_t offset) { + return static_cast(ptr[offset]); +} + +template <> +__device__ inline float kdaLoadAsFloat(const half *ptr, ptrdiff_t offset) { + return __half2float(ptr[offset]); +} + +template <> +__device__ inline float kdaLoadAsFloat<__nv_bfloat16>(const __nv_bfloat16 *ptr, ptrdiff_t offset) { + return __bfloat162float(ptr[offset]); +} + +__device__ inline int64_t kdaLoadOptionalIndex(const void *indices, + bool is_i64, + int idx, + int fallback) { + if (indices == nullptr) { + return static_cast(fallback); + } + return is_i64 + ? static_cast(indices)[idx] + : static_cast(static_cast(indices)[idx]); +} + +__device__ inline float kdaSigmoid(float x) { + if (x >= 0.0f) { + float z = expf(-x); + return 1.0f / (1.0f + z); + } + float z = expf(x); + return z / (1.0f + z); +} + +__device__ inline float kdaBlockReduceSum(float value, float *scratch) { + scratch[threadIdx.x] = value; + __syncthreads(); + for (int stride = blockDim.x / 2; stride > 0; stride >>= 1) { + if (threadIdx.x < stride) { + scratch[threadIdx.x] += scratch[threadIdx.x + stride]; + } + __syncthreads(); + } + return scratch[0]; +} + +template +__global__ void kimiDeltaAttentionDecodeCudaKernel( + Tdata *out, + Tdata *initial_state, + Tdata *final_state, + const Tdata *q, + const Tdata *k, + const Tdata *v, + const Tgate *g, + const Tgate *beta, + const float *A_log, + const float *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + bool cu_seqlens_i64, + bool initial_state_indices_i64, + bool final_state_indices_i64, + bool use_qk_l2norm, + bool has_cu_seqlens, + bool indexed_state_pool, + size_t D, + size_t pool_size, + float scale, + float lower_bound, + ptrdiff_t out_s0, + ptrdiff_t out_s1, + ptrdiff_t out_s2, + ptrdiff_t initial_s0, + ptrdiff_t initial_s1, + ptrdiff_t initial_s2, + ptrdiff_t initial_s3, + ptrdiff_t final_s0, + ptrdiff_t final_s1, + ptrdiff_t final_s2, + ptrdiff_t final_s3, + ptrdiff_t q_s0, + ptrdiff_t q_s1, + ptrdiff_t q_s2, + ptrdiff_t k_s0, + ptrdiff_t k_s1, + ptrdiff_t k_s2, + ptrdiff_t v_s0, + ptrdiff_t v_s1, + ptrdiff_t v_s2, + ptrdiff_t g_s0, + ptrdiff_t g_s1, + ptrdiff_t g_s2, + ptrdiff_t beta_s0, + ptrdiff_t beta_s1, + ptrdiff_t beta_s2, + ptrdiff_t A_log_s0, + ptrdiff_t dt_bias_s0) { + + const int batch_idx = blockIdx.x; + const int head_idx = blockIdx.y; + const int value_dim_idx = blockIdx.z; + + extern __shared__ float scratch[]; + + int64_t token_idx = 0; + if (has_cu_seqlens) { + token_idx = kdaLoadOptionalIndex(cu_seqlens, cu_seqlens_i64, batch_idx, 0); + } + + int64_t read_slot = batch_idx; + int64_t write_slot = batch_idx; + if (indexed_state_pool) { + read_slot = kdaLoadOptionalIndex(initial_state_indices, initial_state_indices_i64, batch_idx, batch_idx); + write_slot = final_state_indices == nullptr + ? static_cast(batch_idx) + : kdaLoadOptionalIndex(final_state_indices, final_state_indices_i64, batch_idx, batch_idx); + if (read_slot < 0 || write_slot < 0 || read_slot >= static_cast(pool_size) || write_slot >= static_cast(pool_size)) { + if (threadIdx.x == 0) { + const int token_batch = has_cu_seqlens ? 0 : batch_idx; + const ptrdiff_t out_base = static_cast(token_batch) * out_s0 + static_cast(token_idx) * out_s1 + static_cast(head_idx) * out_s2; + out[out_base + value_dim_idx] = static_cast(0.0f); + } + return; + } + } + + const int token_batch = has_cu_seqlens ? 0 : batch_idx; + const ptrdiff_t q_base = static_cast(token_batch) * q_s0 + static_cast(token_idx) * q_s1 + static_cast(head_idx) * q_s2; + const ptrdiff_t k_base = static_cast(token_batch) * k_s0 + static_cast(token_idx) * k_s1 + static_cast(head_idx) * k_s2; + const ptrdiff_t v_base = static_cast(token_batch) * v_s0 + static_cast(token_idx) * v_s1 + static_cast(head_idx) * v_s2; + const ptrdiff_t g_base = static_cast(token_batch) * g_s0 + static_cast(token_idx) * g_s1 + static_cast(head_idx) * g_s2; + const ptrdiff_t beta_offset = static_cast(token_batch) * beta_s0 + static_cast(token_idx) * beta_s1 + static_cast(head_idx) * beta_s2; + + float q_sum = 0.0f; + float k_sum = 0.0f; + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + float q_raw = kdaLoadAsFloat(q, q_base + dk); + float k_raw = kdaLoadAsFloat(k, k_base + dk); + q_sum += q_raw * q_raw; + k_sum += k_raw * k_raw; + } + q_sum = kdaBlockReduceSum(q_sum, scratch); + k_sum = kdaBlockReduceSum(k_sum, scratch); + + const float q_scale = use_qk_l2norm ? rsqrtf(q_sum + 1e-6f) * scale : scale; + const float k_scale = use_qk_l2norm ? rsqrtf(k_sum + 1e-6f) : 1.0f; + const float a_log_exp = expf(A_log[static_cast(head_idx) * A_log_s0]); + + const ptrdiff_t initial_base = static_cast(read_slot) * initial_s0 + + static_cast(head_idx) * initial_s1 + + static_cast(value_dim_idx) * initial_s2; + + float kv_mem = 0.0f; + float hq_mem = 0.0f; + float kq_mem = 0.0f; + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + float q_t = kdaLoadAsFloat(q, q_base + dk) * q_scale; + float k_t = kdaLoadAsFloat(k, k_base + dk) * k_scale; + float state = kdaLoadAsFloat(initial_state, initial_base + static_cast(dk) * initial_s3); + float raw_gate = kdaLoadAsFloat(g, g_base + dk) + dt_bias[static_cast(head_idx) * dt_bias_s0 + dk]; + float decay = expf(lower_bound * kdaSigmoid(a_log_exp * raw_gate)); + float decayed_state = state * decay; + kv_mem += decayed_state * k_t; + hq_mem += decayed_state * q_t; + kq_mem += k_t * q_t; + } + kv_mem = kdaBlockReduceSum(kv_mem, scratch); + hq_mem = kdaBlockReduceSum(hq_mem, scratch); + kq_mem = kdaBlockReduceSum(kq_mem, scratch); + + const float beta_t = kdaSigmoid(kdaLoadAsFloat(beta, beta_offset)); + const float v_t = kdaLoadAsFloat(v, v_base + value_dim_idx); + const float delta = (v_t - kv_mem) * beta_t; + + if (threadIdx.x == 0) { + const ptrdiff_t out_base = static_cast(token_batch) * out_s0 + static_cast(token_idx) * out_s1 + static_cast(head_idx) * out_s2; + out[out_base + value_dim_idx] = static_cast(hq_mem + delta * kq_mem); + } + + Tdata *final_state_target = final_state_indices == nullptr ? final_state : initial_state; + const ptrdiff_t final_base = final_state_indices == nullptr + ? static_cast(batch_idx) * final_s0 + + static_cast(head_idx) * final_s1 + + static_cast(value_dim_idx) * final_s2 + : static_cast(write_slot) * initial_s0 + + static_cast(head_idx) * initial_s1 + + static_cast(value_dim_idx) * initial_s2; + const ptrdiff_t final_k_stride = final_state_indices == nullptr ? final_s3 : initial_s3; + + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + float k_t = kdaLoadAsFloat(k, k_base + dk) * k_scale; + float state = kdaLoadAsFloat(initial_state, initial_base + static_cast(dk) * initial_s3); + float raw_gate = kdaLoadAsFloat(g, g_base + dk) + dt_bias[static_cast(head_idx) * dt_bias_s0 + dk]; + float decay = expf(lower_bound * kdaSigmoid(a_log_exp * raw_gate)); + final_state_target[final_base + static_cast(dk) * final_k_stride] = static_cast(state * decay + k_t * delta); + } +} + +template +__global__ void kimiDeltaAttentionRecurrentCudaKernel( + Tdata *out, + Tdata *initial_state, + Tdata *final_state, + const Tdata *q, + const Tdata *k, + const Tdata *v, + const Tgate *g, + const Tgate *beta, + const float *A_log, + const float *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + bool cu_seqlens_i64, + bool initial_state_indices_i64, + bool final_state_indices_i64, + bool use_qk_l2norm, + bool has_cu_seqlens, + bool indexed_state_pool, + size_t T, + size_t D, + size_t pool_size, + float scale, + float lower_bound, + ptrdiff_t out_s0, + ptrdiff_t out_s1, + ptrdiff_t out_s2, + ptrdiff_t initial_s0, + ptrdiff_t initial_s1, + ptrdiff_t initial_s2, + ptrdiff_t initial_s3, + ptrdiff_t final_s0, + ptrdiff_t final_s1, + ptrdiff_t final_s2, + ptrdiff_t final_s3, + ptrdiff_t q_s0, + ptrdiff_t q_s1, + ptrdiff_t q_s2, + ptrdiff_t k_s0, + ptrdiff_t k_s1, + ptrdiff_t k_s2, + ptrdiff_t v_s0, + ptrdiff_t v_s1, + ptrdiff_t v_s2, + ptrdiff_t g_s0, + ptrdiff_t g_s1, + ptrdiff_t g_s2, + ptrdiff_t beta_s0, + ptrdiff_t beta_s1, + ptrdiff_t beta_s2, + ptrdiff_t A_log_s0, + ptrdiff_t dt_bias_s0) { + + const int batch_idx = blockIdx.x; + const int head_idx = blockIdx.y; + const int value_dim_idx = blockIdx.z; + + extern __shared__ float shared[]; + float *state = shared; + float *q_vec = state + D; + float *k_vec = q_vec + D; + float *scratch = k_vec + D; + + int64_t token_begin = 0; + int64_t token_end = static_cast(T); + if (has_cu_seqlens) { + token_begin = kdaLoadOptionalIndex(cu_seqlens, cu_seqlens_i64, batch_idx, 0); + token_end = kdaLoadOptionalIndex(cu_seqlens, cu_seqlens_i64, batch_idx + 1, 0); + if (token_begin < 0 || token_end < token_begin || token_end > static_cast(T)) { + return; + } + } + + int64_t read_slot = batch_idx; + int64_t write_slot = batch_idx; + if (indexed_state_pool) { + read_slot = kdaLoadOptionalIndex(initial_state_indices, initial_state_indices_i64, batch_idx, batch_idx); + write_slot = final_state_indices == nullptr + ? static_cast(batch_idx) + : kdaLoadOptionalIndex(final_state_indices, final_state_indices_i64, batch_idx, batch_idx); + if (read_slot < 0 || write_slot < 0 || read_slot >= static_cast(pool_size) || write_slot >= static_cast(pool_size)) { + return; + } + } + + const ptrdiff_t initial_base = static_cast(read_slot) * initial_s0 + + static_cast(head_idx) * initial_s1 + + static_cast(value_dim_idx) * initial_s2; + + Tdata *final_state_target = final_state_indices == nullptr ? final_state : initial_state; + const ptrdiff_t final_base = final_state_indices == nullptr + ? static_cast(batch_idx) * final_s0 + + static_cast(head_idx) * final_s1 + + static_cast(value_dim_idx) * final_s2 + : static_cast(write_slot) * initial_s0 + + static_cast(head_idx) * initial_s1 + + static_cast(value_dim_idx) * initial_s2; + const ptrdiff_t final_k_stride = final_state_indices == nullptr ? final_s3 : initial_s3; + + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + state[dk] = kdaLoadAsFloat(initial_state, initial_base + static_cast(dk) * initial_s3); + } + __syncthreads(); + + const int token_batch = has_cu_seqlens ? 0 : batch_idx; + const float a_log_exp = expf(A_log[static_cast(head_idx) * A_log_s0]); + + for (int64_t token_idx = token_begin; token_idx < token_end; ++token_idx) { + const ptrdiff_t q_base = static_cast(token_batch) * q_s0 + static_cast(token_idx) * q_s1 + static_cast(head_idx) * q_s2; + const ptrdiff_t k_base = static_cast(token_batch) * k_s0 + static_cast(token_idx) * k_s1 + static_cast(head_idx) * k_s2; + const ptrdiff_t g_base = static_cast(token_batch) * g_s0 + static_cast(token_idx) * g_s1 + static_cast(head_idx) * g_s2; + + float q_sum = 0.0f; + float k_sum = 0.0f; + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + float q_raw = kdaLoadAsFloat(q, q_base + dk); + float k_raw = kdaLoadAsFloat(k, k_base + dk); + q_vec[dk] = q_raw; + k_vec[dk] = k_raw; + q_sum += q_raw * q_raw; + k_sum += k_raw * k_raw; + } + q_sum = kdaBlockReduceSum(q_sum, scratch); + k_sum = kdaBlockReduceSum(k_sum, scratch); + + const float q_scale = use_qk_l2norm ? rsqrtf(q_sum + 1e-6f) * scale : scale; + const float k_scale = use_qk_l2norm ? rsqrtf(k_sum + 1e-6f) : 1.0f; + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + q_vec[dk] *= q_scale; + k_vec[dk] *= k_scale; + } + __syncthreads(); + + float kv_mem = 0.0f; + float hq_mem = 0.0f; + float kq_mem = 0.0f; + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + float raw_gate = kdaLoadAsFloat(g, g_base + dk) + dt_bias[static_cast(head_idx) * dt_bias_s0 + dk]; + float decay = expf(lower_bound * kdaSigmoid(a_log_exp * raw_gate)); + kv_mem += state[dk] * decay * k_vec[dk]; + hq_mem += state[dk] * decay * q_vec[dk]; + kq_mem += k_vec[dk] * q_vec[dk]; + } + kv_mem = kdaBlockReduceSum(kv_mem, scratch); + hq_mem = kdaBlockReduceSum(hq_mem, scratch); + kq_mem = kdaBlockReduceSum(kq_mem, scratch); + + const ptrdiff_t beta_offset = static_cast(token_batch) * beta_s0 + static_cast(token_idx) * beta_s1 + static_cast(head_idx) * beta_s2; + const float beta_t = kdaSigmoid(kdaLoadAsFloat(beta, beta_offset)); + const ptrdiff_t v_base = static_cast(token_batch) * v_s0 + static_cast(token_idx) * v_s1 + static_cast(head_idx) * v_s2; + const float v_t = kdaLoadAsFloat(v, v_base + value_dim_idx); + const float delta = (v_t - kv_mem) * beta_t; + + if (threadIdx.x == 0) { + const ptrdiff_t out_base = static_cast(token_batch) * out_s0 + static_cast(token_idx) * out_s1 + static_cast(head_idx) * out_s2; + out[out_base + value_dim_idx] = static_cast(hq_mem + delta * kq_mem); + } + + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + float raw_gate = kdaLoadAsFloat(g, g_base + dk) + dt_bias[static_cast(head_idx) * dt_bias_s0 + dk]; + float decay = expf(lower_bound * kdaSigmoid(a_log_exp * raw_gate)); + state[dk] = state[dk] * decay + k_vec[dk] * delta; + } + __syncthreads(); + } + + for (int dk = threadIdx.x; dk < static_cast(D); dk += blockDim.x) { + final_state_target[final_base + static_cast(dk) * final_k_stride] = static_cast(state[dk]); + } +} + +#endif diff --git a/src/infiniop/ops/kimi_delta_attention/info.h b/src/infiniop/ops/kimi_delta_attention/info.h new file mode 100644 index 000000000..c75e22cbd --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/info.h @@ -0,0 +1,199 @@ +#ifndef __KIMI_DELTA_ATTENTION_INFO_H__ +#define __KIMI_DELTA_ATTENTION_INFO_H__ + +#include "../../../utils.h" +#include "../../tensor.h" + +#include + +namespace op::kimi_delta_attention { + +class KimiDeltaAttentionInfo { + KimiDeltaAttentionInfo() = default; + +public: + infiniDtype_t data_dtype; + infiniDtype_t gate_dtype; + infiniDtype_t cu_seqlens_dtype; + infiniDtype_t initial_state_indices_dtype; + infiniDtype_t final_state_indices_dtype; + + bool has_cu_seqlens; + bool has_initial_state_indices; + bool has_final_state_indices; + bool indexed_state_pool; + bool is_decode; + bool use_qk_l2norm; + + size_t B, T, total_tokens, H, D, pool_size; + float scale; + float lower_bound; + + std::vector out_strides; + std::vector initial_state_strides; + std::vector final_state_strides; + std::vector q_strides; + std::vector k_strides; + std::vector v_strides; + std::vector g_strides; + std::vector beta_strides; + std::vector A_log_strides; + std::vector dt_bias_strides; + + static utils::Result + create(infiniopTensorDescriptor_t out_desc, + infiniopTensorDescriptor_t initial_state_desc, + infiniopTensorDescriptor_t final_state_desc, + infiniopTensorDescriptor_t q_desc, + infiniopTensorDescriptor_t k_desc, + infiniopTensorDescriptor_t v_desc, + infiniopTensorDescriptor_t g_desc, + infiniopTensorDescriptor_t beta_desc, + infiniopTensorDescriptor_t A_log_desc, + infiniopTensorDescriptor_t dt_bias_desc, + infiniopTensorDescriptor_t cu_seqlens_desc, + infiniopTensorDescriptor_t initial_state_indices_desc, + infiniopTensorDescriptor_t final_state_indices_desc, + float scale, + float lower_bound, + bool use_qk_l2norm) { + + if (out_desc == nullptr || initial_state_desc == nullptr || q_desc == nullptr || k_desc == nullptr || v_desc == nullptr || g_desc == nullptr || beta_desc == nullptr || A_log_desc == nullptr || dt_bias_desc == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + + auto data_dtype = q_desc->dtype(); + CHECK_DTYPE(data_dtype, INFINI_DTYPE_F16, INFINI_DTYPE_BF16, INFINI_DTYPE_F32); + if (k_desc->dtype() != data_dtype || v_desc->dtype() != data_dtype || out_desc->dtype() != data_dtype || initial_state_desc->dtype() != data_dtype || (final_state_desc != nullptr && final_state_desc->dtype() != data_dtype)) { + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } + + auto gate_dtype = g_desc->dtype(); + CHECK_DTYPE(gate_dtype, INFINI_DTYPE_F16, INFINI_DTYPE_BF16, INFINI_DTYPE_F32); + if (beta_desc->dtype() != gate_dtype || A_log_desc->dtype() != INFINI_DTYPE_F32 || dt_bias_desc->dtype() != INFINI_DTYPE_F32) { + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } + + const bool has_cu = cu_seqlens_desc != nullptr; + const bool has_initial_indices = initial_state_indices_desc != nullptr; + const bool has_final_indices = final_state_indices_desc != nullptr; + const bool indexed_pool = has_initial_indices || has_final_indices; + if (has_final_indices && final_state_desc != nullptr) { + return INFINI_STATUS_BAD_PARAM; + } + if (!has_final_indices && final_state_desc == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + + if (q_desc->ndim() != 4 || k_desc->ndim() != 4 || v_desc->ndim() != 4 || g_desc->ndim() != 4 || out_desc->ndim() != 4 || beta_desc->ndim() != 3 || A_log_desc->ndim() != 1 || dt_bias_desc->ndim() != 2 || initial_state_desc->ndim() != 4 || (!has_final_indices && final_state_desc->ndim() != 4)) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + + auto q_shape = q_desc->shape(); + auto k_shape = k_desc->shape(); + auto v_shape = v_desc->shape(); + auto g_shape = g_desc->shape(); + auto out_shape = out_desc->shape(); + auto beta_shape = beta_desc->shape(); + + size_t B = q_shape[0], T = q_shape[1], H = q_shape[2], D = q_shape[3], total_tokens = T; + if (has_cu) { + if (cu_seqlens_desc->ndim() != 1 || cu_seqlens_desc->shape()[0] < 2) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + B = cu_seqlens_desc->shape()[0] - 1; + if (q_shape[0] != 1 || k_shape[0] != 1 || v_shape[0] != 1 || g_shape[0] != 1 || out_shape[0] != 1 || beta_shape[0] != 1) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + total_tokens = q_shape[1]; + T = total_tokens; + } + + if (k_shape != q_shape || v_shape != q_shape || g_shape != q_shape || out_shape != q_shape || beta_shape[0] != q_shape[0] || beta_shape[1] != q_shape[1] || beta_shape[2] != H || H == 0 || D == 0) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + if (A_log_desc->shape()[0] != H || dt_bias_desc->shape()[0] != H || dt_bias_desc->shape()[1] != D) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + if (q_desc->strides()[3] != 1 || k_desc->strides()[3] != 1 || v_desc->strides()[3] != 1 || g_desc->strides()[3] != 1 || out_desc->strides()[3] != 1 || dt_bias_desc->strides()[1] != 1) { + return INFINI_STATUS_BAD_TENSOR_STRIDES; + } + + auto initial_shape = initial_state_desc->shape(); + size_t pool_size = initial_shape[0]; + if (indexed_pool) { + if (initial_shape[1] != H || initial_shape[2] != D || initial_shape[3] != D) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + } else if (initial_shape[0] != B || initial_shape[1] != H || initial_shape[2] != D || initial_shape[3] != D) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + + if (!has_final_indices) { + auto final_shape = final_state_desc->shape(); + if (final_shape[0] != B || final_shape[1] != H || final_shape[2] != D || final_shape[3] != D) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + } + + infiniDtype_t cu_dtype = INFINI_DTYPE_INVALID; + infiniDtype_t initial_indices_dtype = INFINI_DTYPE_INVALID; + infiniDtype_t final_indices_dtype = INFINI_DTYPE_INVALID; + if (has_cu) { + cu_dtype = cu_seqlens_desc->dtype(); + CHECK_DTYPE(cu_dtype, INFINI_DTYPE_I32, INFINI_DTYPE_I64); + } + if (has_initial_indices) { + if (initial_state_indices_desc->ndim() != 1 || initial_state_indices_desc->shape()[0] != B) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + initial_indices_dtype = initial_state_indices_desc->dtype(); + CHECK_DTYPE(initial_indices_dtype, INFINI_DTYPE_I32, INFINI_DTYPE_I64); + } + if (has_final_indices) { + if (final_state_indices_desc->ndim() != 1 || final_state_indices_desc->shape()[0] != B) { + return INFINI_STATUS_BAD_TENSOR_SHAPE; + } + final_indices_dtype = final_state_indices_desc->dtype(); + CHECK_DTYPE(final_indices_dtype, INFINI_DTYPE_I32, INFINI_DTYPE_I64); + } + + KimiDeltaAttentionInfo info; + info.data_dtype = data_dtype; + info.gate_dtype = gate_dtype; + info.cu_seqlens_dtype = cu_dtype; + info.initial_state_indices_dtype = initial_indices_dtype; + info.final_state_indices_dtype = final_indices_dtype; + info.has_cu_seqlens = has_cu; + info.has_initial_state_indices = has_initial_indices; + info.has_final_state_indices = has_final_indices; + info.indexed_state_pool = indexed_pool; + info.is_decode = has_cu ? (total_tokens == B) : (T == 1); + info.use_qk_l2norm = use_qk_l2norm; + info.B = B; + info.T = T; + info.total_tokens = total_tokens; + info.H = H; + info.D = D; + info.pool_size = pool_size; + info.scale = scale; + info.lower_bound = lower_bound; + info.out_strides = out_desc->strides(); + info.initial_state_strides = initial_state_desc->strides(); + if (final_state_desc != nullptr) { + info.final_state_strides = final_state_desc->strides(); + } + info.q_strides = q_desc->strides(); + info.k_strides = k_desc->strides(); + info.v_strides = v_desc->strides(); + info.g_strides = g_desc->strides(); + info.beta_strides = beta_desc->strides(); + info.A_log_strides = A_log_desc->strides(); + info.dt_bias_strides = dt_bias_desc->strides(); + return utils::Result(info); + } +}; + +} // namespace op::kimi_delta_attention + +#endif diff --git a/src/infiniop/ops/kimi_delta_attention/kimi_delta_attention.h b/src/infiniop/ops/kimi_delta_attention/kimi_delta_attention.h new file mode 100644 index 000000000..58727b9da --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/kimi_delta_attention.h @@ -0,0 +1,63 @@ +#ifndef __INFINIOP_KIMI_DELTA_ATTENTION_H__ +#define __INFINIOP_KIMI_DELTA_ATTENTION_H__ + +#include "../../operator.h" +#include "info.h" + +#define DESCRIPTOR(NAMESPACE) \ + \ + namespace op::kimi_delta_attention::NAMESPACE { \ + class Descriptor final : public InfiniopDescriptor { \ + struct Opaque; \ + Opaque *_opaque; \ + KimiDeltaAttentionInfo _info; \ + size_t _workspace_size; \ + \ + Descriptor(Opaque *opaque, \ + KimiDeltaAttentionInfo info, \ + size_t workspace_size, \ + infiniDevice_t device_type, \ + int device_id) \ + : InfiniopDescriptor{device_type, device_id}, \ + _opaque(opaque), \ + _info(info), \ + _workspace_size(workspace_size) {} \ + \ + public: \ + ~Descriptor(); \ + \ + size_t workspaceSize() const { return _workspace_size; } \ + \ + static infiniStatus_t create( \ + infiniopHandle_t handle, \ + Descriptor **desc_ptr, \ + infiniopTensorDescriptor_t out_desc, \ + infiniopTensorDescriptor_t initial_state_desc, \ + infiniopTensorDescriptor_t final_state_desc, \ + infiniopTensorDescriptor_t q_desc, \ + infiniopTensorDescriptor_t k_desc, \ + infiniopTensorDescriptor_t v_desc, \ + infiniopTensorDescriptor_t g_desc, \ + infiniopTensorDescriptor_t beta_desc, \ + infiniopTensorDescriptor_t A_log_desc, \ + infiniopTensorDescriptor_t dt_bias_desc, \ + infiniopTensorDescriptor_t cu_seqlens_desc, \ + infiniopTensorDescriptor_t initial_state_indices_desc, \ + infiniopTensorDescriptor_t final_state_indices_desc, \ + float scale, \ + float lower_bound, \ + bool use_qk_l2norm); \ + \ + infiniStatus_t calculate( \ + void *workspace, size_t workspace_size, \ + void *out, void *initial_state, void *final_state, \ + const void *q, const void *k, const void *v, \ + const void *g, const void *beta, const void *A_log, \ + const void *dt_bias, const void *cu_seqlens, \ + const void *initial_state_indices, \ + const void *final_state_indices, \ + void *stream) const; \ + }; \ + } + +#endif diff --git a/src/infiniop/ops/kimi_delta_attention/metax/kimi_delta_attention_metax.h b/src/infiniop/ops/kimi_delta_attention/metax/kimi_delta_attention_metax.h new file mode 100644 index 000000000..14e069a6a --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/metax/kimi_delta_attention_metax.h @@ -0,0 +1,8 @@ +#ifndef __KIMI_DELTA_ATTENTION_METAX_H__ +#define __KIMI_DELTA_ATTENTION_METAX_H__ + +#include "../kimi_delta_attention.h" + +DESCRIPTOR(metax) + +#endif diff --git a/src/infiniop/ops/kimi_delta_attention/metax/kimi_delta_attention_metax.maca b/src/infiniop/ops/kimi_delta_attention/metax/kimi_delta_attention_metax.maca new file mode 100644 index 000000000..ac138a10e --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/metax/kimi_delta_attention_metax.maca @@ -0,0 +1,273 @@ +#include "../../../devices/metax/metax_common.h" +#include "../../../devices/metax/metax_kernel_common.h" +#include "kimi_delta_attention_metax.h" + +#include "../cuda/kernel.cuh" + +namespace op::kimi_delta_attention::metax { + +struct Descriptor::Opaque { + std::shared_ptr internal; +}; + +Descriptor::~Descriptor() { + delete _opaque; +} + +infiniStatus_t Descriptor::create( + infiniopHandle_t handle, + Descriptor **desc_ptr, + infiniopTensorDescriptor_t out_desc, + infiniopTensorDescriptor_t initial_state_desc, + infiniopTensorDescriptor_t final_state_desc, + infiniopTensorDescriptor_t q_desc, + infiniopTensorDescriptor_t k_desc, + infiniopTensorDescriptor_t v_desc, + infiniopTensorDescriptor_t g_desc, + infiniopTensorDescriptor_t beta_desc, + infiniopTensorDescriptor_t A_log_desc, + infiniopTensorDescriptor_t dt_bias_desc, + infiniopTensorDescriptor_t cu_seqlens_desc, + infiniopTensorDescriptor_t initial_state_indices_desc, + infiniopTensorDescriptor_t final_state_indices_desc, + float scale, + float lower_bound, + bool use_qk_l2norm) { + + auto info = KimiDeltaAttentionInfo::create( + out_desc, + initial_state_desc, + final_state_desc, + q_desc, + k_desc, + v_desc, + g_desc, + beta_desc, + A_log_desc, + dt_bias_desc, + cu_seqlens_desc, + initial_state_indices_desc, + final_state_indices_desc, + scale, + lower_bound, + use_qk_l2norm); + CHECK_RESULT(info); + + *desc_ptr = new Descriptor( + new Opaque{reinterpret_cast(handle)->internal()}, + info.take(), + 0, + handle->device, + handle->device_id); + return INFINI_STATUS_SUCCESS; +} + +template +static infiniStatus_t launch_fallback(const KimiDeltaAttentionInfo &info, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + hcStream_t stream) { + constexpr int threads = 256; + dim3 grid(static_cast(info.B), static_cast(info.H), static_cast(info.D)); + size_t shared = info.is_decode ? threads * sizeof(float) : (info.D * 3 + threads) * sizeof(float); + + // TODO(kimi_delta_attention): Dispatch MoonshotAI FlashKDA SM90 fast path for + // BF16, D=128, contiguous tensors. Original source: + // https://github.com/MoonshotAI/FlashKDA/tree/master/csrc + if (info.is_decode) { + kimiDeltaAttentionDecodeCudaKernel<<>>( + static_cast(out), + static_cast(initial_state), + static_cast(final_state), + static_cast(q), + static_cast(k), + static_cast(v), + static_cast(g), + static_cast(beta), + static_cast(A_log), + static_cast(dt_bias), + cu_seqlens, + initial_state_indices, + final_state_indices, + info.cu_seqlens_dtype == INFINI_DTYPE_I64, + info.initial_state_indices_dtype == INFINI_DTYPE_I64, + info.final_state_indices_dtype == INFINI_DTYPE_I64, + info.use_qk_l2norm, + info.has_cu_seqlens, + info.indexed_state_pool, + info.D, + info.pool_size, + info.scale, + info.lower_bound, + info.out_strides[0], + info.out_strides[1], + info.out_strides[2], + info.initial_state_strides[0], + info.initial_state_strides[1], + info.initial_state_strides[2], + info.initial_state_strides[3], + info.final_state_strides.empty() ? 0 : info.final_state_strides[0], + info.final_state_strides.empty() ? 0 : info.final_state_strides[1], + info.final_state_strides.empty() ? 0 : info.final_state_strides[2], + info.final_state_strides.empty() ? 0 : info.final_state_strides[3], + info.q_strides[0], + info.q_strides[1], + info.q_strides[2], + info.k_strides[0], + info.k_strides[1], + info.k_strides[2], + info.v_strides[0], + info.v_strides[1], + info.v_strides[2], + info.g_strides[0], + info.g_strides[1], + info.g_strides[2], + info.beta_strides[0], + info.beta_strides[1], + info.beta_strides[2], + info.A_log_strides[0], + info.dt_bias_strides[0]); + } else { + kimiDeltaAttentionRecurrentCudaKernel<<>>( + static_cast(out), + static_cast(initial_state), + static_cast(final_state), + static_cast(q), + static_cast(k), + static_cast(v), + static_cast(g), + static_cast(beta), + static_cast(A_log), + static_cast(dt_bias), + cu_seqlens, + initial_state_indices, + final_state_indices, + info.cu_seqlens_dtype == INFINI_DTYPE_I64, + info.initial_state_indices_dtype == INFINI_DTYPE_I64, + info.final_state_indices_dtype == INFINI_DTYPE_I64, + info.use_qk_l2norm, + info.has_cu_seqlens, + info.indexed_state_pool, + info.T, + info.D, + info.pool_size, + info.scale, + info.lower_bound, + info.out_strides[0], + info.out_strides[1], + info.out_strides[2], + info.initial_state_strides[0], + info.initial_state_strides[1], + info.initial_state_strides[2], + info.initial_state_strides[3], + info.final_state_strides.empty() ? 0 : info.final_state_strides[0], + info.final_state_strides.empty() ? 0 : info.final_state_strides[1], + info.final_state_strides.empty() ? 0 : info.final_state_strides[2], + info.final_state_strides.empty() ? 0 : info.final_state_strides[3], + info.q_strides[0], + info.q_strides[1], + info.q_strides[2], + info.k_strides[0], + info.k_strides[1], + info.k_strides[2], + info.v_strides[0], + info.v_strides[1], + info.v_strides[2], + info.g_strides[0], + info.g_strides[1], + info.g_strides[2], + info.beta_strides[0], + info.beta_strides[1], + info.beta_strides[2], + info.A_log_strides[0], + info.dt_bias_strides[0]); + } + return INFINI_STATUS_SUCCESS; +} + +template +static infiniStatus_t launch_for_gate(const KimiDeltaAttentionInfo &info, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + hcStream_t stream) { + switch (info.gate_dtype) { + case INFINI_DTYPE_F16: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_BF16: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_F32: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + default: + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } +} + +infiniStatus_t Descriptor::calculate( + void *workspace, + size_t workspace_size, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + void *stream_) const { + if (workspace_size < _workspace_size) { + return INFINI_STATUS_INSUFFICIENT_WORKSPACE; + } + if (_info.has_cu_seqlens && cu_seqlens == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (_info.has_initial_state_indices && initial_state_indices == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (_info.has_final_state_indices && final_state_indices == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (!_info.has_final_state_indices && final_state == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + + hcStream_t stream = reinterpret_cast(stream_); + switch (_info.data_dtype) { + case INFINI_DTYPE_F16: + return launch_for_gate(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_BF16: + return launch_for_gate<__nv_bfloat16>(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_F32: + return launch_for_gate(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + default: + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } +} + +} // namespace op::kimi_delta_attention::metax diff --git a/src/infiniop/ops/kimi_delta_attention/moore/kimi_delta_attention_moore.h b/src/infiniop/ops/kimi_delta_attention/moore/kimi_delta_attention_moore.h new file mode 100644 index 000000000..dca7e44c5 --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/moore/kimi_delta_attention_moore.h @@ -0,0 +1,8 @@ +#ifndef __KIMI_DELTA_ATTENTION_MOORE_H__ +#define __KIMI_DELTA_ATTENTION_MOORE_H__ + +#include "../kimi_delta_attention.h" + +DESCRIPTOR(moore) + +#endif diff --git a/src/infiniop/ops/kimi_delta_attention/moore/kimi_delta_attention_moore.mu b/src/infiniop/ops/kimi_delta_attention/moore/kimi_delta_attention_moore.mu new file mode 100644 index 000000000..40733353e --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/moore/kimi_delta_attention_moore.mu @@ -0,0 +1,273 @@ +#include "../../../devices/moore/moore_common.h" +#include "../../../devices/moore/moore_kernel_common.h" +#include "kimi_delta_attention_moore.h" + +#include "../cuda/kernel.cuh" + +namespace op::kimi_delta_attention::moore { + +struct Descriptor::Opaque { + std::shared_ptr internal; +}; + +Descriptor::~Descriptor() { + delete _opaque; +} + +infiniStatus_t Descriptor::create( + infiniopHandle_t handle, + Descriptor **desc_ptr, + infiniopTensorDescriptor_t out_desc, + infiniopTensorDescriptor_t initial_state_desc, + infiniopTensorDescriptor_t final_state_desc, + infiniopTensorDescriptor_t q_desc, + infiniopTensorDescriptor_t k_desc, + infiniopTensorDescriptor_t v_desc, + infiniopTensorDescriptor_t g_desc, + infiniopTensorDescriptor_t beta_desc, + infiniopTensorDescriptor_t A_log_desc, + infiniopTensorDescriptor_t dt_bias_desc, + infiniopTensorDescriptor_t cu_seqlens_desc, + infiniopTensorDescriptor_t initial_state_indices_desc, + infiniopTensorDescriptor_t final_state_indices_desc, + float scale, + float lower_bound, + bool use_qk_l2norm) { + + auto info = KimiDeltaAttentionInfo::create( + out_desc, + initial_state_desc, + final_state_desc, + q_desc, + k_desc, + v_desc, + g_desc, + beta_desc, + A_log_desc, + dt_bias_desc, + cu_seqlens_desc, + initial_state_indices_desc, + final_state_indices_desc, + scale, + lower_bound, + use_qk_l2norm); + CHECK_RESULT(info); + + *desc_ptr = new Descriptor( + new Opaque{reinterpret_cast(handle)->internal()}, + info.take(), + 0, + handle->device, + handle->device_id); + return INFINI_STATUS_SUCCESS; +} + +template +static infiniStatus_t launch_fallback(const KimiDeltaAttentionInfo &info, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + musaStream_t stream) { + constexpr int threads = 256; + dim3 grid(static_cast(info.B), static_cast(info.H), static_cast(info.D)); + size_t shared = info.is_decode ? threads * sizeof(float) : (info.D * 3 + threads) * sizeof(float); + + // TODO(kimi_delta_attention): Dispatch MoonshotAI FlashKDA SM90 fast path for + // BF16, D=128, contiguous tensors. Original source: + // https://github.com/MoonshotAI/FlashKDA/tree/master/csrc + if (info.is_decode) { + kimiDeltaAttentionDecodeCudaKernel<<>>( + static_cast(out), + static_cast(initial_state), + static_cast(final_state), + static_cast(q), + static_cast(k), + static_cast(v), + static_cast(g), + static_cast(beta), + static_cast(A_log), + static_cast(dt_bias), + cu_seqlens, + initial_state_indices, + final_state_indices, + info.cu_seqlens_dtype == INFINI_DTYPE_I64, + info.initial_state_indices_dtype == INFINI_DTYPE_I64, + info.final_state_indices_dtype == INFINI_DTYPE_I64, + info.use_qk_l2norm, + info.has_cu_seqlens, + info.indexed_state_pool, + info.D, + info.pool_size, + info.scale, + info.lower_bound, + info.out_strides[0], + info.out_strides[1], + info.out_strides[2], + info.initial_state_strides[0], + info.initial_state_strides[1], + info.initial_state_strides[2], + info.initial_state_strides[3], + info.final_state_strides.empty() ? 0 : info.final_state_strides[0], + info.final_state_strides.empty() ? 0 : info.final_state_strides[1], + info.final_state_strides.empty() ? 0 : info.final_state_strides[2], + info.final_state_strides.empty() ? 0 : info.final_state_strides[3], + info.q_strides[0], + info.q_strides[1], + info.q_strides[2], + info.k_strides[0], + info.k_strides[1], + info.k_strides[2], + info.v_strides[0], + info.v_strides[1], + info.v_strides[2], + info.g_strides[0], + info.g_strides[1], + info.g_strides[2], + info.beta_strides[0], + info.beta_strides[1], + info.beta_strides[2], + info.A_log_strides[0], + info.dt_bias_strides[0]); + } else { + kimiDeltaAttentionRecurrentCudaKernel<<>>( + static_cast(out), + static_cast(initial_state), + static_cast(final_state), + static_cast(q), + static_cast(k), + static_cast(v), + static_cast(g), + static_cast(beta), + static_cast(A_log), + static_cast(dt_bias), + cu_seqlens, + initial_state_indices, + final_state_indices, + info.cu_seqlens_dtype == INFINI_DTYPE_I64, + info.initial_state_indices_dtype == INFINI_DTYPE_I64, + info.final_state_indices_dtype == INFINI_DTYPE_I64, + info.use_qk_l2norm, + info.has_cu_seqlens, + info.indexed_state_pool, + info.T, + info.D, + info.pool_size, + info.scale, + info.lower_bound, + info.out_strides[0], + info.out_strides[1], + info.out_strides[2], + info.initial_state_strides[0], + info.initial_state_strides[1], + info.initial_state_strides[2], + info.initial_state_strides[3], + info.final_state_strides.empty() ? 0 : info.final_state_strides[0], + info.final_state_strides.empty() ? 0 : info.final_state_strides[1], + info.final_state_strides.empty() ? 0 : info.final_state_strides[2], + info.final_state_strides.empty() ? 0 : info.final_state_strides[3], + info.q_strides[0], + info.q_strides[1], + info.q_strides[2], + info.k_strides[0], + info.k_strides[1], + info.k_strides[2], + info.v_strides[0], + info.v_strides[1], + info.v_strides[2], + info.g_strides[0], + info.g_strides[1], + info.g_strides[2], + info.beta_strides[0], + info.beta_strides[1], + info.beta_strides[2], + info.A_log_strides[0], + info.dt_bias_strides[0]); + } + return INFINI_STATUS_SUCCESS; +} + +template +static infiniStatus_t launch_for_gate(const KimiDeltaAttentionInfo &info, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + musaStream_t stream) { + switch (info.gate_dtype) { + case INFINI_DTYPE_F16: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_BF16: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_F32: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + default: + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } +} + +infiniStatus_t Descriptor::calculate( + void *workspace, + size_t workspace_size, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + void *stream_) const { + if (workspace_size < _workspace_size) { + return INFINI_STATUS_INSUFFICIENT_WORKSPACE; + } + if (_info.has_cu_seqlens && cu_seqlens == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (_info.has_initial_state_indices && initial_state_indices == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (_info.has_final_state_indices && final_state_indices == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (!_info.has_final_state_indices && final_state == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + + musaStream_t stream = reinterpret_cast(stream_); + switch (_info.data_dtype) { + case INFINI_DTYPE_F16: + return launch_for_gate(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_BF16: + return launch_for_gate<__nv_bfloat16>(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_F32: + return launch_for_gate(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + default: + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } +} + +} // namespace op::kimi_delta_attention::moore diff --git a/src/infiniop/ops/kimi_delta_attention/nvidia/kimi_delta_attention_nvidia.cu b/src/infiniop/ops/kimi_delta_attention/nvidia/kimi_delta_attention_nvidia.cu new file mode 100644 index 000000000..4f1a8705e --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/nvidia/kimi_delta_attention_nvidia.cu @@ -0,0 +1,277 @@ +#include "../../../devices/nvidia/nvidia_common.cuh" +#include "../../../devices/nvidia/nvidia_handle.cuh" +#include "kimi_delta_attention_nvidia.cuh" + +#include "../cuda/kernel.cuh" + +#include +#include +#include + +namespace op::kimi_delta_attention::nvidia { + +struct Descriptor::Opaque { + std::shared_ptr internal; +}; + +Descriptor::~Descriptor() { + delete _opaque; +} + +infiniStatus_t Descriptor::create( + infiniopHandle_t handle, + Descriptor **desc_ptr, + infiniopTensorDescriptor_t out_desc, + infiniopTensorDescriptor_t initial_state_desc, + infiniopTensorDescriptor_t final_state_desc, + infiniopTensorDescriptor_t q_desc, + infiniopTensorDescriptor_t k_desc, + infiniopTensorDescriptor_t v_desc, + infiniopTensorDescriptor_t g_desc, + infiniopTensorDescriptor_t beta_desc, + infiniopTensorDescriptor_t A_log_desc, + infiniopTensorDescriptor_t dt_bias_desc, + infiniopTensorDescriptor_t cu_seqlens_desc, + infiniopTensorDescriptor_t initial_state_indices_desc, + infiniopTensorDescriptor_t final_state_indices_desc, + float scale, + float lower_bound, + bool use_qk_l2norm) { + + auto info = KimiDeltaAttentionInfo::create( + out_desc, + initial_state_desc, + final_state_desc, + q_desc, + k_desc, + v_desc, + g_desc, + beta_desc, + A_log_desc, + dt_bias_desc, + cu_seqlens_desc, + initial_state_indices_desc, + final_state_indices_desc, + scale, + lower_bound, + use_qk_l2norm); + CHECK_RESULT(info); + + *desc_ptr = new Descriptor( + new Opaque{reinterpret_cast(handle)->internal()}, + info.take(), + 0, + handle->device, + handle->device_id); + return INFINI_STATUS_SUCCESS; +} + +template +static infiniStatus_t launch_fallback(const KimiDeltaAttentionInfo &info, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + cudaStream_t stream) { + constexpr int threads = 256; + dim3 grid(static_cast(info.B), static_cast(info.H), static_cast(info.D)); + size_t shared = info.is_decode ? threads * sizeof(float) : (info.D * 3 + threads) * sizeof(float); + + // TODO(kimi_delta_attention): Dispatch MoonshotAI FlashKDA SM90 fast path for + // BF16, D=128, contiguous tensors. Original source: + // https://github.com/MoonshotAI/FlashKDA/tree/master/csrc + if (info.is_decode) { + kimiDeltaAttentionDecodeCudaKernel<<>>( + static_cast(out), + static_cast(initial_state), + static_cast(final_state), + static_cast(q), + static_cast(k), + static_cast(v), + static_cast(g), + static_cast(beta), + static_cast(A_log), + static_cast(dt_bias), + cu_seqlens, + initial_state_indices, + final_state_indices, + info.cu_seqlens_dtype == INFINI_DTYPE_I64, + info.initial_state_indices_dtype == INFINI_DTYPE_I64, + info.final_state_indices_dtype == INFINI_DTYPE_I64, + info.use_qk_l2norm, + info.has_cu_seqlens, + info.indexed_state_pool, + info.D, + info.pool_size, + info.scale, + info.lower_bound, + info.out_strides[0], + info.out_strides[1], + info.out_strides[2], + info.initial_state_strides[0], + info.initial_state_strides[1], + info.initial_state_strides[2], + info.initial_state_strides[3], + info.final_state_strides.empty() ? 0 : info.final_state_strides[0], + info.final_state_strides.empty() ? 0 : info.final_state_strides[1], + info.final_state_strides.empty() ? 0 : info.final_state_strides[2], + info.final_state_strides.empty() ? 0 : info.final_state_strides[3], + info.q_strides[0], + info.q_strides[1], + info.q_strides[2], + info.k_strides[0], + info.k_strides[1], + info.k_strides[2], + info.v_strides[0], + info.v_strides[1], + info.v_strides[2], + info.g_strides[0], + info.g_strides[1], + info.g_strides[2], + info.beta_strides[0], + info.beta_strides[1], + info.beta_strides[2], + info.A_log_strides[0], + info.dt_bias_strides[0]); + } else { + kimiDeltaAttentionRecurrentCudaKernel<<>>( + static_cast(out), + static_cast(initial_state), + static_cast(final_state), + static_cast(q), + static_cast(k), + static_cast(v), + static_cast(g), + static_cast(beta), + static_cast(A_log), + static_cast(dt_bias), + cu_seqlens, + initial_state_indices, + final_state_indices, + info.cu_seqlens_dtype == INFINI_DTYPE_I64, + info.initial_state_indices_dtype == INFINI_DTYPE_I64, + info.final_state_indices_dtype == INFINI_DTYPE_I64, + info.use_qk_l2norm, + info.has_cu_seqlens, + info.indexed_state_pool, + info.T, + info.D, + info.pool_size, + info.scale, + info.lower_bound, + info.out_strides[0], + info.out_strides[1], + info.out_strides[2], + info.initial_state_strides[0], + info.initial_state_strides[1], + info.initial_state_strides[2], + info.initial_state_strides[3], + info.final_state_strides.empty() ? 0 : info.final_state_strides[0], + info.final_state_strides.empty() ? 0 : info.final_state_strides[1], + info.final_state_strides.empty() ? 0 : info.final_state_strides[2], + info.final_state_strides.empty() ? 0 : info.final_state_strides[3], + info.q_strides[0], + info.q_strides[1], + info.q_strides[2], + info.k_strides[0], + info.k_strides[1], + info.k_strides[2], + info.v_strides[0], + info.v_strides[1], + info.v_strides[2], + info.g_strides[0], + info.g_strides[1], + info.g_strides[2], + info.beta_strides[0], + info.beta_strides[1], + info.beta_strides[2], + info.A_log_strides[0], + info.dt_bias_strides[0]); + } + return INFINI_STATUS_SUCCESS; +} + +template +static infiniStatus_t launch_for_gate(const KimiDeltaAttentionInfo &info, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + cudaStream_t stream) { + switch (info.gate_dtype) { + case INFINI_DTYPE_F16: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_BF16: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_F32: + return launch_fallback(info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + default: + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } +} + +infiniStatus_t Descriptor::calculate( + void *workspace, + size_t workspace_size, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + void *stream_) const { + if (workspace_size < _workspace_size) { + return INFINI_STATUS_INSUFFICIENT_WORKSPACE; + } + if (_info.has_cu_seqlens && cu_seqlens == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (_info.has_initial_state_indices && initial_state_indices == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (_info.has_final_state_indices && final_state_indices == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + if (!_info.has_final_state_indices && final_state == nullptr) { + return INFINI_STATUS_NULL_POINTER; + } + + cudaStream_t stream = reinterpret_cast(stream_); + switch (_info.data_dtype) { + case INFINI_DTYPE_F16: + return launch_for_gate(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_BF16: + return launch_for_gate<__nv_bfloat16>(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + case INFINI_DTYPE_F32: + return launch_for_gate(_info, out, initial_state, final_state, q, k, v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, final_state_indices, stream); + default: + return INFINI_STATUS_BAD_TENSOR_DTYPE; + } +} + +} // namespace op::kimi_delta_attention::nvidia diff --git a/src/infiniop/ops/kimi_delta_attention/nvidia/kimi_delta_attention_nvidia.cuh b/src/infiniop/ops/kimi_delta_attention/nvidia/kimi_delta_attention_nvidia.cuh new file mode 100644 index 000000000..993074957 --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/nvidia/kimi_delta_attention_nvidia.cuh @@ -0,0 +1,8 @@ +#ifndef __KIMI_DELTA_ATTENTION_NVIDIA_CUH__ +#define __KIMI_DELTA_ATTENTION_NVIDIA_CUH__ + +#include "../kimi_delta_attention.h" + +DESCRIPTOR(nvidia) + +#endif diff --git a/src/infiniop/ops/kimi_delta_attention/operator.cc b/src/infiniop/ops/kimi_delta_attention/operator.cc new file mode 100644 index 000000000..eb1944dfb --- /dev/null +++ b/src/infiniop/ops/kimi_delta_attention/operator.cc @@ -0,0 +1,203 @@ +#include "../../operator.h" +#include "../../handle.h" +#include "infiniop/ops/kimi_delta_attention.h" + +#if defined(ENABLE_NVIDIA_API) || defined(ENABLE_QY_API) || defined(ENABLE_ALI_API) || defined(ENABLE_ILUVATAR_API) || defined(ENABLE_HYGON_API) +#include "nvidia/kimi_delta_attention_nvidia.cuh" +#endif +#ifdef ENABLE_METAX_API +#include "metax/kimi_delta_attention_metax.h" +#endif +#ifdef ENABLE_MOORE_API +#include "moore/kimi_delta_attention_moore.h" +#endif + +__INFINI_C __export infiniStatus_t infiniopCreateKimiDeltaAttentionDescriptor( + infiniopHandle_t handle, + infiniopKimiDeltaAttentionDescriptor_t *desc_ptr, + infiniopTensorDescriptor_t out_desc, + infiniopTensorDescriptor_t initial_state_desc, + infiniopTensorDescriptor_t final_state_desc, + infiniopTensorDescriptor_t q_desc, + infiniopTensorDescriptor_t k_desc, + infiniopTensorDescriptor_t v_desc, + infiniopTensorDescriptor_t g_desc, + infiniopTensorDescriptor_t beta_desc, + infiniopTensorDescriptor_t A_log_desc, + infiniopTensorDescriptor_t dt_bias_desc, + infiniopTensorDescriptor_t cu_seqlens_desc, + infiniopTensorDescriptor_t initial_state_indices_desc, + infiniopTensorDescriptor_t final_state_indices_desc, + float scale, + float lower_bound, + bool use_qk_l2norm) { + +#define CREATE(CASE, NAMESPACE) \ + case CASE: \ + return op::kimi_delta_attention::NAMESPACE::Descriptor::create( \ + handle, reinterpret_cast(desc_ptr), \ + out_desc, initial_state_desc, final_state_desc, q_desc, k_desc, v_desc, g_desc, \ + beta_desc, A_log_desc, dt_bias_desc, cu_seqlens_desc, initial_state_indices_desc, \ + final_state_indices_desc, scale, lower_bound, use_qk_l2norm) + + switch (handle->device) { +#ifdef ENABLE_NVIDIA_API + CREATE(INFINI_DEVICE_NVIDIA, nvidia); +#endif +#ifdef ENABLE_QY_API + CREATE(INFINI_DEVICE_QY, nvidia); +#endif +#ifdef ENABLE_ALI_API + CREATE(INFINI_DEVICE_ALI, nvidia); +#endif +#ifdef ENABLE_ILUVATAR_API + CREATE(INFINI_DEVICE_ILUVATAR, nvidia); +#endif +#ifdef ENABLE_HYGON_API + CREATE(INFINI_DEVICE_HYGON, nvidia); +#endif +#ifdef ENABLE_METAX_API + CREATE(INFINI_DEVICE_METAX, metax); +#endif +#ifdef ENABLE_MOORE_API + CREATE(INFINI_DEVICE_MOORE, moore); +#endif + default: + return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED; + } + +#undef CREATE +} + +__INFINI_C __export infiniStatus_t infiniopGetKimiDeltaAttentionWorkspaceSize( + infiniopKimiDeltaAttentionDescriptor_t desc, + size_t *size) { + +#define GET(CASE, NAMESPACE) \ + case CASE: \ + *size = reinterpret_cast(desc) \ + ->workspaceSize(); \ + return INFINI_STATUS_SUCCESS + + switch (desc->device_type) { +#ifdef ENABLE_NVIDIA_API + GET(INFINI_DEVICE_NVIDIA, nvidia); +#endif +#ifdef ENABLE_QY_API + GET(INFINI_DEVICE_QY, nvidia); +#endif +#ifdef ENABLE_ALI_API + GET(INFINI_DEVICE_ALI, nvidia); +#endif +#ifdef ENABLE_ILUVATAR_API + GET(INFINI_DEVICE_ILUVATAR, nvidia); +#endif +#ifdef ENABLE_HYGON_API + GET(INFINI_DEVICE_HYGON, nvidia); +#endif +#ifdef ENABLE_METAX_API + GET(INFINI_DEVICE_METAX, metax); +#endif +#ifdef ENABLE_MOORE_API + GET(INFINI_DEVICE_MOORE, moore); +#endif + default: + return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED; + } + +#undef GET +} + +__INFINI_C __export infiniStatus_t infiniopKimiDeltaAttention( + infiniopKimiDeltaAttentionDescriptor_t desc, + void *workspace, + size_t workspace_size, + void *out, + void *initial_state, + void *final_state, + const void *q, + const void *k, + const void *v, + const void *g, + const void *beta, + const void *A_log, + const void *dt_bias, + const void *cu_seqlens, + const void *initial_state_indices, + const void *final_state_indices, + void *stream) { + +#define CALCULATE(CASE, NAMESPACE) \ + case CASE: \ + return reinterpret_cast( \ + desc) \ + ->calculate(workspace, workspace_size, out, initial_state, final_state, q, k, \ + v, g, beta, A_log, dt_bias, cu_seqlens, initial_state_indices, \ + final_state_indices, stream) + + switch (desc->device_type) { +#ifdef ENABLE_NVIDIA_API + CALCULATE(INFINI_DEVICE_NVIDIA, nvidia); +#endif +#ifdef ENABLE_QY_API + CALCULATE(INFINI_DEVICE_QY, nvidia); +#endif +#ifdef ENABLE_ALI_API + CALCULATE(INFINI_DEVICE_ALI, nvidia); +#endif +#ifdef ENABLE_ILUVATAR_API + CALCULATE(INFINI_DEVICE_ILUVATAR, nvidia); +#endif +#ifdef ENABLE_HYGON_API + CALCULATE(INFINI_DEVICE_HYGON, nvidia); +#endif +#ifdef ENABLE_METAX_API + CALCULATE(INFINI_DEVICE_METAX, metax); +#endif +#ifdef ENABLE_MOORE_API + CALCULATE(INFINI_DEVICE_MOORE, moore); +#endif + default: + return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED; + } + +#undef CALCULATE +} + +__INFINI_C __export infiniStatus_t infiniopDestroyKimiDeltaAttentionDescriptor( + infiniopKimiDeltaAttentionDescriptor_t desc) { + +#define DELETE(CASE, NAMESPACE) \ + case CASE: \ + delete reinterpret_cast( \ + desc); \ + return INFINI_STATUS_SUCCESS + + switch (desc->device_type) { +#ifdef ENABLE_NVIDIA_API + DELETE(INFINI_DEVICE_NVIDIA, nvidia); +#endif +#ifdef ENABLE_QY_API + DELETE(INFINI_DEVICE_QY, nvidia); +#endif +#ifdef ENABLE_ALI_API + DELETE(INFINI_DEVICE_ALI, nvidia); +#endif +#ifdef ENABLE_ILUVATAR_API + DELETE(INFINI_DEVICE_ILUVATAR, nvidia); +#endif +#ifdef ENABLE_HYGON_API + DELETE(INFINI_DEVICE_HYGON, nvidia); +#endif +#ifdef ENABLE_METAX_API + DELETE(INFINI_DEVICE_METAX, metax); +#endif +#ifdef ENABLE_MOORE_API + DELETE(INFINI_DEVICE_MOORE, moore); +#endif + default: + return INFINI_STATUS_DEVICE_TYPE_NOT_SUPPORTED; + } + +#undef DELETE +} diff --git a/test/infinicore/ops/kimi_delta_attention.py b/test/infinicore/ops/kimi_delta_attention.py new file mode 100644 index 000000000..4be7611c2 --- /dev/null +++ b/test/infinicore/ops/kimi_delta_attention.py @@ -0,0 +1,189 @@ +import os +import sys + +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) + +import torch +import infinicore +from framework import BaseOperatorTest, GenericTestRunner, TensorSpec, TestCase, TensorInitializer + + +_TENSOR_DTYPES = [infinicore.float16, infinicore.bfloat16, infinicore.float32] +_TOLERANCE_MAP = { + infinicore.float16: {"atol": 3e-2, "rtol": 3e-2}, + infinicore.bfloat16: {"atol": 4e-2, "rtol": 4e-2}, + infinicore.float32: {"atol": 2e-4, "rtol": 2e-4}, +} + + +def _l2norm(x): + return x * torch.rsqrt((x * x).sum(dim=-1, keepdim=True) + 1e-6) + + +def torch_kimi_delta_attention_ref( + q, + k, + v, + g, + beta, + A_log, + dt_bias, + initial_state, + cu_seqlens=None, + initial_state_indices=None, + final_state_indices=None, + scale=1.0, + lower_bound=-5.0, + use_qk_l2norm=True, +): + initial_dtype = q.dtype + qf = q.float() + kf = k.float() + vf = v.float() + gf = g.float() + betaf = beta.float() + state_pool = initial_state.float().clone() + out = torch.empty_like(vf) + + if cu_seqlens is None: + batch = q.shape[0] + ranges = [(b, 0, q.shape[1], b) for b in range(batch)] + else: + cu = cu_seqlens.cpu().tolist() + ranges = [(i, cu[i], cu[i + 1], 0) for i in range(len(cu) - 1)] + + for req_idx, begin, end, token_batch in ranges: + read_slot = ( + int(initial_state_indices[req_idx].item()) + if initial_state_indices is not None + else req_idx + ) + write_slot = ( + int(final_state_indices[req_idx].item()) + if final_state_indices is not None + else req_idx + ) + if read_slot < 0 or write_slot < 0: + out[token_batch, begin:end].zero_() + continue + + for h in range(q.shape[2]): + state = state_pool[read_slot, h].clone() + a_log_exp = A_log[h].float().exp() + for t in range(begin, end): + q_t = qf[token_batch, t, h] + k_t = kf[token_batch, t, h] + if use_qk_l2norm: + q_t = _l2norm(q_t) + k_t = _l2norm(k_t) + q_t = q_t * scale + + gate = lower_bound * torch.sigmoid( + a_log_exp * (gf[token_batch, t, h] + dt_bias[h].float()) + ) + decay = gate.exp() + beta_t = betaf[token_batch, t, h].sigmoid() + + decayed_state = state * decay.view(1, -1) + kv_mem = (decayed_state * k_t.view(1, -1)).sum(dim=-1) + delta = (vf[token_batch, t, h] - kv_mem) * beta_t + state = decayed_state + delta.view(-1, 1) * k_t.view(1, -1) + out[token_batch, t, h] = (state * q_t.view(1, -1)).sum(dim=-1) + state_pool[write_slot, h].copy_(state) + + return out.to(initial_dtype) + + +def parse_test_cases(): + tests = [] + for dtype in _TENSOR_DTYPES: + tol = _TOLERANCE_MAP[dtype] + for shape, cu, indexed in [ + ((2, 1, 2, 8), None, False), + ((2, 1, 2, 8), None, True), + ((2, 3, 2, 8), None, False), + ((2, 3, 2, 8), None, True), + ((1, 2, 2, 8), torch.tensor([0, 1, 2], dtype=torch.int64), True), + ((1, 5, 2, 8), torch.tensor([0, 2, 5], dtype=torch.int64), False), + ((1, 5, 2, 8), torch.tensor([0, 2, 5], dtype=torch.int64), True), + ]: + B_state = shape[0] if cu is None else cu.numel() - 1 + pool_size = 4 if indexed else B_state + H = shape[2] + D = shape[3] + q = TensorSpec.from_tensor(shape, None, dtype) + k = TensorSpec.from_tensor(shape, None, dtype) + v = TensorSpec.from_tensor(shape, None, dtype) + g = TensorSpec.from_tensor(shape, None, dtype) + beta = TensorSpec.from_tensor(shape[:3], None, dtype) + A_log = TensorSpec.from_tensor((H,), None, infinicore.float32) + dt_bias = TensorSpec.from_tensor((H, D), None, infinicore.float32) + initial_state = TensorSpec.from_tensor((pool_size, H, D, D), None, dtype) + kwargs = { + "scale": D**-0.5, + "lower_bound": -5.0, + "use_qk_l2norm": True, + } + if cu is not None: + kwargs["cu_seqlens"] = TensorSpec.from_tensor( + tuple(cu.shape), + None, + infinicore.int64, + init_mode=TensorInitializer.MANUAL, + set_tensor=cu, + ) + if indexed: + initial_indices = torch.tensor([2, 0], dtype=torch.int64) + final_indices = torch.tensor([1, 3], dtype=torch.int64) + kwargs["initial_state_indices"] = TensorSpec.from_tensor( + tuple(initial_indices.shape), + None, + infinicore.int64, + init_mode=TensorInitializer.MANUAL, + set_tensor=initial_indices, + ) + kwargs["final_state_indices"] = TensorSpec.from_tensor( + tuple(final_indices.shape), + None, + infinicore.int64, + init_mode=TensorInitializer.MANUAL, + set_tensor=final_indices, + ) + tests.append( + TestCase( + inputs=[q, k, v, g, beta, A_log, dt_bias, initial_state], + kwargs=kwargs, + output_spec=None, + comparison_target=None, + tolerance=tol, + description=( + "KimiDeltaAttention" + + (" indexed-pool" if indexed else "") + + (" varlen" if cu is not None else "") + ), + ) + ) + return tests + + +class OpTest(BaseOperatorTest): + def __init__(self): + super().__init__("KimiDeltaAttention") + + def get_test_cases(self): + return parse_test_cases() + + def torch_operator(self, *args, **kwargs): + return torch_kimi_delta_attention_ref(*args, **kwargs) + + def infinicore_operator(self, *args, **kwargs): + return infinicore.nn.functional.kimi_delta_attention(*args, **kwargs) + + +def main(): + runner = GenericTestRunner(OpTest) + runner.run_and_exit() + + +if __name__ == "__main__": + main()