chore: import upstream snapshot with attribution
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
This commit is contained in:
@@ -0,0 +1,280 @@
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.layers.attention.fla.index import (
|
||||
prepare_chunk_indices,
|
||||
prepare_chunk_offsets,
|
||||
)
|
||||
from sglang.srt.layers.attention.fla.op import exp, make_tensor_descriptor, safe_exp
|
||||
from sglang.srt.layers.attention.fla.utils import (
|
||||
autotune_cache_kwargs,
|
||||
)
|
||||
|
||||
CHUNK_SIZE = 64
|
||||
|
||||
|
||||
# This kernel handles K blocks in a for loop to minimize register spills.
|
||||
# Time is the OUTER loop; K blocks are processed in two inner phases per step:
|
||||
# Phase 1: store h to output, accumulate v_correction = sum_k(w_k @ h_k^T)
|
||||
# Phase 2: update h = gate * h + k^T @ v_gated, save to scratch (initial_state)
|
||||
@triton.autotune(
|
||||
configs=[triton.Config({"BV": 64}, num_warps=8, num_stages=2)],
|
||||
key=["H", "K", "V", "BT", "USE_GK", "USE_INITIAL_STATE", "NT_BUCKET"],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop(
|
||||
k,
|
||||
v,
|
||||
w,
|
||||
v_new,
|
||||
g,
|
||||
gk,
|
||||
h,
|
||||
initial_state,
|
||||
initial_state_indices,
|
||||
cu_seqlens,
|
||||
chunk_offsets,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
Hg: tl.constexpr,
|
||||
K: tl.constexpr,
|
||||
V: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BV: tl.constexpr,
|
||||
USE_G: tl.constexpr,
|
||||
USE_GK: tl.constexpr,
|
||||
USE_INITIAL_STATE: tl.constexpr,
|
||||
INPLACE_UPDATE: tl.constexpr,
|
||||
SAVE_NEW_VALUE: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
NT_BUCKET: tl.constexpr, # this arg is kept to align with the triton kernel for CUDA
|
||||
):
|
||||
i_v, i_nh = tl.program_id(0), tl.program_id(1)
|
||||
i_n, i_h = i_nh // H, i_nh % H
|
||||
if IS_VARLEN:
|
||||
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(
|
||||
cu_seqlens + i_n + 1
|
||||
).to(tl.int32)
|
||||
T = eos - bos
|
||||
NT = tl.cdiv(T, BT)
|
||||
boh = tl.load(chunk_offsets + i_n).to(tl.int32)
|
||||
else:
|
||||
bos, eos = i_n * T, i_n * T + T
|
||||
NT = tl.cdiv(T, BT)
|
||||
boh = i_n * NT
|
||||
|
||||
# calculate offset
|
||||
h += ((boh * H + i_h) * V * K).to(tl.int64)
|
||||
v += ((bos * H + i_h) * V).to(tl.int64)
|
||||
k += ((bos * Hg + i_h // (H // Hg)) * K).to(tl.int64)
|
||||
w += ((bos * H + i_h) * K).to(tl.int64)
|
||||
if SAVE_NEW_VALUE:
|
||||
v_new += ((bos * H + i_h) * V).to(tl.int64)
|
||||
stride_v = H * V
|
||||
stride_h = H * V * K
|
||||
stride_k = Hg * K
|
||||
stride_w = H * K
|
||||
|
||||
w_desc = make_tensor_descriptor(
|
||||
base=w,
|
||||
shape=(T, K),
|
||||
strides=(stride_w, 1),
|
||||
block_shape=(BT, 64),
|
||||
)
|
||||
v_desc = make_tensor_descriptor(
|
||||
base=v,
|
||||
shape=(T, V),
|
||||
strides=(stride_v, 1),
|
||||
block_shape=(BT, BV),
|
||||
)
|
||||
k_desc = make_tensor_descriptor(
|
||||
base=k,
|
||||
shape=(T, K),
|
||||
strides=(stride_k, 1),
|
||||
block_shape=(BT, 64),
|
||||
)
|
||||
if SAVE_NEW_VALUE:
|
||||
v_new_desc = make_tensor_descriptor(
|
||||
base=v_new,
|
||||
shape=(T, V),
|
||||
strides=(stride_v, 1),
|
||||
block_shape=(BT, BV),
|
||||
)
|
||||
|
||||
index = tl.load(initial_state_indices + i_n).to(tl.int32)
|
||||
h0 = initial_state + index * stride_h
|
||||
ht = initial_state + index * stride_h
|
||||
if USE_INITIAL_STATE:
|
||||
h0 = h0 + i_h * V * K
|
||||
if INPLACE_UPDATE:
|
||||
ht = ht + i_h * V * K
|
||||
|
||||
# main recurrence — time is the outer loop
|
||||
for i_t in range(NT):
|
||||
########################################################################
|
||||
# Phase 1: store h to output, compute v_new = u - sum_k(w_k @ h_k^T)
|
||||
########################################################################
|
||||
b_v_corr = tl.zeros([BT, BV], dtype=tl.float32)
|
||||
for k_blk in range(0, K, 64):
|
||||
# Load h: from initial_state (i_t==0) or scratch (i_t>0)
|
||||
if i_t == 0:
|
||||
if USE_INITIAL_STATE:
|
||||
p_hs = tl.make_block_ptr(
|
||||
h0, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
|
||||
)
|
||||
b_h = tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
|
||||
else:
|
||||
b_h = tl.zeros([BV, 64], dtype=tl.float32)
|
||||
else:
|
||||
p_hs = tl.make_block_ptr(
|
||||
ht, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
|
||||
)
|
||||
b_h = tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
|
||||
|
||||
# Store pre-update h to output
|
||||
p_ho = tl.make_block_ptr(
|
||||
h + i_t * stride_h,
|
||||
(V, K),
|
||||
(K, 1),
|
||||
(i_v * BV, k_blk),
|
||||
(BV, 64),
|
||||
(1, 0),
|
||||
)
|
||||
tl.store(p_ho, b_h.to(p_ho.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
# Accumulate correction: w_k @ h_k^T
|
||||
b_w = w_desc.load([i_t * BT, k_blk])
|
||||
b_v_corr += tl.dot(b_w, tl.trans(b_h).to(b_w.dtype))
|
||||
|
||||
# v_new = u - correction
|
||||
b_v = v_desc.load([i_t * BT, i_v * BV]) - b_v_corr
|
||||
|
||||
if SAVE_NEW_VALUE:
|
||||
v_new_desc.store([i_t * BT, i_v * BV], b_v.to(v_new.dtype.element_ty))
|
||||
|
||||
# Apply gate to v
|
||||
last_idx = min((i_t + 1) * BT, T) - 1
|
||||
if USE_G:
|
||||
b_g_last = tl.load(g + bos * H + last_idx * H + i_h)
|
||||
p_g = tl.make_block_ptr(
|
||||
g + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,)
|
||||
)
|
||||
b_g = tl.load(p_g, boundary_check=(0,))
|
||||
b_v = b_v * tl.expand_dims(safe_exp(b_g_last - b_g), 1)
|
||||
b_g_last = exp(b_g_last)
|
||||
|
||||
b_v = b_v.to(k.dtype.element_ty)
|
||||
|
||||
########################################################################
|
||||
# Phase 2: reload h, apply gate, update h += k^T @ v, save to scratch
|
||||
########################################################################
|
||||
for k_blk in range(0, K, 64):
|
||||
# Reload h (same source as Phase 1)
|
||||
if i_t == 0:
|
||||
if USE_INITIAL_STATE:
|
||||
p_hs = tl.make_block_ptr(
|
||||
h0, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
|
||||
)
|
||||
b_h = tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
|
||||
else:
|
||||
b_h = tl.zeros([BV, 64], dtype=tl.float32)
|
||||
else:
|
||||
p_hs = tl.make_block_ptr(
|
||||
ht, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
|
||||
)
|
||||
b_h = tl.load(p_hs, boundary_check=(0, 1)).to(tl.float32)
|
||||
|
||||
# Gate decay on h
|
||||
if USE_G:
|
||||
b_h = b_h * b_g_last
|
||||
|
||||
if USE_GK:
|
||||
o_k1 = tl.arange(0, 64) + k_blk
|
||||
b_gk_last1 = tl.load(
|
||||
gk + (bos + last_idx) * H * K + i_h * K + o_k1,
|
||||
mask=(o_k1 < K),
|
||||
other=0.0,
|
||||
)
|
||||
b_h *= tl.expand_dims(exp(b_gk_last1), 0)
|
||||
|
||||
# Delta update: h += k^T @ v
|
||||
b_k = tl.trans(k_desc.load([i_t * BT, k_blk]))
|
||||
b_h += tl.trans(tl.dot(b_k, b_v))
|
||||
|
||||
# Save updated h to scratch (initial_state) for next time step
|
||||
if INPLACE_UPDATE:
|
||||
p_hs = tl.make_block_ptr(
|
||||
ht, (V, K), (K, 1), (i_v * BV, k_blk), (BV, 64), (1, 0)
|
||||
)
|
||||
tl.store(p_hs, b_h.to(p_hs.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
|
||||
def chunk_gated_delta_rule_fwd_h(
|
||||
k: torch.Tensor,
|
||||
w: torch.Tensor,
|
||||
u: torch.Tensor,
|
||||
g: Optional[torch.Tensor] = None,
|
||||
gk: Optional[torch.Tensor] = None,
|
||||
initial_state: Optional[torch.Tensor] = None,
|
||||
initial_state_indices: Optional[torch.Tensor] = None,
|
||||
save_new_value: bool = True,
|
||||
cu_seqlens: Optional[torch.LongTensor] = None,
|
||||
chunk_indices: Optional[torch.LongTensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
B, T, Hg, K, V = *k.shape, u.shape[-1]
|
||||
H = u.shape[-2]
|
||||
BT = CHUNK_SIZE
|
||||
|
||||
if chunk_indices is None and cu_seqlens is not None:
|
||||
chunk_indices = prepare_chunk_indices(cu_seqlens, CHUNK_SIZE)
|
||||
# N: the actual number of sequences in the batch with either equal or variable lengths
|
||||
if cu_seqlens is None:
|
||||
N, NT, chunk_offsets = B, triton.cdiv(T, BT), None
|
||||
else:
|
||||
N, NT, chunk_offsets = (
|
||||
len(cu_seqlens) - 1,
|
||||
len(chunk_indices),
|
||||
prepare_chunk_offsets(cu_seqlens, BT),
|
||||
)
|
||||
assert K <= 256, "current kernel does not support head dimension larger than 256."
|
||||
|
||||
h = k.new_empty(B, NT, H, V, K)
|
||||
|
||||
v_new = torch.empty_like(u) if save_new_value else None
|
||||
|
||||
def grid(meta):
|
||||
return (triton.cdiv(V, meta["BV"]), N * H)
|
||||
|
||||
kernel = chunk_gated_delta_rule_fwd_kernel_h_blockdim64_k_loop
|
||||
|
||||
kernel[grid](
|
||||
k=k,
|
||||
v=u,
|
||||
w=w,
|
||||
v_new=v_new,
|
||||
g=g,
|
||||
gk=gk,
|
||||
h=h,
|
||||
initial_state=initial_state,
|
||||
initial_state_indices=initial_state_indices,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_offsets=chunk_offsets,
|
||||
T=T,
|
||||
H=H,
|
||||
Hg=Hg,
|
||||
K=K,
|
||||
V=V,
|
||||
BT=BT,
|
||||
USE_G=g is not None,
|
||||
USE_GK=gk is not None,
|
||||
USE_INITIAL_STATE=initial_state is not None,
|
||||
INPLACE_UPDATE=True,
|
||||
SAVE_NEW_VALUE=v_new is not None,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
NT_BUCKET=(0 if NT <= 32 else (1 if NT <= 128 else 2)),
|
||||
)
|
||||
return h, v_new
|
||||
@@ -0,0 +1,325 @@
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.layers.attention.fla.index import prepare_chunk_indices
|
||||
from sglang.srt.layers.attention.fla.op import safe_exp
|
||||
from sglang.srt.layers.attention.fla.utils import (
|
||||
autotune_cache_kwargs,
|
||||
)
|
||||
from sglang.srt.layers.attention.fla.wy_fast import recompute_w_u_fwd
|
||||
|
||||
_MERGE_DOT_PRECISION = tl.constexpr("ieee")
|
||||
|
||||
|
||||
@triton.heuristics(
|
||||
{
|
||||
"USE_G": lambda args: args["g"] is not None,
|
||||
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
|
||||
}
|
||||
)
|
||||
@triton.autotune(
|
||||
configs=[
|
||||
triton.Config({"BK": BK}, num_warps=num_warps)
|
||||
for BK in [16, 32, 64]
|
||||
for num_warps in [2, 4, 8, 16, 32]
|
||||
],
|
||||
key=["H", "Hg", "K", "BC", "BK", "USE_G", "IS_VARLEN"],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def chunk_gated_delta_rule_fwd_kkt_solve_kernel_low_reg(
|
||||
k,
|
||||
g,
|
||||
beta,
|
||||
A,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
Hg: tl.constexpr,
|
||||
K: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BC: tl.constexpr,
|
||||
BK: tl.constexpr,
|
||||
USE_G: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
Low-reg version: one [BC,BC] accumulator at a time to minimise register pressure.
|
||||
|
||||
Pass 1: loop over 4 diagonal blocks (tl.static_range unrolls to 4 K-loops).
|
||||
Pass 2: nested loop over off-diagonal distance d=1,2,3 and column j.
|
||||
d=1 (nearest): Ai_{ij} = -Ai_ii @ A_ij_raw @ Ai_jj
|
||||
d>1 (farther): Ai_{ij} = -(Ai_ii @ A_ij_raw
|
||||
+ sum_{m=j+1}^{i-1} Ai_im @ A_mj_raw) @ Ai_jj
|
||||
Each K-loop holds exactly one [BC,BC] accumulator. Raw blocks needed by
|
||||
later correction terms are spilled to upper-triangular scratch slots in A
|
||||
(see _KKT_SCRATCH_COL for the layout; boundary_check makes out-of-bounds
|
||||
stores/loads safe so no runtime `if i_tcX < T` guards are needed).
|
||||
"""
|
||||
i_t, i_bh = tl.program_id(0), tl.program_id(1)
|
||||
i_b, i_h = i_bh // H, i_bh % H
|
||||
|
||||
if IS_VARLEN:
|
||||
i_n, i_t = tl.load(chunk_indices + i_t * 2).to(tl.int32), tl.load(
|
||||
chunk_indices + i_t * 2 + 1
|
||||
).to(tl.int32)
|
||||
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(
|
||||
cu_seqlens + i_n + 1
|
||||
).to(tl.int32)
|
||||
T = eos - bos
|
||||
else:
|
||||
bos, eos = i_b * T, i_b * T + T
|
||||
|
||||
if i_t * BT >= T:
|
||||
return
|
||||
|
||||
i_tc0 = i_t * BT
|
||||
k += (bos * Hg + i_h // (H // Hg)) * K
|
||||
A += (bos * H + i_h) * BT
|
||||
|
||||
o_i = tl.arange(0, BC)
|
||||
m_d = o_i[:, None] > o_i[None, :]
|
||||
m_I = o_i[:, None] == o_i[None, :]
|
||||
|
||||
############################################################################
|
||||
# Pass 1: diagonal blocks — one K-loop per sub-chunk (tl.static_range → 4)
|
||||
############################################################################
|
||||
for i_b in tl.static_range(4):
|
||||
i_tci = i_tc0 + i_b * BC
|
||||
m_tci = (i_tci + o_i) < T
|
||||
p_bi = tl.make_block_ptr(
|
||||
beta + bos * H + i_h, (T,), (H,), (i_tci,), (BC,), (0,)
|
||||
)
|
||||
b_bi = tl.load(p_bi, boundary_check=(0,)).to(tl.float32)
|
||||
if USE_G:
|
||||
p_gi = tl.make_block_ptr(
|
||||
g + bos * H + i_h, (T,), (H,), (i_tci,), (BC,), (0,)
|
||||
)
|
||||
b_gi = tl.load(p_gi, boundary_check=(0,)).to(tl.float32)
|
||||
b_A = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
for i_k in range(tl.cdiv(K, BK)):
|
||||
p_k = tl.make_block_ptr(
|
||||
k, (T, K), (Hg * K, 1), (i_tci, i_k * BK), (BC, BK), (1, 0)
|
||||
)
|
||||
b_k = tl.load(p_k, boundary_check=(0, 1))
|
||||
b_A += tl.dot(b_k, tl.trans(b_k))
|
||||
if USE_G:
|
||||
b_A *= safe_exp(b_gi[:, None] - b_gi[None, :])
|
||||
b_A = (
|
||||
tl.where(m_d & (m_tci[:, None] & m_tci[None, :]), b_A, 0.0) * b_bi[:, None]
|
||||
)
|
||||
# Forward substitution: solve (I + A_diag) x = I column by column.
|
||||
# Extra iterations for out-of-bounds rows are no-ops (b_A rows are zero).
|
||||
b_Ai = -b_A
|
||||
for i in range(2, BC):
|
||||
b_a = tl.sum(tl.where((o_i == i)[:, None], -b_A, 0.0), 0)
|
||||
b_a = tl.where(o_i < i, b_a, 0.0)
|
||||
b_a = b_a + tl.sum(b_a[:, None] * b_Ai, 0)
|
||||
b_Ai = tl.where((o_i == i)[:, None], b_a, b_Ai)
|
||||
b_Ai += m_I
|
||||
p_Aii = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tci, i_b * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
tl.store(p_Aii, b_Ai.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
############################################################################
|
||||
# Pass 2: off-diagonal blocks — one K-loop per (i,j) pair.
|
||||
# Outer loop: d = i-j (diagonal distance) from 1 to 3.
|
||||
# Inner loop: j from 0 to 3-d (so i = j+d).
|
||||
# Processing order ensures Ai_im (m in j+1..i-1) is already in A when needed.
|
||||
# Formula: Ai_ij = -(Ai_ii @ A_ij_raw
|
||||
# + sum_{m=j+1}^{i-1} Ai_im @ A_mj_raw) @ Ai_jj
|
||||
############################################################################
|
||||
for d in tl.static_range(1, 4):
|
||||
for j in tl.static_range(0, 4 - d):
|
||||
i = j + d # compile-time ints from static_range
|
||||
i_tci = i_tc0 + i * BC
|
||||
i_tcj = i_tc0 + j * BC
|
||||
p_bi = tl.make_block_ptr(
|
||||
beta + bos * H + i_h, (T,), (H,), (i_tci,), (BC,), (0,)
|
||||
)
|
||||
b_bi = tl.load(p_bi, boundary_check=(0,)).to(tl.float32)
|
||||
if USE_G:
|
||||
p_gi = tl.make_block_ptr(
|
||||
g + bos * H + i_h, (T,), (H,), (i_tci,), (BC,), (0,)
|
||||
)
|
||||
p_gj = tl.make_block_ptr(
|
||||
g + bos * H + i_h, (T,), (H,), (i_tcj,), (BC,), (0,)
|
||||
)
|
||||
b_gi = tl.load(p_gi, boundary_check=(0,)).to(tl.float32)
|
||||
b_gj = tl.load(p_gj, boundary_check=(0,)).to(tl.float32)
|
||||
# K-loop: accumulate k_i @ k_j^T
|
||||
b_A = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
for i_k in range(tl.cdiv(K, BK)):
|
||||
p_ki = tl.make_block_ptr(
|
||||
k, (T, K), (Hg * K, 1), (i_tci, i_k * BK), (BC, BK), (1, 0)
|
||||
)
|
||||
p_kj = tl.make_block_ptr(
|
||||
k, (T, K), (Hg * K, 1), (i_tcj, i_k * BK), (BC, BK), (1, 0)
|
||||
)
|
||||
b_A += tl.dot(
|
||||
tl.load(p_ki, boundary_check=(0, 1)),
|
||||
tl.trans(tl.load(p_kj, boundary_check=(0, 1))),
|
||||
)
|
||||
if USE_G:
|
||||
b_A *= safe_exp(b_gi[:, None] - b_gj[None, :])
|
||||
b_A *= b_bi[:, None]
|
||||
# Scratch this raw block if a later row in this column needs it as
|
||||
# a correction term. i < 3 means rows i+1..3 exist and will use it.
|
||||
# Scratch column layout (compile-time ternary, evaluated at trace time):
|
||||
# (i=1,j=0)->col1 (i=2,j=1)->col2 (i=2,j=0)->col3
|
||||
if i < 3:
|
||||
sc = 1 if (i == 1 and j == 0) else (2 if (i == 2 and j == 1) else 3)
|
||||
p_s = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tc0, sc * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
tl.store(p_s, b_A.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
# Correction sum: sum_{m=j+1}^{i-1} Ai_im @ A_mj_raw
|
||||
# Unrolled manually (d is a compile-time Python int from tl.static_range):
|
||||
# d=1: no corrections; d=2: m=j+1; d=3: m=j+1 then m=j+2
|
||||
b_corr = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
if d >= 2:
|
||||
m1 = j + 1
|
||||
sc_m1j = (
|
||||
1 if (m1 == 1 and j == 0) else (2 if (m1 == 2 and j == 1) else 3)
|
||||
)
|
||||
p_s_m1j = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tc0, sc_m1j * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
b_A_m1j = tl.load(p_s_m1j, boundary_check=(0, 1)).to(tl.float32)
|
||||
p_Ai_im1 = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tci, m1 * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
b_corr += tl.dot(
|
||||
tl.load(p_Ai_im1, boundary_check=(0, 1)).to(tl.float32),
|
||||
b_A_m1j,
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
if d >= 3:
|
||||
m2 = j + 2
|
||||
sc_m2j = (
|
||||
1 if (m2 == 1 and j == 0) else (2 if (m2 == 2 and j == 1) else 3)
|
||||
)
|
||||
p_s_m2j = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tc0, sc_m2j * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
b_A_m2j = tl.load(p_s_m2j, boundary_check=(0, 1)).to(tl.float32)
|
||||
p_Ai_im2 = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tci, m2 * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
b_corr += tl.dot(
|
||||
tl.load(p_Ai_im2, boundary_check=(0, 1)).to(tl.float32),
|
||||
b_A_m2j,
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
# Compute and store Ai_ij
|
||||
p_Ai_ii = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tci, i * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
p_Ai_jj = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tcj, j * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
b_Ai_ii = tl.load(p_Ai_ii, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_Ai_jj = tl.load(p_Ai_jj, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_Ai_ij = -tl.dot(
|
||||
tl.dot(b_Ai_ii, b_A, input_precision=_MERGE_DOT_PRECISION) + b_corr,
|
||||
b_Ai_jj,
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
p_Ai_ij = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tci, j * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
tl.store(p_Ai_ij, b_Ai_ij.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
# Clean up scratch slots: Pass 2 stored raw A_ij blocks in the upper-triangular
|
||||
# part of row i_tc0 (cols BC..3*BC). These must be zeroed because
|
||||
# recompute_w_u_fwd reads the full BT×BT block.
|
||||
b_zero = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
for sc in tl.static_range(1, BT // BC):
|
||||
p_scratch = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tc0, sc * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
tl.store(p_scratch, b_zero.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
|
||||
def chunk_gated_delta_rule_fwd_intra(
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
g: torch.Tensor | None = None,
|
||||
beta: torch.Tensor | None = None,
|
||||
cu_seqlens: torch.LongTensor | None = None,
|
||||
chunk_size: int = 64,
|
||||
chunk_indices: torch.LongTensor | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
r"""
|
||||
GDN intra-chunk forward: fused kkt + solve_tril + recompute_w_u.
|
||||
|
||||
Equivalent to:
|
||||
A = chunk_scaled_dot_kkt_fwd(k, g, beta, ...) # kernel 1
|
||||
A = solve_tril(A, ...) # kernel 2
|
||||
w, u = recompute_w_u_fwd(k, v, beta, A, g, ...) # kernel 3
|
||||
|
||||
Fuses kernels 1+2 into a single kernel, reducing from 3 to 2 kernel launches
|
||||
and eliminating the HBM round-trip for the intermediate A matrix.
|
||||
|
||||
Args:
|
||||
k (torch.Tensor):
|
||||
The key tensor of shape `[B, T, H, K]`.
|
||||
v (torch.Tensor):
|
||||
The value tensor of shape `[B, T, H, V]`.
|
||||
g (torch.Tensor):
|
||||
The cumulative sum of the gate tensor of shape `[B, T, H]`. Default: `None`.
|
||||
beta (torch.Tensor):
|
||||
The beta tensor of shape `[B, T, H]`.
|
||||
cu_seqlens (torch.LongTensor):
|
||||
The cumulative sequence lengths. Default: `None`.
|
||||
chunk_size (int):
|
||||
The chunk size. Default: 64.
|
||||
chunk_indices (torch.LongTensor):
|
||||
Precomputed chunk indices. Default: `None`.
|
||||
|
||||
Returns:
|
||||
w (torch.Tensor): shape `[B, T, H, K]`
|
||||
u (torch.Tensor): shape `[B, T, H, V]`
|
||||
A (torch.Tensor): shape `[B, T, H, BT]`, the solved (I+A)^{-1} matrix
|
||||
"""
|
||||
B, T, Hg, K = k.shape
|
||||
H = beta.shape[-1]
|
||||
BT = chunk_size
|
||||
BC = 16
|
||||
|
||||
if chunk_indices is None and cu_seqlens is not None:
|
||||
chunk_indices = prepare_chunk_indices(cu_seqlens, BT)
|
||||
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
|
||||
|
||||
# Step 1: fused kkt + solve_tril
|
||||
A = torch.zeros(B, T, H, BT, device=k.device, dtype=k.dtype)
|
||||
kernel = chunk_gated_delta_rule_fwd_kkt_solve_kernel_low_reg
|
||||
kernel[(NT, B * H)](
|
||||
k=k,
|
||||
g=g,
|
||||
beta=beta,
|
||||
A=A,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
H=H,
|
||||
Hg=Hg,
|
||||
K=K,
|
||||
BT=BT,
|
||||
BC=BC,
|
||||
)
|
||||
|
||||
# Step 2: recompute_w_u
|
||||
w, u = recompute_w_u_fwd(
|
||||
k=k,
|
||||
v=v,
|
||||
beta=beta,
|
||||
A=A,
|
||||
g_cumsum=g,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
)
|
||||
return w, u, A
|
||||
@@ -0,0 +1,128 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.srt.layers.attention.fla.fused_sigmoid_gating_recurrent import (
|
||||
fused_sigmoid_gating_delta_rule_update_kernel,
|
||||
)
|
||||
|
||||
|
||||
def fused_sigmoid_gating_delta_rule_update(
|
||||
A_log: torch.Tensor,
|
||||
a: torch.Tensor,
|
||||
dt_bias: torch.Tensor,
|
||||
softplus_beta: float,
|
||||
softplus_threshold: float,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
initial_state_source: torch.Tensor,
|
||||
initial_state_indices: torch.Tensor,
|
||||
scale: Optional[float] = None,
|
||||
use_qk_l2norm_in_kernel: bool = False,
|
||||
cu_seqlens: Optional[torch.Tensor] = None,
|
||||
is_kda: bool = False,
|
||||
# Optional parameters for target_verify support
|
||||
disable_state_update: bool = False,
|
||||
intermediate_states_buffer: Optional[torch.Tensor] = None,
|
||||
intermediate_state_indices: Optional[torch.Tensor] = None,
|
||||
cache_steps: Optional[int] = None,
|
||||
retrieve_parent_token: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""
|
||||
Fused triton implementation of sigmoid gating delta rule update.
|
||||
This function uses a single fused kernel that combines both sigmoid gating computation
|
||||
and the recurrent delta rule update for better performance.
|
||||
|
||||
Supports both decode and target_verify modes:
|
||||
- decode: standard single-step update with state write-back
|
||||
- target_verify: multi-step with intermediate state caching, optional tree attention,
|
||||
and optional state update disable
|
||||
"""
|
||||
B, T, H, K, V = *k.shape, v.shape[-1]
|
||||
stride_q = q.stride()[1]
|
||||
stride_k = k.stride()[1]
|
||||
stride_v = v.stride()[1]
|
||||
stride_b = b.stride()[-2]
|
||||
# Both paths (KDA/GDN) advance p_a once per token, so use the token-axis stride.
|
||||
# For 2D a ([T, ...]) this is stride(0); for 3D a ([B, T, ...]) this is stride(1).
|
||||
# Using stride()[-2] covers GDN [T, HV] and KDA layouts ([T, HV*K] / [B, T, HV*K]).
|
||||
stride_a = a.stride()[-2]
|
||||
HV = v.shape[2]
|
||||
N = B if cu_seqlens is None else len(cu_seqlens) - 1
|
||||
BK, BV = triton.next_power_of_2(K), min(
|
||||
triton.next_power_of_2(V), 16
|
||||
) # use 16 here to reduce register pressure
|
||||
NK, NV = triton.cdiv(K, BK), triton.cdiv(V, BV)
|
||||
assert NK == 1, "NK > 1 is not supported yet"
|
||||
num_stages = 3
|
||||
num_warps = 1
|
||||
|
||||
if scale is None:
|
||||
scale = k.shape[-1] ** -0.5
|
||||
else:
|
||||
assert scale > 0, "scale must be positive"
|
||||
|
||||
o = q.new_empty(NK, *v.shape)
|
||||
|
||||
# Prepare retrieve_parent_token strides
|
||||
if retrieve_parent_token is not None:
|
||||
stride_retrieve_parent_token_seq = retrieve_parent_token.stride(0)
|
||||
stride_retrieve_parent_token_token = retrieve_parent_token.stride(1)
|
||||
else:
|
||||
stride_retrieve_parent_token_seq = 0
|
||||
stride_retrieve_parent_token_token = 0
|
||||
|
||||
NP2_T = triton.next_power_of_2(T)
|
||||
|
||||
grid = (NK, NV, N * HV)
|
||||
|
||||
fused_sigmoid_gating_delta_rule_update_kernel[grid](
|
||||
A_log=A_log,
|
||||
a=a,
|
||||
dt_bias=dt_bias,
|
||||
softplus_beta=softplus_beta,
|
||||
softplus_threshold=softplus_threshold,
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
b=b,
|
||||
o=o,
|
||||
h0_source=initial_state_source,
|
||||
h0_indices=initial_state_indices,
|
||||
cu_seqlens=cu_seqlens,
|
||||
intermediate_states_buffer=intermediate_states_buffer,
|
||||
intermediate_state_indices=intermediate_state_indices,
|
||||
cache_steps=0 if cache_steps is None else cache_steps,
|
||||
retrieve_parent_token_ptr=retrieve_parent_token,
|
||||
stride_retrieve_parent_token_seq=stride_retrieve_parent_token_seq,
|
||||
stride_retrieve_parent_token_token=stride_retrieve_parent_token_token,
|
||||
scale=scale,
|
||||
T=T,
|
||||
stride_a=stride_a,
|
||||
stride_q=stride_q,
|
||||
stride_k=stride_k,
|
||||
stride_v=stride_v,
|
||||
stride_b=stride_b,
|
||||
NP2_T=NP2_T,
|
||||
B=B,
|
||||
H=H,
|
||||
HV=HV,
|
||||
K=K,
|
||||
V=V,
|
||||
BK=BK,
|
||||
BV=BV,
|
||||
USE_INITIAL_STATE=initial_state_source is not None,
|
||||
USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
IS_KDA=is_kda,
|
||||
DISABLE_STATE_UPDATE=disable_state_update,
|
||||
CACHE_INTERMEDIATE_STATES=intermediate_states_buffer is not None,
|
||||
HAS_EAGLE_TREE_CUSTOM_ATTN_MASK=retrieve_parent_token is not None,
|
||||
num_warps=num_warps,
|
||||
num_stages=num_stages,
|
||||
)
|
||||
o = o.squeeze(0)
|
||||
return o
|
||||
Reference in New Issue
Block a user