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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions include/infinicore/ops.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
63 changes: 63 additions & 0 deletions include/infinicore/ops/kimi_delta_attention.hpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,63 @@
#pragma once

#include "infinicore.h"

#include "../device.hpp"
#include "../graph/graph.hpp"
#include "common/op.hpp"

#include <optional>

namespace infinicore::op {

INFINICORE_GRAPH_OP_CLASS(KimiDeltaAttention,
Tensor,
Tensor,
std::optional<Tensor>,
const Tensor &,
const Tensor &,
const Tensor &,
const Tensor &,
const Tensor &,
const Tensor &,
const Tensor &,
std::optional<Tensor>,
std::optional<Tensor>,
std::optional<Tensor>,
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<Tensor> cu_seqlens = std::nullopt,
std::optional<Tensor> initial_state_indices = std::nullopt,
std::optional<Tensor> 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<Tensor> 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<Tensor> cu_seqlens,
std::optional<Tensor> initial_state_indices,
std::optional<Tensor> final_state_indices,
float scale = 1.0f,
float lower_bound = -5.0f,
bool use_qk_l2norm = true);

} // namespace infinicore::op
1 change: 1 addition & 0 deletions include/infiniop.h
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
54 changes: 54 additions & 0 deletions include/infiniop/ops/kimi_delta_attention.h
Original file line number Diff line number Diff line change
@@ -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
2 changes: 2 additions & 0 deletions python/infinicore/nn/functional/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -56,6 +57,7 @@
"fused_gated_delta_net_gating",
"gaussian_nll_loss",
"interpolate",
"kimi_delta_attention",
"linear",
"binary_cross_entropy_with_logits",
"random_sample",
Expand Down
66 changes: 66 additions & 0 deletions python/infinicore/nn/functional/kimi_delta_attention.py
Original file line number Diff line number Diff line change
@@ -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,
)
)
Loading
Loading