[AMD] Perf Kimi-K3 fuse ROCm KDA decode boundary (#34198)
Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
@@ -0,0 +1,166 @@
|
||||
"""Fail-closed adapter for AITER's gfx950 Kimi-K3 fused KDA decode."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.utils import is_hip
|
||||
|
||||
_HEADS = 12
|
||||
_DIM = 128
|
||||
_CHANNELS = 3 * _HEADS * _DIM
|
||||
_WARMED: set[tuple[int, float, float]] = set()
|
||||
|
||||
|
||||
def enabled() -> bool:
|
||||
return os.environ.get("SGLANG_K3_KDA_FUSED_BACKEND", "").lower() == "aiter"
|
||||
|
||||
|
||||
def _ops():
|
||||
try:
|
||||
from sglang.kernels.ops.kimi_k3.flydsl.source import load_module
|
||||
|
||||
module = load_module(
|
||||
"sglang.kernels.ops.kimi_k3.flydsl.kimi_k3_kda_decode",
|
||||
"aiter.ops.flydsl.kimi_k3_kda_decode",
|
||||
)
|
||||
except (ImportError, ModuleNotFoundError):
|
||||
return None, None
|
||||
return (
|
||||
module.flydsl_kimi_k3_kda_decode_with_f_b,
|
||||
module.is_flydsl_kimi_k3_kda_decode_supported,
|
||||
)
|
||||
|
||||
|
||||
def available(device: torch.device | None = None) -> bool:
|
||||
if not is_hip() or not enabled() or not torch.cuda.is_available():
|
||||
return False
|
||||
_, supported = _ops()
|
||||
if supported is None:
|
||||
return False
|
||||
return bool(supported(device))
|
||||
|
||||
|
||||
def covered(
|
||||
f_a: torch.Tensor,
|
||||
f_b_weight: torch.Tensor,
|
||||
mixed_qkv: torch.Tensor,
|
||||
raw_beta: torch.Tensor,
|
||||
conv_state: torch.Tensor,
|
||||
state: torch.Tensor,
|
||||
state_indices: torch.Tensor,
|
||||
output_gate: torch.Tensor,
|
||||
norm_weight: torch.Tensor,
|
||||
) -> bool:
|
||||
if not available(f_a.device) or f_a.ndim != 2:
|
||||
return False
|
||||
batch = f_a.shape[0]
|
||||
return (
|
||||
batch > 0
|
||||
and f_a.shape == (batch, _DIM)
|
||||
and f_a.dtype == torch.bfloat16
|
||||
and f_a.stride(-1) == 1
|
||||
and f_b_weight.shape == (_HEADS, _DIM, _DIM)
|
||||
and f_b_weight.dtype == torch.bfloat16
|
||||
and f_b_weight.stride()[-2:] == (_DIM, 1)
|
||||
and mixed_qkv.shape == (batch, _CHANNELS)
|
||||
and mixed_qkv.dtype == torch.bfloat16
|
||||
and mixed_qkv.stride(-1) == 1
|
||||
and raw_beta.shape == (1, batch, _HEADS)
|
||||
and raw_beta.dtype == torch.bfloat16
|
||||
and conv_state.ndim == 3
|
||||
and conv_state.shape[1:] == (_CHANNELS, 3)
|
||||
and conv_state.dtype == torch.bfloat16
|
||||
and state.ndim == 4
|
||||
and state.shape[1:] == (_HEADS, _DIM, _DIM)
|
||||
and state.dtype == torch.float32
|
||||
and state.stride()[-3:] == (_DIM * _DIM, _DIM, 1)
|
||||
and state_indices.shape == (batch,)
|
||||
and state_indices.dtype == torch.int32
|
||||
and state_indices.stride(0) == 1
|
||||
and output_gate.shape == (batch, _HEADS, _DIM)
|
||||
and output_gate.dtype == torch.bfloat16
|
||||
and norm_weight.shape == (_DIM,)
|
||||
and norm_weight.dtype == torch.bfloat16
|
||||
)
|
||||
|
||||
|
||||
def run(
|
||||
*,
|
||||
f_a: torch.Tensor,
|
||||
f_b_weight: torch.Tensor,
|
||||
mixed_qkv: torch.Tensor,
|
||||
conv_weight: torch.Tensor,
|
||||
conv_state: torch.Tensor,
|
||||
raw_beta: torch.Tensor,
|
||||
A_log: torch.Tensor,
|
||||
dt_bias: torch.Tensor,
|
||||
lower_bound: float,
|
||||
state: torch.Tensor,
|
||||
state_indices: torch.Tensor,
|
||||
output_gate: torch.Tensor,
|
||||
norm_weight: torch.Tensor,
|
||||
norm_eps: float,
|
||||
out: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
op, _ = _ops()
|
||||
if op is None:
|
||||
raise RuntimeError("AITER Kimi-K3 fused KDA decode is unavailable")
|
||||
return op(
|
||||
f_a=f_a,
|
||||
f_b_weight=f_b_weight,
|
||||
x=mixed_qkv,
|
||||
conv_weight=conv_weight,
|
||||
conv_bias=None,
|
||||
conv_state=conv_state,
|
||||
raw_beta=raw_beta,
|
||||
A_log=A_log,
|
||||
dt_bias=dt_bias,
|
||||
lower_bound=lower_bound,
|
||||
state=state,
|
||||
state_indices=state_indices,
|
||||
output_gate=output_gate,
|
||||
norm_weight=norm_weight,
|
||||
norm_eps=norm_eps,
|
||||
out=out,
|
||||
)
|
||||
|
||||
|
||||
def warmup(
|
||||
*,
|
||||
f_b_weight: torch.Tensor,
|
||||
conv_weight: torch.Tensor,
|
||||
A_log: torch.Tensor,
|
||||
dt_bias: torch.Tensor,
|
||||
lower_bound: float,
|
||||
norm_weight: torch.Tensor,
|
||||
norm_eps: float,
|
||||
) -> None:
|
||||
if not available(f_b_weight.device):
|
||||
return
|
||||
device_index = -1 if f_b_weight.device.index is None else f_b_weight.device.index
|
||||
key = (device_index, float(norm_eps), float(lower_bound))
|
||||
if key in _WARMED:
|
||||
return
|
||||
|
||||
device = f_b_weight.device
|
||||
run(
|
||||
f_a=torch.zeros(1, _DIM, dtype=torch.bfloat16, device=device),
|
||||
f_b_weight=f_b_weight,
|
||||
mixed_qkv=torch.zeros(1, _CHANNELS, dtype=torch.bfloat16, device=device),
|
||||
conv_weight=conv_weight,
|
||||
conv_state=torch.zeros(1, _CHANNELS, 3, dtype=torch.bfloat16, device=device),
|
||||
raw_beta=torch.zeros(1, 1, _HEADS, dtype=torch.bfloat16, device=device),
|
||||
A_log=A_log,
|
||||
dt_bias=dt_bias,
|
||||
lower_bound=lower_bound,
|
||||
state=torch.zeros(1, _HEADS, _DIM, _DIM, dtype=torch.float32, device=device),
|
||||
state_indices=torch.zeros(1, dtype=torch.int32, device=device),
|
||||
output_gate=torch.zeros(1, _HEADS, _DIM, dtype=torch.bfloat16, device=device),
|
||||
norm_weight=norm_weight,
|
||||
norm_eps=norm_eps,
|
||||
)
|
||||
torch.cuda.synchronize(device)
|
||||
_WARMED.add(key)
|
||||
@@ -0,0 +1,17 @@
|
||||
"""SGLang-maintained Kimi-K3 FlyDSL specializations."""
|
||||
|
||||
# AITER owns the FlyDSL toolchain bootstrap and shared tensor/buffer shims.
|
||||
# Import it before local kernel modules so its vendored FlyDSL path is active.
|
||||
import aiter as _aiter # noqa: F401
|
||||
|
||||
from .kimi_k3_kda_decode import (
|
||||
flydsl_kimi_k3_kda_decode,
|
||||
flydsl_kimi_k3_kda_decode_with_f_b,
|
||||
is_flydsl_kimi_k3_kda_decode_supported,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"flydsl_kimi_k3_kda_decode",
|
||||
"flydsl_kimi_k3_kda_decode_with_f_b",
|
||||
"is_flydsl_kimi_k3_kda_decode_supported",
|
||||
]
|
||||
@@ -0,0 +1,564 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.
|
||||
|
||||
"""FlyDSL kernel for the fused Kimi-K3 KDA decode path on gfx950."""
|
||||
|
||||
import functools
|
||||
import math
|
||||
|
||||
import flydsl.compiler as flyc
|
||||
import flydsl.expr as fx
|
||||
from aiter.ops.flydsl.kernels import vector
|
||||
from aiter.ops.flydsl.kernels.tensor_shim import GTensor, _to_raw
|
||||
from flydsl._mlir import ir
|
||||
from flydsl._mlir.dialects import gpu as mlir_gpu
|
||||
from flydsl._mlir.dialects import scf
|
||||
from flydsl._mlir.dialects import vector as mlir_vector
|
||||
from flydsl.expr import range_constexpr
|
||||
from flydsl.expr.typing import T
|
||||
|
||||
_HEADS = 12
|
||||
_DIM = 128
|
||||
_LOG2E = math.log2(math.e)
|
||||
_SCALE = _DIM**-0.5
|
||||
_BLOCK_THREADS = 256
|
||||
_NUM_WARPS = 4
|
||||
_WARP_SIZE = 64
|
||||
_WARP_THREADS_K = 8
|
||||
_VALUES_PER_THREAD_K = 4
|
||||
_WARP_TILE_K = _WARP_THREADS_K * _VALUES_PER_THREAD_K
|
||||
_K_ITERS = _DIM // _WARP_TILE_K
|
||||
_WARP_THREADS_V = _WARP_SIZE // _WARP_THREADS_K
|
||||
_V_GROUP_TILE = _NUM_WARPS * _WARP_THREADS_V
|
||||
_V_ITERS = _DIM // _V_GROUP_TILE
|
||||
_WAVES_PER_EU = 3
|
||||
|
||||
|
||||
@functools.cache
|
||||
def create_kimi_k3_kda_decode_kernel(norm_eps: float, lower_bound: float):
|
||||
"""Build the fixed gfx950 BF16 Kimi-K3 decode specialization."""
|
||||
|
||||
@fx.struct
|
||||
class SharedStorage:
|
||||
q: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
k: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
v: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
recurrent_out: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
norm_partial: fx.Array[fx.Float32, 2, 16]
|
||||
|
||||
@flyc.kernel(
|
||||
name="kimi_k3_kda_decode_bf16_gfx950",
|
||||
known_block_size=[_BLOCK_THREADS, 1, 1],
|
||||
)
|
||||
def kernel(
|
||||
x_mem: fx.Tensor,
|
||||
weight_mem: fx.Tensor,
|
||||
conv_state_mem: fx.Tensor,
|
||||
raw_g_mem: fx.Tensor,
|
||||
raw_beta_mem: fx.Tensor,
|
||||
A_log_mem: fx.Tensor,
|
||||
dt_bias_mem: fx.Tensor,
|
||||
state_mem: fx.Tensor,
|
||||
state_indices_mem: fx.Tensor,
|
||||
output_gate_mem: fx.Tensor,
|
||||
norm_weight_mem: fx.Tensor,
|
||||
out_mem: fx.Tensor,
|
||||
batch_size: fx.Int32,
|
||||
stride_x_token: fx.Int32,
|
||||
stride_weight_channel: fx.Int32,
|
||||
stride_weight_width: fx.Int32,
|
||||
stride_conv_slot: fx.Int32,
|
||||
stride_conv_channel: fx.Int32,
|
||||
stride_conv_width: fx.Int32,
|
||||
stride_g_token: fx.Int32,
|
||||
stride_beta_token: fx.Int32,
|
||||
stride_state_slot: fx.Int32,
|
||||
stride_gate_token: fx.Int32,
|
||||
stride_gate_head: fx.Int32,
|
||||
stride_out_token: fx.Int32,
|
||||
stride_out_head: fx.Int32,
|
||||
):
|
||||
del batch_size
|
||||
|
||||
x = GTensor(x_mem, dtype=T.bf16, shape=(-1,))
|
||||
weight = GTensor(weight_mem, dtype=T.f32, shape=(-1,))
|
||||
conv_state = GTensor(conv_state_mem, dtype=T.bf16, shape=(-1,))
|
||||
raw_g = GTensor(raw_g_mem, dtype=T.bf16, shape=(-1,))
|
||||
raw_beta = GTensor(raw_beta_mem, dtype=T.bf16, shape=(-1,))
|
||||
A_log = GTensor(A_log_mem, dtype=T.f32, shape=(-1,))
|
||||
dt_bias = GTensor(dt_bias_mem, dtype=T.f32, shape=(-1,))
|
||||
state = GTensor(state_mem, dtype=T.f32, shape=(-1,))
|
||||
state_indices = GTensor(state_indices_mem, dtype=T.i32, shape=(-1,))
|
||||
output_gate = GTensor(output_gate_mem, dtype=T.bf16, shape=(-1,))
|
||||
norm_weight = GTensor(norm_weight_mem, dtype=T.bf16, shape=(-1,))
|
||||
out = GTensor(out_mem, dtype=T.bf16, shape=(-1,))
|
||||
|
||||
shared = fx.SharedAllocator().allocate(SharedStorage).peek()
|
||||
q_lds = shared.q.ptr
|
||||
k_lds = shared.k.ptr
|
||||
v_lds = shared.v.ptr
|
||||
out_lds = shared.recurrent_out.ptr
|
||||
norm_lds = shared.norm_partial.ptr
|
||||
|
||||
tid = fx.thread_idx.x
|
||||
block = fx.block_idx.x
|
||||
batch = block // fx.Int32(_HEADS)
|
||||
head = block % fx.Int32(_HEADS)
|
||||
lane = tid % fx.Int32(_WARP_SIZE)
|
||||
warp = tid // fx.Int32(_WARP_SIZE)
|
||||
lane_k = lane % fx.Int32(_WARP_THREADS_K)
|
||||
|
||||
state_idx = fx.Int32(state_indices[batch])
|
||||
valid = state_idx > fx.Int32(0)
|
||||
|
||||
valid_if = scf.IfOp(_to_raw(valid), results_=[], has_else=True)
|
||||
with ir.InsertionPoint(valid_if.then_block):
|
||||
# A workgroup exclusively owns all three convolution channels for
|
||||
# its (batch, head), so every cache entry is shifted exactly once.
|
||||
conv_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(conv_if.then_block):
|
||||
channel_local = tid
|
||||
q_channel = head * fx.Int32(_DIM) + channel_local
|
||||
k_channel = (
|
||||
fx.Int32(_HEADS * _DIM) + head * fx.Int32(_DIM) + channel_local
|
||||
)
|
||||
v_channel = (
|
||||
fx.Int32(2 * _HEADS * _DIM) + head * fx.Int32(_DIM) + channel_local
|
||||
)
|
||||
|
||||
def convolve_channel(channel):
|
||||
cs_base = (
|
||||
state_idx * stride_conv_slot + channel * stride_conv_channel
|
||||
)
|
||||
c0 = fx.Float32(conv_state[cs_base])
|
||||
c1 = fx.Float32(conv_state[cs_base + stride_conv_width])
|
||||
c2 = fx.Float32(
|
||||
conv_state[cs_base + fx.Int32(2) * stride_conv_width]
|
||||
)
|
||||
current = fx.BFloat16(x[batch * stride_x_token + channel])
|
||||
current_f32 = fx.Float32(current)
|
||||
w_base = channel * stride_weight_channel
|
||||
acc = c0 * fx.Float32(weight[w_base])
|
||||
acc = acc + c1 * fx.Float32(weight[w_base + stride_weight_width])
|
||||
acc = acc + c2 * fx.Float32(
|
||||
weight[w_base + fx.Int32(2) * stride_weight_width]
|
||||
)
|
||||
acc = acc + current_f32 * fx.Float32(
|
||||
weight[w_base + fx.Int32(3) * stride_weight_width]
|
||||
)
|
||||
silu = acc / (
|
||||
fx.Float32(1.0) + fx.math.exp2(-acc * fx.Float32(_LOG2E))
|
||||
)
|
||||
conv_state.store(
|
||||
cs_base,
|
||||
fx.BFloat16(c1),
|
||||
)
|
||||
conv_state.store(
|
||||
cs_base + stride_conv_width,
|
||||
fx.BFloat16(c2),
|
||||
)
|
||||
conv_state.store(
|
||||
cs_base + fx.Int32(2) * stride_conv_width,
|
||||
current,
|
||||
)
|
||||
return silu.to(fx.BFloat16)
|
||||
|
||||
q_conv = convolve_channel(q_channel)
|
||||
k_conv = convolve_channel(k_channel)
|
||||
v_conv = convolve_channel(v_channel)
|
||||
fx.ptr_store(q_conv, q_lds + tid)
|
||||
fx.ptr_store(k_conv, k_lds + tid)
|
||||
fx.ptr_store(v_conv, v_lds + tid)
|
||||
scf.YieldOp([])
|
||||
|
||||
fx.gpu.barrier()
|
||||
|
||||
# Four waves split V into 32-row groups. Eight-lane subgroups
|
||||
# reduce K; each lane issues one aligned f32x4 state transaction.
|
||||
k_vec_start = lane_k * fx.Int32(_VALUES_PER_THREAD_K)
|
||||
global_v_start = warp * fx.Int32(_WARP_THREADS_V) + lane // fx.Int32(
|
||||
_WARP_THREADS_K
|
||||
)
|
||||
vec_f32 = T.vec(_VALUES_PER_THREAD_K, T.f32)
|
||||
vec_bf16 = T.vec(_VALUES_PER_THREAD_K, T.bf16)
|
||||
zero_vec = fx.full(
|
||||
_VALUES_PER_THREAD_K,
|
||||
0.0,
|
||||
fx.Float32,
|
||||
)
|
||||
|
||||
q_vecs = []
|
||||
k_vecs = []
|
||||
decay_vecs = []
|
||||
sum_q_partial = fx.Float32(0.0)
|
||||
sum_k_partial = fx.Float32(0.0)
|
||||
a = fx.math.exp2(fx.Float32(A_log[head]) * fx.Float32(_LOG2E))
|
||||
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
k_base = k_vec_start + fx.Int32(ki * _WARP_TILE_K)
|
||||
q_bf16 = fx.ptr_load(
|
||||
q_lds + k_base,
|
||||
result_type=vec_bf16,
|
||||
)
|
||||
k_bf16 = fx.ptr_load(
|
||||
k_lds + k_base,
|
||||
result_type=vec_bf16,
|
||||
)
|
||||
q_f32 = q_bf16.extf(vec_f32)
|
||||
k_f32 = k_bf16.extf(vec_f32)
|
||||
q_vecs.append(q_f32)
|
||||
k_vecs.append(k_f32)
|
||||
sum_q_vec = q_f32 * q_f32
|
||||
sum_k_vec = k_f32 * k_f32
|
||||
sum_q_partial = (
|
||||
sum_q_partial
|
||||
+ mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
sum_q_vec,
|
||||
).dest
|
||||
)
|
||||
sum_k_partial = (
|
||||
sum_k_partial
|
||||
+ mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
sum_k_vec,
|
||||
).dest
|
||||
)
|
||||
|
||||
gate_bf16 = raw_g.vec_load(
|
||||
(batch * stride_g_token + head * fx.Int32(_DIM) + k_base,),
|
||||
_VALUES_PER_THREAD_K,
|
||||
)
|
||||
gate_f32 = gate_bf16.extf(vec_f32)
|
||||
dt = dt_bias.vec_load(
|
||||
(head * fx.Int32(_DIM) + k_base,),
|
||||
_VALUES_PER_THREAD_K,
|
||||
)
|
||||
sigmoid_arg = (gate_f32 + dt) * a
|
||||
gate = fx.Float32(lower_bound) / (
|
||||
fx.Float32(1.0) + fx.math.exp2(-sigmoid_arg * fx.Float32(_LOG2E))
|
||||
)
|
||||
decay_vecs.append(fx.math.exp2(gate * fx.Float32(_LOG2E)))
|
||||
|
||||
width = fx.Int32(_WARP_SIZE)
|
||||
for offset in (1, 2, 4):
|
||||
sum_q_partial = (
|
||||
sum_q_partial
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_q_partial),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
sum_k_partial = (
|
||||
sum_k_partial
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_k_partial),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
|
||||
subgroup_leader = (lane // fx.Int32(_WARP_THREADS_K)) * fx.Int32(
|
||||
_WARP_THREADS_K
|
||||
)
|
||||
norm_q = mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_q_partial),
|
||||
_to_raw(subgroup_leader),
|
||||
_to_raw(width),
|
||||
mode="idx",
|
||||
).shuffleResult
|
||||
norm_k = mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_k_partial),
|
||||
_to_raw(subgroup_leader),
|
||||
_to_raw(width),
|
||||
mode="idx",
|
||||
).shuffleResult
|
||||
inv_q = fx.math.rsqrt(fx.Float32(norm_q) + fx.Float32(1e-6))
|
||||
inv_k = fx.math.rsqrt(fx.Float32(norm_k) + fx.Float32(1e-6))
|
||||
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
q_vecs[ki] = q_vecs[ki] * fx.Float32(inv_q) * fx.Float32(_SCALE)
|
||||
k_vecs[ki] = k_vecs[ki] * fx.Float32(inv_k)
|
||||
|
||||
dot_kq_vec = zero_vec
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
dot_kq_vec = mlir_vector.FMAOp(
|
||||
k_vecs[ki],
|
||||
q_vecs[ki],
|
||||
dot_kq_vec,
|
||||
).result
|
||||
dot_kq = mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
dot_kq_vec,
|
||||
).dest
|
||||
for offset in (1, 2, 4):
|
||||
dot_kq = (
|
||||
dot_kq
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(dot_kq),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
|
||||
beta_value = fx.Float32(raw_beta[batch * stride_beta_token + head])
|
||||
beta = fx.Float32(1.0) / (
|
||||
fx.Float32(1.0) + fx.math.exp2(-beta_value * fx.Float32(_LOG2E))
|
||||
)
|
||||
state_head_base = state_idx * stride_state_slot + head * fx.Int32(
|
||||
_DIM * _DIM
|
||||
)
|
||||
|
||||
state_vecs = []
|
||||
for vi in range_constexpr(_V_ITERS):
|
||||
global_v = global_v_start + fx.Int32(vi * _V_GROUP_TILE)
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
k_base = k_vec_start + fx.Int32(ki * _WARP_TILE_K)
|
||||
state_off = state_head_base + global_v * fx.Int32(_DIM) + k_base
|
||||
state_vecs.append(
|
||||
state.vec_load(
|
||||
(state_off,),
|
||||
_VALUES_PER_THREAD_K,
|
||||
)
|
||||
)
|
||||
|
||||
for vi in range_constexpr(_V_ITERS):
|
||||
global_v = global_v_start + fx.Int32(vi * _V_GROUP_TILE)
|
||||
sum_hk_vec = zero_vec
|
||||
sum_hq_vec = zero_vec
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
state_pos = vi * _K_ITERS + ki
|
||||
decayed = state_vecs[state_pos] * decay_vecs[ki]
|
||||
state_vecs[state_pos] = decayed
|
||||
sum_hk_vec = mlir_vector.FMAOp(
|
||||
decayed,
|
||||
k_vecs[ki],
|
||||
sum_hk_vec,
|
||||
).result
|
||||
sum_hq_vec = mlir_vector.FMAOp(
|
||||
decayed,
|
||||
q_vecs[ki],
|
||||
sum_hq_vec,
|
||||
).result
|
||||
|
||||
sum_hk = mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
sum_hk_vec,
|
||||
).dest
|
||||
sum_hq = mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
sum_hq_vec,
|
||||
).dest
|
||||
for offset in (1, 2, 4):
|
||||
sum_hk = (
|
||||
sum_hk
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_hk),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
sum_hq = (
|
||||
sum_hq
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_hq),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
|
||||
conv_v = fx.Float32(fx.ptr_load(v_lds + global_v))
|
||||
v_new = (conv_v - fx.Float32(sum_hk)) * beta
|
||||
v_new = mlir_gpu.ShuffleOp(
|
||||
_to_raw(v_new),
|
||||
_to_raw(subgroup_leader),
|
||||
_to_raw(width),
|
||||
mode="idx",
|
||||
).shuffleResult
|
||||
recurrent_value = fx.Float32(sum_hq) + fx.Float32(v_new) * fx.Float32(
|
||||
dot_kq
|
||||
)
|
||||
v_new_vec = mlir_vector.BroadcastOp(
|
||||
vec_f32,
|
||||
_to_raw(v_new),
|
||||
).vector
|
||||
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
state_pos = vi * _K_ITERS + ki
|
||||
updated = mlir_vector.FMAOp(
|
||||
k_vecs[ki],
|
||||
v_new_vec,
|
||||
state_vecs[state_pos],
|
||||
).result
|
||||
k_base = k_vec_start + fx.Int32(ki * _WARP_TILE_K)
|
||||
state_off = state_head_base + global_v * fx.Int32(_DIM) + k_base
|
||||
state.vec_store(
|
||||
(state_off,),
|
||||
updated,
|
||||
_VALUES_PER_THREAD_K,
|
||||
)
|
||||
|
||||
if lane_k == fx.Int32(0):
|
||||
fx.ptr_store(
|
||||
fx.BFloat16(recurrent_value),
|
||||
out_lds + global_v,
|
||||
)
|
||||
|
||||
fx.gpu.barrier()
|
||||
|
||||
# Preserve the model's BF16 boundary before RMSNorm and gating.
|
||||
output_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(output_if.then_block):
|
||||
recurrent_bf16 = fx.ptr_load(out_lds + tid)
|
||||
recurrent_f32 = fx.Float32(recurrent_bf16)
|
||||
square = recurrent_f32 * recurrent_f32
|
||||
for offset in (32, 16, 8, 4, 2, 1):
|
||||
square = (
|
||||
square
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(square),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
if lane == fx.Int32(0):
|
||||
fx.ptr_store(square, norm_lds + warp)
|
||||
scf.YieldOp([])
|
||||
|
||||
fx.gpu.barrier()
|
||||
|
||||
output_store_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(output_store_if.then_block):
|
||||
norm_sum = fx.Float32(fx.ptr_load(norm_lds))
|
||||
norm_sum = norm_sum + fx.Float32(fx.ptr_load(norm_lds + fx.Int32(1)))
|
||||
inv_rms = fx.math.rsqrt(
|
||||
norm_sum * fx.Float32(1.0 / _DIM) + fx.Float32(norm_eps)
|
||||
)
|
||||
recurrent_f32 = fx.Float32(fx.ptr_load(out_lds + tid))
|
||||
norm_w = fx.Float32(norm_weight[tid])
|
||||
gate_value = fx.Float32(
|
||||
output_gate[
|
||||
batch * stride_gate_token + head * stride_gate_head + tid
|
||||
]
|
||||
)
|
||||
output_sigmoid = fx.Float32(1.0) / (
|
||||
fx.Float32(1.0) + fx.math.exp2(-gate_value * fx.Float32(_LOG2E))
|
||||
)
|
||||
result = recurrent_f32 * inv_rms * norm_w * output_sigmoid
|
||||
out.store(
|
||||
batch * stride_out_token + head * stride_out_head + tid,
|
||||
result.to(fx.BFloat16),
|
||||
)
|
||||
scf.YieldOp([])
|
||||
scf.YieldOp([])
|
||||
with ir.InsertionPoint(valid_if.else_block):
|
||||
zero_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(zero_if.then_block):
|
||||
out.store(
|
||||
batch * stride_out_token + head * stride_out_head + tid,
|
||||
fx.BFloat16(0.0),
|
||||
)
|
||||
scf.YieldOp([])
|
||||
scf.YieldOp([])
|
||||
|
||||
@flyc.jit
|
||||
def launch(
|
||||
x_mem: fx.Tensor,
|
||||
weight_mem: fx.Tensor,
|
||||
conv_state_mem: fx.Tensor,
|
||||
raw_g_mem: fx.Tensor,
|
||||
raw_beta_mem: fx.Tensor,
|
||||
A_log_mem: fx.Tensor,
|
||||
dt_bias_mem: fx.Tensor,
|
||||
state_mem: fx.Tensor,
|
||||
state_indices_mem: fx.Tensor,
|
||||
output_gate_mem: fx.Tensor,
|
||||
norm_weight_mem: fx.Tensor,
|
||||
out_mem: fx.Tensor,
|
||||
batch_size: fx.Int32,
|
||||
stride_x_token: fx.Int32,
|
||||
stride_weight_channel: fx.Int32,
|
||||
stride_weight_width: fx.Int32,
|
||||
stride_conv_slot: fx.Int32,
|
||||
stride_conv_channel: fx.Int32,
|
||||
stride_conv_width: fx.Int32,
|
||||
stride_g_token: fx.Int32,
|
||||
stride_beta_token: fx.Int32,
|
||||
stride_state_slot: fx.Int32,
|
||||
stride_gate_token: fx.Int32,
|
||||
stride_gate_head: fx.Int32,
|
||||
stride_out_token: fx.Int32,
|
||||
stride_out_head: fx.Int32,
|
||||
stream: fx.Stream = fx.Stream(None), # noqa: B008
|
||||
):
|
||||
kernel(
|
||||
x_mem,
|
||||
weight_mem,
|
||||
conv_state_mem,
|
||||
raw_g_mem,
|
||||
raw_beta_mem,
|
||||
A_log_mem,
|
||||
dt_bias_mem,
|
||||
state_mem,
|
||||
state_indices_mem,
|
||||
output_gate_mem,
|
||||
norm_weight_mem,
|
||||
out_mem,
|
||||
batch_size,
|
||||
stride_x_token,
|
||||
stride_weight_channel,
|
||||
stride_weight_width,
|
||||
stride_conv_slot,
|
||||
stride_conv_channel,
|
||||
stride_conv_width,
|
||||
stride_g_token,
|
||||
stride_beta_token,
|
||||
stride_state_slot,
|
||||
stride_gate_token,
|
||||
stride_gate_head,
|
||||
stride_out_token,
|
||||
stride_out_head,
|
||||
).launch(
|
||||
grid=(batch_size * fx.Int32(_HEADS), 1, 1),
|
||||
block=(_BLOCK_THREADS, 1, 1),
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
launch.compile_hints = {
|
||||
"waves_per_eu": _WAVES_PER_EU,
|
||||
"llvm_options": {
|
||||
"amdgpu-expert-scheduling-mode": True,
|
||||
},
|
||||
}
|
||||
return launch
|
||||
|
||||
|
||||
__all__ = ["create_kimi_k3_kda_decode_kernel"]
|
||||
@@ -0,0 +1,779 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.
|
||||
|
||||
"""FlyDSL Kimi-K3 KDA decode with a fused head-local f_b projection."""
|
||||
|
||||
import functools
|
||||
import math
|
||||
|
||||
import flydsl.compiler as flyc
|
||||
import flydsl.expr as fx
|
||||
from aiter.ops.flydsl.kernels import vector
|
||||
from aiter.ops.flydsl.kernels.tensor_shim import (
|
||||
AITER_FLYDSL_KERNARG_PRELOAD,
|
||||
AITER_FLYDSL_KERNARG_PRELOAD_COUNT,
|
||||
GTensor,
|
||||
_to_raw,
|
||||
)
|
||||
from flydsl._mlir import ir
|
||||
from flydsl._mlir.dialects import gpu as mlir_gpu
|
||||
from flydsl._mlir.dialects import llvm, scf
|
||||
from flydsl._mlir.dialects import vector as mlir_vector
|
||||
from flydsl.expr import arith, const_expr, range_constexpr
|
||||
from flydsl.expr.arith import ArithValue
|
||||
from flydsl.expr.typing import T
|
||||
|
||||
_HEADS = 12
|
||||
_DIM = 128
|
||||
_LOG2E = math.log2(math.e)
|
||||
_SCALE = _DIM**-0.5
|
||||
_BLOCK_THREADS = 256
|
||||
_NUM_WARPS = 4
|
||||
_WARP_SIZE = 64
|
||||
_WARP_THREADS_K = 8
|
||||
_VALUES_PER_THREAD_K = 4
|
||||
_WARP_TILE_K = _WARP_THREADS_K * _VALUES_PER_THREAD_K
|
||||
_K_ITERS = _DIM // _WARP_TILE_K
|
||||
_WARP_THREADS_V = _WARP_SIZE // _WARP_THREADS_K
|
||||
_V_GROUP_TILE = _NUM_WARPS * _WARP_THREADS_V
|
||||
_V_ITERS = _DIM // _V_GROUP_TILE
|
||||
_PROJECTION_VECTOR = 4
|
||||
_PROJECTION_ITERS = _DIM // _PROJECTION_VECTOR
|
||||
_DEFAULT_WAVES_PER_EU = 2
|
||||
|
||||
|
||||
@functools.cache
|
||||
def create_kimi_k3_kda_decode_fb_kernel(
|
||||
norm_eps: float,
|
||||
lower_bound: float,
|
||||
*,
|
||||
waves_per_eu: int = _DEFAULT_WAVES_PER_EU,
|
||||
cooperative_f_a: bool = False,
|
||||
parallel_front: bool = False,
|
||||
fused_norm_reduce: bool = False,
|
||||
projection_fdot2: bool = False,
|
||||
):
|
||||
"""Build the fixed gfx950 BF16 f_b plus KDA decode specialization."""
|
||||
conv_tid_offset = _DIM if parallel_front else 0
|
||||
conv_tid_upper = 2 * _DIM if parallel_front else _DIM
|
||||
|
||||
if cooperative_f_a:
|
||||
|
||||
@fx.struct
|
||||
class SharedStorage:
|
||||
f_a: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
q: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
k: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
v: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
gate: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
recurrent_out: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
norm_partial: fx.Array[fx.Float32, 4, 16]
|
||||
|
||||
else:
|
||||
|
||||
@fx.struct
|
||||
class SharedStorage:
|
||||
q: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
k: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
v: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
gate: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
recurrent_out: fx.Array[fx.BFloat16, _DIM, 16]
|
||||
norm_partial: fx.Array[fx.Float32, 4, 16]
|
||||
|
||||
kernel_name = "kimi_k3_kda_decode_fb_bf16_gfx950"
|
||||
if (
|
||||
cooperative_f_a
|
||||
or parallel_front
|
||||
or fused_norm_reduce
|
||||
or projection_fdot2
|
||||
or waves_per_eu != _DEFAULT_WAVES_PER_EU
|
||||
):
|
||||
kernel_name += (
|
||||
f"_wpe{waves_per_eu}_cfa{int(cooperative_f_a)}"
|
||||
f"_pf{int(parallel_front)}"
|
||||
f"_fnr{int(fused_norm_reduce)}"
|
||||
f"_fd2{int(projection_fdot2)}"
|
||||
)
|
||||
|
||||
@flyc.kernel(
|
||||
name=kernel_name,
|
||||
known_block_size=[_BLOCK_THREADS, 1, 1],
|
||||
)
|
||||
def kernel(
|
||||
f_a_mem: fx.Tensor,
|
||||
f_b_weight_mem: fx.Tensor,
|
||||
x_mem: fx.Tensor,
|
||||
weight_mem: fx.Tensor,
|
||||
conv_state_mem: fx.Tensor,
|
||||
raw_beta_mem: fx.Tensor,
|
||||
A_log_mem: fx.Tensor,
|
||||
dt_bias_mem: fx.Tensor,
|
||||
state_mem: fx.Tensor,
|
||||
state_indices_mem: fx.Tensor,
|
||||
output_gate_mem: fx.Tensor,
|
||||
norm_weight_mem: fx.Tensor,
|
||||
out_mem: fx.Tensor,
|
||||
batch_size: fx.Int32,
|
||||
stride_f_a_token: fx.Int32,
|
||||
stride_f_b_head: fx.Int32,
|
||||
stride_f_b_output: fx.Int32,
|
||||
stride_x_token: fx.Int32,
|
||||
stride_weight_channel: fx.Int32,
|
||||
stride_weight_width: fx.Int32,
|
||||
stride_conv_slot: fx.Int32,
|
||||
stride_conv_channel: fx.Int32,
|
||||
stride_conv_width: fx.Int32,
|
||||
stride_beta_token: fx.Int32,
|
||||
stride_state_slot: fx.Int32,
|
||||
stride_gate_token: fx.Int32,
|
||||
stride_gate_head: fx.Int32,
|
||||
stride_out_token: fx.Int32,
|
||||
stride_out_head: fx.Int32,
|
||||
):
|
||||
del batch_size
|
||||
|
||||
f_a = GTensor(f_a_mem, dtype=T.bf16, shape=(-1,))
|
||||
f_b_weight = GTensor(f_b_weight_mem, dtype=T.bf16, shape=(-1,))
|
||||
x = GTensor(x_mem, dtype=T.bf16, shape=(-1,))
|
||||
weight = GTensor(weight_mem, dtype=T.f32, shape=(-1,))
|
||||
conv_state = GTensor(conv_state_mem, dtype=T.bf16, shape=(-1,))
|
||||
raw_beta = GTensor(raw_beta_mem, dtype=T.bf16, shape=(-1,))
|
||||
A_log = GTensor(A_log_mem, dtype=T.f32, shape=(-1,))
|
||||
dt_bias = GTensor(dt_bias_mem, dtype=T.f32, shape=(-1,))
|
||||
state = GTensor(state_mem, dtype=T.f32, shape=(-1,))
|
||||
state_indices = GTensor(state_indices_mem, dtype=T.i32, shape=(-1,))
|
||||
output_gate = GTensor(output_gate_mem, dtype=T.bf16, shape=(-1,))
|
||||
norm_weight = GTensor(norm_weight_mem, dtype=T.bf16, shape=(-1,))
|
||||
out = GTensor(out_mem, dtype=T.bf16, shape=(-1,))
|
||||
|
||||
shared = fx.SharedAllocator().allocate(SharedStorage).peek()
|
||||
f_a_lds = shared.f_a.ptr if cooperative_f_a else shared.q.ptr
|
||||
q_lds = shared.q.ptr
|
||||
k_lds = shared.k.ptr
|
||||
v_lds = shared.v.ptr
|
||||
gate_lds = shared.gate.ptr
|
||||
out_lds = shared.recurrent_out.ptr
|
||||
norm_lds = shared.norm_partial.ptr
|
||||
|
||||
tid = fx.thread_idx.x
|
||||
block = fx.block_idx.x
|
||||
batch = block // fx.Int32(_HEADS)
|
||||
head = block % fx.Int32(_HEADS)
|
||||
lane = tid % fx.Int32(_WARP_SIZE)
|
||||
warp = tid // fx.Int32(_WARP_SIZE)
|
||||
lane_k = lane % fx.Int32(_WARP_THREADS_K)
|
||||
|
||||
state_idx = fx.Int32(state_indices[batch])
|
||||
valid = state_idx > fx.Int32(0)
|
||||
|
||||
valid_if = scf.IfOp(_to_raw(valid), results_=[], has_else=True)
|
||||
with ir.InsertionPoint(valid_if.then_block):
|
||||
if const_expr(cooperative_f_a):
|
||||
f_a_load_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(f_a_load_if.then_block):
|
||||
fx.ptr_store(
|
||||
fx.BFloat16(f_a[batch * stride_f_a_token + tid]),
|
||||
f_a_lds + tid,
|
||||
)
|
||||
scf.YieldOp([])
|
||||
fx.gpu.barrier()
|
||||
|
||||
# Threads 0..127 own one output each. Accumulation is FP32 and the
|
||||
# single BF16 store is the same numerical boundary as F.linear.
|
||||
projection_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(projection_if.then_block):
|
||||
i1 = ir.IntegerType.get_signless(1)
|
||||
vec_f32_projection = T.vec(_PROJECTION_VECTOR, T.f32)
|
||||
vec2_bf16 = T.vec(2, T.bf16)
|
||||
accum = fx.full(
|
||||
_PROJECTION_VECTOR,
|
||||
0.0,
|
||||
fx.Float32,
|
||||
)
|
||||
local_dot = fx.Float32(0.0)
|
||||
f_a_base = batch * stride_f_a_token
|
||||
f_b_base = head * stride_f_b_head + tid * stride_f_b_output
|
||||
for projection_iter in range_constexpr(_PROJECTION_ITERS):
|
||||
projection_offset = fx.Int32(projection_iter * _PROJECTION_VECTOR)
|
||||
if const_expr(cooperative_f_a):
|
||||
f_a_values = fx.ptr_load(
|
||||
f_a_lds + projection_offset,
|
||||
result_type=T.vec(_PROJECTION_VECTOR, T.bf16),
|
||||
).extf(vec_f32_projection)
|
||||
else:
|
||||
f_a_values = f_a.vec_load(
|
||||
(f_a_base + projection_offset,),
|
||||
_PROJECTION_VECTOR,
|
||||
).extf(vec_f32_projection)
|
||||
weight_values = f_b_weight.vec_load(
|
||||
(f_b_base + projection_offset,),
|
||||
_PROJECTION_VECTOR,
|
||||
).extf(vec_f32_projection)
|
||||
if const_expr(projection_fdot2):
|
||||
f_a_bf16 = f_a_values.truncf(T.vec(_PROJECTION_VECTOR, T.bf16))
|
||||
weight_bf16 = weight_values.truncf(
|
||||
T.vec(_PROJECTION_VECTOR, T.bf16)
|
||||
)
|
||||
for pair_index in range_constexpr(_PROJECTION_VECTOR // 2):
|
||||
f_a_pair = vector.from_elements(
|
||||
vec2_bf16,
|
||||
[
|
||||
vector.extract(
|
||||
f_a_bf16,
|
||||
static_position=[pair_index * 2],
|
||||
dynamic_position=[],
|
||||
),
|
||||
vector.extract(
|
||||
f_a_bf16,
|
||||
static_position=[pair_index * 2 + 1],
|
||||
dynamic_position=[],
|
||||
),
|
||||
],
|
||||
)
|
||||
weight_pair = vector.from_elements(
|
||||
vec2_bf16,
|
||||
[
|
||||
vector.extract(
|
||||
weight_bf16,
|
||||
static_position=[pair_index * 2],
|
||||
dynamic_position=[],
|
||||
),
|
||||
vector.extract(
|
||||
weight_bf16,
|
||||
static_position=[pair_index * 2 + 1],
|
||||
dynamic_position=[],
|
||||
),
|
||||
],
|
||||
)
|
||||
local_dot = ArithValue(
|
||||
llvm.call_intrinsic(
|
||||
T.f32,
|
||||
"llvm.amdgcn.fdot2.f32.bf16",
|
||||
[
|
||||
f_a_pair,
|
||||
weight_pair,
|
||||
_to_raw(local_dot),
|
||||
arith.constant(False, type=i1),
|
||||
],
|
||||
[],
|
||||
[],
|
||||
)
|
||||
)
|
||||
else:
|
||||
accum = mlir_vector.FMAOp(
|
||||
f_a_values,
|
||||
weight_values,
|
||||
accum,
|
||||
).result
|
||||
if const_expr(projection_fdot2):
|
||||
projected = local_dot
|
||||
else:
|
||||
projected = mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
accum,
|
||||
).dest
|
||||
fx.ptr_store(
|
||||
fx.BFloat16(projected),
|
||||
gate_lds + tid,
|
||||
)
|
||||
scf.YieldOp([])
|
||||
|
||||
# A workgroup exclusively owns all three convolution channels for
|
||||
# its (batch, head), so every cache entry is shifted exactly once.
|
||||
conv_if = scf.IfOp(
|
||||
_to_raw(
|
||||
(tid >= fx.Int32(conv_tid_offset))
|
||||
& (tid < fx.Int32(conv_tid_upper))
|
||||
),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(conv_if.then_block):
|
||||
channel_local = tid - fx.Int32(conv_tid_offset)
|
||||
q_channel = head * fx.Int32(_DIM) + channel_local
|
||||
k_channel = (
|
||||
fx.Int32(_HEADS * _DIM) + head * fx.Int32(_DIM) + channel_local
|
||||
)
|
||||
v_channel = (
|
||||
fx.Int32(2 * _HEADS * _DIM) + head * fx.Int32(_DIM) + channel_local
|
||||
)
|
||||
|
||||
def convolve_channel(channel):
|
||||
cs_base = (
|
||||
state_idx * stride_conv_slot + channel * stride_conv_channel
|
||||
)
|
||||
c0 = fx.Float32(conv_state[cs_base])
|
||||
c1 = fx.Float32(conv_state[cs_base + stride_conv_width])
|
||||
c2 = fx.Float32(
|
||||
conv_state[cs_base + fx.Int32(2) * stride_conv_width]
|
||||
)
|
||||
current = fx.BFloat16(x[batch * stride_x_token + channel])
|
||||
current_f32 = fx.Float32(current)
|
||||
w_base = channel * stride_weight_channel
|
||||
acc = c0 * fx.Float32(weight[w_base])
|
||||
acc = acc + c1 * fx.Float32(weight[w_base + stride_weight_width])
|
||||
acc = acc + c2 * fx.Float32(
|
||||
weight[w_base + fx.Int32(2) * stride_weight_width]
|
||||
)
|
||||
acc = acc + current_f32 * fx.Float32(
|
||||
weight[w_base + fx.Int32(3) * stride_weight_width]
|
||||
)
|
||||
silu = acc / (
|
||||
fx.Float32(1.0) + fx.math.exp2(-acc * fx.Float32(_LOG2E))
|
||||
)
|
||||
conv_state.store(
|
||||
cs_base,
|
||||
fx.BFloat16(c1),
|
||||
)
|
||||
conv_state.store(
|
||||
cs_base + stride_conv_width,
|
||||
fx.BFloat16(c2),
|
||||
)
|
||||
conv_state.store(
|
||||
cs_base + fx.Int32(2) * stride_conv_width,
|
||||
current,
|
||||
)
|
||||
return silu.to(fx.BFloat16)
|
||||
|
||||
q_conv = convolve_channel(q_channel)
|
||||
k_conv = convolve_channel(k_channel)
|
||||
v_conv = convolve_channel(v_channel)
|
||||
fx.ptr_store(q_conv, q_lds + channel_local)
|
||||
fx.ptr_store(k_conv, k_lds + channel_local)
|
||||
fx.ptr_store(v_conv, v_lds + channel_local)
|
||||
scf.YieldOp([])
|
||||
|
||||
# Both projection and convolution LDS values must be visible
|
||||
# before the recurrent core begins.
|
||||
fx.gpu.barrier()
|
||||
|
||||
# Four waves split V into 32-row groups. Eight-lane subgroups
|
||||
# reduce K; each lane issues one aligned f32x4 state transaction.
|
||||
k_vec_start = lane_k * fx.Int32(_VALUES_PER_THREAD_K)
|
||||
global_v_start = warp * fx.Int32(_WARP_THREADS_V) + lane // fx.Int32(
|
||||
_WARP_THREADS_K
|
||||
)
|
||||
vec_f32 = T.vec(_VALUES_PER_THREAD_K, T.f32)
|
||||
vec_bf16 = T.vec(_VALUES_PER_THREAD_K, T.bf16)
|
||||
zero_vec = fx.full(
|
||||
_VALUES_PER_THREAD_K,
|
||||
0.0,
|
||||
fx.Float32,
|
||||
)
|
||||
|
||||
q_vecs = []
|
||||
k_vecs = []
|
||||
decay_vecs = []
|
||||
sum_q_partial = fx.Float32(0.0)
|
||||
sum_k_partial = fx.Float32(0.0)
|
||||
a = fx.math.exp2(fx.Float32(A_log[head]) * fx.Float32(_LOG2E))
|
||||
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
k_base = k_vec_start + fx.Int32(ki * _WARP_TILE_K)
|
||||
q_bf16 = fx.ptr_load(
|
||||
q_lds + k_base,
|
||||
result_type=vec_bf16,
|
||||
)
|
||||
k_bf16 = fx.ptr_load(
|
||||
k_lds + k_base,
|
||||
result_type=vec_bf16,
|
||||
)
|
||||
q_f32 = q_bf16.extf(vec_f32)
|
||||
k_f32 = k_bf16.extf(vec_f32)
|
||||
q_vecs.append(q_f32)
|
||||
k_vecs.append(k_f32)
|
||||
sum_q_vec = q_f32 * q_f32
|
||||
sum_k_vec = k_f32 * k_f32
|
||||
sum_q_partial = (
|
||||
sum_q_partial
|
||||
+ mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
sum_q_vec,
|
||||
).dest
|
||||
)
|
||||
sum_k_partial = (
|
||||
sum_k_partial
|
||||
+ mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
sum_k_vec,
|
||||
).dest
|
||||
)
|
||||
|
||||
# The projection is rounded in LDS before the lower-bound gate.
|
||||
gate_bf16 = fx.ptr_load(
|
||||
gate_lds + k_base,
|
||||
result_type=vec_bf16,
|
||||
)
|
||||
gate_f32 = gate_bf16.extf(vec_f32)
|
||||
dt = dt_bias.vec_load(
|
||||
(head * fx.Int32(_DIM) + k_base,),
|
||||
_VALUES_PER_THREAD_K,
|
||||
)
|
||||
sigmoid_arg = (gate_f32 + dt) * a
|
||||
gate = fx.Float32(lower_bound) / (
|
||||
fx.Float32(1.0) + fx.math.exp2(-sigmoid_arg * fx.Float32(_LOG2E))
|
||||
)
|
||||
decay_vecs.append(fx.math.exp2(gate * fx.Float32(_LOG2E)))
|
||||
|
||||
width = fx.Int32(_WARP_SIZE)
|
||||
for offset in (1, 2, 4):
|
||||
sum_q_partial = (
|
||||
sum_q_partial
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_q_partial),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
sum_k_partial = (
|
||||
sum_k_partial
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_k_partial),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
|
||||
subgroup_leader = (lane // fx.Int32(_WARP_THREADS_K)) * fx.Int32(
|
||||
_WARP_THREADS_K
|
||||
)
|
||||
norm_q = mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_q_partial),
|
||||
_to_raw(subgroup_leader),
|
||||
_to_raw(width),
|
||||
mode="idx",
|
||||
).shuffleResult
|
||||
norm_k = mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_k_partial),
|
||||
_to_raw(subgroup_leader),
|
||||
_to_raw(width),
|
||||
mode="idx",
|
||||
).shuffleResult
|
||||
inv_q = fx.math.rsqrt(fx.Float32(norm_q) + fx.Float32(1e-6))
|
||||
inv_k = fx.math.rsqrt(fx.Float32(norm_k) + fx.Float32(1e-6))
|
||||
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
q_vecs[ki] = q_vecs[ki] * fx.Float32(inv_q) * fx.Float32(_SCALE)
|
||||
k_vecs[ki] = k_vecs[ki] * fx.Float32(inv_k)
|
||||
|
||||
dot_kq_vec = zero_vec
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
dot_kq_vec = mlir_vector.FMAOp(
|
||||
k_vecs[ki],
|
||||
q_vecs[ki],
|
||||
dot_kq_vec,
|
||||
).result
|
||||
dot_kq = mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
dot_kq_vec,
|
||||
).dest
|
||||
for offset in (1, 2, 4):
|
||||
dot_kq = (
|
||||
dot_kq
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(dot_kq),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
|
||||
beta_value = fx.Float32(raw_beta[batch * stride_beta_token + head])
|
||||
beta = fx.Float32(1.0) / (
|
||||
fx.Float32(1.0) + fx.math.exp2(-beta_value * fx.Float32(_LOG2E))
|
||||
)
|
||||
state_head_base = state_idx * stride_state_slot + head * fx.Int32(
|
||||
_DIM * _DIM
|
||||
)
|
||||
|
||||
def process_state_row(vi, row_state_vecs):
|
||||
global_v = global_v_start + fx.Int32(vi * _V_GROUP_TILE)
|
||||
sum_hk_vec = zero_vec
|
||||
sum_hq_vec = zero_vec
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
decayed = row_state_vecs[ki] * decay_vecs[ki]
|
||||
row_state_vecs[ki] = decayed
|
||||
sum_hk_vec = mlir_vector.FMAOp(
|
||||
decayed,
|
||||
k_vecs[ki],
|
||||
sum_hk_vec,
|
||||
).result
|
||||
sum_hq_vec = mlir_vector.FMAOp(
|
||||
decayed,
|
||||
q_vecs[ki],
|
||||
sum_hq_vec,
|
||||
).result
|
||||
|
||||
sum_hk = mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
sum_hk_vec,
|
||||
).dest
|
||||
sum_hq = mlir_vector.ReductionOp(
|
||||
T.f32,
|
||||
vector.CombiningKind.ADD,
|
||||
sum_hq_vec,
|
||||
).dest
|
||||
for offset in (1, 2, 4):
|
||||
sum_hk = (
|
||||
sum_hk
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_hk),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
sum_hq = (
|
||||
sum_hq
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(sum_hq),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
|
||||
conv_v = fx.Float32(fx.ptr_load(v_lds + global_v))
|
||||
v_new = (conv_v - fx.Float32(sum_hk)) * beta
|
||||
v_new = mlir_gpu.ShuffleOp(
|
||||
_to_raw(v_new),
|
||||
_to_raw(subgroup_leader),
|
||||
_to_raw(width),
|
||||
mode="idx",
|
||||
).shuffleResult
|
||||
recurrent_value = fx.Float32(sum_hq) + fx.Float32(v_new) * fx.Float32(
|
||||
dot_kq
|
||||
)
|
||||
v_new_vec = mlir_vector.BroadcastOp(
|
||||
vec_f32,
|
||||
_to_raw(v_new),
|
||||
).vector
|
||||
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
updated = mlir_vector.FMAOp(
|
||||
k_vecs[ki],
|
||||
v_new_vec,
|
||||
row_state_vecs[ki],
|
||||
).result
|
||||
k_base = k_vec_start + fx.Int32(ki * _WARP_TILE_K)
|
||||
state_off = state_head_base + global_v * fx.Int32(_DIM) + k_base
|
||||
state.vec_store(
|
||||
(state_off,),
|
||||
updated,
|
||||
_VALUES_PER_THREAD_K,
|
||||
)
|
||||
|
||||
if lane_k == fx.Int32(0):
|
||||
fx.ptr_store(
|
||||
fx.BFloat16(recurrent_value),
|
||||
out_lds + global_v,
|
||||
)
|
||||
rounded = fx.BFloat16(recurrent_value)
|
||||
rounded_f32 = fx.Float32(rounded)
|
||||
return rounded_f32 * rounded_f32
|
||||
|
||||
norm_accum = fx.Float32(0.0)
|
||||
state_vecs = []
|
||||
for vi in range_constexpr(_V_ITERS):
|
||||
global_v = global_v_start + fx.Int32(vi * _V_GROUP_TILE)
|
||||
for ki in range_constexpr(_K_ITERS):
|
||||
k_base = k_vec_start + fx.Int32(ki * _WARP_TILE_K)
|
||||
state_off = state_head_base + global_v * fx.Int32(_DIM) + k_base
|
||||
state_vecs.append(state.vec_load((state_off,), 4))
|
||||
for vi in range_constexpr(_V_ITERS):
|
||||
norm_accum = norm_accum + process_state_row(
|
||||
vi,
|
||||
state_vecs[vi * _K_ITERS : (vi + 1) * _K_ITERS],
|
||||
)
|
||||
|
||||
if const_expr(fused_norm_reduce):
|
||||
for offset in (32, 16, 8, 4, 2, 1):
|
||||
norm_accum = (
|
||||
norm_accum
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(norm_accum),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
if lane == fx.Int32(0):
|
||||
fx.ptr_store(
|
||||
norm_accum * fx.Float32(1.0 / _WARP_THREADS_K),
|
||||
norm_lds + warp,
|
||||
)
|
||||
|
||||
fx.gpu.barrier()
|
||||
|
||||
# Preserve the model's BF16 boundary before RMSNorm and gating.
|
||||
if const_expr(not fused_norm_reduce):
|
||||
output_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(output_if.then_block):
|
||||
recurrent_bf16 = fx.ptr_load(out_lds + tid)
|
||||
recurrent_f32 = fx.Float32(recurrent_bf16)
|
||||
square = recurrent_f32 * recurrent_f32
|
||||
for offset in (32, 16, 8, 4, 2, 1):
|
||||
square = (
|
||||
square
|
||||
+ mlir_gpu.ShuffleOp(
|
||||
_to_raw(square),
|
||||
_to_raw(fx.Int32(offset)),
|
||||
_to_raw(width),
|
||||
mode="xor",
|
||||
).shuffleResult
|
||||
)
|
||||
if lane == fx.Int32(0):
|
||||
fx.ptr_store(square, norm_lds + warp)
|
||||
scf.YieldOp([])
|
||||
|
||||
fx.gpu.barrier()
|
||||
|
||||
output_store_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(output_store_if.then_block):
|
||||
norm_sum = fx.Float32(fx.ptr_load(norm_lds))
|
||||
norm_sum = norm_sum + fx.Float32(fx.ptr_load(norm_lds + fx.Int32(1)))
|
||||
if const_expr(fused_norm_reduce):
|
||||
norm_sum = norm_sum + fx.Float32(
|
||||
fx.ptr_load(norm_lds + fx.Int32(2))
|
||||
)
|
||||
norm_sum = norm_sum + fx.Float32(
|
||||
fx.ptr_load(norm_lds + fx.Int32(3))
|
||||
)
|
||||
inv_rms = fx.math.rsqrt(
|
||||
norm_sum * fx.Float32(1.0 / _DIM) + fx.Float32(norm_eps)
|
||||
)
|
||||
recurrent_f32 = fx.Float32(fx.ptr_load(out_lds + tid))
|
||||
norm_w = fx.Float32(norm_weight[tid])
|
||||
gate_value = fx.Float32(
|
||||
output_gate[
|
||||
batch * stride_gate_token + head * stride_gate_head + tid
|
||||
]
|
||||
)
|
||||
output_sigmoid = fx.Float32(1.0) / (
|
||||
fx.Float32(1.0) + fx.math.exp2(-gate_value * fx.Float32(_LOG2E))
|
||||
)
|
||||
result = recurrent_f32 * inv_rms * norm_w * output_sigmoid
|
||||
out.store(
|
||||
batch * stride_out_token + head * stride_out_head + tid,
|
||||
result.to(fx.BFloat16),
|
||||
)
|
||||
scf.YieldOp([])
|
||||
scf.YieldOp([])
|
||||
with ir.InsertionPoint(valid_if.else_block):
|
||||
zero_if = scf.IfOp(
|
||||
_to_raw(tid < fx.Int32(_DIM)),
|
||||
results_=[],
|
||||
has_else=False,
|
||||
)
|
||||
with ir.InsertionPoint(zero_if.then_block):
|
||||
out.store(
|
||||
batch * stride_out_token + head * stride_out_head + tid,
|
||||
fx.BFloat16(0.0),
|
||||
)
|
||||
scf.YieldOp([])
|
||||
scf.YieldOp([])
|
||||
|
||||
@flyc.jit
|
||||
def launch(
|
||||
f_a_mem: fx.Tensor,
|
||||
f_b_weight_mem: fx.Tensor,
|
||||
x_mem: fx.Tensor,
|
||||
weight_mem: fx.Tensor,
|
||||
conv_state_mem: fx.Tensor,
|
||||
raw_beta_mem: fx.Tensor,
|
||||
A_log_mem: fx.Tensor,
|
||||
dt_bias_mem: fx.Tensor,
|
||||
state_mem: fx.Tensor,
|
||||
state_indices_mem: fx.Tensor,
|
||||
output_gate_mem: fx.Tensor,
|
||||
norm_weight_mem: fx.Tensor,
|
||||
out_mem: fx.Tensor,
|
||||
batch_size: fx.Int32,
|
||||
stride_f_a_token: fx.Int32,
|
||||
stride_f_b_head: fx.Int32,
|
||||
stride_f_b_output: fx.Int32,
|
||||
stride_x_token: fx.Int32,
|
||||
stride_weight_channel: fx.Int32,
|
||||
stride_weight_width: fx.Int32,
|
||||
stride_conv_slot: fx.Int32,
|
||||
stride_conv_channel: fx.Int32,
|
||||
stride_conv_width: fx.Int32,
|
||||
stride_beta_token: fx.Int32,
|
||||
stride_state_slot: fx.Int32,
|
||||
stride_gate_token: fx.Int32,
|
||||
stride_gate_head: fx.Int32,
|
||||
stride_out_token: fx.Int32,
|
||||
stride_out_head: fx.Int32,
|
||||
stream: fx.Stream = fx.Stream(None), # noqa: B008
|
||||
):
|
||||
kernel(
|
||||
f_a_mem,
|
||||
f_b_weight_mem,
|
||||
x_mem,
|
||||
weight_mem,
|
||||
conv_state_mem,
|
||||
raw_beta_mem,
|
||||
A_log_mem,
|
||||
dt_bias_mem,
|
||||
state_mem,
|
||||
state_indices_mem,
|
||||
output_gate_mem,
|
||||
norm_weight_mem,
|
||||
out_mem,
|
||||
batch_size,
|
||||
stride_f_a_token,
|
||||
stride_f_b_head,
|
||||
stride_f_b_output,
|
||||
stride_x_token,
|
||||
stride_weight_channel,
|
||||
stride_weight_width,
|
||||
stride_conv_slot,
|
||||
stride_conv_channel,
|
||||
stride_conv_width,
|
||||
stride_beta_token,
|
||||
stride_state_slot,
|
||||
stride_gate_token,
|
||||
stride_gate_head,
|
||||
stride_out_token,
|
||||
stride_out_head,
|
||||
).launch(
|
||||
grid=(batch_size * fx.Int32(_HEADS), 1, 1),
|
||||
block=(_BLOCK_THREADS, 1, 1),
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
launch.compile_hints = {
|
||||
"waves_per_eu": waves_per_eu,
|
||||
"llvm_options": {
|
||||
"amdgpu-expert-scheduling-mode": True,
|
||||
"amdgpu-kernarg-preload": AITER_FLYDSL_KERNARG_PRELOAD,
|
||||
"amdgpu-kernarg-preload-count": AITER_FLYDSL_KERNARG_PRELOAD_COUNT,
|
||||
},
|
||||
}
|
||||
return launch
|
||||
|
||||
|
||||
__all__ = ["create_kimi_k3_kda_decode_fb_kernel"]
|
||||
@@ -0,0 +1,477 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.
|
||||
|
||||
"""High-level API for the fused Kimi-K3 KDA decode specialization."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import functools
|
||||
from collections.abc import Iterable
|
||||
|
||||
import torch
|
||||
from aiter.ops.flydsl.kernels.tensor_shim import _run_compiled
|
||||
|
||||
from .kernels.kimi_k3_kda_decode import (
|
||||
create_kimi_k3_kda_decode_kernel,
|
||||
)
|
||||
from .kernels.kimi_k3_kda_decode_fb import (
|
||||
create_kimi_k3_kda_decode_fb_kernel,
|
||||
)
|
||||
|
||||
_HEADS = 12
|
||||
_DIM = 128
|
||||
_CONV_CHANNELS = 3 * _HEADS * _DIM
|
||||
_CONV_WIDTH = 4
|
||||
|
||||
|
||||
def _fb_build_options(batch: int) -> dict[str, int | bool]:
|
||||
"""Use the validated gfx950 winner only for the exact-C2 bucket."""
|
||||
if batch != 2:
|
||||
return {}
|
||||
return {
|
||||
"waves_per_eu": 3,
|
||||
"cooperative_f_a": True,
|
||||
"parallel_front": True,
|
||||
"fused_norm_reduce": True,
|
||||
"projection_fdot2": True,
|
||||
}
|
||||
|
||||
|
||||
@functools.cache
|
||||
def _rocm_arch(device: torch.device) -> str | None:
|
||||
properties = torch.cuda.get_device_properties(device)
|
||||
arch = getattr(properties, "gcnArchName", None)
|
||||
return arch.split(":", 1)[0] if arch is not None else None
|
||||
|
||||
|
||||
def is_flydsl_kimi_k3_kda_decode_supported(
|
||||
device: torch.device | str | int | None = None,
|
||||
) -> bool:
|
||||
"""Return whether ``device`` can run this gfx950-only specialization."""
|
||||
if not torch.cuda.is_available():
|
||||
return False
|
||||
try:
|
||||
resolved = torch.device(
|
||||
"cuda",
|
||||
torch.cuda.current_device(),
|
||||
)
|
||||
if device is not None:
|
||||
resolved = (
|
||||
torch.device("cuda", device)
|
||||
if isinstance(device, int)
|
||||
else torch.device(device)
|
||||
)
|
||||
if resolved.type != "cuda":
|
||||
return False
|
||||
if resolved.index is None:
|
||||
resolved = torch.device(
|
||||
"cuda",
|
||||
torch.cuda.current_device(),
|
||||
)
|
||||
return _rocm_arch(resolved) == "gfx950"
|
||||
except (AssertionError, RuntimeError, TypeError, ValueError):
|
||||
return False
|
||||
|
||||
|
||||
def _check_tensor(
|
||||
name: str,
|
||||
tensor: torch.Tensor,
|
||||
*,
|
||||
shape: tuple[int, ...],
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
inner_strides: tuple[int, ...] = (),
|
||||
) -> None:
|
||||
if tensor.shape != shape:
|
||||
raise ValueError(
|
||||
f"`{name}` must have shape {list(shape)}, got {list(tensor.shape)}."
|
||||
)
|
||||
if tensor.dtype != dtype:
|
||||
raise ValueError(f"`{name}` must have dtype {dtype}, got {tensor.dtype}.")
|
||||
if tensor.device != device:
|
||||
raise ValueError(f"`{name}` must be on {device}, got {tensor.device}.")
|
||||
if inner_strides and tensor.stride()[-len(inner_strides) :] != inner_strides:
|
||||
raise ValueError(
|
||||
f"`{name}` must have inner strides {inner_strides}, got {tensor.stride()}."
|
||||
)
|
||||
|
||||
|
||||
def _check_same_device(
|
||||
tensors: Iterable[tuple[str, torch.Tensor]],
|
||||
device: torch.device,
|
||||
) -> None:
|
||||
for name, tensor in tensors:
|
||||
if not tensor.is_cuda:
|
||||
raise ValueError(f"`{name}` must be a CUDA tensor.")
|
||||
if tensor.device != device:
|
||||
raise ValueError(f"`{name}` must be on {device}, got {tensor.device}.")
|
||||
|
||||
|
||||
def _validate_kda_inputs(
|
||||
*,
|
||||
api_name: str,
|
||||
batch_source: str,
|
||||
device: torch.device,
|
||||
batch: int,
|
||||
x: torch.Tensor,
|
||||
conv_weight: torch.Tensor,
|
||||
conv_bias: torch.Tensor | None,
|
||||
conv_state: torch.Tensor,
|
||||
raw_beta: torch.Tensor,
|
||||
A_log: torch.Tensor,
|
||||
dt_bias: torch.Tensor,
|
||||
lower_bound: float | None,
|
||||
state: torch.Tensor,
|
||||
state_indices: torch.Tensor,
|
||||
output_gate: torch.Tensor,
|
||||
norm_weight: torch.Tensor,
|
||||
out: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
"""Validate operands shared by both explicit KDA specializations."""
|
||||
if not is_flydsl_kimi_k3_kda_decode_supported(device):
|
||||
raise RuntimeError(f"`{api_name}` requires a gfx950 GPU.")
|
||||
if batch <= 0:
|
||||
raise ValueError(f"`{batch_source}` must have a non-empty batch dimension.")
|
||||
if conv_bias is not None:
|
||||
raise ValueError("This specialization requires `conv_bias=None`.")
|
||||
if lower_bound is None:
|
||||
raise ValueError("This specialization requires the KDA lower-bound gate.")
|
||||
|
||||
_check_same_device(
|
||||
(
|
||||
("x", x),
|
||||
("conv_weight", conv_weight),
|
||||
("conv_state", conv_state),
|
||||
("raw_beta", raw_beta),
|
||||
("A_log", A_log),
|
||||
("dt_bias", dt_bias),
|
||||
("state", state),
|
||||
("state_indices", state_indices),
|
||||
("output_gate", output_gate),
|
||||
("norm_weight", norm_weight),
|
||||
),
|
||||
device,
|
||||
)
|
||||
_check_tensor(
|
||||
"x",
|
||||
x,
|
||||
shape=(batch, _CONV_CHANNELS),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
_check_tensor(
|
||||
"conv_weight",
|
||||
conv_weight,
|
||||
shape=(_CONV_CHANNELS, _CONV_WIDTH),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
if conv_state.ndim != 3 or conv_state.shape[1:] != (
|
||||
_CONV_CHANNELS,
|
||||
_CONV_WIDTH - 1,
|
||||
):
|
||||
raise ValueError(
|
||||
"`conv_state` must have shape [cache, 4608, 3], "
|
||||
f"got {list(conv_state.shape)}."
|
||||
)
|
||||
if conv_state.dtype != torch.bfloat16:
|
||||
raise ValueError("`conv_state` must have dtype torch.bfloat16.")
|
||||
if state.ndim != 4 or state.shape[1:] != (
|
||||
_HEADS,
|
||||
_DIM,
|
||||
_DIM,
|
||||
):
|
||||
raise ValueError(
|
||||
f"`state` must have shape [cache, 12, 128, 128], got {list(state.shape)}."
|
||||
)
|
||||
if state.dtype != torch.float32:
|
||||
raise ValueError("`state` must have dtype torch.float32.")
|
||||
if state.stride()[-3:] != (_DIM * _DIM, _DIM, 1):
|
||||
raise ValueError("`state` must be contiguous within each cache slot.")
|
||||
_check_tensor(
|
||||
"raw_beta",
|
||||
raw_beta,
|
||||
shape=(1, batch, _HEADS),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
_check_tensor(
|
||||
"A_log",
|
||||
A_log,
|
||||
shape=(_HEADS,),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
_check_tensor(
|
||||
"dt_bias",
|
||||
dt_bias,
|
||||
shape=(_HEADS * _DIM,),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
_check_tensor(
|
||||
"state_indices",
|
||||
state_indices,
|
||||
shape=(batch,),
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
_check_tensor(
|
||||
"output_gate",
|
||||
output_gate,
|
||||
shape=(batch, _HEADS, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
_check_tensor(
|
||||
"norm_weight",
|
||||
norm_weight,
|
||||
shape=(_DIM,),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
|
||||
if out is None:
|
||||
return torch.empty(
|
||||
(1, batch, _HEADS, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
)
|
||||
_check_same_device((("out", out),), device)
|
||||
_check_tensor(
|
||||
"out",
|
||||
out,
|
||||
shape=(1, batch, _HEADS, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def flydsl_kimi_k3_kda_decode(
|
||||
x: torch.Tensor,
|
||||
conv_weight: torch.Tensor,
|
||||
conv_bias: torch.Tensor | None,
|
||||
conv_state: torch.Tensor,
|
||||
raw_g: torch.Tensor,
|
||||
raw_beta: torch.Tensor,
|
||||
A_log: torch.Tensor,
|
||||
dt_bias: torch.Tensor,
|
||||
lower_bound: float | None,
|
||||
state: torch.Tensor,
|
||||
state_indices: torch.Tensor,
|
||||
output_gate: torch.Tensor,
|
||||
norm_weight: torch.Tensor,
|
||||
norm_eps: float,
|
||||
out: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run fused Kimi-K3 KDA decode on MI350-series GPUs.
|
||||
|
||||
This pure-decode specialization fuses the packed width-4 Q/K/V causal
|
||||
convolution, the FP32 recurrent-state update, and the BF16
|
||||
RMSNorm/sigmoid output gate. Slot zero is reserved: non-positive
|
||||
``state_indices`` produce zero output without modifying either cache.
|
||||
|
||||
The layout is fixed to Kimi-K3 TP8: 12 local heads and 128-dimensional
|
||||
key/value state. Call
|
||||
:func:`is_flydsl_kimi_k3_kda_decode_supported` before dispatching from a
|
||||
model implementation.
|
||||
"""
|
||||
if x.ndim != 2:
|
||||
raise ValueError(f"`x` must have rank 2, got rank {x.ndim}.")
|
||||
if not x.is_cuda:
|
||||
raise ValueError("`x` must be a CUDA tensor.")
|
||||
device = x.device
|
||||
batch = x.shape[0]
|
||||
out = _validate_kda_inputs(
|
||||
api_name="flydsl_kimi_k3_kda_decode",
|
||||
batch_source="x",
|
||||
device=device,
|
||||
batch=batch,
|
||||
x=x,
|
||||
conv_weight=conv_weight,
|
||||
conv_bias=conv_bias,
|
||||
conv_state=conv_state,
|
||||
raw_beta=raw_beta,
|
||||
A_log=A_log,
|
||||
dt_bias=dt_bias,
|
||||
lower_bound=lower_bound,
|
||||
state=state,
|
||||
state_indices=state_indices,
|
||||
output_gate=output_gate,
|
||||
norm_weight=norm_weight,
|
||||
out=out,
|
||||
)
|
||||
_check_same_device((("raw_g", raw_g),), device)
|
||||
_check_tensor(
|
||||
"raw_g",
|
||||
raw_g,
|
||||
shape=(1, batch, _HEADS, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
inner_strides=(_DIM, 1),
|
||||
)
|
||||
|
||||
executable = create_kimi_k3_kda_decode_kernel(
|
||||
float(norm_eps),
|
||||
float(lower_bound),
|
||||
)
|
||||
with torch.cuda.device(device):
|
||||
stream = torch.cuda.current_stream(device)
|
||||
_run_compiled(
|
||||
executable,
|
||||
x,
|
||||
conv_weight,
|
||||
conv_state,
|
||||
raw_g,
|
||||
raw_beta,
|
||||
A_log,
|
||||
dt_bias,
|
||||
state,
|
||||
state_indices,
|
||||
output_gate,
|
||||
norm_weight,
|
||||
out,
|
||||
batch,
|
||||
x.stride(0),
|
||||
conv_weight.stride(0),
|
||||
conv_weight.stride(1),
|
||||
conv_state.stride(0),
|
||||
conv_state.stride(1),
|
||||
conv_state.stride(2),
|
||||
raw_g.stride(1),
|
||||
raw_beta.stride(1),
|
||||
state.stride(0),
|
||||
output_gate.stride(0),
|
||||
output_gate.stride(1),
|
||||
out.stride(1),
|
||||
out.stride(2),
|
||||
stream,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def flydsl_kimi_k3_kda_decode_with_f_b(
|
||||
f_a: torch.Tensor,
|
||||
f_b_weight: torch.Tensor,
|
||||
x: torch.Tensor,
|
||||
conv_weight: torch.Tensor,
|
||||
conv_bias: torch.Tensor | None,
|
||||
conv_state: torch.Tensor,
|
||||
raw_beta: torch.Tensor,
|
||||
A_log: torch.Tensor,
|
||||
dt_bias: torch.Tensor,
|
||||
lower_bound: float | None,
|
||||
state: torch.Tensor,
|
||||
state_indices: torch.Tensor,
|
||||
output_gate: torch.Tensor,
|
||||
norm_weight: torch.Tensor,
|
||||
norm_eps: float,
|
||||
out: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Run the explicit gfx950 Kimi-K3 f_b plus KDA decode specialization.
|
||||
|
||||
The kernel consumes ``f_a`` and the head-local ``f_b_weight`` directly,
|
||||
accumulates the projection in FP32, and rounds once to BF16 before the KDA
|
||||
lower-bound decay gate. It does not materialize the projected raw-g tensor
|
||||
in global memory.
|
||||
"""
|
||||
if f_a.ndim != 2:
|
||||
raise ValueError(f"`f_a` must have rank 2, got rank {f_a.ndim}.")
|
||||
if not f_a.is_cuda:
|
||||
raise ValueError("`f_a` must be a CUDA tensor.")
|
||||
device = f_a.device
|
||||
batch = f_a.shape[0]
|
||||
_check_same_device((("f_b_weight", f_b_weight),), device)
|
||||
_check_tensor(
|
||||
"f_a",
|
||||
f_a,
|
||||
shape=(batch, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
inner_strides=(1,),
|
||||
)
|
||||
_check_tensor(
|
||||
"f_b_weight",
|
||||
f_b_weight,
|
||||
shape=(_HEADS, _DIM, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=device,
|
||||
inner_strides=(_DIM, 1),
|
||||
)
|
||||
out = _validate_kda_inputs(
|
||||
api_name="flydsl_kimi_k3_kda_decode_with_f_b",
|
||||
batch_source="f_a",
|
||||
device=device,
|
||||
batch=batch,
|
||||
x=x,
|
||||
conv_weight=conv_weight,
|
||||
conv_bias=conv_bias,
|
||||
conv_state=conv_state,
|
||||
raw_beta=raw_beta,
|
||||
A_log=A_log,
|
||||
dt_bias=dt_bias,
|
||||
lower_bound=lower_bound,
|
||||
state=state,
|
||||
state_indices=state_indices,
|
||||
output_gate=output_gate,
|
||||
norm_weight=norm_weight,
|
||||
out=out,
|
||||
)
|
||||
|
||||
executable = create_kimi_k3_kda_decode_fb_kernel(
|
||||
float(norm_eps),
|
||||
float(lower_bound),
|
||||
**_fb_build_options(batch),
|
||||
)
|
||||
with torch.cuda.device(device):
|
||||
stream = torch.cuda.current_stream(device)
|
||||
_run_compiled(
|
||||
executable,
|
||||
f_a,
|
||||
f_b_weight,
|
||||
x,
|
||||
conv_weight,
|
||||
conv_state,
|
||||
raw_beta,
|
||||
A_log,
|
||||
dt_bias,
|
||||
state,
|
||||
state_indices,
|
||||
output_gate,
|
||||
norm_weight,
|
||||
out,
|
||||
batch,
|
||||
f_a.stride(0),
|
||||
f_b_weight.stride(0),
|
||||
f_b_weight.stride(1),
|
||||
x.stride(0),
|
||||
conv_weight.stride(0),
|
||||
conv_weight.stride(1),
|
||||
conv_state.stride(0),
|
||||
conv_state.stride(1),
|
||||
conv_state.stride(2),
|
||||
raw_beta.stride(1),
|
||||
state.stride(0),
|
||||
output_gate.stride(0),
|
||||
output_gate.stride(1),
|
||||
out.stride(1),
|
||||
out.stride(2),
|
||||
stream,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
__all__ = [
|
||||
"flydsl_kimi_k3_kda_decode",
|
||||
"flydsl_kimi_k3_kda_decode_with_f_b",
|
||||
"is_flydsl_kimi_k3_kda_decode_supported",
|
||||
]
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Select SGLang-vendored or upstream AITER Kimi-K3 FlyDSL operators."""
|
||||
|
||||
import importlib
|
||||
import os
|
||||
|
||||
|
||||
def load_module(local_module: str, aiter_module: str):
|
||||
mode = os.environ.get("SGLANG_K3_FLYDSL_SOURCE", "auto").lower()
|
||||
if mode not in ("auto", "sglang", "aiter"):
|
||||
raise ValueError(
|
||||
"SGLANG_K3_FLYDSL_SOURCE must be one of auto, sglang, or aiter"
|
||||
)
|
||||
|
||||
candidates = (
|
||||
((local_module, "sglang"), (aiter_module, "aiter"))
|
||||
if mode in ("auto", "sglang")
|
||||
else ((aiter_module, "aiter"),)
|
||||
)
|
||||
errors = []
|
||||
for module_name, source in candidates:
|
||||
if mode == "sglang" and source != "sglang":
|
||||
continue
|
||||
try:
|
||||
return importlib.import_module(module_name)
|
||||
except (ImportError, ModuleNotFoundError) as error:
|
||||
errors.append(f"{source}: {error}")
|
||||
raise ImportError("Kimi-K3 FlyDSL source unavailable: " + "; ".join(errors))
|
||||
@@ -3,7 +3,7 @@ from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention import kda_fused_decode
|
||||
from sglang.kernels.ops.attention import kda_fused_decode, kda_fused_decode_aiter_hip
|
||||
from sglang.kernels.ops.mamba.causal_conv1d_triton import (
|
||||
causal_conv1d_fn,
|
||||
causal_conv1d_update,
|
||||
@@ -20,7 +20,7 @@ from sglang.srt.layers.attention.linear.utils import (
|
||||
)
|
||||
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
|
||||
from sglang.srt.utils import is_cpu, is_cuda, is_npu
|
||||
from sglang.srt.utils.common import rank0_log
|
||||
from sglang.srt.utils.common import is_gfx95_supported, rank0_log
|
||||
|
||||
# KDA always uses the triton causal_conv1d_fn (no CUDA override).
|
||||
# Only causal_conv1d_update needs platform-specific overrides for decode.
|
||||
@@ -555,6 +555,92 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
||||
replayssm_k = layer_cache.replayssm_k
|
||||
replayssm_g = layer_cache.replayssm_g
|
||||
|
||||
deferred_f_b = bool(getattr(layer, "_k3_deferred_f_b", False))
|
||||
if replayssm_d is None and deferred_f_b and is_gfx95_supported():
|
||||
fused_static = getattr(layer, "_k3_hip_fused_decode_args", None)
|
||||
fused_backend = getattr(layer, "_k3_hip_fused_decode_backend", "")
|
||||
onorm_gate = getattr(layer, "_k3_onorm_gate", None)
|
||||
if fused_static is not None and onorm_gate is not None:
|
||||
f_b_weight, norm_weight, norm_eps, a_log = fused_static
|
||||
conv_state_view = conv_states.transpose(-1, -2)
|
||||
output_gate = onorm_gate.view(
|
||||
onorm_gate.shape[0], layer.num_v_heads, layer.head_v_dim
|
||||
)
|
||||
out = mixed_qkv.new_empty(
|
||||
(1, mixed_qkv.shape[0], layer.num_v_heads, layer.head_v_dim)
|
||||
)
|
||||
if fused_backend == "aiter" and kda_fused_decode_aiter_hip.covered(
|
||||
a,
|
||||
f_b_weight,
|
||||
mixed_qkv,
|
||||
b,
|
||||
conv_state_view,
|
||||
ssm_states,
|
||||
cache_indices,
|
||||
output_gate,
|
||||
norm_weight,
|
||||
):
|
||||
core_attn_out = kda_fused_decode_aiter_hip.run(
|
||||
f_a=a,
|
||||
f_b_weight=f_b_weight,
|
||||
mixed_qkv=mixed_qkv,
|
||||
conv_weight=layer.conv_weights,
|
||||
conv_state=conv_state_view,
|
||||
raw_beta=b,
|
||||
A_log=a_log,
|
||||
dt_bias=layer.dt_bias,
|
||||
lower_bound=float(layer.lower_bound),
|
||||
state=ssm_states,
|
||||
state_indices=cache_indices,
|
||||
output_gate=output_gate,
|
||||
norm_weight=norm_weight,
|
||||
norm_eps=norm_eps,
|
||||
out=out,
|
||||
)
|
||||
else:
|
||||
core_attn_out = None
|
||||
if not getattr(KDAAttnBackend, "_hip_fused_reject_logged", False):
|
||||
KDAAttnBackend._hip_fused_reject_logged = True
|
||||
rank0_log(
|
||||
"K3 HIP fused KDA rejected: "
|
||||
f"backend={fused_backend}, "
|
||||
f"f_a={tuple(a.shape)}/{a.dtype}/{a.stride()}, "
|
||||
f"mixed={tuple(mixed_qkv.shape)}/{mixed_qkv.dtype}/"
|
||||
f"{mixed_qkv.stride()}, beta={tuple(b.shape)}/{b.dtype}/"
|
||||
f"{b.stride()}, conv={tuple(conv_state_view.shape)}/"
|
||||
f"{conv_state_view.dtype}/{conv_state_view.stride()}, "
|
||||
f"state={tuple(ssm_states.shape)}/{ssm_states.dtype}/"
|
||||
f"{ssm_states.stride()}, indices={cache_indices.dtype}/"
|
||||
f"{cache_indices.stride()}, gate={tuple(output_gate.shape)}/"
|
||||
f"{output_gate.dtype}/{output_gate.stride()}, "
|
||||
f"norm={norm_weight.dtype}/{norm_weight.stride()}"
|
||||
)
|
||||
|
||||
if core_attn_out is not None:
|
||||
layer._k3_onorm_consumed = True
|
||||
self._track_mamba_state_decode(
|
||||
forward_batch,
|
||||
conv_states,
|
||||
ssm_states,
|
||||
cache_indices,
|
||||
layer.layer_id,
|
||||
)
|
||||
return core_attn_out
|
||||
|
||||
# The model deferred f_b only after publishing static fallback
|
||||
# weights. Materialize the original gate before entering the
|
||||
# unchanged conv + packed-KDA fallback chain.
|
||||
from sglang.kernels.ops.kimi_k3 import kimi_k3_tiny_gemm
|
||||
|
||||
if fused_static is None:
|
||||
raise RuntimeError("K3 deferred f_b is missing fallback weights")
|
||||
a = kimi_k3_tiny_gemm(
|
||||
a,
|
||||
fused_static[0].view(
|
||||
layer.num_v_heads * layer.head_v_dim, layer.head_v_dim
|
||||
),
|
||||
)
|
||||
|
||||
# Fully fused decode step: conv1d update + delta-rule recurrence +
|
||||
# gated RMSNorm in one kernel. Engages only when the model handed off
|
||||
# the output-norm gate for this forward (attempt-and-verify stash,
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
# - Full-rank KDA gate (use_full_rank_gate)
|
||||
|
||||
import logging
|
||||
import os
|
||||
from collections.abc import Iterable
|
||||
from functools import cached_property
|
||||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||
@@ -1695,6 +1696,7 @@ class KimiK3DeltaAttention(nn.Module):
|
||||
self.attn.lower_bound = config.linear_attn_config.get("gate_lower_bound", None)
|
||||
# Set by _prepare_fused_decode() once weights are loaded.
|
||||
self._kda_fused_decode_ready = False
|
||||
self._kda_hip_fused_decode_ready = False
|
||||
|
||||
def forward_qkvbfg(self, hidden_states: torch.Tensor):
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
@@ -1746,7 +1748,51 @@ class KimiK3DeltaAttention(nn.Module):
|
||||
unfused chain. Called once from load_weights (after all weights are
|
||||
loaded, before cuda graph capture)."""
|
||||
if _is_hip:
|
||||
# The fused KDA decode kernel is NVIDIA-only
|
||||
from sglang.kernels.ops.attention import kda_fused_decode_aiter_hip
|
||||
|
||||
layer = self.attn
|
||||
w = layer.conv_weights
|
||||
f_b_weight = self.f_b_proj.weight
|
||||
backend = os.environ.get("SGLANG_K3_KDA_FUSED_BACKEND", "").lower()
|
||||
backend_available = (
|
||||
backend == "aiter"
|
||||
and kda_fused_decode_aiter_hip.available(f_b_weight.device)
|
||||
)
|
||||
if (
|
||||
backend_available
|
||||
and w is not None
|
||||
and tuple(w.shape) == (3 * 12 * 128, 4)
|
||||
and w.dtype == torch.float32
|
||||
and f_b_weight.shape == (12 * 128, 128)
|
||||
and f_b_weight.dtype == torch.bfloat16
|
||||
and layer.A_log is not None
|
||||
and layer.A_log.numel() == 12
|
||||
and layer.A_log.dtype == torch.float32
|
||||
and layer.dt_bias is not None
|
||||
and tuple(layer.dt_bias.shape) == (12 * 128,)
|
||||
and layer.dt_bias.dtype == torch.float32
|
||||
and layer.lower_bound is not None
|
||||
):
|
||||
norm_weight = self.o_norm.weight.data.to(torch.bfloat16).contiguous()
|
||||
f_b_weight = f_b_weight.view(12, 128, 128).contiguous()
|
||||
a_log = layer.A_log.detach().reshape(-1).contiguous()
|
||||
layer._k3_hip_fused_decode_args = (
|
||||
f_b_weight,
|
||||
norm_weight,
|
||||
float(self.o_norm.eps),
|
||||
a_log,
|
||||
)
|
||||
kda_fused_decode_aiter_hip.warmup(
|
||||
f_b_weight=f_b_weight,
|
||||
conv_weight=w,
|
||||
A_log=a_log,
|
||||
dt_bias=layer.dt_bias,
|
||||
lower_bound=float(layer.lower_bound),
|
||||
norm_weight=norm_weight,
|
||||
norm_eps=float(self.o_norm.eps),
|
||||
)
|
||||
layer._k3_hip_fused_decode_backend = backend
|
||||
self._kda_hip_fused_decode_ready = True
|
||||
return
|
||||
layer = self.attn
|
||||
w = layer.conv_weights
|
||||
@@ -1792,7 +1838,9 @@ class KimiK3DeltaAttention(nn.Module):
|
||||
)
|
||||
self._kda_fused_decode_ready = True
|
||||
|
||||
def forward_qkvbfg_fused(self, hidden_states: torch.Tensor):
|
||||
def forward_qkvbfg_fused(
|
||||
self, hidden_states: torch.Tensor, defer_f_b: bool = False
|
||||
):
|
||||
if self.use_full_rank_gate:
|
||||
if self._bfa_w is not None:
|
||||
w = self._bfa_w
|
||||
@@ -1813,7 +1861,11 @@ class KimiK3DeltaAttention(nn.Module):
|
||||
alt.wait_stream(cur)
|
||||
with torch.cuda.stream(alt):
|
||||
bfa = gemm(hidden_states, w)
|
||||
forget_gate = gemm(bfa[..., :n_fa], self._bfa_f_b_w)
|
||||
forget_gate = (
|
||||
bfa[..., :n_fa]
|
||||
if defer_f_b
|
||||
else gemm(bfa[..., :n_fa], self._bfa_f_b_w)
|
||||
)
|
||||
beta = bfa[..., n_fa : n_fa + n_b]
|
||||
fused_states, _ = self.fused_qkvg_proj(hidden_states)
|
||||
qkv, g_proj_states = torch.split(
|
||||
@@ -1825,13 +1877,18 @@ class KimiK3DeltaAttention(nn.Module):
|
||||
fused_states, _ = self.fused_qkvg_proj(hidden_states)
|
||||
qkv, g_proj_states = torch.split(fused_states, self.split_sizes, dim=-1)
|
||||
bfa = gemm(hidden_states, w)
|
||||
forget_gate = gemm(bfa[..., :n_fa], self._bfa_f_b_w)
|
||||
forget_gate = (
|
||||
bfa[..., :n_fa]
|
||||
if defer_f_b
|
||||
else gemm(bfa[..., :n_fa], self._bfa_f_b_w)
|
||||
)
|
||||
beta = bfa[..., n_fa : n_fa + n_b]
|
||||
else:
|
||||
fused_states, _ = self.fused_qkvg_proj(hidden_states)
|
||||
qkv, g_proj_states = torch.split(fused_states, self.split_sizes, dim=-1)
|
||||
beta = self.b_proj(hidden_states)[0]
|
||||
forget_gate = self.f_b_proj(self.f_a_proj(hidden_states)[0])[0]
|
||||
f_a = self.f_a_proj(hidden_states)[0]
|
||||
forget_gate = f_a if defer_f_b else self.f_b_proj(f_a)[0]
|
||||
else:
|
||||
fused_states = self.fused_qkvbfg_a_proj(hidden_states)
|
||||
qkv, beta, fg_a_states = torch.split(fused_states, self.split_sizes, dim=-1)
|
||||
@@ -1847,9 +1904,12 @@ class KimiK3DeltaAttention(nn.Module):
|
||||
forward_batch: ForwardBatch,
|
||||
zero_allocator: BumpAllocator,
|
||||
) -> torch.Tensor:
|
||||
defer_f_b = (
|
||||
self._kda_hip_fused_decode_ready and forward_batch.forward_mode.is_decode()
|
||||
)
|
||||
if self.do_fuse_qkvbfg:
|
||||
mixed_qkv, beta, forget_gate, g_proj_states = self.forward_qkvbfg_fused(
|
||||
hidden_states
|
||||
hidden_states, defer_f_b=defer_f_b
|
||||
)
|
||||
else:
|
||||
mixed_qkv, beta, forget_gate, g_proj_states = self.forward_qkvbfg(
|
||||
@@ -1870,13 +1930,15 @@ class KimiK3DeltaAttention(nn.Module):
|
||||
# into the recurrence kernel. If the backend leaves the stash
|
||||
# unconsumed (env off or shape not covered), apply o_norm here as
|
||||
# before.
|
||||
fused_onorm = self._kda_fused_decode_ready and (
|
||||
fused_onorm = (self._kda_fused_decode_ready or defer_f_b) and (
|
||||
forward_batch.forward_mode.is_decode()
|
||||
or forward_batch.forward_mode.is_target_verify()
|
||||
)
|
||||
if fused_onorm:
|
||||
self.attn._k3_onorm_gate = g_proj_states
|
||||
self.attn._k3_onorm_consumed = False
|
||||
if defer_f_b:
|
||||
self.attn._k3_deferred_f_b = True
|
||||
|
||||
core_attn_out = self.attn(
|
||||
forward_batch,
|
||||
@@ -1888,6 +1950,8 @@ class KimiK3DeltaAttention(nn.Module):
|
||||
if fused_onorm:
|
||||
self.attn._k3_onorm_gate = None
|
||||
fused_onorm = self.attn._k3_onorm_consumed
|
||||
if defer_f_b:
|
||||
self.attn._k3_deferred_f_b = False
|
||||
if not fused_onorm:
|
||||
norm_gate = g_proj_states.unflatten(-1, (-1, self.head_dim))
|
||||
core_attn_out = self.o_norm(core_attn_out, norm_gate)
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
"""Allocation-free microbenchmark for the Kimi-K3 fused f_b + KDA decode."""
|
||||
|
||||
import argparse
|
||||
import runpy
|
||||
import statistics
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_amd_ci
|
||||
from sglang.utils import is_in_ci
|
||||
|
||||
register_amd_ci(est_time=120, stage="jit-kernel-benchmark", runner_config="amd")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--batch", type=int, default=2)
|
||||
parser.add_argument("--warmup", type=int, default=100)
|
||||
parser.add_argument("--iters", type=int, default=1000)
|
||||
parser.add_argument("--trials", type=int, default=21)
|
||||
parser.add_argument("--mode", choices=("eager", "graph"), default="graph")
|
||||
args = parser.parse_args()
|
||||
|
||||
test = runpy.run_path(str(Path(__file__).with_name("test_kimi_k3_kda_decode.py")))
|
||||
f_a, f_b_weight, inputs = test["_make_fb_inputs"](args.batch)
|
||||
out = torch.empty((1, args.batch, 12, 128), dtype=torch.bfloat16, device="cuda")
|
||||
kwargs = dict(
|
||||
f_a=f_a,
|
||||
f_b_weight=f_b_weight,
|
||||
x=inputs.x,
|
||||
conv_weight=inputs.conv_weight,
|
||||
conv_bias=None,
|
||||
conv_state=inputs.conv_state,
|
||||
raw_beta=inputs.raw_beta,
|
||||
A_log=inputs.A_log,
|
||||
dt_bias=inputs.dt_bias,
|
||||
lower_bound=test["_LOWER_BOUND"],
|
||||
state=inputs.state,
|
||||
state_indices=inputs.state_indices,
|
||||
output_gate=inputs.output_gate,
|
||||
norm_weight=inputs.norm_weight,
|
||||
norm_eps=test["_NORM_EPS"],
|
||||
out=out,
|
||||
)
|
||||
fn = test["flydsl_kimi_k3_kda_decode_with_f_b"]
|
||||
for _ in range(args.warmup):
|
||||
fn(**kwargs)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
if args.mode == "graph":
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
fn(**kwargs)
|
||||
replay = graph.replay
|
||||
else:
|
||||
replay = lambda: fn(**kwargs)
|
||||
|
||||
for _ in range(args.warmup):
|
||||
replay()
|
||||
torch.cuda.synchronize()
|
||||
samples = []
|
||||
for _ in range(args.trials):
|
||||
start = torch.cuda.Event(enable_timing=True)
|
||||
end = torch.cuda.Event(enable_timing=True)
|
||||
start.record()
|
||||
for _ in range(args.iters):
|
||||
replay()
|
||||
end.record()
|
||||
end.synchronize()
|
||||
samples.append(start.elapsed_time(end) * 1000 / args.iters)
|
||||
|
||||
print(
|
||||
f"batch={args.batch} mode={args.mode} "
|
||||
f"p50_us={statistics.median(samples):.4f} "
|
||||
f"mean_us={statistics.mean(samples):.4f} "
|
||||
f"p10_us={sorted(samples)[max(0, args.trials // 10 - 1)]:.4f} "
|
||||
f"p90_us={sorted(samples)[min(args.trials - 1, 9 * args.trials // 10)]:.4f}"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if is_in_ci():
|
||||
print("Skipping bench_kimi_k3_kda_decode.py in CI")
|
||||
raise SystemExit(0)
|
||||
main()
|
||||
@@ -0,0 +1,531 @@
|
||||
# SPDX-License-Identifier: MIT
|
||||
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.
|
||||
|
||||
"""Correctness tests for the fused FlyDSL Kimi-K3 KDA decode path."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
pytest.importorskip("flydsl")
|
||||
from aiter.jit.utils.chip_info import get_gfx
|
||||
from aiter.ops.flydsl.utils import is_flydsl_available
|
||||
|
||||
from sglang.kernels.ops.kimi_k3.flydsl.kimi_k3_kda_decode import (
|
||||
_fb_build_options,
|
||||
flydsl_kimi_k3_kda_decode,
|
||||
flydsl_kimi_k3_kda_decode_with_f_b,
|
||||
is_flydsl_kimi_k3_kda_decode_supported,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_amd_ci
|
||||
|
||||
register_amd_ci(est_time=120, suite="stage-b-test-1-gpu-small-amd-mi35x")
|
||||
|
||||
|
||||
def _gfx950_flydsl_available() -> bool:
|
||||
if not torch.cuda.is_available() or not is_flydsl_available():
|
||||
return False
|
||||
try:
|
||||
return get_gfx() == "gfx950"
|
||||
except (AssertionError, KeyError, RuntimeError):
|
||||
return False
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not _gfx950_flydsl_available(),
|
||||
reason="gfx950 FlyDSL required",
|
||||
)
|
||||
|
||||
_DEVICE = torch.device("cuda")
|
||||
_HEADS = 12
|
||||
_DIM = 128
|
||||
_CHANNELS = 3 * _HEADS * _DIM
|
||||
_CONV_WIDTH = 4
|
||||
_LOWER_BOUND = -5.0
|
||||
_NORM_EPS = 1e-5
|
||||
|
||||
|
||||
@dataclass
|
||||
class Inputs:
|
||||
x: torch.Tensor
|
||||
conv_weight: torch.Tensor
|
||||
conv_state: torch.Tensor
|
||||
raw_g: torch.Tensor
|
||||
raw_beta: torch.Tensor
|
||||
A_log: torch.Tensor
|
||||
dt_bias: torch.Tensor
|
||||
state: torch.Tensor
|
||||
state_indices: torch.Tensor
|
||||
output_gate: torch.Tensor
|
||||
norm_weight: torch.Tensor
|
||||
|
||||
|
||||
def _make_inputs(batch: int, seed: int = 20260728) -> Inputs:
|
||||
generator = torch.Generator(device=_DEVICE).manual_seed(seed + batch)
|
||||
slots = batch + 2
|
||||
|
||||
x_storage = torch.randn(
|
||||
(batch, _CHANNELS + 17),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
)
|
||||
x = x_storage[:, :_CHANNELS]
|
||||
conv_weight = 0.1 * torch.randn(
|
||||
(_CHANNELS, _CONV_WIDTH),
|
||||
dtype=torch.float32,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
)
|
||||
|
||||
# Kimi's hybrid cache can pad the slot stride. Exercise that layout
|
||||
# explicitly while keeping each slot's inner dimensions contiguous.
|
||||
conv_storage = torch.randn(
|
||||
(slots, _CHANNELS * (_CONV_WIDTH - 1) + 19),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
)
|
||||
conv_state = conv_storage[:, : _CHANNELS * (_CONV_WIDTH - 1)].view(
|
||||
slots, _CHANNELS, _CONV_WIDTH - 1
|
||||
)
|
||||
state_storage = 0.01 * torch.randn(
|
||||
(slots, _HEADS * _DIM * _DIM + 23),
|
||||
dtype=torch.float32,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
)
|
||||
state = state_storage[:, : _HEADS * _DIM * _DIM].view(slots, _HEADS, _DIM, _DIM)
|
||||
|
||||
raw_beta_storage = torch.randn(
|
||||
(1, batch, _HEADS + 1),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
)
|
||||
output_gate_storage = torch.randn(
|
||||
(batch, _HEADS * _DIM + 7),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
)
|
||||
return Inputs(
|
||||
x=x,
|
||||
conv_weight=conv_weight,
|
||||
conv_state=conv_state,
|
||||
raw_g=torch.randn(
|
||||
(1, batch, _HEADS, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
),
|
||||
raw_beta=raw_beta_storage[:, :, :_HEADS],
|
||||
A_log=0.5
|
||||
* torch.randn(
|
||||
(_HEADS,),
|
||||
dtype=torch.float32,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
),
|
||||
dt_bias=0.1
|
||||
* torch.randn(
|
||||
(_HEADS * _DIM,),
|
||||
dtype=torch.float32,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
),
|
||||
state=state,
|
||||
state_indices=torch.arange(
|
||||
1,
|
||||
batch + 1,
|
||||
dtype=torch.int32,
|
||||
device=_DEVICE,
|
||||
),
|
||||
output_gate=output_gate_storage[:, : _HEADS * _DIM].view(batch, _HEADS, _DIM),
|
||||
norm_weight=torch.randn(
|
||||
(_DIM,),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _copy_inputs(inputs: Inputs) -> Inputs:
|
||||
def clone_preserving_strides(tensor: torch.Tensor) -> torch.Tensor:
|
||||
clone = torch.empty_strided(
|
||||
tensor.shape,
|
||||
tensor.stride(),
|
||||
dtype=tensor.dtype,
|
||||
device=tensor.device,
|
||||
)
|
||||
clone.copy_(tensor)
|
||||
return clone
|
||||
|
||||
return Inputs(
|
||||
**{
|
||||
name: clone_preserving_strides(getattr(inputs, name))
|
||||
for name in Inputs.__dataclass_fields__
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _reference(inputs: Inputs) -> torch.Tensor:
|
||||
batch = inputs.x.shape[0]
|
||||
output = torch.zeros(
|
||||
(1, batch, _HEADS, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
)
|
||||
dt_bias = inputs.dt_bias.view(_HEADS, _DIM)
|
||||
|
||||
for batch_idx in range(batch):
|
||||
state_idx = int(inputs.state_indices[batch_idx])
|
||||
if state_idx <= 0:
|
||||
continue
|
||||
|
||||
history = inputs.conv_state[state_idx]
|
||||
conv_values = torch.cat(
|
||||
(
|
||||
history.float(),
|
||||
inputs.x[batch_idx, :, None].float(),
|
||||
),
|
||||
dim=-1,
|
||||
)
|
||||
packed_qkv = F.silu((conv_values * inputs.conv_weight).sum(dim=-1)).to(
|
||||
torch.bfloat16
|
||||
)
|
||||
history[:, 0].copy_(history[:, 1])
|
||||
history[:, 1].copy_(history[:, 2])
|
||||
history[:, 2].copy_(inputs.x[batch_idx])
|
||||
|
||||
q, k, v = packed_qkv.view(
|
||||
3,
|
||||
_HEADS,
|
||||
_DIM,
|
||||
).unbind(0)
|
||||
for head_idx in range(_HEADS):
|
||||
q_head = q[head_idx].float()
|
||||
k_head = k[head_idx].float()
|
||||
q_head = q_head * torch.rsqrt(q_head.square().sum() + 1e-6)
|
||||
q_head = q_head * (_DIM**-0.5)
|
||||
k_head = k_head * torch.rsqrt(k_head.square().sum() + 1e-6)
|
||||
|
||||
a = inputs.A_log[head_idx].exp()
|
||||
decay = (
|
||||
_LOWER_BOUND
|
||||
* torch.sigmoid(
|
||||
(
|
||||
inputs.raw_g[
|
||||
0,
|
||||
batch_idx,
|
||||
head_idx,
|
||||
].float()
|
||||
+ dt_bias[head_idx]
|
||||
)
|
||||
* a
|
||||
)
|
||||
).exp()
|
||||
decayed_state = inputs.state[state_idx, head_idx] * decay[None, :]
|
||||
state_dot_k = decayed_state @ k_head
|
||||
state_dot_q = decayed_state @ q_head
|
||||
k_dot_q = torch.dot(k_head, q_head)
|
||||
beta = torch.sigmoid(
|
||||
inputs.raw_beta[
|
||||
0,
|
||||
batch_idx,
|
||||
head_idx,
|
||||
].float()
|
||||
)
|
||||
v_new = (v[head_idx].float() - state_dot_k) * beta
|
||||
inputs.state[state_idx, head_idx].copy_(
|
||||
decayed_state + v_new[:, None] * k_head[None, :]
|
||||
)
|
||||
|
||||
# The model materializes recurrent output in BF16 before the
|
||||
# normalization/gating operation.
|
||||
recurrent = (state_dot_q + v_new * k_dot_q).to(torch.bfloat16)
|
||||
recurrent_f32 = recurrent.float()
|
||||
inv_rms = torch.rsqrt(recurrent_f32.square().mean() + _NORM_EPS)
|
||||
output[0, batch_idx, head_idx] = (
|
||||
recurrent_f32
|
||||
* inv_rms
|
||||
* inputs.norm_weight.float()
|
||||
* torch.sigmoid(
|
||||
inputs.output_gate[
|
||||
batch_idx,
|
||||
head_idx,
|
||||
].float()
|
||||
)
|
||||
).to(torch.bfloat16)
|
||||
return output
|
||||
|
||||
|
||||
def _relative_rmse(
|
||||
reference: torch.Tensor,
|
||||
actual: torch.Tensor,
|
||||
) -> float:
|
||||
delta = actual.float() - reference.float()
|
||||
return float(
|
||||
delta.square().mean().sqrt() / (reference.float().square().mean().sqrt() + 1e-8)
|
||||
)
|
||||
|
||||
|
||||
def _run(inputs: Inputs) -> torch.Tensor:
|
||||
return flydsl_kimi_k3_kda_decode(
|
||||
x=inputs.x,
|
||||
conv_weight=inputs.conv_weight,
|
||||
conv_bias=None,
|
||||
conv_state=inputs.conv_state,
|
||||
raw_g=inputs.raw_g,
|
||||
raw_beta=inputs.raw_beta,
|
||||
A_log=inputs.A_log,
|
||||
dt_bias=inputs.dt_bias,
|
||||
lower_bound=_LOWER_BOUND,
|
||||
state=inputs.state,
|
||||
state_indices=inputs.state_indices,
|
||||
output_gate=inputs.output_gate,
|
||||
norm_weight=inputs.norm_weight,
|
||||
norm_eps=_NORM_EPS,
|
||||
)
|
||||
|
||||
|
||||
def _make_fb_inputs(
|
||||
batch: int,
|
||||
seed: int = 20260728,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, Inputs]:
|
||||
generator = torch.Generator(device=_DEVICE).manual_seed(seed + 10_000 + batch)
|
||||
f_a_storage = torch.randn(
|
||||
(batch, _DIM + 5),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
)
|
||||
f_a = f_a_storage[:, :_DIM]
|
||||
f_b_weight = (
|
||||
0.05
|
||||
* torch.randn(
|
||||
(_HEADS, _DIM, _DIM),
|
||||
dtype=torch.bfloat16,
|
||||
device=_DEVICE,
|
||||
generator=generator,
|
||||
)
|
||||
).to(torch.bfloat16)
|
||||
inputs = _make_inputs(batch, seed)
|
||||
projected = F.linear(
|
||||
f_a.float(),
|
||||
f_b_weight.view(_HEADS * _DIM, _DIM).float(),
|
||||
).to(torch.bfloat16)
|
||||
inputs.raw_g.copy_(projected.view(1, batch, _HEADS, _DIM))
|
||||
return f_a, f_b_weight, inputs
|
||||
|
||||
|
||||
def _run_with_f_b(
|
||||
f_a: torch.Tensor,
|
||||
f_b_weight: torch.Tensor,
|
||||
inputs: Inputs,
|
||||
) -> torch.Tensor:
|
||||
return flydsl_kimi_k3_kda_decode_with_f_b(
|
||||
f_a=f_a,
|
||||
f_b_weight=f_b_weight,
|
||||
x=inputs.x,
|
||||
conv_weight=inputs.conv_weight,
|
||||
conv_bias=None,
|
||||
conv_state=inputs.conv_state,
|
||||
raw_beta=inputs.raw_beta,
|
||||
A_log=inputs.A_log,
|
||||
dt_bias=inputs.dt_bias,
|
||||
lower_bound=_LOWER_BOUND,
|
||||
state=inputs.state,
|
||||
state_indices=inputs.state_indices,
|
||||
output_gate=inputs.output_gate,
|
||||
norm_weight=inputs.norm_weight,
|
||||
norm_eps=_NORM_EPS,
|
||||
)
|
||||
|
||||
|
||||
def test_public_api_and_support_predicate() -> None:
|
||||
import sglang.kernels.ops.kimi_k3.flydsl as flydsl_ops
|
||||
|
||||
assert flydsl_ops.flydsl_kimi_k3_kda_decode is flydsl_kimi_k3_kda_decode
|
||||
assert (
|
||||
flydsl_ops.is_flydsl_kimi_k3_kda_decode_supported
|
||||
is is_flydsl_kimi_k3_kda_decode_supported
|
||||
)
|
||||
assert is_flydsl_kimi_k3_kda_decode_supported(0)
|
||||
assert not is_flydsl_kimi_k3_kda_decode_supported("cpu")
|
||||
|
||||
|
||||
def test_f_b_public_api() -> None:
|
||||
import sglang.kernels.ops.kimi_k3.flydsl as flydsl_ops
|
||||
|
||||
assert (
|
||||
flydsl_ops.flydsl_kimi_k3_kda_decode_with_f_b
|
||||
is flydsl_kimi_k3_kda_decode_with_f_b
|
||||
)
|
||||
|
||||
|
||||
def test_f_b_batch_two_dispatch_is_guarded() -> None:
|
||||
assert _fb_build_options(1) == {}
|
||||
assert _fb_build_options(4) == {}
|
||||
options = _fb_build_options(2)
|
||||
assert options == {
|
||||
"waves_per_eu": 3,
|
||||
"cooperative_f_a": True,
|
||||
"parallel_front": True,
|
||||
"fused_norm_reduce": True,
|
||||
"projection_fdot2": True,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch", [1, 8, 16])
|
||||
def test_kimi_k3_kda_decode_matches_reference(batch: int) -> None:
|
||||
seed = _make_inputs(batch)
|
||||
reference_inputs = _copy_inputs(seed)
|
||||
actual_inputs = _copy_inputs(seed)
|
||||
|
||||
reference = _reference(reference_inputs)
|
||||
actual = _run(actual_inputs)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
assert is_flydsl_kimi_k3_kda_decode_supported(_DEVICE)
|
||||
assert not torch.isnan(actual).any()
|
||||
assert _relative_rmse(reference, actual) < 1e-3
|
||||
assert (
|
||||
_relative_rmse(
|
||||
reference_inputs.state,
|
||||
actual_inputs.state,
|
||||
)
|
||||
< 1e-3
|
||||
)
|
||||
assert torch.equal(
|
||||
reference_inputs.conv_state,
|
||||
actual_inputs.conv_state,
|
||||
)
|
||||
|
||||
|
||||
def test_non_positive_slots_do_not_modify_caches() -> None:
|
||||
inputs = _make_inputs(batch=2)
|
||||
inputs.state_indices.copy_(torch.tensor([0, -1], dtype=torch.int32, device=_DEVICE))
|
||||
conv_before = inputs.conv_state.clone()
|
||||
state_before = inputs.state.clone()
|
||||
|
||||
actual = _run(inputs)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
assert torch.count_nonzero(actual) == 0
|
||||
assert torch.equal(inputs.conv_state, conv_before)
|
||||
assert torch.equal(inputs.state, state_before)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch", [1, 8, 16])
|
||||
def test_kimi_k3_kda_decode_with_f_b_matches_reference(batch: int) -> None:
|
||||
f_a, f_b_weight, seed = _make_fb_inputs(batch)
|
||||
reference_inputs = _copy_inputs(seed)
|
||||
actual_inputs = _copy_inputs(seed)
|
||||
|
||||
reference = _reference(reference_inputs)
|
||||
actual = _run_with_f_b(f_a, f_b_weight, actual_inputs)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
assert not torch.isnan(actual).any()
|
||||
assert _relative_rmse(reference, actual) < 1e-3
|
||||
assert _relative_rmse(reference_inputs.state, actual_inputs.state) < 1e-3
|
||||
assert torch.equal(reference_inputs.conv_state, actual_inputs.conv_state)
|
||||
|
||||
|
||||
def test_kimi_k3_kda_decode_with_f_b_recurrent_sequence() -> None:
|
||||
f_a, f_b_weight, seed = _make_fb_inputs(batch=2)
|
||||
reference_inputs = _copy_inputs(seed)
|
||||
actual_inputs = _copy_inputs(seed)
|
||||
for _ in range(32):
|
||||
reference = _reference(reference_inputs)
|
||||
actual = _run_with_f_b(f_a, f_b_weight, actual_inputs)
|
||||
torch.cuda.synchronize()
|
||||
assert _relative_rmse(reference, actual) < 1e-3
|
||||
assert _relative_rmse(reference_inputs.state, actual_inputs.state) < 1e-3
|
||||
assert torch.equal(reference_inputs.conv_state, actual_inputs.conv_state)
|
||||
|
||||
|
||||
def test_kimi_k3_kda_decode_with_f_b_graph_replay() -> None:
|
||||
f_a, f_b_weight, inputs = _make_fb_inputs(batch=2)
|
||||
out = torch.empty((1, 2, _HEADS, _DIM), dtype=torch.bfloat16, device=_DEVICE)
|
||||
_run_with_f_b(f_a, f_b_weight, inputs)
|
||||
torch.cuda.synchronize()
|
||||
graph = torch.cuda.CUDAGraph()
|
||||
with torch.cuda.graph(graph):
|
||||
flydsl_kimi_k3_kda_decode_with_f_b(
|
||||
f_a=f_a,
|
||||
f_b_weight=f_b_weight,
|
||||
x=inputs.x,
|
||||
conv_weight=inputs.conv_weight,
|
||||
conv_bias=None,
|
||||
conv_state=inputs.conv_state,
|
||||
raw_beta=inputs.raw_beta,
|
||||
A_log=inputs.A_log,
|
||||
dt_bias=inputs.dt_bias,
|
||||
lower_bound=_LOWER_BOUND,
|
||||
state=inputs.state,
|
||||
state_indices=inputs.state_indices,
|
||||
output_gate=inputs.output_gate,
|
||||
norm_weight=inputs.norm_weight,
|
||||
norm_eps=_NORM_EPS,
|
||||
out=out,
|
||||
)
|
||||
graph.replay()
|
||||
torch.cuda.synchronize()
|
||||
assert not torch.isnan(out).any()
|
||||
|
||||
|
||||
def test_f_b_non_positive_slots_do_not_modify_caches() -> None:
|
||||
f_a, f_b_weight, inputs = _make_fb_inputs(batch=2)
|
||||
inputs.state_indices.copy_(torch.tensor([0, -1], dtype=torch.int32, device=_DEVICE))
|
||||
conv_before = inputs.conv_state.clone()
|
||||
state_before = inputs.state.clone()
|
||||
|
||||
actual = _run_with_f_b(f_a, f_b_weight, inputs)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
assert torch.count_nonzero(actual) == 0
|
||||
assert torch.equal(inputs.conv_state, conv_before)
|
||||
assert torch.equal(inputs.state, state_before)
|
||||
|
||||
|
||||
def test_f_b_api_rejects_invalid_projection_inputs() -> None:
|
||||
f_a, f_b_weight, inputs = _make_fb_inputs(batch=1)
|
||||
|
||||
with pytest.raises(ValueError, match="`f_a` must have rank 2"):
|
||||
_run_with_f_b(f_a.unsqueeze(0), f_b_weight, inputs)
|
||||
with pytest.raises(ValueError, match="`f_a` must have dtype"):
|
||||
_run_with_f_b(f_a.float(), f_b_weight, inputs)
|
||||
with pytest.raises(ValueError, match="`f_b_weight` must have shape"):
|
||||
_run_with_f_b(f_a, f_b_weight[:, :, :-1], inputs)
|
||||
with pytest.raises(ValueError, match="`f_b_weight` must have inner strides"):
|
||||
_run_with_f_b(f_a, f_b_weight.transpose(1, 2), inputs)
|
||||
|
||||
|
||||
def test_decode_api_rejects_invalid_input_rank() -> None:
|
||||
inputs = _make_inputs(batch=1)
|
||||
inputs.x = inputs.x.unsqueeze(0)
|
||||
|
||||
with pytest.raises(ValueError, match="`x` must have rank 2"):
|
||||
_run(inputs)
|
||||
|
||||
|
||||
def test_fused_kda_backend_is_opt_in(monkeypatch):
|
||||
from sglang.kernels.ops.attention import kda_fused_decode_aiter_hip
|
||||
|
||||
monkeypatch.delenv("SGLANG_K3_KDA_FUSED_BACKEND", raising=False)
|
||||
assert not kda_fused_decode_aiter_hip.enabled()
|
||||
monkeypatch.setenv("SGLANG_K3_KDA_FUSED_BACKEND", "aiter")
|
||||
assert kda_fused_decode_aiter_hip.enabled()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import sys
|
||||
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
Reference in New Issue
Block a user