Skip to content
Merged
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
34 changes: 26 additions & 8 deletions aiter/ops/triton/_triton_kernels/chunk_delta_attn/flash_kda.py
Original file line number Diff line number Diff line change
Expand Up @@ -494,17 +494,30 @@ def _flash_kda_segment_kernel(
if STORE_FINAL: # noqa: SIM102
if tl.load(seg_is_last + i_seg) == 1:
i_n = tl.load(seg_seq + i_seg).to(tl.int64)
dt_s = final_state.dtype.element_ty
if STATE_V_FIRST:
f_off = (i_n * H + i_h) * V * K + o_w[None, :] * K
tl.store(final_state + f_off + o_k1[:, None], b_h1, mask=m_w[None, :])
tl.store(final_state + f_off + o_k2[:, None], b_h2, mask=m_w[None, :])
tl.store(
final_state + f_off + o_k1[:, None],
b_h1.to(dt_s),
mask=m_w[None, :],
)
tl.store(
final_state + f_off + o_k2[:, None],
b_h2.to(dt_s),
mask=m_w[None, :],
)
else:
f_off = (i_n * H + i_h) * K * V + o_w[None, :]
tl.store(
final_state + f_off + o_k1[:, None] * V, b_h1, mask=m_w[None, :]
final_state + f_off + o_k1[:, None] * V,
b_h1.to(dt_s),
mask=m_w[None, :],
)
tl.store(
final_state + f_off + o_k2[:, None] * V, b_h2, mask=m_w[None, :]
final_state + f_off + o_k2[:, None] * V,
b_h2.to(dt_s),
mask=m_w[None, :],
)


Expand Down Expand Up @@ -541,8 +554,12 @@ def _flash_kda_seg_scan_kernel(

if HAS_H0:
base = (i_n * H + i_h) * K * V + o_v[None, :]
b_h1 = tl.load(h0 + base + o_k1[:, None] * V, mask=m_v[None, :], other=0.0)
b_h2 = tl.load(h0 + base + o_k2[:, None] * V, mask=m_v[None, :], other=0.0)
b_h1 = tl.load(h0 + base + o_k1[:, None] * V, mask=m_v[None, :], other=0.0).to(
tl.float32
)
b_h2 = tl.load(h0 + base + o_k2[:, None] * V, mask=m_v[None, :], other=0.0).to(
tl.float32
Comment thread
XiaobingSuper marked this conversation as resolved.
)
else:
b_h1 = tl.zeros([64, BV], dtype=tl.float32)
b_h2 = tl.zeros([64, BV], dtype=tl.float32)
Expand Down Expand Up @@ -831,13 +848,14 @@ def flash_kda_fwd(
if h0 is not None:
if state_v_first:
h0 = h0.transpose(-1, -2)
h0 = h0.to(torch.float32).contiguous()
h0 = h0.contiguous()
Comment thread
XiaobingSuper marked this conversation as resolved.

o = torch.empty_like(v)
final_state = None
if output_final_state:
shape = (N, H, V, K) if state_v_first else (N, H, K, V)
final_state = torch.empty(shape, dtype=torch.float32, device=dev)
state_dtype = h0.dtype if h0 is not None else torch.float32
final_state = torch.empty(shape, dtype=state_dtype, device=dev)

Comment on lines 855 to 859

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The store was already correct (Triton implicitly converts to the pointer element type, and fp16/bf16 both round-trip in the numerics table above), but the surrounding out/h_out stores cast explicitly, so I made these match.

common = {
"ws_kd": ws_kd,
Expand Down
14 changes: 7 additions & 7 deletions aiter/ops/triton/kimi_delta_attn/chunk_delta_attn.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,11 +96,13 @@ def chunk_kimi_delta_attn(
Scale factor for the attention scores. Default: `1 / sqrt(K)`.
initial_state (torch.Tensor, optional):
Initial state of shape `[N, HV, K, V]` (`[N, HV, V, K]` when
`state_v_first=True`) and dtype fp32, for `N` input sequences. For
`state_v_first=True`), for `N` input sequences. Any float dtype; the
recurrence accumulates in fp32 whatever the state is stored as. For
equal-length inputs `N` equals the batch size `B`. Default: `None`.
output_final_state (bool):
Whether to return the final state, same shape and dtype as
`initial_state`. Default: `False`.
Whether to return the final state, same shape as `initial_state`.
Its dtype is `initial_state`'s on the FlashKDA path and fp32 on the
default pipeline. Default: `False`.
Comment thread
XiaobingSuper marked this conversation as resolved.
use_qk_l2norm_in_kernel (bool):
Whether to L2-normalize `q` and `k` before the recurrence.
use_gate_in_kernel (bool):
Expand Down Expand Up @@ -196,12 +198,10 @@ def chunk_kimi_delta_attn(
f"of input sequences, i.e., {len(cu_seqlens) - 1} rather than "
f"{initial_state.shape[0]}."
)
if initial_state is not None and initial_state.dtype != torch.float32:
if initial_state is not None and not initial_state.is_floating_point():
raise ValueError(
f"`initial_state` must be fp32, got {initial_state.dtype}. The recurrence "
"accumulates in fp32 and the state is read back verbatim."
f"`initial_state` must be a float tensor, got {initial_state.dtype}."
)
Comment thread
XiaobingSuper marked this conversation as resolved.

if use_gate_in_kernel and A_log is None:
raise ValueError("`A_log` must be provided when `use_gate_in_kernel=True`.")
if safe_gate and use_gate_in_kernel:
Expand Down
Loading