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,166 @@
|
||||
"""Microbenchmark: buffered output-only GDN decode (ReplaySSM Part A) vs. the
|
||||
existing packed GDN decode kernel.
|
||||
|
||||
Compares per-step decode latency of
|
||||
``fused_recurrent_gated_delta_rule_packed_decode`` (writes the full recurrent
|
||||
state S every step) against ``fused_recurrent_gdn_replayssm_decode`` at
|
||||
L in {1, 8, 16} (writes the full state only every L steps) across batch sizes
|
||||
{1, 16, 64, 256} for a realistic GDN config (HV=32, K=V=128).
|
||||
|
||||
The win is per-step HBM *state* traffic: the packed kernel reads + writes S
|
||||
(~2 * num_slots * HV * V * K * 4 bytes / step for an fp32 state), while the
|
||||
ReplaySSM kernel reads S every step but writes it only 1-in-L steps, plus a
|
||||
small ring append (d:[HV,V], k:[H,K], g:[HV] per step). The amortized state
|
||||
traffic ratio is reported per L.
|
||||
|
||||
Run::
|
||||
|
||||
python -m sglang.srt.layers.attention.fla.bench_gdn_replayssm_decode
|
||||
|
||||
Requires a GPU (Triton).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.srt.layers.attention.fla.fused_recurrent import (
|
||||
fused_recurrent_gated_delta_rule_packed_decode,
|
||||
)
|
||||
from sglang.srt.layers.attention.fla.fused_recurrent_linear_replayssm import (
|
||||
fused_recurrent_gdn_replayssm_decode,
|
||||
)
|
||||
|
||||
|
||||
def _make_static(B, H, HV, K, V, dtype, device):
|
||||
qk_dim = 2 * H * K
|
||||
v_dim = HV * V
|
||||
mixed_qkv = torch.randn(B, qk_dim + v_dim, device=device, dtype=dtype)
|
||||
a = torch.randn(B, HV, device=device, dtype=dtype) * 0.5
|
||||
b = torch.randn(B, HV, device=device, dtype=dtype)
|
||||
A_log = (torch.randn(HV, device=device, dtype=torch.float32) * 0.3).contiguous()
|
||||
dt_bias = (torch.randn(HV, device=device, dtype=torch.float32) * 0.1).contiguous()
|
||||
return mixed_qkv, a, b, A_log, dt_bias
|
||||
|
||||
|
||||
def _state_bytes_per_step(B, HV, K, V, L, dtype):
|
||||
"""Amortized per-step HBM *state* traffic (bytes), state in fp32.
|
||||
|
||||
packed: read S + write S every step.
|
||||
replay: read S every step; write S once per L steps; append ring records
|
||||
(d:[HV,V] in `dtype`, k:[H,K] in `dtype` shared across HV//H, g:[HV]
|
||||
fp32) every step. We report the dominant fp32-state terms; ring
|
||||
appends are tiny by comparison and shown separately.
|
||||
"""
|
||||
fp32 = 4
|
||||
state_elems = B * HV * V * K # one record per active request slot
|
||||
packed = (state_elems * fp32) * 2 # read + write
|
||||
replay = (state_elems * fp32) * (1 + 1.0 / L) # read every step + write 1/L
|
||||
return packed, replay
|
||||
|
||||
|
||||
def _bench_cfg(B, H, HV, K, V, Ls, dtype, device, num_slots=None, warmup=25, rep=100):
|
||||
num_slots = num_slots or B
|
||||
mixed_qkv, a, b, A_log, dt_bias = _make_static(B, H, HV, K, V, dtype, device)
|
||||
scale = K**-0.5
|
||||
cache_indices = torch.arange(B, device=device, dtype=torch.int32)
|
||||
|
||||
# packed decode
|
||||
state = torch.randn(num_slots, HV, V, K, device=device, dtype=torch.float32)
|
||||
out = mixed_qkv.new_empty(B, 1, HV, V)
|
||||
|
||||
def run_packed():
|
||||
fused_recurrent_gated_delta_rule_packed_decode(
|
||||
mixed_qkv=mixed_qkv,
|
||||
a=a,
|
||||
b=b,
|
||||
A_log=A_log,
|
||||
dt_bias=dt_bias,
|
||||
scale=scale,
|
||||
initial_state=state,
|
||||
out=out,
|
||||
ssm_state_indices=cache_indices,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
)
|
||||
|
||||
t_packed = triton.testing.do_bench(run_packed, warmup=warmup, rep=rep)
|
||||
|
||||
rows = []
|
||||
for L in Ls:
|
||||
rstate = torch.randn(num_slots, HV, V, K, device=device, dtype=torch.float32)
|
||||
d_cache = torch.zeros(num_slots, HV, L, V, device=device, dtype=dtype)
|
||||
k_cache = torch.zeros(num_slots, H, L, K, device=device, dtype=dtype)
|
||||
g_cache = torch.zeros(num_slots, HV, L, device=device, dtype=torch.float32)
|
||||
write_pos = torch.zeros(B, device=device, dtype=torch.int32)
|
||||
rout = mixed_qkv.new_empty(B, 1, HV, V)
|
||||
nk = 1 if L == 1 else 2
|
||||
|
||||
def run_replay():
|
||||
fused_recurrent_gdn_replayssm_decode(
|
||||
mixed_qkv=mixed_qkv,
|
||||
a=a,
|
||||
b=b,
|
||||
A_log=A_log,
|
||||
dt_bias=dt_bias,
|
||||
scale=scale,
|
||||
initial_state=rstate,
|
||||
d_cache=d_cache,
|
||||
k_cache=k_cache,
|
||||
g_cache=g_cache,
|
||||
out=rout,
|
||||
ssm_state_indices=cache_indices,
|
||||
write_pos=write_pos,
|
||||
use_qk_l2norm_in_kernel=True,
|
||||
nk=nk,
|
||||
)
|
||||
|
||||
t_replay = triton.testing.do_bench(run_replay, warmup=warmup, rep=rep)
|
||||
packed_bytes, replay_bytes = _state_bytes_per_step(B, HV, K, V, L, dtype)
|
||||
rows.append((L, t_replay, t_packed / t_replay, replay_bytes / packed_bytes))
|
||||
return t_packed, rows
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--hv", type=int, default=32, help="num value heads")
|
||||
parser.add_argument("--h", type=int, default=16, help="num key/query heads")
|
||||
parser.add_argument("--k", type=int, default=128)
|
||||
parser.add_argument("--v", type=int, default=128)
|
||||
parser.add_argument("--batch-sizes", type=int, nargs="+", default=[1, 16, 64, 256])
|
||||
parser.add_argument("--ls", type=int, nargs="+", default=[1, 8, 16])
|
||||
parser.add_argument("--dtype", choices=["bf16", "fp16", "fp32"], default="bf16")
|
||||
args = parser.parse_args()
|
||||
|
||||
if not torch.cuda.is_available():
|
||||
raise SystemExit("CUDA / Triton required for this microbenchmark.")
|
||||
device = "cuda"
|
||||
dtype = {
|
||||
"bf16": torch.bfloat16,
|
||||
"fp16": torch.float16,
|
||||
"fp32": torch.float32,
|
||||
}[args.dtype]
|
||||
|
||||
print(
|
||||
f"GDN ReplaySSM decode microbench HV={args.hv} H={args.h} "
|
||||
f"K={args.k} V={args.v} dtype={args.dtype}\n"
|
||||
"per-step latency (ms); speedup = packed/replay; "
|
||||
"state-traffic = replay/packed (lower is better)"
|
||||
)
|
||||
for B in args.batch_sizes:
|
||||
t_packed, rows = _bench_cfg(
|
||||
B, args.h, args.hv, args.k, args.v, args.ls, dtype, device
|
||||
)
|
||||
print(f"\nB={B:<4d} packed={t_packed:.4f} ms")
|
||||
for L, t_replay, speedup, traffic_ratio in rows:
|
||||
print(
|
||||
f" L={L:<3d} replay={t_replay:.4f} ms "
|
||||
f"speedup={speedup:5.2f}x "
|
||||
f"state-traffic={traffic_ratio:5.2f}x"
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,262 @@
|
||||
# Adapted from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/gated_delta_rule/chunk.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
|
||||
from sglang.srt.layers.attention.fla.chunk_delta_h import chunk_gated_delta_rule_fwd_h
|
||||
from sglang.srt.layers.attention.fla.chunk_fwd import chunk_gated_delta_rule_fwd_intra
|
||||
from sglang.srt.layers.attention.fla.chunk_o import chunk_fwd_o
|
||||
from sglang.srt.layers.attention.fla.cumsum import chunk_local_cumsum
|
||||
from sglang.srt.layers.attention.fla.index import (
|
||||
prepare_chunk_indices,
|
||||
)
|
||||
from sglang.srt.layers.attention.fla.l2norm import l2norm_fwd
|
||||
from sglang.srt.layers.attention.fla.utils import (
|
||||
SUPPRESS_LEVEL,
|
||||
autocast_custom_fwd,
|
||||
input_guard,
|
||||
is_intel,
|
||||
)
|
||||
|
||||
if is_intel:
|
||||
from sglang.srt.hardware_backend.xpu.kernels.fla.chunk_delta_h import (
|
||||
chunk_gated_delta_rule_fwd_h,
|
||||
)
|
||||
from sglang.srt.hardware_backend.xpu.kernels.fla.chunk_fwd import (
|
||||
chunk_gated_delta_rule_fwd_intra,
|
||||
)
|
||||
|
||||
CHUNK_SIZE = 64
|
||||
|
||||
|
||||
def chunk_gated_delta_rule_fwd(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
beta: torch.Tensor,
|
||||
scale: float,
|
||||
initial_state: torch.Tensor,
|
||||
initial_state_indices: torch.Tensor,
|
||||
cu_seqlens: Optional[torch.LongTensor] = None,
|
||||
chunk_indices: torch.LongTensor | None = None,
|
||||
):
|
||||
g = chunk_local_cumsum(
|
||||
g, chunk_size=CHUNK_SIZE, cu_seqlens=cu_seqlens, chunk_indices=chunk_indices
|
||||
)
|
||||
|
||||
# fused kkt + solve_tril + recompute_w_u
|
||||
w, u, A = chunk_gated_delta_rule_fwd_intra(
|
||||
k=k,
|
||||
v=v,
|
||||
g=g,
|
||||
beta=beta,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
)
|
||||
|
||||
h, v_new = chunk_gated_delta_rule_fwd_h(
|
||||
k=k,
|
||||
w=w,
|
||||
u=u,
|
||||
g=g,
|
||||
initial_state=initial_state,
|
||||
initial_state_indices=initial_state_indices,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
)
|
||||
o = chunk_fwd_o(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v_new,
|
||||
h=h,
|
||||
g=g,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
)
|
||||
if SUPPRESS_LEVEL < 3:
|
||||
return g, o, A, None, h, None
|
||||
elif SUPPRESS_LEVEL >= 3:
|
||||
return g, o, A, w, h, v_new
|
||||
|
||||
|
||||
class ChunkGatedDeltaRuleFunction(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
@input_guard
|
||||
@autocast_custom_fwd
|
||||
def forward(
|
||||
ctx,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
beta: torch.Tensor,
|
||||
scale: float,
|
||||
initial_state: torch.Tensor,
|
||||
initial_state_indices: torch.Tensor,
|
||||
cu_seqlens: Optional[torch.LongTensor] = None,
|
||||
use_qk_l2norm_in_kernel: bool = False,
|
||||
):
|
||||
q_orig = q
|
||||
k_orig = k
|
||||
|
||||
if use_qk_l2norm_in_kernel:
|
||||
q = l2norm_fwd(q)
|
||||
k = l2norm_fwd(k)
|
||||
|
||||
chunk_indices = (
|
||||
prepare_chunk_indices(cu_seqlens, CHUNK_SIZE)
|
||||
if cu_seqlens is not None
|
||||
else None
|
||||
)
|
||||
g, o, A, w, h, v_new = chunk_gated_delta_rule_fwd(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
g=g,
|
||||
beta=beta,
|
||||
scale=scale,
|
||||
initial_state=initial_state,
|
||||
initial_state_indices=initial_state_indices,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
)
|
||||
return o.to(q.dtype), h
|
||||
|
||||
|
||||
@torch.compiler.disable
|
||||
def chunk_gated_delta_rule(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
beta: torch.Tensor,
|
||||
scale: float = None,
|
||||
initial_state: torch.Tensor = None,
|
||||
initial_state_indices: torch.Tensor = None,
|
||||
cu_seqlens: Optional[torch.LongTensor] = None,
|
||||
head_first: bool = False,
|
||||
use_qk_l2norm_in_kernel: bool = False,
|
||||
):
|
||||
r"""
|
||||
Args:
|
||||
q (torch.Tensor):
|
||||
queries of shape `[B, T, H, K]` if `head_first=False` else `[B, H, T, K]`.
|
||||
k (torch.Tensor):
|
||||
keys of shape `[B, T, H, K]` if `head_first=False` else `[B, H, T, K]`.
|
||||
v (torch.Tensor):
|
||||
values of shape `[B, T, H, V]` if `head_first=False` else `[B, H, T, V]`.
|
||||
g (torch.Tensor):
|
||||
(forget) gating tensor (in log space!) of shape `[B, T, H]` if `head_first=False` else `[B, H, T]`.
|
||||
beta (torch.Tensor):
|
||||
betas of shape `[B, T, H]` if `head_first=False` else `[B, H, T]`.
|
||||
scale (Optional[int]):
|
||||
Scale factor for the RetNet attention scores.
|
||||
If not provided, it will default to `1 / sqrt(K)`. Default: `None`.
|
||||
initial_state (Optional[torch.Tensor]):
|
||||
Initial state of shape `[N, H, V, K]` for `N` input sequences.
|
||||
For equal-length input sequences, `N` equals the batch size `B`.
|
||||
Default: `None`.
|
||||
output_final_state (Optional[bool]):
|
||||
Whether to output the final state of shape `[N, H, V, K]`. Default: `False`.
|
||||
cu_seqlens (torch.LongTensor):
|
||||
Cumulative sequence lengths of shape `[N+1]` used for variable-length training,
|
||||
consistent with the FlashAttention API.
|
||||
head_first (Optional[bool]):
|
||||
Whether the inputs are in the head-first format, which is not supported for variable-length inputs.
|
||||
Default: `False`.
|
||||
|
||||
Returns:
|
||||
o (torch.Tensor):
|
||||
Outputs of shape `[B, T, H, V]` if `head_first=False` else `[B, H, T, V]`.
|
||||
final_state (torch.Tensor):
|
||||
Final state of shape `[N, H, V, K]` if `output_final_state=True` else `None`.
|
||||
|
||||
Examples::
|
||||
>>> import torch
|
||||
>>> import torch.nn.functional as F
|
||||
>>> from einops import rearrange
|
||||
>>> from fla.ops.gated_delta_rule import chunk_gated_delta_rule
|
||||
# inputs with equal lengths
|
||||
>>> B, T, H, K, V = 4, 2048, 4, 512, 512
|
||||
>>> q = torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda')
|
||||
>>> k = F.normalize(torch.randn(B, T, H, K, dtype=torch.bfloat16, device='cuda'), p=2, dim=-1)
|
||||
>>> v = torch.randn(B, T, H, V, dtype=torch.bfloat16, device='cuda')
|
||||
>>> beta = torch.rand(B, T, H, dtype=torch.bfloat16, device='cuda').sigmoid()
|
||||
>>> g = F.logsigmoid(torch.rand(B, T, H, dtype=torch.bfloat16, device='cuda'))
|
||||
>>> h0 = torch.randn(B, H, K, V, dtype=torch.bfloat16, device='cuda')
|
||||
>>> o, ht = chunk_gated_delta_rule(
|
||||
q, k, v, g, beta,
|
||||
initial_state=h0,
|
||||
output_final_state=True
|
||||
)
|
||||
# for variable-length inputs, the batch size `B` is expected to be 1 and `cu_seqlens` is required
|
||||
>>> q, k, v, beta, g = map(lambda x: rearrange(x, 'b t ... -> 1 (b t) ...'), (q, k, v, beta, g))
|
||||
# for a batch with 4 sequences, `cu_seqlens` with 5 start/end positions are expected
|
||||
>>> cu_seqlens = q.new_tensor([0, 2048, 4096, 6144, 8192], dtype=torch.long)
|
||||
>>> o_var, ht_var = chunk_gated_delta_rule(
|
||||
q, k, v, g, beta,
|
||||
initial_state=h0,
|
||||
output_final_state=True,
|
||||
cu_seqlens=cu_seqlens
|
||||
)
|
||||
"""
|
||||
assert q.dtype == k.dtype == v.dtype
|
||||
assert (
|
||||
q.dtype != torch.float32
|
||||
), "ChunkGatedDeltaRuleFunction does not support float32. Please use bfloat16."
|
||||
assert (
|
||||
len(beta.shape) == 3
|
||||
), "beta must be of shape [B, T, H] if head_first=False, or [B, H, T] otherwise."
|
||||
|
||||
if head_first:
|
||||
raise DeprecationWarning(
|
||||
"head_first is deprecated and will be removed in a future version. "
|
||||
"Please use head_first=False for now instead."
|
||||
)
|
||||
q, k, v, beta, g = map(
|
||||
lambda x: rearrange(x, "b h t ... -> b t h ..."), (q, k, v, beta, g)
|
||||
)
|
||||
# if not head_first and q.shape[1] < q.shape[2]:
|
||||
# warnings.warn(
|
||||
# f"Input tensor shape suggests potential format mismatch: seq_len ({q.shape[1]}) < num_heads ({q.shape[2]}). "
|
||||
# "This may indicate the inputs were passed in head-first format [B, H, T, ...] "
|
||||
# "when head_first=False was specified. "
|
||||
# "Please verify your input tensor format matches the expected shape [B, T, H, ...]."
|
||||
# )
|
||||
if cu_seqlens is not None:
|
||||
if q.shape[0] != 1:
|
||||
raise ValueError(
|
||||
f"The batch size is expected to be 1 rather than {q.shape[0]} when using `cu_seqlens`."
|
||||
f"Please flatten variable-length inputs before processing."
|
||||
)
|
||||
if (
|
||||
initial_state_indices is not None
|
||||
and initial_state_indices.shape[0] != len(cu_seqlens) - 1
|
||||
):
|
||||
raise ValueError(
|
||||
f"The number of initial states is expected to be equal to the number of input sequences, "
|
||||
f"i.e., {len(cu_seqlens) - 1} rather than {initial_state_indices.shape[0]}."
|
||||
)
|
||||
if scale is None:
|
||||
scale = k.shape[-1] ** -0.5
|
||||
o, h = ChunkGatedDeltaRuleFunction.apply(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
g,
|
||||
beta,
|
||||
scale,
|
||||
initial_state,
|
||||
initial_state_indices,
|
||||
cu_seqlens,
|
||||
use_qk_l2norm_in_kernel,
|
||||
)
|
||||
if head_first:
|
||||
o = rearrange(o, "b t h ... -> b h t ...")
|
||||
return o, None, h
|
||||
@@ -0,0 +1,357 @@
|
||||
# Adapted from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/common/chunk_delta_h.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
import os
|
||||
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, safe_exp
|
||||
from sglang.srt.layers.attention.fla.utils import (
|
||||
autotune_cache_kwargs,
|
||||
is_nvidia_hopper,
|
||||
)
|
||||
|
||||
NUM_WARPS = [2, 4] if is_nvidia_hopper else [2, 4, 8, 16]
|
||||
CHUNK_SIZE = 64
|
||||
GDN_CHUNK_H_BV = int(os.getenv("SGLANG_GDN_CHUNK_H_BV", "32"))
|
||||
GDN_CHUNK_H_NUM_WARPS = int(os.getenv("SGLANG_GDN_CHUNK_H_NUM_WARPS", "4"))
|
||||
GDN_CHUNK_H_NUM_STAGES = int(os.getenv("SGLANG_GDN_CHUNK_H_NUM_STAGES", "2"))
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
# Single hardcoded config. The kernel writes ht (final state) back into
|
||||
# initial_state in-place; with multiple configs, triton's autotune benchmark
|
||||
# phase invokes the kernel many times for timing and corrupts the cache pool,
|
||||
# producing silently wrong output on the first user request. Restoring via
|
||||
# `restore_value=["initial_state"]` works for unit tests but OOMs on
|
||||
# production-scale models (e.g. Kimi-Linear-48B at default mem_fraction)
|
||||
# because cloning the cache pool for each benchmark exceeds available memory.
|
||||
# NT_BUCKET is kept in the autotune key for forward-compatibility (allows
|
||||
# future per-bucket configs once the kernel is refactored to write final
|
||||
# state to a separate output buffer). The env knobs keep this single-config
|
||||
# property while allowing model/hardware-local validation of the selected
|
||||
# tile without corrupting the state pool through multi-config autotune.
|
||||
configs=[
|
||||
triton.Config(
|
||||
{"BV": GDN_CHUNK_H_BV},
|
||||
num_warps=GDN_CHUNK_H_NUM_WARPS,
|
||||
num_stages=GDN_CHUNK_H_NUM_STAGES,
|
||||
)
|
||||
],
|
||||
key=["H", "K", "V", "BT", "USE_GK", "NT_BUCKET"],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def chunk_gated_delta_rule_fwd_kernel_h_blockdim64(
|
||||
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,
|
||||
):
|
||||
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
|
||||
|
||||
# [BV, BK]
|
||||
b_h1 = tl.zeros([BV, 64], dtype=tl.float32)
|
||||
if K > 64:
|
||||
b_h2 = tl.zeros([BV, 64], dtype=tl.float32)
|
||||
if K > 128:
|
||||
b_h3 = tl.zeros([BV, 64], dtype=tl.float32)
|
||||
if K > 192:
|
||||
b_h4 = tl.zeros([BV, 64], dtype=tl.float32)
|
||||
|
||||
# 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
|
||||
|
||||
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
|
||||
|
||||
# load initial state
|
||||
if USE_INITIAL_STATE:
|
||||
p_h0_1 = tl.make_block_ptr(h0, (V, K), (K, 1), (i_v * BV, 0), (BV, 64), (1, 0))
|
||||
b_h1 += tl.load(p_h0_1, boundary_check=(0, 1)).to(tl.float32)
|
||||
if K > 64:
|
||||
p_h0_2 = tl.make_block_ptr(
|
||||
h0, (V, K), (K, 1), (i_v * BV, 64), (BV, 64), (1, 0)
|
||||
)
|
||||
b_h2 += tl.load(p_h0_2, boundary_check=(0, 1)).to(tl.float32)
|
||||
if K > 128:
|
||||
p_h0_3 = tl.make_block_ptr(
|
||||
h0, (V, K), (K, 1), (i_v * BV, 128), (BV, 64), (1, 0)
|
||||
)
|
||||
b_h3 += tl.load(p_h0_3, boundary_check=(0, 1)).to(tl.float32)
|
||||
if K > 192:
|
||||
p_h0_4 = tl.make_block_ptr(
|
||||
h0, (V, K), (K, 1), (i_v * BV, 192), (BV, 64), (1, 0)
|
||||
)
|
||||
b_h4 += tl.load(p_h0_4, boundary_check=(0, 1)).to(tl.float32)
|
||||
|
||||
# main recurrence
|
||||
for i_t in range(NT):
|
||||
p_h1 = tl.make_block_ptr(
|
||||
h + i_t * stride_h, (V, K), (K, 1), (i_v * BV, 0), (BV, 64), (1, 0)
|
||||
)
|
||||
tl.store(p_h1, b_h1.to(p_h1.dtype.element_ty), boundary_check=(0, 1))
|
||||
if K > 64:
|
||||
p_h2 = tl.make_block_ptr(
|
||||
h + i_t * stride_h, (V, K), (K, 1), (i_v * BV, 64), (BV, 64), (1, 0)
|
||||
)
|
||||
tl.store(p_h2, b_h2.to(p_h2.dtype.element_ty), boundary_check=(0, 1))
|
||||
if K > 128:
|
||||
p_h3 = tl.make_block_ptr(
|
||||
h + i_t * stride_h, (V, K), (K, 1), (i_v * BV, 128), (BV, 64), (1, 0)
|
||||
)
|
||||
tl.store(p_h3, b_h3.to(p_h3.dtype.element_ty), boundary_check=(0, 1))
|
||||
if K > 192:
|
||||
p_h4 = tl.make_block_ptr(
|
||||
h + i_t * stride_h, (V, K), (K, 1), (i_v * BV, 192), (BV, 64), (1, 0)
|
||||
)
|
||||
tl.store(p_h4, b_h4.to(p_h4.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
p_w = tl.make_block_ptr(
|
||||
w, (T, K), (stride_w, 1), (i_t * BT, 0), (BT, 64), (1, 0)
|
||||
)
|
||||
b_w = tl.load(p_w, boundary_check=(0, 1))
|
||||
b_v = tl.dot(b_w, tl.trans(b_h1).to(b_w.dtype))
|
||||
if K > 64:
|
||||
p_w = tl.make_block_ptr(
|
||||
w, (T, K), (stride_w, 1), (i_t * BT, 64), (BT, 64), (1, 0)
|
||||
)
|
||||
b_w = tl.load(p_w, boundary_check=(0, 1))
|
||||
b_v += tl.dot(b_w, tl.trans(b_h2).to(b_w.dtype))
|
||||
if K > 128:
|
||||
p_w = tl.make_block_ptr(
|
||||
w, (T, K), (stride_w, 1), (i_t * BT, 128), (BT, 64), (1, 0)
|
||||
)
|
||||
b_w = tl.load(p_w, boundary_check=(0, 1))
|
||||
b_v += tl.dot(b_w, tl.trans(b_h3).to(b_w.dtype))
|
||||
if K > 192:
|
||||
p_w = tl.make_block_ptr(
|
||||
w, (T, K), (stride_w, 1), (i_t * BT, 192), (BT, 64), (1, 0)
|
||||
)
|
||||
b_w = tl.load(p_w, boundary_check=(0, 1))
|
||||
b_v += tl.dot(b_w, tl.trans(b_h4).to(b_w.dtype))
|
||||
p_v = tl.make_block_ptr(
|
||||
v, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0)
|
||||
)
|
||||
b_v = tl.load(p_v, boundary_check=(0, 1)) - b_v
|
||||
|
||||
if SAVE_NEW_VALUE:
|
||||
p_v = tl.make_block_ptr(
|
||||
v_new, (T, V), (stride_v, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0)
|
||||
)
|
||||
tl.store(p_v, b_v.to(p_v.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
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 * safe_exp(b_g_last - b_g)[:, None]
|
||||
b_g_last = exp(b_g_last)
|
||||
b_h1 = b_h1 * b_g_last
|
||||
if K > 64:
|
||||
b_h2 = b_h2 * b_g_last
|
||||
if K > 128:
|
||||
b_h3 = b_h3 * b_g_last
|
||||
if K > 192:
|
||||
b_h4 = b_h4 * b_g_last
|
||||
|
||||
if USE_GK:
|
||||
o_k1 = tl.arange(0, 64)
|
||||
b_gk_last1 = tl.load(
|
||||
gk + (bos + last_idx) * H * K + i_h * K + o_k1,
|
||||
mask=(o_k1 < K),
|
||||
other=0.0,
|
||||
)
|
||||
b_h1 *= exp(b_gk_last1)[None, :]
|
||||
if K > 64:
|
||||
o_k2 = 64 + o_k1
|
||||
b_gk_last2 = tl.load(
|
||||
gk + (bos + last_idx) * H * K + i_h * K + o_k2,
|
||||
mask=(o_k2 < K),
|
||||
other=0.0,
|
||||
)
|
||||
b_h2 *= exp(b_gk_last2)[None, :]
|
||||
if K > 128:
|
||||
o_k3 = 128 + o_k1
|
||||
b_gk_last3 = tl.load(
|
||||
gk + (bos + last_idx) * H * K + i_h * K + o_k3,
|
||||
mask=(o_k3 < K),
|
||||
other=0.0,
|
||||
)
|
||||
b_h3 *= exp(b_gk_last3)[None, :]
|
||||
if K > 192:
|
||||
o_k4 = 192 + o_k1
|
||||
b_gk_last4 = tl.load(
|
||||
gk + (bos + last_idx) * H * K + i_h * K + o_k4,
|
||||
mask=(o_k4 < K),
|
||||
other=0.0,
|
||||
)
|
||||
b_h4 *= exp(b_gk_last4)[None, :]
|
||||
b_v = b_v.to(k.dtype.element_ty)
|
||||
|
||||
p_k = tl.make_block_ptr(
|
||||
k, (K, T), (1, stride_k), (0, i_t * BT), (64, BT), (0, 1)
|
||||
)
|
||||
b_k = tl.load(p_k, boundary_check=(0, 1))
|
||||
b_h1 += tl.trans(tl.dot(b_k, b_v))
|
||||
if K > 64:
|
||||
p_k = tl.make_block_ptr(
|
||||
k, (K, T), (1, stride_k), (64, i_t * BT), (64, BT), (0, 1)
|
||||
)
|
||||
b_k = tl.load(p_k, boundary_check=(0, 1))
|
||||
b_h2 += tl.trans(tl.dot(b_k, b_v))
|
||||
if K > 128:
|
||||
p_k = tl.make_block_ptr(
|
||||
k, (K, T), (1, stride_k), (128, i_t * BT), (64, BT), (0, 1)
|
||||
)
|
||||
b_k = tl.load(p_k, boundary_check=(0, 1))
|
||||
b_h3 += tl.trans(tl.dot(b_k, b_v))
|
||||
if K > 192:
|
||||
p_k = tl.make_block_ptr(
|
||||
k, (K, T), (1, stride_k), (192, i_t * BT), (64, BT), (0, 1)
|
||||
)
|
||||
b_k = tl.load(p_k, boundary_check=(0, 1))
|
||||
b_h4 += tl.trans(tl.dot(b_k, b_v))
|
||||
|
||||
# epilogue
|
||||
if INPLACE_UPDATE:
|
||||
p_ht = tl.make_block_ptr(ht, (V, K), (K, 1), (i_v * BV, 0), (BV, 64), (1, 0))
|
||||
tl.store(p_ht, b_h1.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
|
||||
if K > 64:
|
||||
p_ht = tl.make_block_ptr(
|
||||
ht, (V, K), (K, 1), (i_v * BV, 64), (BV, 64), (1, 0)
|
||||
)
|
||||
tl.store(p_ht, b_h2.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
|
||||
if K > 128:
|
||||
p_ht = tl.make_block_ptr(
|
||||
ht, (V, K), (K, 1), (i_v * BV, 128), (BV, 64), (1, 0)
|
||||
)
|
||||
tl.store(p_ht, b_h3.to(p_ht.dtype.element_ty), boundary_check=(0, 1))
|
||||
if K > 192:
|
||||
p_ht = tl.make_block_ptr(
|
||||
ht, (V, K), (K, 1), (i_v * BV, 192), (BV, 64), (1, 0)
|
||||
)
|
||||
tl.store(p_ht, b_h4.to(p_ht.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)
|
||||
|
||||
chunk_gated_delta_rule_fwd_kernel_h_blockdim64[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,416 @@
|
||||
# Adapted from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/gated_delta_rule/chunk_fwd.py
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
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,
|
||||
is_tf32_supported,
|
||||
)
|
||||
from sglang.srt.layers.attention.fla.wy_fast import recompute_w_u_fwd
|
||||
|
||||
# TF32 for the block-merge dot products (16x16 matmuls) is safe and ~2x faster on SM90.
|
||||
# The numerically sensitive forward-substitution uses scalar ops, not tl.dot.
|
||||
if is_tf32_supported:
|
||||
_MERGE_DOT_PRECISION = tl.constexpr("tf32")
|
||||
else:
|
||||
_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 [32, 64]
|
||||
for num_warps in [1, 2, 4]
|
||||
],
|
||||
key=["H", "Hg", "K", "BC"],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def chunk_gated_delta_rule_fwd_kkt_solve_kernel(
|
||||
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,
|
||||
):
|
||||
"""
|
||||
Fused kernel: compute beta * K @ K^T (lower triangular) + solve_tril (I+A)^{-1} in one pass.
|
||||
|
||||
This kernel fuses chunk_scaled_dot_kkt_fwd and solve_tril into a single kernel,
|
||||
avoiding the HBM round-trip for the intermediate A matrix.
|
||||
|
||||
Steps:
|
||||
1. Compute all 10 lower-triangular [BC, BC] blocks of beta * K @ K^T in registers
|
||||
2. Apply gate and beta scaling
|
||||
3. Forward substitution on diagonal blocks
|
||||
4. Block merge to get full (I+A)^{-1}
|
||||
5. Write result to A (output)
|
||||
"""
|
||||
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
|
||||
i_tc1 = i_t * BT + BC
|
||||
i_tc2 = i_t * BT + 2 * BC
|
||||
i_tc3 = i_t * BT + 3 * BC
|
||||
|
||||
k += (bos * Hg + i_h // (H // Hg)) * K
|
||||
A += (bos * H + i_h) * BT
|
||||
|
||||
o_i = tl.arange(0, BC)
|
||||
m_tc0 = (i_tc0 + o_i) < T
|
||||
m_tc1 = (i_tc1 + o_i) < T
|
||||
m_tc2 = (i_tc2 + o_i) < T
|
||||
m_tc3 = (i_tc3 + o_i) < T
|
||||
|
||||
# load beta for each sub-chunk
|
||||
p_b0 = tl.make_block_ptr(beta + bos * H + i_h, (T,), (H,), (i_tc0,), (BC,), (0,))
|
||||
p_b1 = tl.make_block_ptr(beta + bos * H + i_h, (T,), (H,), (i_tc1,), (BC,), (0,))
|
||||
p_b2 = tl.make_block_ptr(beta + bos * H + i_h, (T,), (H,), (i_tc2,), (BC,), (0,))
|
||||
p_b3 = tl.make_block_ptr(beta + bos * H + i_h, (T,), (H,), (i_tc3,), (BC,), (0,))
|
||||
b_b0 = tl.load(p_b0, boundary_check=(0,)).to(tl.float32)
|
||||
b_b1 = tl.load(p_b1, boundary_check=(0,)).to(tl.float32)
|
||||
b_b2 = tl.load(p_b2, boundary_check=(0,)).to(tl.float32)
|
||||
b_b3 = tl.load(p_b3, boundary_check=(0,)).to(tl.float32)
|
||||
|
||||
# load gate if used
|
||||
if USE_G:
|
||||
p_g0 = tl.make_block_ptr(g + bos * H + i_h, (T,), (H,), (i_tc0,), (BC,), (0,))
|
||||
p_g1 = tl.make_block_ptr(g + bos * H + i_h, (T,), (H,), (i_tc1,), (BC,), (0,))
|
||||
p_g2 = tl.make_block_ptr(g + bos * H + i_h, (T,), (H,), (i_tc2,), (BC,), (0,))
|
||||
p_g3 = tl.make_block_ptr(g + bos * H + i_h, (T,), (H,), (i_tc3,), (BC,), (0,))
|
||||
|
||||
b_g0 = tl.load(p_g0, boundary_check=(0,)).to(tl.float32)
|
||||
b_g1 = tl.load(p_g1, boundary_check=(0,)).to(tl.float32)
|
||||
b_g2 = tl.load(p_g2, boundary_check=(0,)).to(tl.float32)
|
||||
b_g3 = tl.load(p_g3, boundary_check=(0,)).to(tl.float32)
|
||||
|
||||
############################################################################
|
||||
# Step 1: compute all 10 lower-triangular [BC, BC] blocks of K @ K^T
|
||||
############################################################################
|
||||
|
||||
# 4 diagonal blocks
|
||||
b_A00 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
b_A11 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
b_A22 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
b_A33 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
|
||||
# 6 off-diagonal blocks
|
||||
b_A10 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
b_A20 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
b_A21 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
b_A30 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
b_A31 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
b_A32 = tl.zeros([BC, BC], dtype=tl.float32)
|
||||
|
||||
for i_k in range(tl.cdiv(K, BK)):
|
||||
p_k0 = tl.make_block_ptr(
|
||||
k, (T, K), (Hg * K, 1), (i_tc0, i_k * BK), (BC, BK), (1, 0)
|
||||
)
|
||||
b_k0 = tl.load(p_k0, boundary_check=(0, 1))
|
||||
# diagonal block 0
|
||||
b_A00 += tl.dot(b_k0, tl.trans(b_k0))
|
||||
|
||||
if i_tc1 < T:
|
||||
p_k1 = tl.make_block_ptr(
|
||||
k, (T, K), (Hg * K, 1), (i_tc1, i_k * BK), (BC, BK), (1, 0)
|
||||
)
|
||||
b_k1 = tl.load(p_k1, boundary_check=(0, 1))
|
||||
# diagonal block 1
|
||||
b_A11 += tl.dot(b_k1, tl.trans(b_k1))
|
||||
# off-diagonal (1,0)
|
||||
b_A10 += tl.dot(b_k1, tl.trans(b_k0))
|
||||
|
||||
if i_tc2 < T:
|
||||
p_k2 = tl.make_block_ptr(
|
||||
k, (T, K), (Hg * K, 1), (i_tc2, i_k * BK), (BC, BK), (1, 0)
|
||||
)
|
||||
b_k2 = tl.load(p_k2, boundary_check=(0, 1))
|
||||
# diagonal block 2
|
||||
b_A22 += tl.dot(b_k2, tl.trans(b_k2))
|
||||
# off-diagonal (2,0), (2,1)
|
||||
b_A20 += tl.dot(b_k2, tl.trans(b_k0))
|
||||
b_A21 += tl.dot(b_k2, tl.trans(b_k1))
|
||||
|
||||
if i_tc3 < T:
|
||||
p_k3 = tl.make_block_ptr(
|
||||
k, (T, K), (Hg * K, 1), (i_tc3, i_k * BK), (BC, BK), (1, 0)
|
||||
)
|
||||
b_k3 = tl.load(p_k3, boundary_check=(0, 1))
|
||||
# diagonal block 3
|
||||
b_A33 += tl.dot(b_k3, tl.trans(b_k3))
|
||||
# off-diagonal (3,0), (3,1), (3,2)
|
||||
b_A30 += tl.dot(b_k3, tl.trans(b_k0))
|
||||
b_A31 += tl.dot(b_k3, tl.trans(b_k1))
|
||||
b_A32 += tl.dot(b_k3, tl.trans(b_k2))
|
||||
|
||||
############################################################################
|
||||
# Step 2: apply gate and beta scaling
|
||||
############################################################################
|
||||
|
||||
if USE_G:
|
||||
# diagonal blocks: g_diff = g_i - g_j within sub-chunk
|
||||
b_A00 *= safe_exp(b_g0[:, None] - b_g0[None, :])
|
||||
b_A11 *= safe_exp(b_g1[:, None] - b_g1[None, :])
|
||||
b_A22 *= safe_exp(b_g2[:, None] - b_g2[None, :])
|
||||
b_A33 *= safe_exp(b_g3[:, None] - b_g3[None, :])
|
||||
|
||||
# off-diagonal blocks: g_diff = g_row - g_col (cross sub-chunk)
|
||||
b_A10 *= safe_exp(b_g1[:, None] - b_g0[None, :])
|
||||
b_A20 *= safe_exp(b_g2[:, None] - b_g0[None, :])
|
||||
b_A21 *= safe_exp(b_g2[:, None] - b_g1[None, :])
|
||||
b_A30 *= safe_exp(b_g3[:, None] - b_g0[None, :])
|
||||
b_A31 *= safe_exp(b_g3[:, None] - b_g1[None, :])
|
||||
b_A32 *= safe_exp(b_g3[:, None] - b_g2[None, :])
|
||||
|
||||
# apply beta to row dimension and mask
|
||||
m_d = o_i[:, None] > o_i[None, :]
|
||||
m_I = o_i[:, None] == o_i[None, :]
|
||||
|
||||
# diagonal blocks: strictly lower triangular within sub-chunk, scaled by beta
|
||||
b_A00 = (
|
||||
tl.where(m_d & (m_tc0[:, None] & m_tc0[None, :]), b_A00, 0.0) * b_b0[:, None]
|
||||
)
|
||||
b_A11 = (
|
||||
tl.where(m_d & (m_tc1[:, None] & m_tc1[None, :]), b_A11, 0.0) * b_b1[:, None]
|
||||
)
|
||||
b_A22 = (
|
||||
tl.where(m_d & (m_tc2[:, None] & m_tc2[None, :]), b_A22, 0.0) * b_b2[:, None]
|
||||
)
|
||||
b_A33 = (
|
||||
tl.where(m_d & (m_tc3[:, None] & m_tc3[None, :]), b_A33, 0.0) * b_b3[:, None]
|
||||
)
|
||||
|
||||
# off-diagonal blocks: full block, scaled by beta
|
||||
b_A10 = b_A10 * b_b1[:, None]
|
||||
b_A20 = b_A20 * b_b2[:, None]
|
||||
b_A21 = b_A21 * b_b2[:, None]
|
||||
b_A30 = b_A30 * b_b3[:, None]
|
||||
b_A31 = b_A31 * b_b3[:, None]
|
||||
b_A32 = b_A32 * b_b3[:, None]
|
||||
|
||||
############################################################################
|
||||
# Step 3: forward substitution on diagonal blocks -> (I + A_diag)^{-1}
|
||||
#
|
||||
# Same algorithm as solve_tril, but rows are extracted from in-register
|
||||
# [BC, BC] tensor via tl.sum(tl.where(mask, tensor, 0), 0) instead of
|
||||
# tl.load from HBM.
|
||||
############################################################################
|
||||
|
||||
b_Ai00 = -b_A00
|
||||
b_Ai11 = -b_A11
|
||||
b_Ai22 = -b_A22
|
||||
b_Ai33 = -b_A33
|
||||
|
||||
for i in range(2, min(BC, T - i_tc0)):
|
||||
b_a00 = tl.sum(tl.where((o_i == i)[:, None], -b_A00, 0.0), 0)
|
||||
b_a00 = tl.where(o_i < i, b_a00, 0.0)
|
||||
b_a00 = b_a00 + tl.sum(b_a00[:, None] * b_Ai00, 0)
|
||||
b_Ai00 = tl.where((o_i == i)[:, None], b_a00, b_Ai00)
|
||||
for i in range(2, min(BC, T - i_tc1)):
|
||||
b_a11 = tl.sum(tl.where((o_i == i)[:, None], -b_A11, 0.0), 0)
|
||||
b_a11 = tl.where(o_i < i, b_a11, 0.0)
|
||||
b_a11 = b_a11 + tl.sum(b_a11[:, None] * b_Ai11, 0)
|
||||
b_Ai11 = tl.where((o_i == i)[:, None], b_a11, b_Ai11)
|
||||
for i in range(2, min(BC, T - i_tc2)):
|
||||
b_a22 = tl.sum(tl.where((o_i == i)[:, None], -b_A22, 0.0), 0)
|
||||
b_a22 = tl.where(o_i < i, b_a22, 0.0)
|
||||
b_a22 = b_a22 + tl.sum(b_a22[:, None] * b_Ai22, 0)
|
||||
b_Ai22 = tl.where((o_i == i)[:, None], b_a22, b_Ai22)
|
||||
for i in range(2, min(BC, T - i_tc3)):
|
||||
b_a33 = tl.sum(tl.where((o_i == i)[:, None], -b_A33, 0.0), 0)
|
||||
b_a33 = tl.where(o_i < i, b_a33, 0.0)
|
||||
b_a33 = b_a33 + tl.sum(b_a33[:, None] * b_Ai33, 0)
|
||||
b_Ai33 = tl.where((o_i == i)[:, None], b_a33, b_Ai33)
|
||||
|
||||
b_Ai00 += m_I
|
||||
b_Ai11 += m_I
|
||||
b_Ai22 += m_I
|
||||
b_Ai33 += m_I
|
||||
|
||||
############################################################################
|
||||
# Step 4: block merge -> full (I + A)^{-1}
|
||||
############################################################################
|
||||
|
||||
b_Ai10 = -tl.dot(
|
||||
tl.dot(b_Ai11, b_A10, input_precision=_MERGE_DOT_PRECISION),
|
||||
b_Ai00,
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
b_Ai21 = -tl.dot(
|
||||
tl.dot(b_Ai22, b_A21, input_precision=_MERGE_DOT_PRECISION),
|
||||
b_Ai11,
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
b_Ai32 = -tl.dot(
|
||||
tl.dot(b_Ai33, b_A32, input_precision=_MERGE_DOT_PRECISION),
|
||||
b_Ai22,
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
|
||||
b_Ai20 = -tl.dot(
|
||||
b_Ai22,
|
||||
tl.dot(b_A20, b_Ai00, input_precision=_MERGE_DOT_PRECISION)
|
||||
+ tl.dot(b_A21, b_Ai10, input_precision=_MERGE_DOT_PRECISION),
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
b_Ai31 = -tl.dot(
|
||||
b_Ai33,
|
||||
tl.dot(b_A31, b_Ai11, input_precision=_MERGE_DOT_PRECISION)
|
||||
+ tl.dot(b_A32, b_Ai21, input_precision=_MERGE_DOT_PRECISION),
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
b_Ai30 = -tl.dot(
|
||||
b_Ai33,
|
||||
tl.dot(b_A30, b_Ai00, input_precision=_MERGE_DOT_PRECISION)
|
||||
+ tl.dot(b_A31, b_Ai10, input_precision=_MERGE_DOT_PRECISION)
|
||||
+ tl.dot(b_A32, b_Ai20, input_precision=_MERGE_DOT_PRECISION),
|
||||
input_precision=_MERGE_DOT_PRECISION,
|
||||
)
|
||||
|
||||
############################################################################
|
||||
# Step 5: store full (I + A)^{-1} to output A
|
||||
############################################################################
|
||||
|
||||
p_A00 = tl.make_block_ptr(A, (T, BT), (H * BT, 1), (i_tc0, 0), (BC, BC), (1, 0))
|
||||
p_A10 = tl.make_block_ptr(A, (T, BT), (H * BT, 1), (i_tc1, 0), (BC, BC), (1, 0))
|
||||
p_A11 = tl.make_block_ptr(A, (T, BT), (H * BT, 1), (i_tc1, BC), (BC, BC), (1, 0))
|
||||
p_A20 = tl.make_block_ptr(A, (T, BT), (H * BT, 1), (i_tc2, 0), (BC, BC), (1, 0))
|
||||
p_A21 = tl.make_block_ptr(A, (T, BT), (H * BT, 1), (i_tc2, BC), (BC, BC), (1, 0))
|
||||
p_A22 = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tc2, 2 * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
p_A30 = tl.make_block_ptr(A, (T, BT), (H * BT, 1), (i_tc3, 0), (BC, BC), (1, 0))
|
||||
p_A31 = tl.make_block_ptr(A, (T, BT), (H * BT, 1), (i_tc3, BC), (BC, BC), (1, 0))
|
||||
p_A32 = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tc3, 2 * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
p_A33 = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_tc3, 3 * BC), (BC, BC), (1, 0)
|
||||
)
|
||||
|
||||
tl.store(p_A00, b_Ai00.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A10, b_Ai10.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A11, b_Ai11.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A20, b_Ai20.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A21, b_Ai21.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A22, b_Ai22.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A30, b_Ai30.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A31, b_Ai31.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A32, b_Ai32.to(A.dtype.element_ty), boundary_check=(0, 1))
|
||||
tl.store(p_A33, b_Ai33.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)
|
||||
chunk_gated_delta_rule_fwd_kkt_solve_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
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,197 @@
|
||||
# Adapted from flash-linear-attention project.
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
# Token-parallel implementation of KDA intra chunk kernel
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.layers.attention.fla.op import exp2
|
||||
from sglang.srt.layers.attention.fla.utils import autotune_cache_kwargs
|
||||
|
||||
|
||||
@triton.heuristics(
|
||||
{
|
||||
"IS_VARLEN": lambda args: args["cu_seqlens"] is not None,
|
||||
}
|
||||
)
|
||||
@triton.autotune(
|
||||
configs=[
|
||||
triton.Config({"BH": BH}, num_warps=num_warps)
|
||||
for BH in [1, 2, 4, 8]
|
||||
for num_warps in [1, 2, 4, 8]
|
||||
],
|
||||
key=["K", "H"],
|
||||
**autotune_cache_kwargs,
|
||||
)
|
||||
@triton.jit(do_not_specialize=["T", "N"])
|
||||
def chunk_kda_fwd_kernel_intra_token_parallel(
|
||||
q,
|
||||
k,
|
||||
g,
|
||||
beta,
|
||||
Aqk,
|
||||
Akk,
|
||||
scale,
|
||||
cu_seqlens,
|
||||
N,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
K: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BC: tl.constexpr,
|
||||
BK: tl.constexpr,
|
||||
BH: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
i_tg, i_hg = tl.program_id(0), tl.program_id(1)
|
||||
|
||||
if IS_VARLEN:
|
||||
i_n = 0
|
||||
left, right = 0, N
|
||||
|
||||
# Unrolled binary search (max B=2^32)
|
||||
# We can limit iterations based on expected max batch size if needed
|
||||
# 20 iterations covers B=1M, usually enough
|
||||
for _ in range(20):
|
||||
if left < right:
|
||||
mid = (left + right) // 2
|
||||
if i_tg < tl.load(cu_seqlens + mid + 1).to(tl.int32):
|
||||
right = mid
|
||||
else:
|
||||
left = mid + 1
|
||||
i_n = left
|
||||
|
||||
bos, eos = tl.load(cu_seqlens + i_n).to(tl.int32), tl.load(
|
||||
cu_seqlens + i_n + 1
|
||||
).to(tl.int32)
|
||||
T = eos - bos
|
||||
i_t = i_tg - bos
|
||||
else:
|
||||
bos = (i_tg // T) * T
|
||||
i_t = i_tg % T
|
||||
|
||||
if i_t >= T:
|
||||
return
|
||||
|
||||
i_c = i_t // BT
|
||||
i_s = (i_t % BT) // BC
|
||||
i_tc = i_c * BT
|
||||
i_ts = i_tc + i_s * BC
|
||||
|
||||
q += bos * H * K
|
||||
k += bos * H * K
|
||||
g += bos * H * K
|
||||
Aqk += bos * H * BT
|
||||
Akk += bos * H * BC
|
||||
beta += bos * H
|
||||
|
||||
o_h = tl.arange(0, BH)
|
||||
o_k = tl.arange(0, BK)
|
||||
m_h = (i_hg * BH + o_h) < H
|
||||
m_k = o_k < K
|
||||
|
||||
p_q = tl.make_block_ptr(
|
||||
q + i_t * H * K, (H, K), (K, 1), (i_hg * BH, 0), (BH, BK), (1, 0)
|
||||
)
|
||||
p_k = tl.make_block_ptr(
|
||||
k + i_t * H * K, (H, K), (K, 1), (i_hg * BH, 0), (BH, BK), (1, 0)
|
||||
)
|
||||
p_g = tl.make_block_ptr(
|
||||
g + i_t * H * K, (H, K), (K, 1), (i_hg * BH, 0), (BH, BK), (1, 0)
|
||||
)
|
||||
p_beta = tl.make_block_ptr(beta + i_t * H, (H,), (1,), (i_hg * BH,), (BH,), (0,))
|
||||
# [BH, BK]
|
||||
b_q = tl.load(p_q, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_k = tl.load(p_k, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_g = tl.load(p_g, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_k = b_k * tl.load(p_beta, boundary_check=(0,)).to(tl.float32)[:, None]
|
||||
|
||||
for j in range(i_ts, min(i_t + 1, min(T, i_ts + BC))):
|
||||
p_kj = tl.make_block_ptr(
|
||||
k + j * H * K, (H, K), (K, 1), (i_hg * BH, 0), (BH, BK), (1, 0)
|
||||
)
|
||||
p_gj = tl.make_block_ptr(
|
||||
g + j * H * K, (H, K), (K, 1), (i_hg * BH, 0), (BH, BK), (1, 0)
|
||||
)
|
||||
# [BH, BK]
|
||||
b_kj = tl.load(p_kj, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_gj = tl.load(p_gj, boundary_check=(0, 1)).to(tl.float32)
|
||||
|
||||
b_kgj = b_kj * exp2(b_g - b_gj)
|
||||
|
||||
b_kgj = tl.where(m_k[None, :], b_kgj, 0.0)
|
||||
# [BH]
|
||||
b_Aqk = tl.sum(b_q * b_kgj, axis=1) * scale
|
||||
b_Akk = tl.sum(b_k * b_kgj, axis=1) * tl.where(j < i_t, 1.0, 0.0)
|
||||
|
||||
tl.store(
|
||||
Aqk + i_t * H * BT + (i_hg * BH + o_h) * BT + j % BT,
|
||||
b_Aqk.to(Aqk.dtype.element_ty),
|
||||
mask=m_h,
|
||||
)
|
||||
tl.store(
|
||||
Akk + i_t * H * BC + (i_hg * BH + o_h) * BC + j - i_ts,
|
||||
b_Akk.to(Akk.dtype.element_ty),
|
||||
mask=m_h,
|
||||
)
|
||||
|
||||
|
||||
def chunk_kda_fwd_intra_token_parallel(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
gk: torch.Tensor,
|
||||
beta: torch.Tensor,
|
||||
Aqk: torch.Tensor,
|
||||
Akk: torch.Tensor,
|
||||
scale: float,
|
||||
cu_seqlens: torch.LongTensor | None = None,
|
||||
chunk_size: int = 64,
|
||||
sub_chunk_size: int = 16,
|
||||
) -> None:
|
||||
"""
|
||||
Token-parallel implementation: each token gets its own thread block.
|
||||
Supports both fixed-length and variable-length sequences.
|
||||
Reduces wasted computation on padding.
|
||||
|
||||
Writes directly to Aqk and Akk tensors (in-place).
|
||||
|
||||
Args:
|
||||
q: [B, T, H, K]
|
||||
k: [B, T, H, K]
|
||||
gk: [B, T, H, K] cumsum of gates
|
||||
beta: [B, T, H]
|
||||
Aqk: [B, T, H, BT] output tensor to write to
|
||||
Akk: [B, T, H, BC] output tensor for diagonal blocks (fp32)
|
||||
scale: attention scale
|
||||
chunk_size: BT (default 64)
|
||||
sub_chunk_size: BC (default 16)
|
||||
"""
|
||||
B, T, H, K = q.shape
|
||||
N = len(cu_seqlens) - 1 if cu_seqlens is not None else B
|
||||
BT = chunk_size
|
||||
BC = sub_chunk_size
|
||||
|
||||
def grid(meta):
|
||||
return (B * T, triton.cdiv(H, meta["BH"]))
|
||||
|
||||
BK = triton.next_power_of_2(K)
|
||||
|
||||
chunk_kda_fwd_kernel_intra_token_parallel[grid](
|
||||
q=q,
|
||||
k=k,
|
||||
g=gk,
|
||||
beta=beta,
|
||||
Aqk=Aqk,
|
||||
Akk=Akk,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
N=N,
|
||||
T=T,
|
||||
H=H,
|
||||
K=K,
|
||||
BT=BT,
|
||||
BC=BC,
|
||||
BK=BK,
|
||||
)
|
||||
return Aqk, Akk
|
||||
@@ -0,0 +1,174 @@
|
||||
# Adapted from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/common/chunk_o.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
from typing import Optional
|
||||
|
||||
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 exp, safe_exp
|
||||
from sglang.srt.layers.attention.fla.utils import check_shared_mem, is_nvidia_hopper
|
||||
|
||||
BKV_LIST = [64, 128] if check_shared_mem() else [32, 64]
|
||||
NUM_WARPS = [2, 4] if is_nvidia_hopper else [2, 4, 8]
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[
|
||||
# triton.Config({"BK": BK, "BV": BV}, num_warps=num_warps, num_stages=num_stages)
|
||||
# for BK in BKV_LIST
|
||||
# for BV in BKV_LIST
|
||||
# for num_warps in NUM_WARPS
|
||||
# for num_stages in [2, 3, 4]
|
||||
# ],
|
||||
# key=["H", "K", "V", "BT"],
|
||||
# )
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def chunk_fwd_kernel_o(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
h,
|
||||
g,
|
||||
o,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
scale,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
Hg: tl.constexpr,
|
||||
K: tl.constexpr,
|
||||
V: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BK: tl.constexpr,
|
||||
BV: tl.constexpr,
|
||||
USE_G: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
i_v, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
i_b, i_h = i_bh // H, i_bh % H
|
||||
|
||||
if IS_VARLEN:
|
||||
i_tg = i_t
|
||||
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
|
||||
NT = tl.cdiv(T, BT)
|
||||
else:
|
||||
NT = tl.cdiv(T, BT)
|
||||
i_tg = i_b * NT + i_t
|
||||
bos, eos = i_b * T, i_b * T + T
|
||||
|
||||
# offset calculation
|
||||
q += (bos * Hg + i_h // (H // Hg)) * K
|
||||
k += (bos * Hg + i_h // (H // Hg)) * K
|
||||
v += (bos * H + i_h) * V
|
||||
o += (bos * H + i_h) * V
|
||||
h += (i_tg * H + i_h).to(tl.int64) * V * K
|
||||
|
||||
b_o = tl.zeros([BT, BV], dtype=tl.float32)
|
||||
b_A = tl.zeros([BT, BT], dtype=tl.float32)
|
||||
|
||||
for i_k in range(tl.cdiv(K, BK)):
|
||||
p_q = tl.make_block_ptr(
|
||||
q, (T, K), (Hg * K, 1), (i_t * BT, i_k * BK), (BT, BK), (1, 0)
|
||||
)
|
||||
p_k = tl.make_block_ptr(
|
||||
k, (K, T), (1, Hg * K), (i_k * BK, i_t * BT), (BK, BT), (0, 1)
|
||||
)
|
||||
p_h = tl.make_block_ptr(
|
||||
h, (V, K), (K, 1), (i_v * BV, i_k * BK), (BV, BK), (1, 0)
|
||||
)
|
||||
# [BT, BK]
|
||||
b_q = tl.load(p_q, boundary_check=(0, 1))
|
||||
# [BK, BT]
|
||||
b_k = tl.load(p_k, boundary_check=(0, 1))
|
||||
# [BV, BK]
|
||||
b_h = tl.load(p_h, boundary_check=(0, 1))
|
||||
|
||||
# [BT, BK] @ [BK, BV] -> [BT, BV]
|
||||
b_o += tl.dot(b_q, tl.trans(b_h))
|
||||
# [BT, BK] @ [BK, BT] -> [BT, BT]
|
||||
b_A += tl.dot(b_q, b_k)
|
||||
|
||||
if USE_G:
|
||||
g += bos * H + i_h
|
||||
p_g = tl.make_block_ptr(g, (T,), (H,), (i_t * BT,), (BT,), (0,))
|
||||
b_g = tl.load(p_g, boundary_check=(0,))
|
||||
b_o = b_o * exp(b_g)[:, None]
|
||||
b_A = b_A * safe_exp(b_g[:, None] - b_g[None, :])
|
||||
|
||||
o_i = tl.arange(0, BT)
|
||||
m_A = o_i[:, None] >= o_i[None, :]
|
||||
b_A = tl.where(m_A, b_A, 0)
|
||||
|
||||
p_v = tl.make_block_ptr(
|
||||
v, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0)
|
||||
)
|
||||
p_o = tl.make_block_ptr(
|
||||
o, (T, V), (H * V, 1), (i_t * BT, i_v * BV), (BT, BV), (1, 0)
|
||||
)
|
||||
b_v = tl.load(p_v, boundary_check=(0, 1))
|
||||
|
||||
# to fix mma -> mma layout conversion
|
||||
# already solved by triton v3.2 or higher
|
||||
b_o = b_o * scale + tl.dot(b_A.to(b_v.dtype), b_v) * scale
|
||||
tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
|
||||
def chunk_fwd_o(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
h: torch.Tensor,
|
||||
g: Optional[torch.Tensor] = None, # cumsum of log decay
|
||||
scale: Optional[float] = None,
|
||||
cu_seqlens: Optional[torch.LongTensor] = None,
|
||||
chunk_size: int = 64,
|
||||
) -> torch.Tensor:
|
||||
B, T, Hg, K, V = *q.shape, v.shape[-1]
|
||||
H = v.shape[-2]
|
||||
BT = min(chunk_size, max(16, triton.next_power_of_2(T)))
|
||||
chunk_indices = (
|
||||
prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
|
||||
)
|
||||
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
|
||||
if scale is None:
|
||||
scale = k.shape[-1] ** -0.5
|
||||
|
||||
o = torch.zeros_like(v)
|
||||
|
||||
def grid(meta):
|
||||
return (triton.cdiv(V, meta["BV"]), NT, B * H)
|
||||
|
||||
chunk_fwd_kernel_o[grid](
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
h,
|
||||
g,
|
||||
o,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
scale,
|
||||
T=T,
|
||||
H=H,
|
||||
Hg=Hg,
|
||||
K=K,
|
||||
V=V,
|
||||
BT=BT,
|
||||
BK=128,
|
||||
BV=64,
|
||||
USE_G=g is not None,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
num_warps=4,
|
||||
num_stages=2,
|
||||
)
|
||||
return o
|
||||
@@ -0,0 +1,147 @@
|
||||
# Adapted from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/common/chunk_scaled_dot_kkt.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
from typing import Optional
|
||||
|
||||
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
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[
|
||||
# triton.Config({"BK": BK}, num_warps=num_warps, num_stages=num_stages)
|
||||
# for BK in [32, 64, 128]
|
||||
# for num_warps in [2, 4, 8]
|
||||
# for num_stages in [2, 3, 4]
|
||||
# ],
|
||||
# key=["H", "K", "BT", "IS_VARLEN"],
|
||||
# )
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def chunk_scaled_dot_kkt_fwd_kernel(
|
||||
k,
|
||||
beta,
|
||||
g_cumsum,
|
||||
A,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
Hg: tl.constexpr,
|
||||
K: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BK: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
USE_G: tl.constexpr,
|
||||
):
|
||||
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
|
||||
o_t = tl.arange(0, BT)
|
||||
|
||||
p_beta = tl.make_block_ptr(
|
||||
beta + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,)
|
||||
)
|
||||
b_beta = tl.load(p_beta, boundary_check=(0,))
|
||||
|
||||
b_A = tl.zeros([BT, BT], dtype=tl.float32)
|
||||
for i_k in range(tl.cdiv(K, BK)):
|
||||
p_k = tl.make_block_ptr(
|
||||
k + (bos * Hg + i_h // (H // Hg)) * K,
|
||||
(T, K),
|
||||
(Hg * K, 1),
|
||||
(i_t * BT, i_k * BK),
|
||||
(BT, 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:
|
||||
p_g = tl.make_block_ptr(
|
||||
g_cumsum + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,)
|
||||
)
|
||||
b_g = tl.load(p_g, boundary_check=(0,))
|
||||
b_g_diff = b_g[:, None] - b_g[None, :]
|
||||
b_A = b_A * safe_exp(b_g_diff)
|
||||
|
||||
b_A *= b_beta[:, None]
|
||||
b_A = tl.where(o_t[:, None] > o_t[None, :], b_A, 0)
|
||||
p_A = tl.make_block_ptr(
|
||||
A + (bos * H + i_h) * BT, (T, BT), (BT * H, 1), (i_t * BT, 0), (BT, BT), (1, 0)
|
||||
)
|
||||
tl.store(p_A, b_A.to(p_A.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
|
||||
def chunk_scaled_dot_kkt_fwd(
|
||||
k: torch.Tensor,
|
||||
beta: torch.Tensor,
|
||||
g_cumsum: Optional[torch.Tensor] = None,
|
||||
cu_seqlens: Optional[torch.LongTensor] = None,
|
||||
chunk_size: int = 64,
|
||||
output_dtype: torch.dtype = torch.float32,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
Compute beta * K * K^T.
|
||||
|
||||
Args:
|
||||
k (torch.Tensor):
|
||||
The key tensor of shape `[B, T, H, K]`.
|
||||
beta (torch.Tensor):
|
||||
The beta tensor of shape `[B, T, H]`.
|
||||
g_cumsum (torch.Tensor):
|
||||
The cumulative sum of the gate tensor of shape `[B, T, H]`.
|
||||
Default: None
|
||||
cu_seqlens (torch.LongTensor):
|
||||
The cumulative sequence lengths of the input tensor.
|
||||
Default: None
|
||||
chunk_size (int):
|
||||
The chunk size. Default: 64.
|
||||
output_dtype (torch.dtype):
|
||||
The dtype of the output tensor. Default: `torch.float32`
|
||||
|
||||
Returns:
|
||||
beta * K * K^T of shape `[B, T, H, BT]` where `BT` is the chunk size.
|
||||
"""
|
||||
|
||||
B, T, Hg, K = k.shape
|
||||
|
||||
H = beta.shape[-1]
|
||||
BT = chunk_size
|
||||
chunk_indices = (
|
||||
prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
|
||||
)
|
||||
NT = triton.cdiv(T, BT) if cu_seqlens is None else len(chunk_indices)
|
||||
A = torch.empty(B, T, H, BT, device=k.device, dtype=output_dtype)
|
||||
chunk_scaled_dot_kkt_fwd_kernel[(NT, B * H)](
|
||||
k=k,
|
||||
beta=beta,
|
||||
g_cumsum=g_cumsum,
|
||||
A=A,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
H=H,
|
||||
Hg=Hg,
|
||||
K=K,
|
||||
BT=BT,
|
||||
BK=64,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
USE_G=g_cumsum is not None,
|
||||
num_warps=8,
|
||||
num_stages=3,
|
||||
)
|
||||
return A
|
||||
@@ -0,0 +1,294 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/utils/cumsum.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
from typing import Optional
|
||||
|
||||
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.utils import check_shared_mem, input_guard
|
||||
|
||||
BS_LIST = [32, 64] if check_shared_mem() else [16, 32]
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[triton.Config({}, num_warps=num_warps) for num_warps in [1, 2, 4, 8]],
|
||||
# key=["B", "H", "BT", "IS_VARLEN", "REVERSE"],
|
||||
# )
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def chunk_local_cumsum_scalar_kernel(
|
||||
s,
|
||||
o,
|
||||
scale,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
B: tl.constexpr,
|
||||
H: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
REVERSE: tl.constexpr,
|
||||
HAS_SCALE: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
HEAD_FIRST: tl.constexpr,
|
||||
):
|
||||
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 HEAD_FIRST:
|
||||
p_s = tl.make_block_ptr(
|
||||
s + bos * H + i_h * T, (T,), (1,), (i_t * BT,), (BT,), (0,)
|
||||
)
|
||||
p_o = tl.make_block_ptr(
|
||||
o + bos * H + i_h * T, (T,), (1,), (i_t * BT,), (BT,), (0,)
|
||||
)
|
||||
else:
|
||||
p_s = tl.make_block_ptr(s + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
|
||||
p_o = tl.make_block_ptr(o + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,))
|
||||
# [BT]
|
||||
b_s = tl.load(p_s, boundary_check=(0,)).to(tl.float32)
|
||||
b_o = tl.cumsum(b_s, axis=0)
|
||||
if REVERSE:
|
||||
b_z = tl.sum(b_s, axis=0)
|
||||
b_o = -b_o + b_z[None] + b_s
|
||||
if HAS_SCALE:
|
||||
b_o *= scale
|
||||
tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0,))
|
||||
|
||||
|
||||
@triton.autotune(
|
||||
configs=[
|
||||
triton.Config({"BS": BS}, num_warps=num_warps, num_stages=num_stages)
|
||||
for BS in BS_LIST
|
||||
for num_warps in [2, 4, 8]
|
||||
for num_stages in [2, 3, 4]
|
||||
],
|
||||
key=["B", "H", "S", "BT", "IS_VARLEN", "REVERSE", "HAS_SCALE"],
|
||||
)
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def chunk_local_cumsum_vector_kernel(
|
||||
s,
|
||||
o,
|
||||
scale,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
B: tl.constexpr,
|
||||
H: tl.constexpr,
|
||||
S: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BS: tl.constexpr,
|
||||
REVERSE: tl.constexpr,
|
||||
HAS_SCALE: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
HEAD_FIRST: tl.constexpr,
|
||||
):
|
||||
i_s, i_t, i_bh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
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
|
||||
|
||||
o_i = tl.arange(0, BT)
|
||||
if REVERSE:
|
||||
m_s = tl.where(o_i[:, None] <= o_i[None, :], 1.0, 0.0)
|
||||
else:
|
||||
m_s = tl.where(o_i[:, None] >= o_i[None, :], 1.0, 0.0)
|
||||
|
||||
if HEAD_FIRST:
|
||||
p_s = tl.make_block_ptr(
|
||||
s + (bos * H + i_h * T) * S,
|
||||
(T, S),
|
||||
(S, 1),
|
||||
(i_t * BT, i_s * BS),
|
||||
(BT, BS),
|
||||
(1, 0),
|
||||
)
|
||||
p_o = tl.make_block_ptr(
|
||||
o + (bos * H + i_h * T) * S,
|
||||
(T, S),
|
||||
(S, 1),
|
||||
(i_t * BT, i_s * BS),
|
||||
(BT, BS),
|
||||
(1, 0),
|
||||
)
|
||||
else:
|
||||
p_s = tl.make_block_ptr(
|
||||
s + (bos * H + i_h) * S,
|
||||
(T, S),
|
||||
(H * S, 1),
|
||||
(i_t * BT, i_s * BS),
|
||||
(BT, BS),
|
||||
(1, 0),
|
||||
)
|
||||
p_o = tl.make_block_ptr(
|
||||
o + (bos * H + i_h) * S,
|
||||
(T, S),
|
||||
(H * S, 1),
|
||||
(i_t * BT, i_s * BS),
|
||||
(BT, BS),
|
||||
(1, 0),
|
||||
)
|
||||
# [BT, BS]
|
||||
b_s = tl.load(p_s, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_o = tl.dot(m_s, b_s, allow_tf32=False)
|
||||
if HAS_SCALE:
|
||||
b_o *= scale
|
||||
tl.store(p_o, b_o.to(p_o.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
|
||||
def chunk_local_cumsum_scalar(
|
||||
g: torch.Tensor,
|
||||
chunk_size: int,
|
||||
reverse: bool = False,
|
||||
scale: float = None,
|
||||
cu_seqlens: Optional[torch.Tensor] = None,
|
||||
head_first: bool = False,
|
||||
output_dtype: Optional[torch.dtype] = torch.float,
|
||||
chunk_indices: Optional[torch.LongTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if head_first:
|
||||
B, H, T = g.shape
|
||||
else:
|
||||
B, T, H = g.shape
|
||||
assert chunk_size == 2 ** (
|
||||
chunk_size.bit_length() - 1
|
||||
), "chunk_size must be a power of 2"
|
||||
BT = chunk_size
|
||||
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)
|
||||
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
|
||||
grid = (NT, B * H)
|
||||
chunk_local_cumsum_scalar_kernel[grid](
|
||||
s=g_org,
|
||||
o=g,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
B=B,
|
||||
H=H,
|
||||
BT=BT,
|
||||
HEAD_FIRST=head_first,
|
||||
REVERSE=reverse,
|
||||
HAS_SCALE=scale is not None,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
num_warps=8,
|
||||
num_stages=3,
|
||||
)
|
||||
return g
|
||||
|
||||
|
||||
def chunk_local_cumsum_vector(
|
||||
g: torch.Tensor,
|
||||
chunk_size: int,
|
||||
reverse: bool = False,
|
||||
scale: float = None,
|
||||
cu_seqlens: Optional[torch.Tensor] = None,
|
||||
head_first: bool = False,
|
||||
output_dtype: Optional[torch.dtype] = torch.float,
|
||||
chunk_indices: Optional[torch.LongTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if head_first:
|
||||
B, H, T, S = g.shape
|
||||
else:
|
||||
B, T, H, S = g.shape
|
||||
BT = chunk_size
|
||||
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)
|
||||
assert chunk_size == 2 ** (
|
||||
chunk_size.bit_length() - 1
|
||||
), "chunk_size must be a power of 2"
|
||||
|
||||
g_org, g = g, torch.empty_like(g, dtype=output_dtype or g.dtype)
|
||||
|
||||
def grid(meta):
|
||||
return (triton.cdiv(meta["S"], meta["BS"]), NT, B * H)
|
||||
|
||||
# keep cumulative normalizer in fp32
|
||||
# this kernel is equivalent to
|
||||
# g = g.view(B, H, NT, BT, -1).cumsum(-2).view(B, H, T, -1)
|
||||
chunk_local_cumsum_vector_kernel[grid](
|
||||
s=g_org,
|
||||
o=g,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
B=B,
|
||||
H=H,
|
||||
S=S,
|
||||
BT=BT,
|
||||
HEAD_FIRST=head_first,
|
||||
REVERSE=reverse,
|
||||
HAS_SCALE=scale is not None,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
)
|
||||
return g
|
||||
|
||||
|
||||
@input_guard
|
||||
def chunk_local_cumsum(
|
||||
g: torch.Tensor,
|
||||
chunk_size: int,
|
||||
reverse: bool = False,
|
||||
scale: float = None,
|
||||
cu_seqlens: Optional[torch.Tensor] = None,
|
||||
head_first: bool = False,
|
||||
output_dtype: Optional[torch.dtype] = torch.float,
|
||||
chunk_indices: Optional[torch.LongTensor] = None,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
if cu_seqlens is not None:
|
||||
assert (
|
||||
g.shape[0] == 1
|
||||
), "Only batch size 1 is supported when cu_seqlens are provided"
|
||||
if len(g.shape) == 3:
|
||||
return chunk_local_cumsum_scalar(
|
||||
g=g,
|
||||
chunk_size=chunk_size,
|
||||
reverse=reverse,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
head_first=head_first,
|
||||
output_dtype=output_dtype,
|
||||
chunk_indices=chunk_indices,
|
||||
)
|
||||
elif len(g.shape) == 4:
|
||||
return chunk_local_cumsum_vector(
|
||||
g=g,
|
||||
chunk_size=chunk_size,
|
||||
reverse=reverse,
|
||||
scale=scale,
|
||||
cu_seqlens=cu_seqlens,
|
||||
head_first=head_first,
|
||||
output_dtype=output_dtype,
|
||||
chunk_indices=chunk_indices,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported input shape {g.shape}, "
|
||||
f"which should be (B, T, H, D) if `head_first=False` "
|
||||
f"or (B, H, T, D) otherwise"
|
||||
)
|
||||
@@ -0,0 +1,75 @@
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
# g = -self.A_log.float().exp() * F.softplus(a.float() + self.dt_bias)
|
||||
# beta_output = b.sigmoid()
|
||||
@triton.jit
|
||||
def fused_gdn_gating_kernel(
|
||||
g,
|
||||
beta_output,
|
||||
A_log,
|
||||
a,
|
||||
b,
|
||||
dt_bias,
|
||||
seq_len,
|
||||
stride_a,
|
||||
stride_b,
|
||||
NUM_HEADS: tl.constexpr,
|
||||
beta: tl.constexpr,
|
||||
threshold: tl.constexpr,
|
||||
BLK_HEADS: tl.constexpr,
|
||||
):
|
||||
i_b, i_s, i_d = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
head_off = i_d * BLK_HEADS + tl.arange(0, BLK_HEADS)
|
||||
off = i_b * seq_len * NUM_HEADS + i_s * NUM_HEADS + head_off
|
||||
mask = head_off < NUM_HEADS
|
||||
blk_A_log = tl.load(A_log + head_off, mask=mask)
|
||||
blk_a = tl.load(a + i_b * stride_a + head_off, mask=mask)
|
||||
blk_b = tl.load(b + i_b * stride_b + head_off, mask=mask)
|
||||
blk_bias = tl.load(dt_bias + head_off, mask=mask)
|
||||
x = blk_a.to(tl.float32) + blk_bias.to(tl.float32)
|
||||
softplus_x = tl.where(
|
||||
beta * x <= threshold, (1 / beta) * tl.log(1 + tl.exp(beta * x)), x
|
||||
)
|
||||
blk_g = -tl.exp(blk_A_log.to(tl.float32)) * softplus_x
|
||||
tl.store(g + off, blk_g.to(g.dtype.element_ty), mask=mask)
|
||||
blk_beta_output = tl.sigmoid(blk_b.to(tl.float32))
|
||||
tl.store(beta_output + off, blk_beta_output.to(b.dtype.element_ty), mask=mask)
|
||||
|
||||
|
||||
def fused_gdn_gating(
|
||||
A_log: torch.Tensor,
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
dt_bias: torch.Tensor,
|
||||
beta: float = 1.0,
|
||||
threshold: float = 20.0,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
batch, num_heads = a.shape
|
||||
seq_len = 1
|
||||
stride_a = a.stride(0)
|
||||
stride_b = b.stride(0)
|
||||
grid = (batch, seq_len, triton.cdiv(num_heads, 8))
|
||||
g = torch.empty(1, batch, num_heads, dtype=torch.float32, device=a.device)
|
||||
beta_output = torch.empty(1, batch, num_heads, dtype=torch.float32, device=b.device)
|
||||
fused_gdn_gating_kernel[grid](
|
||||
g,
|
||||
beta_output,
|
||||
A_log,
|
||||
a,
|
||||
b,
|
||||
dt_bias,
|
||||
seq_len,
|
||||
stride_a,
|
||||
stride_b,
|
||||
num_heads,
|
||||
beta,
|
||||
threshold,
|
||||
8,
|
||||
num_warps=1,
|
||||
)
|
||||
return g, beta_output
|
||||
@@ -0,0 +1,396 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/modules/fused_norm_gate.py
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.utils import (
|
||||
cdiv,
|
||||
cpu_has_amx_support,
|
||||
is_cpu,
|
||||
is_npu,
|
||||
next_power_of_2,
|
||||
)
|
||||
|
||||
_is_npu = is_npu()
|
||||
_use_cpu = is_cpu() and cpu_has_amx_support()
|
||||
|
||||
# Maximum rows per Triton block for layernorm gated kernel
|
||||
MAX_ROWS_PER_BLOCK = 4
|
||||
|
||||
|
||||
@triton.jit
|
||||
def layer_norm_gated_fwd_kernel(
|
||||
x, # pointer to the input
|
||||
g, # pointer to the gate
|
||||
y, # pointer to the output
|
||||
w, # pointer to the weights
|
||||
b, # pointer to the biases
|
||||
residual, # pointer to the residual
|
||||
residual_out, # pointer to the residual
|
||||
mean, # pointer to the mean
|
||||
rstd, # pointer to the 1/std
|
||||
eps, # epsilon to avoid division by zero
|
||||
T, # number of rows in x
|
||||
D: tl.constexpr, # number of columns in x
|
||||
BT: tl.constexpr,
|
||||
BD: tl.constexpr,
|
||||
ACTIVATION: tl.constexpr,
|
||||
IS_RMS_NORM: tl.constexpr,
|
||||
STORE_RESIDUAL_OUT: tl.constexpr,
|
||||
HAS_RESIDUAL: tl.constexpr,
|
||||
HAS_WEIGHT: tl.constexpr,
|
||||
HAS_BIAS: tl.constexpr,
|
||||
):
|
||||
i_t = tl.program_id(0)
|
||||
|
||||
o_d = tl.arange(0, BD)
|
||||
m_d = o_d < D
|
||||
|
||||
p_x = tl.make_block_ptr(x, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
|
||||
b_x = tl.load(p_x, boundary_check=(0, 1)).to(tl.float32)
|
||||
if HAS_RESIDUAL:
|
||||
p_res = tl.make_block_ptr(
|
||||
residual, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0)
|
||||
)
|
||||
b_x += tl.load(p_res, boundary_check=(0, 1)).to(tl.float32)
|
||||
if STORE_RESIDUAL_OUT:
|
||||
p_res_out = tl.make_block_ptr(
|
||||
residual_out, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0)
|
||||
)
|
||||
tl.store(p_res_out, b_x.to(p_res_out.dtype.element_ty), boundary_check=(0, 1))
|
||||
if not IS_RMS_NORM:
|
||||
b_mean = tl.sum(b_x, axis=1) / D
|
||||
p_mean = tl.make_block_ptr(mean, (T,), (1,), (i_t * BT,), (BT,), (0,))
|
||||
tl.store(p_mean, b_mean.to(p_mean.dtype.element_ty), boundary_check=(0,))
|
||||
b_xbar = tl.where(m_d[None, :], b_x - b_mean[:, None], 0.0)
|
||||
b_var = tl.sum(b_xbar * b_xbar, axis=1) / D
|
||||
else:
|
||||
b_xbar = tl.where(m_d[None, :], b_x, 0.0)
|
||||
b_var = tl.sum(b_xbar * b_xbar, axis=1) / D
|
||||
b_rstd = 1 / tl.sqrt(b_var + eps)
|
||||
|
||||
p_rstd = tl.make_block_ptr(rstd, (T,), (1,), (i_t * BT,), (BT,), (0,))
|
||||
tl.store(p_rstd, b_rstd.to(p_rstd.dtype.element_ty), boundary_check=(0,))
|
||||
|
||||
if HAS_WEIGHT:
|
||||
b_w = tl.load(w + o_d, mask=m_d).to(tl.float32)
|
||||
if HAS_BIAS:
|
||||
b_b = tl.load(b + o_d, mask=m_d).to(tl.float32)
|
||||
b_x_hat = (
|
||||
(b_x - b_mean[:, None]) * b_rstd[:, None]
|
||||
if not IS_RMS_NORM
|
||||
else b_x * b_rstd[:, None]
|
||||
)
|
||||
b_y = b_x_hat * b_w[None, :] if HAS_WEIGHT else b_x_hat
|
||||
if HAS_BIAS:
|
||||
b_y = b_y + b_b[None, :]
|
||||
|
||||
# swish/sigmoid output gate
|
||||
p_g = tl.make_block_ptr(g, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
|
||||
b_g = tl.load(p_g, boundary_check=(0, 1)).to(tl.float32)
|
||||
if ACTIVATION == "swish" or ACTIVATION == "silu":
|
||||
b_y = b_y * b_g * tl.sigmoid(b_g)
|
||||
elif ACTIVATION == "sigmoid":
|
||||
b_y = b_y * tl.sigmoid(b_g)
|
||||
|
||||
# Write output
|
||||
p_y = tl.make_block_ptr(y, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
|
||||
tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
|
||||
@triton.jit
|
||||
def layer_norm_gated_fwd_kernel1(
|
||||
x, # pointer to the input
|
||||
g, # pointer to the gate
|
||||
y, # pointer to the output
|
||||
w, # pointer to the weights
|
||||
b, # pointer to the biases
|
||||
residual, # pointer to the residual
|
||||
residual_out, # pointer to the residual
|
||||
mean, # pointer to the mean
|
||||
rstd, # pointer to the 1/std
|
||||
eps, # epsilon to avoid division by zero
|
||||
D: tl.constexpr, # number of columns in x
|
||||
BD: tl.constexpr,
|
||||
ACTIVATION: tl.constexpr,
|
||||
IS_RMS_NORM: tl.constexpr,
|
||||
STORE_RESIDUAL_OUT: tl.constexpr,
|
||||
HAS_RESIDUAL: tl.constexpr,
|
||||
HAS_WEIGHT: tl.constexpr,
|
||||
HAS_BIAS: tl.constexpr,
|
||||
):
|
||||
i_t = tl.program_id(0)
|
||||
x += i_t * D
|
||||
y += i_t * D
|
||||
g += i_t * D
|
||||
if HAS_RESIDUAL:
|
||||
residual += i_t * D
|
||||
if STORE_RESIDUAL_OUT:
|
||||
residual_out += i_t * D
|
||||
|
||||
o_d = tl.arange(0, BD)
|
||||
m_d = o_d < D
|
||||
b_x = tl.load(x + o_d, mask=m_d, other=0.0).to(tl.float32)
|
||||
if HAS_RESIDUAL:
|
||||
b_x += tl.load(residual + o_d, mask=m_d, other=0.0).to(tl.float32)
|
||||
if STORE_RESIDUAL_OUT:
|
||||
tl.store(residual_out + o_d, b_x, mask=m_d)
|
||||
if not IS_RMS_NORM:
|
||||
b_mean = tl.sum(b_x, axis=0) / D
|
||||
tl.store(mean + i_t, b_mean)
|
||||
b_xbar = tl.where(m_d, b_x - b_mean, 0.0)
|
||||
b_var = tl.sum(b_xbar * b_xbar, axis=0) / D
|
||||
else:
|
||||
b_xbar = tl.where(m_d, b_x, 0.0)
|
||||
b_var = tl.sum(b_xbar * b_xbar, axis=0) / D
|
||||
b_rstd = 1 / tl.sqrt(b_var + eps)
|
||||
tl.store(rstd + i_t, b_rstd)
|
||||
|
||||
if HAS_WEIGHT:
|
||||
b_w = tl.load(w + o_d, mask=m_d).to(tl.float32)
|
||||
if HAS_BIAS:
|
||||
b_b = tl.load(b + o_d, mask=m_d).to(tl.float32)
|
||||
b_x_hat = (b_x - b_mean) * b_rstd if not IS_RMS_NORM else b_x * b_rstd
|
||||
b_y = b_x_hat * b_w if HAS_WEIGHT else b_x_hat
|
||||
if HAS_BIAS:
|
||||
b_y = b_y + b_b
|
||||
|
||||
# swish/sigmoid output gate
|
||||
b_g = tl.load(g + o_d, mask=m_d, other=0.0).to(tl.float32)
|
||||
if ACTIVATION == "swish" or ACTIVATION == "silu":
|
||||
b_y = b_y * b_g * tl.sigmoid(b_g)
|
||||
elif ACTIVATION == "sigmoid":
|
||||
b_y = b_y * tl.sigmoid(b_g)
|
||||
|
||||
# Write output
|
||||
tl.store(y + o_d, b_y, mask=m_d)
|
||||
|
||||
|
||||
def layer_norm_gated_fwd(
|
||||
x: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
activation: str = "swish",
|
||||
eps: float = 1e-5,
|
||||
residual: torch.Tensor = None,
|
||||
out_dtype: torch.dtype = None,
|
||||
residual_dtype: torch.dtype = None,
|
||||
is_rms_norm: bool = False,
|
||||
):
|
||||
if residual is not None:
|
||||
residual_dtype = residual.dtype
|
||||
T, D = x.shape
|
||||
if residual is not None:
|
||||
assert residual.shape == (T, D)
|
||||
if weight is not None:
|
||||
assert weight.shape == (D,)
|
||||
if bias is not None:
|
||||
assert bias.shape == (D,)
|
||||
# allocate output
|
||||
y = x if out_dtype is None else torch.empty_like(x, dtype=out_dtype)
|
||||
if residual is not None or (
|
||||
residual_dtype is not None and residual_dtype != x.dtype
|
||||
):
|
||||
residual_out = torch.empty(T, D, device=x.device, dtype=residual_dtype)
|
||||
else:
|
||||
residual_out = None
|
||||
mean = (
|
||||
torch.empty((T,), dtype=torch.float, device=x.device)
|
||||
if not is_rms_norm
|
||||
else None
|
||||
)
|
||||
rstd = torch.empty((T,), dtype=torch.float, device=x.device)
|
||||
# Less than 64KB per feature: enqueue fused kernel
|
||||
MAX_FUSED_SIZE = 65536 // x.element_size()
|
||||
BD = min(MAX_FUSED_SIZE, next_power_of_2(D))
|
||||
if D > BD:
|
||||
raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
|
||||
# heuristics for number of warps
|
||||
|
||||
if D <= 512:
|
||||
BT = 32
|
||||
layer_norm_gated_fwd_kernel[(cdiv(T, BT),)](
|
||||
x=x,
|
||||
g=g,
|
||||
y=y,
|
||||
w=weight,
|
||||
b=bias,
|
||||
residual=residual,
|
||||
residual_out=residual_out,
|
||||
mean=mean,
|
||||
rstd=rstd,
|
||||
eps=eps,
|
||||
T=T,
|
||||
D=D,
|
||||
BD=BD,
|
||||
BT=BT,
|
||||
ACTIVATION=activation,
|
||||
IS_RMS_NORM=is_rms_norm,
|
||||
STORE_RESIDUAL_OUT=residual_out is not None,
|
||||
HAS_RESIDUAL=residual is not None,
|
||||
HAS_WEIGHT=weight is not None,
|
||||
HAS_BIAS=bias is not None,
|
||||
num_warps=4,
|
||||
)
|
||||
else:
|
||||
layer_norm_gated_fwd_kernel1[(T,)](
|
||||
x=x,
|
||||
g=g,
|
||||
y=y,
|
||||
w=weight,
|
||||
b=bias,
|
||||
residual=residual,
|
||||
residual_out=residual_out,
|
||||
mean=mean,
|
||||
rstd=rstd,
|
||||
eps=eps,
|
||||
D=D,
|
||||
BD=BD,
|
||||
ACTIVATION=activation,
|
||||
IS_RMS_NORM=is_rms_norm,
|
||||
STORE_RESIDUAL_OUT=residual_out is not None,
|
||||
HAS_RESIDUAL=residual is not None,
|
||||
HAS_WEIGHT=weight is not None,
|
||||
HAS_BIAS=bias is not None,
|
||||
num_warps=4,
|
||||
)
|
||||
# residual_out is None if residual is None and residual_dtype == input_dtype
|
||||
return y, mean, rstd, residual_out if residual_out is not None else x
|
||||
|
||||
|
||||
class LayerNormGatedFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
x: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
activation: str,
|
||||
residual: torch.Tensor | None = None,
|
||||
eps: float = 1e-6,
|
||||
prenorm: bool = False,
|
||||
residual_in_fp32: bool = False,
|
||||
is_rms_norm: bool = False,
|
||||
):
|
||||
x_shape_og = x.shape
|
||||
g_shape_og = g.shape
|
||||
# reshape input data into 2D tensor
|
||||
x = x.reshape(-1, x.shape[-1])
|
||||
g = g.reshape(-1, g.shape[-1])
|
||||
if residual is not None:
|
||||
assert residual.shape == x_shape_og
|
||||
residual = residual.reshape(-1, residual.shape[-1])
|
||||
residual_dtype = (
|
||||
residual.dtype
|
||||
if residual is not None
|
||||
else (torch.float if residual_in_fp32 else None)
|
||||
)
|
||||
y, mean, rstd, residual_out = layer_norm_gated_fwd(
|
||||
x=x,
|
||||
g=g,
|
||||
weight=weight,
|
||||
bias=bias,
|
||||
activation=activation,
|
||||
eps=eps,
|
||||
residual=residual,
|
||||
residual_dtype=residual_dtype,
|
||||
is_rms_norm=is_rms_norm,
|
||||
)
|
||||
ctx.save_for_backward(residual_out, g, weight, bias, mean, rstd)
|
||||
ctx.x_shape_og = x_shape_og
|
||||
ctx.g_shape_og = g_shape_og
|
||||
ctx.activation = activation
|
||||
ctx.eps = eps
|
||||
ctx.is_rms_norm = is_rms_norm
|
||||
ctx.has_residual = residual is not None
|
||||
ctx.prenorm = prenorm
|
||||
ctx.x_dtype = x.dtype
|
||||
y = y.reshape(x_shape_og)
|
||||
return y if not prenorm else (y, residual_out.reshape(x_shape_og))
|
||||
|
||||
|
||||
def rms_norm_gated(
|
||||
x: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
bias: torch.Tensor,
|
||||
activation: str = "swish",
|
||||
residual: torch.Tensor | None = None,
|
||||
prenorm: bool = False,
|
||||
residual_in_fp32: bool = False,
|
||||
eps: float = 1e-6,
|
||||
):
|
||||
return LayerNormGatedFunction.apply(
|
||||
x,
|
||||
g,
|
||||
weight,
|
||||
bias,
|
||||
activation,
|
||||
residual,
|
||||
eps,
|
||||
prenorm,
|
||||
residual_in_fp32,
|
||||
True,
|
||||
)
|
||||
|
||||
|
||||
class FusedRMSNormGated(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
elementwise_affine: bool = True,
|
||||
eps: float = 1e-5,
|
||||
activation: str = "swish",
|
||||
device: torch.device | None = None,
|
||||
dtype: torch.dtype | None = None,
|
||||
) -> None:
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
|
||||
self.hidden_size = hidden_size
|
||||
self.elementwise_affine = elementwise_affine
|
||||
self.eps = eps
|
||||
self.activation = activation
|
||||
|
||||
if self.activation not in ["swish", "silu", "sigmoid"]:
|
||||
raise ValueError(f"Unsupported activation: {self.activation}")
|
||||
|
||||
if elementwise_affine:
|
||||
self.weight = nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
|
||||
else:
|
||||
self.register_parameter("weight", None)
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
g: torch.Tensor,
|
||||
residual: torch.Tensor | None = None,
|
||||
prenorm: bool = False,
|
||||
residual_in_fp32: bool = False,
|
||||
) -> torch.Tensor:
|
||||
if _use_cpu:
|
||||
assert (
|
||||
self.activation == "silu"
|
||||
), "CPU rmsnorm_gated currently only supports activation silu"
|
||||
return torch.ops.sgl_kernel.fused_rmsnorm_gated_cpu(
|
||||
x, self.weight, g, self.eps
|
||||
)
|
||||
else:
|
||||
return rms_norm_gated(
|
||||
x,
|
||||
g,
|
||||
self.weight,
|
||||
self.bias,
|
||||
self.activation,
|
||||
residual=residual,
|
||||
eps=self.eps,
|
||||
prenorm=prenorm,
|
||||
residual_in_fp32=residual_in_fp32,
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,628 @@
|
||||
# Buffered output-only linear-attention decode (ReplaySSM Part A), ported to
|
||||
# SGLang. Covers BOTH gate granularities with one kernel:
|
||||
# * GDN (``IS_KDA=False``): per-head SCALAR gate ``alpha = exp(g)``.
|
||||
# * KDA (``IS_KDA=True``): per-K-channel gate ``alpha[k] = exp(g[k])`` —
|
||||
# the state decays column-wise, ``S' = S . Diag(alpha) + d k^T``.
|
||||
# GDN is the special case of KDA with all per-K decays equal; a single
|
||||
# ``IS_KDA`` constexpr selects the gate path and the GDN path is bit-for-bit
|
||||
# the original (no regression to the validated GDN kernel).
|
||||
#
|
||||
# This is a STANDALONE increment: kernel + wrapper only. It is NOT yet wired
|
||||
# into the memory pool / radix cache / scheduler / backend dispatch. The caller
|
||||
# (currently the correctness test and the microbenchmark) owns the ring tensors.
|
||||
#
|
||||
# Idea (vs. ``fused_recurrent_gated_delta_rule_packed_decode`` / ``..._kda_...``):
|
||||
# The plain packed decode reads the full recurrent state S [HV, V, K] from
|
||||
# HBM and writes it back *every* decode step (~8*d*n bytes/step of state
|
||||
# traffic for an fp32 state, read+write). ReplaySSM keeps a small per-slot
|
||||
# ring buffer of the last L steps' (d, k, g) and only WRITES the full state
|
||||
# every L steps (a "flush"); on non-flush steps it appends a tiny
|
||||
# (d, k, g) record and reconstructs the readout from the checkpoint S0 plus
|
||||
# the buffer. S0 is still READ every step, so per-step state traffic drops
|
||||
# from read+write (~8*d*n) to read-only (~4*d*n) -> roughly halved.
|
||||
#
|
||||
# Math (single head, single step; matches the packed decode kernels exactly).
|
||||
# Let ``a = exp(g)`` be the decay (scalar for GDN, per-K vector for KDA) and
|
||||
# ``S`` the state *before* this token:
|
||||
# d_cur = beta * (v - (S . Diag(a)) . k) = beta * (v - S . (a (.) k))
|
||||
# o = (S . Diag(a)) . q + d_cur*(k^T q) = S . (a (.) q) + d_cur*(k^T q)
|
||||
# S_new = S . Diag(a) + d_cur k^T # only persisted on flush
|
||||
# where ``(.)`` is elementwise over K. For GDN ``a`` is scalar so
|
||||
# ``S . (a (.) q) = a * (S . q)`` (the cheap scalar post-multiply); for KDA the
|
||||
# per-K ``a`` folds into q/k before the matvec. k^T q uses the RAW current
|
||||
# k/q (the rank-1 term), so it is identical for both gate types.
|
||||
#
|
||||
# Buffered reconstruction: with buffered steps j=0..m-1 holding (d_j, k_j, g_j),
|
||||
# the state *before* the current token is
|
||||
# S = Diag(A) . S0 + sum_j d_j (W_j (.) k_j)^T (per-K form)
|
||||
# with A[c] = exp(sum_j g_j[c]) (total decay, per-K)
|
||||
# W_j[c] = exp(sum_i g_i[c] - cumsum_inclusive_j[c]) = prod_{i>j} a_i[c].
|
||||
# For GDN A and W_j are scalars (g is K-independent) and W_j folds onto d_j
|
||||
# instead of k_j (either factor works for a scalar). S is reconstructed in
|
||||
# K-tiles and immediately read with q (and k) -> the [V,K] state tile is never
|
||||
# fully materialized to HBM on a non-flush step.
|
||||
#
|
||||
# At L=1 the ring is always empty and ``write_pos == L-1`` every step, so the
|
||||
# reconstruction term is zero, the total decay is 1, and this kernel reduces
|
||||
# *algebraically* to the corresponding packed-decode kernel.
|
||||
#
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Ported from vllm/model_executor/layers/fla/ops/fused_recurrent_replayssm.py
|
||||
# (ReplaySSM, commit 3c85112) and adapted to SGLang's packed GDN/KDA decode
|
||||
# layout.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def fused_recurrent_linear_replayssm_decode_kernel(
|
||||
mixed_qkv, # [B, 2*H*K + HV*V] packed (q | k | v) after conv1d
|
||||
a, # GDN: [B, HV] gate input ; KDA: [B, HV, K] per-K gate input
|
||||
b, # [B, HV] beta input b
|
||||
A_log, # [HV] log-space decay parameter (per-head scalar, both gate types)
|
||||
dt_bias, # GDN: [HV] ; KDA: [HV, K] time-step bias
|
||||
o, # [B, HV, V] output (written every step)
|
||||
h0, # [num_slots, HV, V, K] checkpoint state (read every step)
|
||||
ht, # [num_slots, HV, V, K] checkpoint state (written only on flush; == h0)
|
||||
d_cache, # [num_slots, HV, L, V] ring: corrected delta vectors
|
||||
k_cache, # [num_slots, H, L, K] ring: (normed/scaled) keys
|
||||
g_cache, # GDN: [num_slots, HV, L] ; KDA: [num_slots, HV, L, K] log-decay gates (fp32)
|
||||
ssm_state_indices, # [B] physical state slot per decode row
|
||||
write_pos, # [B] int32 per-row ring cursor (0..L-1)
|
||||
force_flush, # [B] int32: !=0 forces a flush this step (radix track boundary)
|
||||
scale,
|
||||
stride_mixed_qkv_tok: tl.constexpr,
|
||||
stride_a_tok: tl.constexpr,
|
||||
stride_b_tok: tl.constexpr,
|
||||
stride_init_state_token: tl.constexpr,
|
||||
stride_final_state_token: tl.constexpr,
|
||||
stride_indices_seq: tl.constexpr,
|
||||
H: tl.constexpr,
|
||||
HV: tl.constexpr,
|
||||
K: tl.constexpr,
|
||||
V: tl.constexpr,
|
||||
BK: tl.constexpr,
|
||||
BV: tl.constexpr,
|
||||
BC: tl.constexpr,
|
||||
NK: tl.constexpr,
|
||||
BKT: tl.constexpr,
|
||||
MAX_CACHE_LEN: tl.constexpr,
|
||||
SOFTPLUS_THRESHOLD: tl.constexpr,
|
||||
USE_QK_L2NORM_IN_KERNEL: tl.constexpr,
|
||||
HAS_FORCE_FLUSH: tl.constexpr,
|
||||
IS_KDA: tl.constexpr,
|
||||
):
|
||||
i_v = tl.program_id(0)
|
||||
i_n = tl.program_id(1)
|
||||
i_hv = tl.program_id(2)
|
||||
i_h = i_hv // (HV // H)
|
||||
|
||||
o_v = i_v * BV + tl.arange(0, BV)
|
||||
o_c = tl.arange(0, BC)
|
||||
mask_v = o_v < V
|
||||
|
||||
# Resolve the physical state slot; zero the output and bail for padded rows.
|
||||
state_idx = tl.load(ssm_state_indices + i_n * stride_indices_seq).to(tl.int64)
|
||||
p_o = o + (i_n * HV + i_hv) * V + o_v
|
||||
if state_idx < 0:
|
||||
tl.store(
|
||||
p_o,
|
||||
tl.zeros([BV], dtype=tl.float32).to(p_o.dtype.element_ty),
|
||||
mask=mask_v,
|
||||
)
|
||||
return
|
||||
|
||||
# Per-row buffer cursor and flush flag (device-side branch, no host branch),
|
||||
# plus the set of valid (already committed) cache positions.
|
||||
b_write_pos = tl.load(write_pos + i_n).to(tl.int64)
|
||||
b_is_flush = b_write_pos == MAX_CACHE_LEN - 1
|
||||
if HAS_FORCE_FLUSH:
|
||||
# A radix track-boundary (or any caller-forced) flush folds the partial
|
||||
# ring (the real `write_pos` entries, NOT L-1) + current token into the
|
||||
# checkpoint so an external snapshot reads an up-to-date state. cache_valid
|
||||
# below still uses the true write_pos, so only committed entries are read.
|
||||
b_is_flush = b_is_flush | (tl.load(force_flush + i_n) != 0)
|
||||
cache_valid = o_c < b_write_pos
|
||||
|
||||
# Gate for the current token. beta is a per-head scalar for both gate
|
||||
# types; A_log is a per-head scalar for both. The decay g/alpha is a
|
||||
# per-head scalar for GDN (computed here) and a per-K vector for KDA
|
||||
# (computed per K-tile inside the loop, since it is K-indexed).
|
||||
# g = -exp(A_log) * softplus(a + dt_bias); alpha = exp(g); beta = sigmoid(b)
|
||||
A_log_val = tl.load(A_log + i_hv).to(tl.float32)
|
||||
b_val = tl.load(b + i_n * stride_b_tok + i_hv).to(tl.float32)
|
||||
beta_val = tl.sigmoid(b_val).to(b.dtype.element_ty).to(tl.float32)
|
||||
if not IS_KDA:
|
||||
a_val = tl.load(a + i_n * stride_a_tok + i_hv).to(tl.float32)
|
||||
dt_bias_val = tl.load(dt_bias + i_hv).to(tl.float32)
|
||||
x = a_val + dt_bias_val
|
||||
softplus_x = tl.where(x <= SOFTPLUS_THRESHOLD, tl.log(1.0 + tl.exp(x)), x)
|
||||
g_val = -tl.exp(A_log_val) * softplus_x
|
||||
alpha_val = tl.exp(g_val)
|
||||
|
||||
# Replay decay over the committed cache, from the cached per-step gates.
|
||||
# b_replay_decay[j] = exp(sum_i g_i - cumsum_inclusive_j) = prod_{i>j} alpha_i
|
||||
p_g_main = g_cache + (state_idx * HV + i_hv) * MAX_CACHE_LEN + o_c
|
||||
b_g_all = tl.load(p_g_main, mask=cache_valid, other=0.0).to(tl.float32)
|
||||
b_g_prefix = tl.cumsum(b_g_all, axis=0)
|
||||
b_g_total = tl.sum(b_g_all, axis=0)
|
||||
b_replay_decay = tl.where(cache_valid, tl.exp(b_g_total - b_g_prefix), 0.0)
|
||||
b_total_decay = tl.exp(b_g_total)
|
||||
|
||||
# Cached corrected-delta vectors d (K-independent). Layout
|
||||
# d_cache[slot, hv, L, V] -> index [V, BC] tile.
|
||||
p_d_main = d_cache + (
|
||||
((state_idx * HV + i_hv) * MAX_CACHE_LEN + o_c[None, :]) * V + o_v[:, None]
|
||||
)
|
||||
b_d_all = tl.load(
|
||||
p_d_main, mask=mask_v[:, None] & cache_valid[None, :], other=0
|
||||
).to(tl.float32)
|
||||
# Cast the (d, k) reconstruction-dot operands to the I/O dtype so tl.dot
|
||||
# runs on TENSOR CORES (bf16 tensor cores for bf16, TF32 for fp32). This is
|
||||
# the ReplaySSM reference path and is performance-critical: the L-deep
|
||||
# reconstruction is L x the baseline's rank-1 compute, which only stays
|
||||
# hidden under the (reduced) memory traffic if it runs on tensor cores. An
|
||||
# IEEE fp32 dot here disables tensor cores and makes the kernel ~10-20x
|
||||
# slower. Tensor-core precision (~4e-4 TF32 / ~1e-3 bf16) is benign
|
||||
# end-to-end (ReplaySSM bf16 GSM8K parity); the unit test uses
|
||||
# tensor-core-realistic tolerances. At L=1 the buffer is empty so this dot
|
||||
# is identically zero and the path stays bit-exact regardless of precision.
|
||||
# GDN folds the (scalar) replay decay onto d here; KDA folds the (per-K)
|
||||
# replay decay onto the cached keys inside the K-tile loop instead.
|
||||
if not IS_KDA:
|
||||
b_d_tc = (b_d_all * b_replay_decay[None, :]).to(
|
||||
p_o.dtype.element_ty
|
||||
) # [BV, BC]
|
||||
else:
|
||||
b_d_tc = b_d_all.to(p_o.dtype.element_ty) # [BV, BC]
|
||||
|
||||
# Current token value (for the delta-rule update).
|
||||
v_off = (2 * H * K) + i_hv * V + o_v
|
||||
b_v = tl.load(
|
||||
mixed_qkv + i_n * stride_mixed_qkv_tok + v_off, mask=mask_v, other=0
|
||||
).to(tl.float32)
|
||||
|
||||
# Optional q/k L2 norm: full-vector reciprocal norms (computed, not kept).
|
||||
if USE_QK_L2NORM_IN_KERNEL:
|
||||
o_kf = tl.arange(0, BK)
|
||||
mask_kf = o_kf < K
|
||||
p_mix = mixed_qkv + i_n * stride_mixed_qkv_tok
|
||||
qf = tl.load(p_mix + i_h * K + o_kf, mask=mask_kf, other=0).to(tl.float32)
|
||||
kf = tl.load(p_mix + H * K + i_h * K + o_kf, mask=mask_kf, other=0).to(
|
||||
tl.float32
|
||||
)
|
||||
q_rnorm = 1.0 / tl.sqrt(tl.sum(qf * qf) + 1e-6)
|
||||
k_rnorm = 1.0 / tl.sqrt(tl.sum(kf * kf) + 1e-6)
|
||||
else:
|
||||
q_rnorm = 1.0
|
||||
k_rnorm = 1.0
|
||||
|
||||
# Reconstruct S from the checkpoint + cached (d, k) in K-tiles and read it
|
||||
# with the current (scaled) q and k. K-tiling keeps the per-program tile
|
||||
# small so the full [V, K] state is never materialized. Also append the
|
||||
# current key chunk to the ring cache (non-flush only).
|
||||
b_state_q = tl.zeros([BV], dtype=tl.float32)
|
||||
b_state_k = tl.zeros([BV], dtype=tl.float32)
|
||||
cur_kq = tl.zeros([1], dtype=tl.float32)
|
||||
write_k = (not b_is_flush) and (i_v == 0) and (i_hv == i_h * (HV // H))
|
||||
write_g_kda = IS_KDA and (not b_is_flush) and (i_v == 0)
|
||||
for kk in range(NK):
|
||||
o_kt = kk * BKT + tl.arange(0, BKT)
|
||||
mask_kt = o_kt < K
|
||||
p_mix = mixed_qkv + i_n * stride_mixed_qkv_tok
|
||||
q_c = (
|
||||
tl.load(p_mix + i_h * K + o_kt, mask=mask_kt, other=0).to(tl.float32)
|
||||
* q_rnorm
|
||||
)
|
||||
k_c = (
|
||||
tl.load(p_mix + H * K + i_h * K + o_kt, mask=mask_kt, other=0).to(
|
||||
tl.float32
|
||||
)
|
||||
* k_rnorm
|
||||
)
|
||||
q_cs = q_c * scale
|
||||
# Rank-1 output term uses the RAW current k/q (gate-independent).
|
||||
cur_kq += tl.sum(k_c * q_cs)
|
||||
|
||||
# This K-tile of the state: S_tile = Diag(A_tile) S0_tile + d (.) (W (.) k_cache).
|
||||
p_h0_c = (
|
||||
h0
|
||||
+ state_idx * stride_init_state_token
|
||||
+ i_hv * V * K
|
||||
+ o_v[:, None] * K
|
||||
+ o_kt[None, :]
|
||||
)
|
||||
b_h0_c = tl.load(p_h0_c, mask=mask_v[:, None] & mask_kt[None, :], other=0).to(
|
||||
tl.float32
|
||||
)
|
||||
p_k_c = (
|
||||
k_cache
|
||||
+ ((state_idx * H + i_h) * MAX_CACHE_LEN + o_c[:, None]) * K
|
||||
+ o_kt[None, :]
|
||||
)
|
||||
|
||||
if not IS_KDA:
|
||||
b_k_all_c = tl.load(
|
||||
p_k_c, mask=cache_valid[:, None] & mask_kt[None, :], other=0
|
||||
).to(p_o.dtype.element_ty)
|
||||
b_h_c = b_h0_c * b_total_decay + tl.dot(b_d_tc, b_k_all_c).to(tl.float32)
|
||||
# GDN: scalar current-token decay applied after the loop.
|
||||
q_eff = q_cs
|
||||
k_eff = k_c
|
||||
else:
|
||||
# KDA per-K decay: load this tile's cached gates [BC, BKT], form the
|
||||
# per-K total / replay decay, fold the replay decay onto the cached
|
||||
# keys and the total decay onto S0.
|
||||
p_g_c = (
|
||||
g_cache
|
||||
+ ((state_idx * HV + i_hv) * MAX_CACHE_LEN + o_c[:, None]) * K
|
||||
+ o_kt[None, :]
|
||||
)
|
||||
b_g_all_c = tl.load(
|
||||
p_g_c, mask=cache_valid[:, None] & mask_kt[None, :], other=0.0
|
||||
).to(tl.float32)
|
||||
b_g_prefix_c = tl.cumsum(b_g_all_c, axis=0) # [BC, BKT]
|
||||
b_g_total_c = tl.sum(b_g_all_c, axis=0) # [BKT]
|
||||
b_replay_decay_c = tl.where(
|
||||
cache_valid[:, None],
|
||||
tl.exp(b_g_total_c[None, :] - b_g_prefix_c),
|
||||
0.0,
|
||||
) # [BC, BKT]
|
||||
b_total_decay_c = tl.exp(b_g_total_c) # [BKT]
|
||||
b_k_all_c = tl.load(
|
||||
p_k_c, mask=cache_valid[:, None] & mask_kt[None, :], other=0.0
|
||||
).to(tl.float32)
|
||||
b_k_scaled = (b_k_all_c * b_replay_decay_c).to(p_o.dtype.element_ty)
|
||||
b_h_c = b_h0_c * b_total_decay_c[None, :] + tl.dot(b_d_tc, b_k_scaled).to(
|
||||
tl.float32
|
||||
)
|
||||
# KDA: current-token per-K decay folds into q/k for the readout.
|
||||
p_a_c = a + i_n * stride_a_tok + i_hv * K + o_kt
|
||||
p_dt_c = dt_bias + i_hv * K + o_kt
|
||||
b_a_c = tl.load(p_a_c, mask=mask_kt, other=0.0).to(tl.float32)
|
||||
b_dt_c = tl.load(p_dt_c, mask=mask_kt, other=0.0).to(tl.float32)
|
||||
x_c = b_a_c + b_dt_c
|
||||
softplus_c = tl.where(
|
||||
x_c <= SOFTPLUS_THRESHOLD, tl.log(1.0 + tl.exp(x_c)), x_c
|
||||
)
|
||||
g_cur_c = -tl.exp(A_log_val) * softplus_c # [BKT]
|
||||
alpha_cur_c = tl.exp(g_cur_c)
|
||||
q_eff = q_cs * alpha_cur_c
|
||||
k_eff = k_c * alpha_cur_c
|
||||
|
||||
# Append this tile of the current gate to the ring (non-flush only).
|
||||
if write_g_kda:
|
||||
p_cur_g = (
|
||||
g_cache
|
||||
+ ((state_idx * HV + i_hv) * MAX_CACHE_LEN + b_write_pos) * K
|
||||
+ o_kt
|
||||
)
|
||||
tl.store(p_cur_g, g_cur_c, mask=mask_kt & (b_write_pos < MAX_CACHE_LEN))
|
||||
|
||||
# Read the state with the (gate-folded) q and k, accumulated across tiles.
|
||||
b_state_q += tl.sum(b_h_c * q_eff[None, :], axis=1)
|
||||
b_state_k += tl.sum(b_h_c * k_eff[None, :], axis=1)
|
||||
|
||||
if write_k:
|
||||
p_cur_k = (
|
||||
k_cache
|
||||
+ ((state_idx * H + i_h) * MAX_CACHE_LEN + b_write_pos) * K
|
||||
+ o_kt
|
||||
)
|
||||
tl.store(
|
||||
p_cur_k,
|
||||
k_c.to(p_o.dtype.element_ty),
|
||||
mask=mask_kt & (b_write_pos < MAX_CACHE_LEN),
|
||||
)
|
||||
|
||||
# Current-token output: (S . Diag(a)) q + d_cur*(k . q), with the new
|
||||
# corrected delta-rule vector d_cur = beta * (v - (S . Diag(a)) k).
|
||||
# For GDN the per-head scalar decay is applied here; for KDA it was already
|
||||
# folded into q/k above.
|
||||
if not IS_KDA:
|
||||
b_state_q *= alpha_val
|
||||
b_state_k *= alpha_val
|
||||
b_d_cur = beta_val * (b_v - b_state_k)
|
||||
b_o = b_state_q + b_d_cur * tl.sum(cur_kq)
|
||||
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v)
|
||||
|
||||
if b_is_flush:
|
||||
# Flush: fold the current token into the checkpoint, S_new = S.Diag(a) +
|
||||
# d_cur k^T, and persist it. Re-walk K chunks to rebuild S, then apply
|
||||
# the update. After this the ring is logically cleared (the caller
|
||||
# resets write_pos to 0 on the next step).
|
||||
for kk in range(NK):
|
||||
o_kt = kk * BKT + tl.arange(0, BKT)
|
||||
mask_kt = o_kt < K
|
||||
p_mix = mixed_qkv + i_n * stride_mixed_qkv_tok
|
||||
k_c = (
|
||||
tl.load(p_mix + H * K + i_h * K + o_kt, mask=mask_kt, other=0).to(
|
||||
tl.float32
|
||||
)
|
||||
* k_rnorm
|
||||
)
|
||||
p_h0_c = (
|
||||
h0
|
||||
+ state_idx * stride_init_state_token
|
||||
+ i_hv * V * K
|
||||
+ o_v[:, None] * K
|
||||
+ o_kt[None, :]
|
||||
)
|
||||
b_h0_c = tl.load(
|
||||
p_h0_c, mask=mask_v[:, None] & mask_kt[None, :], other=0
|
||||
).to(tl.float32)
|
||||
p_k_c = (
|
||||
k_cache
|
||||
+ ((state_idx * H + i_h) * MAX_CACHE_LEN + o_c[:, None]) * K
|
||||
+ o_kt[None, :]
|
||||
)
|
||||
if not IS_KDA:
|
||||
b_k_all_c = tl.load(
|
||||
p_k_c, mask=cache_valid[:, None] & mask_kt[None, :], other=0
|
||||
).to(p_o.dtype.element_ty)
|
||||
b_h_c = b_h0_c * b_total_decay + tl.dot(b_d_tc, b_k_all_c).to(
|
||||
tl.float32
|
||||
)
|
||||
b_h_new_c = alpha_val * b_h_c + b_d_cur[:, None] * k_c[None, :]
|
||||
else:
|
||||
p_g_c = (
|
||||
g_cache
|
||||
+ ((state_idx * HV + i_hv) * MAX_CACHE_LEN + o_c[:, None]) * K
|
||||
+ o_kt[None, :]
|
||||
)
|
||||
b_g_all_c = tl.load(
|
||||
p_g_c, mask=cache_valid[:, None] & mask_kt[None, :], other=0.0
|
||||
).to(tl.float32)
|
||||
b_g_prefix_c = tl.cumsum(b_g_all_c, axis=0)
|
||||
b_g_total_c = tl.sum(b_g_all_c, axis=0)
|
||||
b_replay_decay_c = tl.where(
|
||||
cache_valid[:, None],
|
||||
tl.exp(b_g_total_c[None, :] - b_g_prefix_c),
|
||||
0.0,
|
||||
)
|
||||
b_total_decay_c = tl.exp(b_g_total_c)
|
||||
b_k_all_c = tl.load(
|
||||
p_k_c, mask=cache_valid[:, None] & mask_kt[None, :], other=0.0
|
||||
).to(tl.float32)
|
||||
b_k_scaled = (b_k_all_c * b_replay_decay_c).to(p_o.dtype.element_ty)
|
||||
b_h_c = b_h0_c * b_total_decay_c[None, :] + tl.dot(
|
||||
b_d_tc, b_k_scaled
|
||||
).to(tl.float32)
|
||||
p_a_c = a + i_n * stride_a_tok + i_hv * K + o_kt
|
||||
p_dt_c = dt_bias + i_hv * K + o_kt
|
||||
b_a_c = tl.load(p_a_c, mask=mask_kt, other=0.0).to(tl.float32)
|
||||
b_dt_c = tl.load(p_dt_c, mask=mask_kt, other=0.0).to(tl.float32)
|
||||
x_c = b_a_c + b_dt_c
|
||||
softplus_c = tl.where(
|
||||
x_c <= SOFTPLUS_THRESHOLD, tl.log(1.0 + tl.exp(x_c)), x_c
|
||||
)
|
||||
alpha_cur_c = tl.exp(-tl.exp(A_log_val) * softplus_c) # [BKT]
|
||||
b_h_new_c = (
|
||||
b_h_c * alpha_cur_c[None, :] + b_d_cur[:, None] * k_c[None, :]
|
||||
)
|
||||
p_ht_c = (
|
||||
ht
|
||||
+ state_idx * stride_final_state_token
|
||||
+ i_hv * V * K
|
||||
+ o_v[:, None] * K
|
||||
+ o_kt[None, :]
|
||||
)
|
||||
tl.store(
|
||||
p_ht_c,
|
||||
b_h_new_c.to(p_ht_c.dtype.element_ty),
|
||||
mask=mask_v[:, None] & mask_kt[None, :],
|
||||
)
|
||||
else:
|
||||
# Non-flush: append the current token's corrected delta d to the cache
|
||||
# (k chunks were written inside the loop; KDA's g chunks too). GDN's
|
||||
# scalar g is appended here.
|
||||
p_cur_d = (
|
||||
d_cache + ((state_idx * HV + i_hv) * MAX_CACHE_LEN + b_write_pos) * V + o_v
|
||||
)
|
||||
tl.store(
|
||||
p_cur_d,
|
||||
b_d_cur.to(p_cur_d.dtype.element_ty),
|
||||
mask=mask_v & (b_write_pos < MAX_CACHE_LEN),
|
||||
)
|
||||
if (not IS_KDA) and (i_v == 0):
|
||||
p_cur_g = g_cache + (state_idx * HV + i_hv) * MAX_CACHE_LEN + b_write_pos
|
||||
tl.store(p_cur_g, g_val, mask=b_write_pos < MAX_CACHE_LEN)
|
||||
|
||||
|
||||
def fused_recurrent_linear_replayssm_decode(
|
||||
mixed_qkv: torch.Tensor,
|
||||
a: torch.Tensor,
|
||||
b: torch.Tensor,
|
||||
A_log: torch.Tensor,
|
||||
dt_bias: torch.Tensor,
|
||||
scale: float,
|
||||
initial_state: torch.Tensor,
|
||||
d_cache: torch.Tensor,
|
||||
k_cache: torch.Tensor,
|
||||
g_cache: torch.Tensor,
|
||||
out: torch.Tensor,
|
||||
ssm_state_indices: torch.Tensor,
|
||||
write_pos: torch.Tensor,
|
||||
force_flush: torch.Tensor | None = None,
|
||||
use_qk_l2norm_in_kernel: bool = False,
|
||||
is_kda: bool = False,
|
||||
block_v: int | None = None,
|
||||
num_warps: int = 1,
|
||||
num_stages: int = 3,
|
||||
nk: int = 2,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Buffered output-only linear-attention autoregressive decode (1 token/seq).
|
||||
|
||||
One kernel for both gate granularities, selected by ``is_kda``:
|
||||
* ``is_kda=False`` (GDN): per-head SCALAR gate. ``a``=[B, HV],
|
||||
``dt_bias``=[HV], ``g_cache``=[num_slots, HV, L].
|
||||
* ``is_kda=True`` (KDA): per-K-channel gate. ``a``=[B, HV, K],
|
||||
``dt_bias``=[HV, K], ``g_cache``=[num_slots, HV, L, K].
|
||||
``A_log`` is [HV] (per-head scalar) for both.
|
||||
|
||||
Same call surface as the packed decode plus the three ring caches
|
||||
(``d_cache`` / ``k_cache`` / ``g_cache``) and the per-decode-row
|
||||
``write_pos`` cursor. ``initial_state`` is both the checkpoint read (h0)
|
||||
and the (flush-only) checkpoint write (ht), in place.
|
||||
|
||||
Allocates nothing persistent: the caller owns the ring tensors and is
|
||||
responsible for advancing / resetting ``write_pos`` (e.g. ``(write_pos+1) %
|
||||
L`` after each step). This is a STANDALONE kernel; the memory-pool / cache
|
||||
integration is a later phase.
|
||||
"""
|
||||
if mixed_qkv.ndim != 2:
|
||||
raise ValueError(f"`mixed_qkv` must be 2D (got ndim={mixed_qkv.ndim}).")
|
||||
if mixed_qkv.stride(-1) != 1:
|
||||
raise ValueError("`mixed_qkv` must be contiguous in the last dim.")
|
||||
if b.ndim != 2:
|
||||
raise ValueError(f"`b` must be 2D (got b.ndim={b.ndim}).")
|
||||
if A_log.ndim != 1:
|
||||
raise ValueError("`A_log` must be a 1D tensor.")
|
||||
if initial_state.ndim != 4:
|
||||
raise ValueError(f"`initial_state` must be 4D (got ndim={initial_state.ndim}).")
|
||||
if not out.is_contiguous():
|
||||
raise ValueError("`out` must be contiguous.")
|
||||
if write_pos.ndim != 1 or write_pos.dtype != torch.int32:
|
||||
raise ValueError("`write_pos` must be a 1D int32 tensor.")
|
||||
if force_flush is not None and (
|
||||
force_flush.ndim != 1 or force_flush.dtype != torch.int32
|
||||
):
|
||||
raise ValueError("`force_flush` must be a 1D int32 tensor or None.")
|
||||
|
||||
B = mixed_qkv.shape[0]
|
||||
num_state_slots, HV, V, K = initial_state.shape
|
||||
qkv_dim = mixed_qkv.shape[1]
|
||||
q_dim = (qkv_dim - HV * V) // 2
|
||||
if q_dim <= 0 or q_dim % K != 0:
|
||||
raise ValueError(
|
||||
f"Invalid packed `mixed_qkv` last dim={qkv_dim} for HV={HV}, V={V}, K={K}."
|
||||
)
|
||||
H = q_dim // K
|
||||
if H <= 0 or HV % H != 0:
|
||||
raise ValueError(
|
||||
f"Invalid head config inferred from mixed_qkv: H={H}, HV={HV}."
|
||||
)
|
||||
max_cache_len = d_cache.shape[2]
|
||||
|
||||
# Gate-shape sanity: GDN scalar gate vs KDA per-K gate.
|
||||
if is_kda:
|
||||
if a.ndim != 3 or tuple(a.shape) != (B, HV, K):
|
||||
raise ValueError(
|
||||
f"KDA `a` must have shape {(B, HV, K)} (got {tuple(a.shape)})."
|
||||
)
|
||||
if dt_bias.ndim != 2 or tuple(dt_bias.shape) != (HV, K):
|
||||
raise ValueError(
|
||||
f"KDA `dt_bias` must have shape {(HV, K)} (got {tuple(dt_bias.shape)})."
|
||||
)
|
||||
if not a.is_contiguous() or not dt_bias.is_contiguous():
|
||||
raise ValueError("KDA `a`/`dt_bias` must be contiguous.")
|
||||
g_expect = (HV, max_cache_len, K)
|
||||
else:
|
||||
if a.ndim != 2 or tuple(a.shape) != (B, HV):
|
||||
raise ValueError(
|
||||
f"GDN `a` must have shape {(B, HV)} (got {tuple(a.shape)})."
|
||||
)
|
||||
if dt_bias.ndim != 1 or dt_bias.shape[0] != HV:
|
||||
raise ValueError(
|
||||
f"GDN `dt_bias` must have shape {(HV,)} (got {tuple(dt_bias.shape)})."
|
||||
)
|
||||
g_expect = (HV, max_cache_len)
|
||||
|
||||
# Cache shape sanity (per state slot): d=(HV, L, V), k=(H, L, K).
|
||||
if tuple(d_cache.shape[1:]) != (HV, max_cache_len, V):
|
||||
raise ValueError(
|
||||
f"`d_cache` per-slot shape must be {(HV, max_cache_len, V)} "
|
||||
f"(got {tuple(d_cache.shape[1:])})."
|
||||
)
|
||||
if tuple(k_cache.shape[1:]) != (H, max_cache_len, K):
|
||||
raise ValueError(
|
||||
f"`k_cache` per-slot shape must be {(H, max_cache_len, K)} "
|
||||
f"(got {tuple(k_cache.shape[1:])})."
|
||||
)
|
||||
if tuple(g_cache.shape[1:]) != g_expect:
|
||||
raise ValueError(
|
||||
f"`g_cache` per-slot shape must be {g_expect} "
|
||||
f"(got {tuple(g_cache.shape[1:])})."
|
||||
)
|
||||
if g_cache.dtype != torch.float32:
|
||||
raise ValueError(f"`g_cache` must be float32 (got {g_cache.dtype}).")
|
||||
if out.shape != (B, 1, HV, V):
|
||||
raise ValueError(
|
||||
f"`out` must have shape {(B, 1, HV, V)} (got {tuple(out.shape)})."
|
||||
)
|
||||
if write_pos.shape[0] != B or ssm_state_indices.shape[0] != B:
|
||||
raise ValueError(
|
||||
"`write_pos` and `ssm_state_indices` must both have length B="
|
||||
f"{B} (got {write_pos.shape[0]}, {ssm_state_indices.shape[0]})."
|
||||
)
|
||||
|
||||
BK = triton.next_power_of_2(K)
|
||||
if triton.cdiv(K, BK) != 1:
|
||||
raise ValueError(
|
||||
f"Cached decode kernel only supports NK_global=1 (got K={K}, BK={BK})."
|
||||
)
|
||||
if BK % nk != 0:
|
||||
raise ValueError(f"nk={nk} must divide BK={BK}.")
|
||||
BKT = BK // nk
|
||||
if BKT < 16:
|
||||
raise ValueError(f"BKT={BKT} must be >=16 for tl.dot (nk={nk}, BK={BK}).")
|
||||
# K-tiling keeps the per-program tile small enough that a larger BV fits
|
||||
# without register spilling -> NV=1 -> half the grid -> fewer redundant
|
||||
# cache / metadata loads.
|
||||
BV = block_v if block_v is not None else min(triton.next_power_of_2(V), 64)
|
||||
BC = max(16, triton.next_power_of_2(max_cache_len))
|
||||
|
||||
grid = (triton.cdiv(V, BV), B, HV)
|
||||
fused_recurrent_linear_replayssm_decode_kernel[grid](
|
||||
mixed_qkv=mixed_qkv,
|
||||
a=a,
|
||||
b=b,
|
||||
A_log=A_log,
|
||||
dt_bias=dt_bias,
|
||||
o=out,
|
||||
h0=initial_state,
|
||||
ht=initial_state,
|
||||
d_cache=d_cache,
|
||||
k_cache=k_cache,
|
||||
g_cache=g_cache,
|
||||
ssm_state_indices=ssm_state_indices,
|
||||
write_pos=write_pos,
|
||||
force_flush=force_flush if force_flush is not None else write_pos,
|
||||
scale=scale,
|
||||
stride_mixed_qkv_tok=mixed_qkv.stride(0),
|
||||
stride_a_tok=a.stride(0),
|
||||
stride_b_tok=b.stride(0),
|
||||
stride_init_state_token=initial_state.stride(0),
|
||||
stride_final_state_token=initial_state.stride(0),
|
||||
stride_indices_seq=ssm_state_indices.stride(0),
|
||||
H=H,
|
||||
HV=HV,
|
||||
K=K,
|
||||
V=V,
|
||||
BK=BK,
|
||||
BV=BV,
|
||||
BC=BC,
|
||||
NK=nk,
|
||||
BKT=BKT,
|
||||
MAX_CACHE_LEN=max_cache_len,
|
||||
SOFTPLUS_THRESHOLD=20.0,
|
||||
USE_QK_L2NORM_IN_KERNEL=use_qk_l2norm_in_kernel,
|
||||
HAS_FORCE_FLUSH=force_flush is not None,
|
||||
IS_KDA=is_kda,
|
||||
num_warps=num_warps,
|
||||
num_stages=num_stages,
|
||||
)
|
||||
return out, initial_state
|
||||
|
||||
|
||||
# Backwards-compatible aliases: the original GDN-only names. Existing callers
|
||||
# (backend dispatch, tests, microbench) keep working; ``is_kda`` defaults to
|
||||
# False so these are the GDN path unchanged.
|
||||
fused_recurrent_gdn_replayssm_decode_kernel = (
|
||||
fused_recurrent_linear_replayssm_decode_kernel
|
||||
)
|
||||
fused_recurrent_gdn_replayssm_decode = fused_recurrent_linear_replayssm_decode
|
||||
@@ -0,0 +1,366 @@
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def fused_sigmoid_gating_delta_rule_update_kernel(
|
||||
A_log,
|
||||
a,
|
||||
dt_bias,
|
||||
softplus_beta,
|
||||
softplus_threshold,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
b,
|
||||
o,
|
||||
h0_source,
|
||||
h0_indices,
|
||||
cu_seqlens,
|
||||
# Parameters for target_verify support (unused for decode)
|
||||
intermediate_states_buffer,
|
||||
intermediate_state_indices,
|
||||
cache_steps,
|
||||
retrieve_parent_token_ptr,
|
||||
stride_retrieve_parent_token_seq: tl.constexpr,
|
||||
stride_retrieve_parent_token_token: tl.constexpr,
|
||||
# ================================================
|
||||
scale,
|
||||
T,
|
||||
stride_a,
|
||||
stride_q,
|
||||
stride_k,
|
||||
stride_v,
|
||||
stride_b,
|
||||
NP2_T: tl.constexpr,
|
||||
B: tl.constexpr,
|
||||
H: tl.constexpr,
|
||||
HV: tl.constexpr,
|
||||
K: tl.constexpr,
|
||||
V: tl.constexpr,
|
||||
BK: tl.constexpr,
|
||||
BV: tl.constexpr,
|
||||
USE_INITIAL_STATE: tl.constexpr,
|
||||
USE_QK_L2NORM_IN_KERNEL: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
IS_KDA: tl.constexpr,
|
||||
# Optional flags for target_verify support (default False for decode)
|
||||
DISABLE_STATE_UPDATE: tl.constexpr = False,
|
||||
CACHE_INTERMEDIATE_STATES: tl.constexpr = False,
|
||||
HAS_EAGLE_TREE_CUSTOM_ATTN_MASK: tl.constexpr = False,
|
||||
):
|
||||
"""
|
||||
Fused kernel that combines sigmoid gating computation with recurrent delta rule update.
|
||||
"""
|
||||
i_k, i_v, i_nh = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
i_n, i_hv = i_nh // HV, i_nh % HV
|
||||
i_h = i_hv // (HV // H)
|
||||
|
||||
if IS_VARLEN:
|
||||
bos, eos = (
|
||||
tl.load(cu_seqlens + i_n).to(tl.int64),
|
||||
tl.load(cu_seqlens + i_n + 1).to(tl.int64),
|
||||
)
|
||||
all = T
|
||||
T = eos - bos
|
||||
else:
|
||||
bos, eos = i_n * T, i_n * T + T
|
||||
all = B * T
|
||||
|
||||
o_k = i_k * BK + tl.arange(0, BK)
|
||||
o_v = i_v * BV + tl.arange(0, BV)
|
||||
|
||||
p_q = q + bos * stride_q + i_h * K + o_k
|
||||
p_k = k + bos * stride_k + i_h * K + o_k
|
||||
p_v = v + bos * stride_v + i_hv * V + o_v
|
||||
p_b = b + bos * stride_b + i_hv
|
||||
p_o = o + ((i_k * all + bos) * HV + i_hv) * V + o_v
|
||||
|
||||
# Gating computation pointers
|
||||
p_A_log = A_log + i_hv
|
||||
if IS_KDA:
|
||||
p_a = a + bos * stride_a + i_hv * K + o_k
|
||||
p_dt_bias = dt_bias + i_hv * K + o_k
|
||||
else:
|
||||
p_a = a + bos * stride_a + i_hv
|
||||
p_dt_bias = dt_bias + i_hv
|
||||
|
||||
mask_k = o_k < K
|
||||
mask_v = o_v < V
|
||||
mask_h = mask_k[:, None] & mask_v[None, :]
|
||||
|
||||
b_h = tl.zeros([BK, BV], dtype=tl.float32)
|
||||
if USE_INITIAL_STATE:
|
||||
idx = tl.load(h0_indices + i_n)
|
||||
if idx >= 0:
|
||||
p_h0 = (
|
||||
h0_source
|
||||
+ idx * HV * K * V
|
||||
+ i_hv * K * V
|
||||
+ o_v[None, :] * K
|
||||
+ o_k[:, None]
|
||||
)
|
||||
b_h += tl.load(p_h0, mask=mask_h, other=0).to(tl.float32)
|
||||
|
||||
# Preload tree attention data if needed
|
||||
if HAS_EAGLE_TREE_CUSTOM_ATTN_MASK:
|
||||
token_indices = tl.arange(0, NP2_T)
|
||||
mask_retrieve = token_indices < T
|
||||
retrieve_parent_token_base = (
|
||||
retrieve_parent_token_ptr
|
||||
+ (i_n * stride_retrieve_parent_token_seq)
|
||||
+ token_indices * stride_retrieve_parent_token_token
|
||||
)
|
||||
parent_idx_tokens = tl.load(
|
||||
retrieve_parent_token_base, mask=mask_retrieve, other=0
|
||||
)
|
||||
|
||||
# Prepare intermediate state cache index if enabled
|
||||
cache_idx = -1
|
||||
if CACHE_INTERMEDIATE_STATES:
|
||||
cache_idx = tl.load(intermediate_state_indices + i_n)
|
||||
|
||||
step_idx = 0
|
||||
for _ in range(0, T):
|
||||
# Tree attention: load parent's cached state
|
||||
if HAS_EAGLE_TREE_CUSTOM_ATTN_MASK:
|
||||
# step_idx == 0 uses b_h from USE_INITIAL_STATE
|
||||
if step_idx != 0 and cache_idx >= 0:
|
||||
parent_step_idx = tl.sum(
|
||||
tl.where(token_indices == step_idx, parent_idx_tokens, 0)
|
||||
)
|
||||
step_offset = parent_step_idx * HV * K * V
|
||||
cache_ptr = (
|
||||
intermediate_states_buffer
|
||||
+ cache_idx * cache_steps * HV * K * V
|
||||
+ step_offset
|
||||
+ i_hv * K * V
|
||||
+ o_v[None, :] * K
|
||||
+ o_k[:, None]
|
||||
)
|
||||
b_h = tl.load(cache_ptr, mask=mask_h, other=0).to(tl.float32)
|
||||
|
||||
# Load inputs
|
||||
b_q = tl.load(p_q, mask=mask_k, other=0).to(tl.float32)
|
||||
b_k = tl.load(p_k, mask=mask_k, other=0).to(tl.float32)
|
||||
b_v = tl.load(p_v, mask=mask_v, other=0).to(tl.float32)
|
||||
b_b = tl.load(p_b).to(tl.float32)
|
||||
|
||||
# Compute sigmoid gating
|
||||
# Load gating parameters
|
||||
b_A_log = tl.load(p_A_log).to(tl.float32)
|
||||
if IS_KDA:
|
||||
b_a = tl.load(p_a, mask=mask_k, other=0).to(tl.float32)
|
||||
b_dt_bias = tl.load(p_dt_bias, mask=mask_k, other=0).to(tl.float32)
|
||||
else:
|
||||
b_a = tl.load(p_a).to(tl.float32)
|
||||
b_dt_bias = tl.load(p_dt_bias).to(tl.float32)
|
||||
|
||||
# Compute g = -exp(A_log) * softplus(a + dt_bias)
|
||||
x = b_a + b_dt_bias
|
||||
beta_x = softplus_beta * x
|
||||
# Apply softplus with numerical stability
|
||||
softplus_x = tl.where(
|
||||
beta_x <= softplus_threshold,
|
||||
(1.0 / softplus_beta) * tl.log(1.0 + tl.exp(beta_x)),
|
||||
x,
|
||||
)
|
||||
b_g = -tl.exp(b_A_log) * softplus_x
|
||||
|
||||
# Compute beta = sigmoid(b)
|
||||
b_beta = 1.0 / (1.0 + tl.exp(-b_b))
|
||||
|
||||
# Apply L2 normalization if enabled
|
||||
if USE_QK_L2NORM_IN_KERNEL:
|
||||
b_q = b_q / (tl.sqrt(tl.sum(b_q * b_q) + 1e-6))
|
||||
b_k = b_k / (tl.sqrt(tl.sum(b_k * b_k) + 1e-6))
|
||||
|
||||
b_q = b_q * scale
|
||||
|
||||
# Apply gating to hidden state: h *= exp(g)
|
||||
if IS_KDA:
|
||||
b_h *= tl.exp(b_g[:, None])
|
||||
else:
|
||||
b_h *= tl.exp(b_g)
|
||||
|
||||
# Delta rule: v -= sum(h * k, dim=0)
|
||||
b_v -= tl.sum(b_h * b_k[:, None], 0)
|
||||
|
||||
# Apply beta gating: v *= beta
|
||||
b_v *= b_beta
|
||||
|
||||
# Update hidden state: h += k[:, None] * v[None, :]
|
||||
b_h += b_k[:, None] * b_v[None, :]
|
||||
|
||||
# Compute output: o = sum(h * q, dim=0)
|
||||
b_o = tl.sum(b_h * b_q[:, None], 0)
|
||||
tl.store(p_o, b_o.to(p_o.dtype.element_ty), mask=mask_v)
|
||||
|
||||
# Cache intermediate states if enabled
|
||||
if CACHE_INTERMEDIATE_STATES:
|
||||
if cache_idx >= 0:
|
||||
step_offset = step_idx * HV * K * V
|
||||
cache_ptr = (
|
||||
intermediate_states_buffer
|
||||
+ cache_idx * cache_steps * HV * K * V
|
||||
+ step_offset
|
||||
+ i_hv * K * V
|
||||
+ o_v[None, :] * K
|
||||
+ o_k[:, None]
|
||||
)
|
||||
tl.store(cache_ptr, b_h.to(cache_ptr.dtype.element_ty), mask=mask_h)
|
||||
|
||||
step_idx += 1
|
||||
|
||||
# Update pointers for next timestep
|
||||
p_q += stride_q
|
||||
p_k += stride_k
|
||||
p_v += stride_v
|
||||
p_b += stride_b
|
||||
p_o += HV * V
|
||||
p_a += stride_a
|
||||
|
||||
# Store final state back to h0_source with bounds checking
|
||||
if not DISABLE_STATE_UPDATE:
|
||||
if USE_INITIAL_STATE:
|
||||
idx = tl.load(h0_indices + i_n)
|
||||
if idx >= 0:
|
||||
p_h0 = (
|
||||
h0_source
|
||||
+ idx * HV * K * V
|
||||
+ i_hv * K * V
|
||||
+ o_v[None, :] * K
|
||||
+ o_k[:, None]
|
||||
)
|
||||
tl.store(p_h0, b_h.to(p_h0.dtype.element_ty), mask=mask_h)
|
||||
|
||||
|
||||
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, # kept for API compat; stride is derived from ``intermediate_states_buffer.shape[1]``
|
||||
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), 32)
|
||||
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)
|
||||
|
||||
# Per-req stride must match the buffer's allocated dim, not runtime steps
|
||||
# (they can differ under --speculative-adaptive).
|
||||
cache_stride_steps = (
|
||||
intermediate_states_buffer.shape[1]
|
||||
if intermediate_states_buffer is not None
|
||||
else 0
|
||||
)
|
||||
|
||||
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=cache_stride_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
|
||||
@@ -0,0 +1,35 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/utils/index.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.srt.layers.attention.fla.utils import tensor_cache
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_lens(cu_seqlens: torch.LongTensor) -> torch.LongTensor:
|
||||
return cu_seqlens[1:] - cu_seqlens[:-1]
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_chunk_indices(
|
||||
cu_seqlens: torch.LongTensor, chunk_size: int
|
||||
) -> torch.LongTensor:
|
||||
indices = torch.cat(
|
||||
[
|
||||
torch.arange(n)
|
||||
for n in triton.cdiv(prepare_lens(cu_seqlens), chunk_size).tolist()
|
||||
]
|
||||
)
|
||||
return torch.stack([indices.eq(0).cumsum(0) - 1, indices], 1).to(cu_seqlens)
|
||||
|
||||
|
||||
@tensor_cache
|
||||
def prepare_chunk_offsets(
|
||||
cu_seqlens: torch.LongTensor, chunk_size: int
|
||||
) -> torch.LongTensor:
|
||||
return torch.cat(
|
||||
[cu_seqlens.new_tensor([0]), triton.cdiv(prepare_lens(cu_seqlens), chunk_size)]
|
||||
).cumsum(-1)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,150 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/modules/l2norm.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.layers.attention.fla.utils import input_guard
|
||||
|
||||
BT_LIST = [8, 16, 32, 64, 128]
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[
|
||||
# triton.Config({}, num_warps=num_warps) for num_warps in [1, 2, 4, 8, 16, 32]
|
||||
# ],
|
||||
# key=["D"],
|
||||
# )
|
||||
@triton.jit
|
||||
def l2norm_fwd_kernel1(
|
||||
x,
|
||||
y,
|
||||
D,
|
||||
BD: tl.constexpr,
|
||||
eps,
|
||||
):
|
||||
i_t = tl.program_id(0)
|
||||
x += i_t * D
|
||||
y += i_t * D
|
||||
# Compute mean and variance
|
||||
cols = tl.arange(0, BD)
|
||||
mask = cols < D
|
||||
b_x = tl.load(x + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
b_var = tl.sum(b_x * b_x, axis=0)
|
||||
b_rstd = 1 / tl.sqrt(b_var + eps)
|
||||
# tl.store(Rstd + i_t, rstd)
|
||||
# Normalize and apply linear transformation
|
||||
b_y = b_x * b_rstd
|
||||
tl.store(y + cols, b_y, mask=mask)
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[
|
||||
# triton.Config({"BT": BT}, num_warps=num_warps)
|
||||
# for num_warps in [1, 2, 4, 8, 16]
|
||||
# for BT in BT_LIST
|
||||
# ],
|
||||
# key=["D", "NB"],
|
||||
# )
|
||||
@triton.jit
|
||||
def l2norm_fwd_kernel(
|
||||
x,
|
||||
y,
|
||||
eps,
|
||||
NB: tl.constexpr,
|
||||
T: tl.constexpr,
|
||||
D: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BD: tl.constexpr,
|
||||
):
|
||||
i_t = tl.program_id(0)
|
||||
p_x = tl.make_block_ptr(x, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
|
||||
b_x = tl.load(p_x, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_var = tl.sum(b_x * b_x, axis=1)
|
||||
b_y = b_x / tl.sqrt(b_var + eps)[:, None]
|
||||
p_y = tl.make_block_ptr(y, (T, D), (D, 1), (i_t * BT, 0), (BT, BD), (1, 0))
|
||||
tl.store(p_y, b_y.to(p_y.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
|
||||
def l2norm_fwd(
|
||||
x: torch.Tensor, eps: float = 1e-6, output_dtype: Optional[torch.dtype] = None
|
||||
):
|
||||
x_shape_og = x.shape
|
||||
x = x.view(-1, x.shape[-1])
|
||||
# allocate output
|
||||
if output_dtype is None:
|
||||
y = torch.empty_like(x)
|
||||
else:
|
||||
y = torch.empty_like(x, dtype=output_dtype)
|
||||
assert y.stride(-1) == 1
|
||||
T, D = x.shape[0], x.shape[-1]
|
||||
# rstd = torch.empty((T,), dtype=torch.float32, device=x.device)
|
||||
# Less than 64KB per feature: enqueue fused kernel
|
||||
MAX_FUSED_SIZE = 65536 // x.element_size()
|
||||
BD = min(MAX_FUSED_SIZE, triton.next_power_of_2(D))
|
||||
if D > BD:
|
||||
raise RuntimeError("This layer doesn't support feature dim >= 64KB.")
|
||||
|
||||
if D <= 512:
|
||||
NB = triton.cdiv(T, 2048)
|
||||
|
||||
def grid(meta):
|
||||
return (triton.cdiv(T, meta["BT"]),)
|
||||
|
||||
l2norm_fwd_kernel[grid](
|
||||
x,
|
||||
y,
|
||||
eps,
|
||||
NB=NB,
|
||||
T=T,
|
||||
D=D,
|
||||
BD=BD,
|
||||
BT=16,
|
||||
num_warps=8,
|
||||
num_stages=3,
|
||||
)
|
||||
else:
|
||||
l2norm_fwd_kernel1[(T,)](
|
||||
x,
|
||||
y,
|
||||
eps=eps,
|
||||
D=D,
|
||||
BD=BD,
|
||||
num_warps=8,
|
||||
num_stages=3,
|
||||
)
|
||||
|
||||
return y.view(x_shape_og)
|
||||
|
||||
|
||||
class L2NormFunction(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
@input_guard
|
||||
def forward(ctx, x, eps=1e-6, output_dtype=None):
|
||||
return l2norm_fwd(x, eps, output_dtype)
|
||||
|
||||
|
||||
def l2norm(
|
||||
x: torch.Tensor, eps: float = 1e-6, output_dtype: Optional[torch.dtype] = None
|
||||
) -> torch.Tensor:
|
||||
return L2NormFunction.apply(x, eps, output_dtype)
|
||||
|
||||
|
||||
l2_norm = l2norm
|
||||
|
||||
|
||||
class L2Norm(nn.Module):
|
||||
|
||||
def __init__(self, eps: float = 1e-6, output_dtype: Optional[torch.dtype] = None):
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.output_dtype = output_dtype
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return l2norm(x, self.eps, self.output_dtype)
|
||||
@@ -0,0 +1,483 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/modules/layernorm_gated.py
|
||||
# Copyright (c) 2024, Tri Dao.
|
||||
# Based on the Triton LayerNorm tutorial: https://triton-lang.org/main/getting-started/tutorials/05-layer-norm.html
|
||||
# For the backward pass, we keep weight_grad and bias_grad in registers and accumulate.
|
||||
# This backward pass is faster for dimensions up to 8k, but after that it's much slower due to register spilling.
|
||||
# The models we train have hidden dim up to 8k anyway (e.g. Llama 70B), so this is fine.
|
||||
|
||||
|
||||
from contextlib import nullcontext
|
||||
from functools import lru_cache
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from einops import rearrange
|
||||
|
||||
from sglang.jit_kernel.utils import is_arch_support_pdl
|
||||
from sglang.srt.batch_invariant_ops import is_batch_invariant_mode_enabled
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Backend,
|
||||
Phase,
|
||||
check_cuda_graph_backend,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
cdiv,
|
||||
cpu_has_amx_support,
|
||||
device_context,
|
||||
is_cpu,
|
||||
is_npu,
|
||||
next_power_of_2,
|
||||
)
|
||||
|
||||
_is_npu = is_npu()
|
||||
_use_cpu = is_cpu() and cpu_has_amx_support()
|
||||
|
||||
# Maximum rows per Triton block for layernorm gated kernel
|
||||
MAX_ROWS_PER_BLOCK = 4
|
||||
|
||||
|
||||
def rms_norm_ref(
|
||||
x,
|
||||
weight,
|
||||
bias,
|
||||
z=None,
|
||||
eps=1e-6,
|
||||
group_size=None,
|
||||
norm_before_gate=True,
|
||||
upcast=True,
|
||||
):
|
||||
dtype = x.dtype
|
||||
N = x.shape[-1]
|
||||
weight = weight.float()
|
||||
bias = bias.float() if bias is not None else None
|
||||
if upcast:
|
||||
x = x.float()
|
||||
z = z.float() if z is not None else z
|
||||
if z is not None and not norm_before_gate:
|
||||
x = x * F.silu(z)
|
||||
if group_size is None:
|
||||
rstd = 1 / torch.sqrt((x.square()).mean(dim=-1, keepdim=True) + eps)
|
||||
out = (x * rstd * weight) + bias if bias is not None else (x * rstd * weight)
|
||||
else:
|
||||
x_group = rearrange(x, "... (g d) -> ... g d", d=group_size)
|
||||
rstd = 1 / torch.sqrt((x_group.square()).mean(dim=-1, keepdim=True) + eps)
|
||||
out = rearrange(x_group * rstd, "... g d -> ... (g d)") * weight
|
||||
if bias is not None:
|
||||
out = out + bias
|
||||
if z is not None and norm_before_gate:
|
||||
out *= F.silu(z)
|
||||
return out.to(dtype)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _layer_norm_fwd_1pass_kernel(
|
||||
X, # pointer to the input
|
||||
Y, # pointer to the output
|
||||
W, # pointer to the weights
|
||||
B, # pointer to the biases
|
||||
Z, # pointer to the other branch
|
||||
Mean, # pointer to the mean
|
||||
Rstd, # pointer to the 1/std
|
||||
stride_x_row, # how much to increase the pointer when moving by 1 row
|
||||
stride_y_row,
|
||||
stride_z_row,
|
||||
M, # number of rows in X
|
||||
N: tl.constexpr, # number of columns in X
|
||||
eps, # epsilon to avoid division by zero
|
||||
BLOCK_N: tl.constexpr,
|
||||
ROWS_PER_BLOCK: tl.constexpr,
|
||||
HAS_BIAS: tl.constexpr,
|
||||
HAS_Z: tl.constexpr,
|
||||
NORM_BEFORE_GATE: tl.constexpr,
|
||||
IS_RMS_NORM: tl.constexpr,
|
||||
ACTIVATION: tl.constexpr,
|
||||
USE_GDC: tl.constexpr = False,
|
||||
):
|
||||
if USE_GDC:
|
||||
tl.extra.cuda.gdc_wait()
|
||||
|
||||
# Map the program id to the starting row of X and Y it should compute.
|
||||
row_start = tl.program_id(0) * ROWS_PER_BLOCK
|
||||
group = tl.program_id(1)
|
||||
|
||||
# Create 2D tile: [ROWS_PER_BLOCK, BLOCK_N]
|
||||
rows = row_start + tl.arange(0, ROWS_PER_BLOCK)
|
||||
cols = tl.arange(0, BLOCK_N)
|
||||
|
||||
# Compute offsets for 2D tile
|
||||
row_offsets = rows[:, None] * stride_x_row
|
||||
col_offsets = cols[None, :] + group * N
|
||||
|
||||
# Base pointers
|
||||
X_base = X + row_offsets + col_offsets
|
||||
Y_base = Y + rows[:, None] * stride_y_row + col_offsets
|
||||
|
||||
# Create mask for valid rows and columns
|
||||
row_mask = rows[:, None] < M
|
||||
col_mask = cols[None, :] < N
|
||||
mask = row_mask & col_mask
|
||||
|
||||
# Load input data with 2D tile
|
||||
x = tl.load(X_base, mask=mask, other=0.0).to(tl.float32)
|
||||
|
||||
if HAS_Z and not NORM_BEFORE_GATE:
|
||||
Z_base = Z + rows[:, None] * stride_z_row + col_offsets
|
||||
z = tl.load(Z_base, mask=mask, other=0.0).to(tl.float32)
|
||||
if ACTIVATION == "swish" or ACTIVATION == "silu":
|
||||
x *= z * tl.sigmoid(z)
|
||||
elif ACTIVATION == "sigmoid":
|
||||
x *= tl.sigmoid(z)
|
||||
|
||||
# Compute mean and variance per row (reduce along axis 1)
|
||||
if not IS_RMS_NORM:
|
||||
mean = tl.sum(x, axis=1) / N # Shape: [ROWS_PER_BLOCK]
|
||||
# Store mean for each row
|
||||
mean_offsets = group * M + rows
|
||||
mean_mask = rows < M
|
||||
tl.store(Mean + mean_offsets, mean, mask=mean_mask)
|
||||
# Broadcast mean back to 2D for subtraction
|
||||
xbar = tl.where(mask, x - mean[:, None], 0.0)
|
||||
var = tl.sum(xbar * xbar, axis=1) / N # Shape: [ROWS_PER_BLOCK]
|
||||
else:
|
||||
xbar = tl.where(mask, x, 0.0)
|
||||
var = tl.sum(xbar * xbar, axis=1) / N # Shape: [ROWS_PER_BLOCK]
|
||||
mean = 0.0 # Placeholder for RMS norm
|
||||
|
||||
rstd = tl.rsqrt(var + eps) # Shape: [ROWS_PER_BLOCK]
|
||||
|
||||
# Store rstd for each row
|
||||
rstd_offsets = group * M + rows
|
||||
rstd_mask = rows < M
|
||||
tl.store(Rstd + rstd_offsets, rstd, mask=rstd_mask)
|
||||
|
||||
# Load weights and biases (broadcast across rows)
|
||||
w_offsets = cols + group * N
|
||||
w_mask = cols < N
|
||||
w = tl.load(W + w_offsets, mask=w_mask, other=0.0).to(tl.float32)
|
||||
|
||||
if HAS_BIAS:
|
||||
b = tl.load(B + w_offsets, mask=w_mask, other=0.0).to(tl.float32)
|
||||
|
||||
# Normalize and apply linear transformation
|
||||
if not IS_RMS_NORM:
|
||||
x_hat = (x - mean[:, None]) * rstd[:, None]
|
||||
else:
|
||||
x_hat = x * rstd[:, None]
|
||||
|
||||
y = x_hat * w[None, :] + b[None, :] if HAS_BIAS else x_hat * w[None, :]
|
||||
|
||||
if HAS_Z and NORM_BEFORE_GATE:
|
||||
Z_base = Z + rows[:, None] * stride_z_row + col_offsets
|
||||
z = tl.load(Z_base, mask=mask, other=0.0).to(tl.float32)
|
||||
if ACTIVATION == "swish" or ACTIVATION == "silu":
|
||||
y *= z * tl.sigmoid(z)
|
||||
elif ACTIVATION == "sigmoid":
|
||||
y *= tl.sigmoid(z)
|
||||
|
||||
# Write output
|
||||
tl.store(Y_base, y, mask=mask)
|
||||
|
||||
if USE_GDC:
|
||||
tl.extra.cuda.gdc_launch_dependents()
|
||||
|
||||
|
||||
@lru_cache
|
||||
def _get_sm_count(device: torch.device) -> int:
|
||||
"""Get and cache the SM count for a given device."""
|
||||
if device.type == "xpu":
|
||||
assert torch.xpu.is_available(), "XPU device is not available"
|
||||
return torch.xpu.get_device_properties(device).gpu_subslice_count
|
||||
props = torch.cuda.get_device_properties(device)
|
||||
return props.multi_processor_count
|
||||
|
||||
|
||||
def calc_rows_per_block(M: int, device: torch.device) -> int:
|
||||
# Use a constant value when the row count must not affect kernel numerics.
|
||||
if is_batch_invariant_mode_enabled() or check_cuda_graph_backend(
|
||||
Phase.PREFILL, Backend.TC_PIECEWISE
|
||||
):
|
||||
return MAX_ROWS_PER_BLOCK
|
||||
sm_count = _get_sm_count(device)
|
||||
rows_per_block = next_power_of_2(cdiv(M, 2 * sm_count))
|
||||
rows_per_block = min(rows_per_block, MAX_ROWS_PER_BLOCK)
|
||||
return rows_per_block
|
||||
|
||||
|
||||
def _layer_norm_fwd(
|
||||
x,
|
||||
weight,
|
||||
bias,
|
||||
eps,
|
||||
z=None,
|
||||
out=None,
|
||||
group_size=None,
|
||||
norm_before_gate=True,
|
||||
is_rms_norm=False,
|
||||
activation: str = "swish",
|
||||
):
|
||||
M, N = x.shape
|
||||
if group_size is None:
|
||||
group_size = N
|
||||
assert N % group_size == 0
|
||||
ngroups = N // group_size
|
||||
assert x.stride(-1) == 1
|
||||
if z is not None:
|
||||
assert z.stride(-1) == 1
|
||||
assert z.shape == (M, N)
|
||||
assert weight.shape == (N,)
|
||||
assert weight.stride(-1) == 1
|
||||
if bias is not None:
|
||||
assert bias.stride(-1) == 1
|
||||
assert bias.shape == (N,)
|
||||
# allocate output
|
||||
if out is not None:
|
||||
assert out.shape == x.shape
|
||||
else:
|
||||
out = torch.empty_like(x)
|
||||
assert out.stride(-1) == 1
|
||||
mean = (
|
||||
torch.empty((ngroups * M,), dtype=torch.float32, device=x.device)
|
||||
if not is_rms_norm
|
||||
else None
|
||||
)
|
||||
rstd = torch.empty((ngroups * M,), dtype=torch.float32, device=x.device)
|
||||
# Less than 64KB per feature: enqueue fused kernel
|
||||
MAX_FUSED_SIZE = 65536 // x.element_size()
|
||||
BLOCK_N = min(MAX_FUSED_SIZE, triton.next_power_of_2(group_size))
|
||||
if group_size > BLOCK_N:
|
||||
raise RuntimeError("This layer norm doesn't support feature dim >= 64KB.")
|
||||
# heuristics for number of warps
|
||||
num_warps = min(max(BLOCK_N // 256, 1), 8)
|
||||
# Calculate rows per block based on SM count
|
||||
rows_per_block = calc_rows_per_block(M, x.device)
|
||||
# Update grid to use rows_per_block
|
||||
grid = (cdiv(M, rows_per_block), ngroups)
|
||||
pdl_kwargs = {"USE_GDC": True, "launch_pdl": True} if is_arch_support_pdl() else {}
|
||||
# Workaround for PyTorch <= 2.12: torch.xpu.device is not Dynamo-compatible
|
||||
# in that release — it creates a DynamoConfigPatchProxy that
|
||||
# SourcelessBuilder cannot wrap, causing a hard error under
|
||||
# torch.compile(fullgraph=True). The device context is a functional no-op
|
||||
# for Triton kernel launches (device is determined by the tensor, not the
|
||||
# surrounding context), so we simply skip it when Dynamo is tracing.
|
||||
# PyTorch main already has the proper fix (XPUDeviceVariable registered in
|
||||
# torch/_dynamo/variables/ctx_manager.py analogous to CUDADeviceVariable).
|
||||
# TODO: remove this branch once we upgrade from PyTorch 2.12.
|
||||
device_ctx = (
|
||||
nullcontext()
|
||||
if x.device.type == "xpu" and torch.compiler.is_compiling()
|
||||
else device_context(x.device)
|
||||
)
|
||||
with device_ctx:
|
||||
_layer_norm_fwd_1pass_kernel[grid](
|
||||
x,
|
||||
out,
|
||||
weight,
|
||||
bias,
|
||||
z,
|
||||
mean,
|
||||
rstd,
|
||||
x.stride(0),
|
||||
out.stride(0),
|
||||
z.stride(0) if z is not None else 0,
|
||||
M,
|
||||
group_size,
|
||||
eps,
|
||||
BLOCK_N=BLOCK_N,
|
||||
ROWS_PER_BLOCK=rows_per_block,
|
||||
HAS_BIAS=bias is not None,
|
||||
HAS_Z=z is not None,
|
||||
NORM_BEFORE_GATE=norm_before_gate,
|
||||
IS_RMS_NORM=is_rms_norm,
|
||||
num_warps=num_warps,
|
||||
ACTIVATION=activation,
|
||||
**pdl_kwargs,
|
||||
)
|
||||
return out, mean, rstd
|
||||
|
||||
|
||||
if _is_npu:
|
||||
from sgl_kernel_npu.fla.layernorm_gated import layer_norm_fwd_npu as _layer_norm_fwd
|
||||
|
||||
|
||||
def rms_norm_gated(
|
||||
*,
|
||||
x,
|
||||
weight,
|
||||
bias,
|
||||
z=None,
|
||||
eps=1e-6,
|
||||
group_size=None,
|
||||
norm_before_gate=True,
|
||||
is_rms_norm=False,
|
||||
activation: str = "swish",
|
||||
):
|
||||
"""If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))"""
|
||||
|
||||
x_shape_og = x.shape
|
||||
# reshape input data into 2D tensor
|
||||
x = x.reshape(-1, x.shape[-1])
|
||||
if x.stride(-1) != 1:
|
||||
x = x.contiguous()
|
||||
if z is not None:
|
||||
assert z.shape == x_shape_og
|
||||
z = z.reshape(-1, z.shape[-1])
|
||||
if z.stride(-1) != 1:
|
||||
z = z.contiguous()
|
||||
weight = weight.contiguous()
|
||||
if bias is not None:
|
||||
bias = bias.contiguous()
|
||||
if _is_npu:
|
||||
assert activation == "swish", "NPU only supports swish activation"
|
||||
y, mean, rstd = _layer_norm_fwd(
|
||||
x,
|
||||
weight,
|
||||
bias,
|
||||
eps,
|
||||
z=z,
|
||||
group_size=group_size,
|
||||
norm_before_gate=norm_before_gate,
|
||||
is_rms_norm=is_rms_norm,
|
||||
activation=activation,
|
||||
)
|
||||
return y.reshape(x_shape_og)
|
||||
|
||||
|
||||
class LayerNormFn(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
x,
|
||||
weight,
|
||||
bias,
|
||||
z=None,
|
||||
eps=1e-6,
|
||||
group_size=None,
|
||||
norm_before_gate=True,
|
||||
is_rms_norm=False,
|
||||
activation: str = "swish",
|
||||
):
|
||||
return rms_norm_gated(
|
||||
x=x,
|
||||
weight=weight,
|
||||
bias=bias,
|
||||
eps=eps,
|
||||
z=z,
|
||||
group_size=group_size,
|
||||
norm_before_gate=norm_before_gate,
|
||||
is_rms_norm=is_rms_norm,
|
||||
activation=activation,
|
||||
)
|
||||
|
||||
|
||||
def layernorm_fn(
|
||||
x,
|
||||
weight,
|
||||
bias,
|
||||
z=None,
|
||||
eps=1e-6,
|
||||
group_size=None,
|
||||
norm_before_gate=True,
|
||||
is_rms_norm=False,
|
||||
activation: str = "swish",
|
||||
):
|
||||
return LayerNormFn.apply(
|
||||
x, weight, bias, z, eps, group_size, norm_before_gate, is_rms_norm, activation
|
||||
)
|
||||
|
||||
|
||||
class LayerNorm(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
eps=1e-5,
|
||||
group_size=None,
|
||||
norm_before_gate=True,
|
||||
device=None,
|
||||
dtype=None,
|
||||
):
|
||||
"""If group_size is not None, we do GroupNorm with each group having group_size elements.
|
||||
group_size=None is equivalent to group_size=hidden_size (i.e. there's only 1 group).
|
||||
"""
|
||||
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.weight = torch.nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
|
||||
self.bias = torch.nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
|
||||
self.group_size = group_size
|
||||
self.norm_before_gate = norm_before_gate
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self):
|
||||
torch.nn.init.ones_(self.weight)
|
||||
torch.nn.init.zeros_(self.bias)
|
||||
|
||||
def forward(self, x, z=None):
|
||||
"""If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))"""
|
||||
return layernorm_fn(
|
||||
x,
|
||||
self.weight,
|
||||
self.bias,
|
||||
z=z,
|
||||
group_size=self.group_size,
|
||||
eps=self.eps,
|
||||
norm_before_gate=self.norm_before_gate,
|
||||
is_rms_norm=False,
|
||||
)
|
||||
|
||||
|
||||
class RMSNorm(torch.nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
eps=1e-5,
|
||||
group_size=None,
|
||||
norm_before_gate=True,
|
||||
device=None,
|
||||
dtype=None,
|
||||
activation: str = "swish",
|
||||
):
|
||||
"""If group_size is not None, we do GroupNorm with each group having group_size elements.
|
||||
group_size=None is equivalent to group_size=hidden_size (i.e. there's only 1 group).
|
||||
"""
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.eps = eps
|
||||
self.activation = activation
|
||||
self.weight = torch.nn.Parameter(torch.empty(hidden_size, **factory_kwargs))
|
||||
self.register_parameter("bias", None)
|
||||
self.group_size = group_size
|
||||
self.norm_before_gate = norm_before_gate
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self):
|
||||
torch.nn.init.ones_(self.weight)
|
||||
|
||||
def forward(self, x, z=None):
|
||||
"""If z is not None, we do norm(x) * silu(z) if norm_before_gate, else norm(x * silu(z))"""
|
||||
if _use_cpu:
|
||||
assert (
|
||||
self.norm_before_gate
|
||||
and self.group_size is None
|
||||
and self.activation == "swish"
|
||||
), "CPU rmsnorm_gated currently only supports norm before gate without group size or activation other than swish"
|
||||
return torch.ops.sgl_kernel.fused_rmsnorm_gated_cpu(
|
||||
x, self.weight, z, self.eps
|
||||
)
|
||||
else:
|
||||
return layernorm_fn(
|
||||
x,
|
||||
self.weight,
|
||||
self.bias,
|
||||
z=z,
|
||||
eps=self.eps,
|
||||
group_size=self.group_size,
|
||||
norm_before_gate=self.norm_before_gate,
|
||||
is_rms_norm=True,
|
||||
activation=self.activation,
|
||||
)
|
||||
@@ -0,0 +1,66 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/utils/op.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
import os
|
||||
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import triton.language.extra.libdevice as tldevice
|
||||
|
||||
from sglang.srt.layers.attention.fla.utils import is_gather_supported
|
||||
|
||||
if os.environ.get("FLA_USE_FAST_OPS", "0") == "1":
|
||||
exp = tldevice.fast_expf
|
||||
exp2 = tldevice.exp2
|
||||
log = tldevice.fast_logf
|
||||
log2 = tldevice.fast_log2f
|
||||
else:
|
||||
exp = tl.exp
|
||||
exp2 = tl.math.exp2
|
||||
log = tl.log
|
||||
log2 = tl.log2
|
||||
|
||||
|
||||
@triton.jit
|
||||
def safe_exp(x):
|
||||
return exp(tl.where(x <= 0, x, float("-inf")))
|
||||
|
||||
|
||||
if not is_gather_supported:
|
||||
|
||||
@triton.jit
|
||||
def gather(src, index, axis, _builder=None):
|
||||
"""
|
||||
Gather operation that works when tl.gather is not supported.
|
||||
This is a fallback implementation that returns None.
|
||||
Just to make triton compiler happy.
|
||||
"""
|
||||
return None
|
||||
|
||||
else:
|
||||
gather = tl.gather
|
||||
|
||||
|
||||
if hasattr(triton.language, "_experimental_make_tensor_descriptor"):
|
||||
# For Triton 3.3.x
|
||||
make_tensor_descriptor = triton.language._experimental_make_tensor_descriptor
|
||||
elif hasattr(triton.language, "make_tensor_descriptor"):
|
||||
# For Triton 3.4.x and later
|
||||
make_tensor_descriptor = triton.language.make_tensor_descriptor
|
||||
else:
|
||||
"""
|
||||
Fallback implementation when TMA is not supported.
|
||||
Returns None to indicate TMA descriptors are unavailable.
|
||||
Just make triton compiler happy.
|
||||
"""
|
||||
|
||||
@triton.jit
|
||||
def make_tensor_descriptor(
|
||||
base,
|
||||
shape,
|
||||
strides,
|
||||
block_shape,
|
||||
_builder=None,
|
||||
):
|
||||
return None
|
||||
@@ -0,0 +1,464 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/utils/solve_tril.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
from typing import Optional
|
||||
|
||||
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.utils import input_guard
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[
|
||||
# triton.Config({}, num_warps=num_warps, num_stages=num_stages)
|
||||
# for num_warps in [1, 2, 4, 8]
|
||||
# for num_stages in [2, 3, 4, 5]
|
||||
# ],
|
||||
# key=["BT"],
|
||||
# )
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def solve_tril_16x16_kernel(
|
||||
A,
|
||||
Ad,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
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
|
||||
|
||||
A = A + (bos * H + i_h) * BT
|
||||
Ad = Ad + (bos * H + i_h) * 16
|
||||
|
||||
offset = (i_t * 16) % BT
|
||||
p_A = tl.make_block_ptr(
|
||||
A, (T, BT), (H * BT, 1), (i_t * 16, offset), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai = tl.make_block_ptr(Ad, (T, 16), (H * 16, 1), (i_t * 16, 0), (16, 16), (1, 0))
|
||||
b_A = tl.load(p_A, boundary_check=(0, 1)).to(tl.float32)
|
||||
b_A = -tl.where(tl.arange(0, 16)[:, None] > tl.arange(0, 16)[None, :], b_A, 0)
|
||||
|
||||
o_i = tl.arange(0, 16)
|
||||
for i in range(1, min(16, T - i_t * 16)):
|
||||
b_a = -tl.load(A + (i_t * 16 + i) * H * BT + o_i + offset)
|
||||
b_a = b_a + tl.sum(b_a[:, None] * b_A, 0)
|
||||
mask = o_i == i
|
||||
b_A = tl.where(mask[:, None], b_a, b_A)
|
||||
b_A += o_i[:, None] == o_i[None, :]
|
||||
tl.store(
|
||||
p_Ai,
|
||||
b_A.to(p_Ai.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[
|
||||
# triton.Config({}, num_warps=num_warps, num_stages=num_stages)
|
||||
# for num_warps in [1, 2, 4, 8]
|
||||
# for num_stages in [2, 3, 4, 5]
|
||||
# ],
|
||||
# key=["H", "BT", "IS_VARLEN"],
|
||||
# )
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def merge_16x16_to_32x32_inverse_kernel(
|
||||
A,
|
||||
Ad,
|
||||
Ai,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
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
|
||||
|
||||
A += (bos * H + i_h) * 32
|
||||
Ad += (bos * H + i_h) * 16
|
||||
Ai += (bos * H + i_h) * 32
|
||||
|
||||
p_A_21 = tl.make_block_ptr(
|
||||
A, (T, 32), (H * 32, 1), (i_t * 32 + 16, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ad_11 = tl.make_block_ptr(
|
||||
Ad, (T, 16), (H * 16, 1), (i_t * 32, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ad_22 = tl.make_block_ptr(
|
||||
Ad, (T, 16), (H * 16, 1), (i_t * 32 + 16, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_11 = tl.make_block_ptr(
|
||||
Ai, (T, 32), (H * 32, 1), (i_t * 32, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_22 = tl.make_block_ptr(
|
||||
Ai, (T, 32), (H * 32, 1), (i_t * 32 + 16, 16), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_21 = tl.make_block_ptr(
|
||||
Ai, (T, 32), (H * 32, 1), (i_t * 32 + 16, 0), (16, 16), (1, 0)
|
||||
)
|
||||
|
||||
A_21 = tl.load(p_A_21, boundary_check=(0, 1)).to(tl.float32)
|
||||
Ai_11 = tl.load(p_Ad_11, boundary_check=(0, 1)).to(tl.float32)
|
||||
Ai_22 = tl.load(p_Ad_22, boundary_check=(0, 1)).to(tl.float32)
|
||||
Ai_21 = -tl.dot(
|
||||
tl.dot(Ai_22, A_21, input_precision="ieee"), Ai_11, input_precision="ieee"
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_11,
|
||||
Ai_11.to(p_Ai_11.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_22,
|
||||
Ai_22.to(p_Ai_22.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_21,
|
||||
Ai_21.to(p_Ai_21.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[
|
||||
# triton.Config({}, num_warps=num_warps, num_stages=num_stages)
|
||||
# for num_warps in [2, 4, 8]
|
||||
# for num_stages in [2, 3, 4, 5]
|
||||
# ],
|
||||
# key=["H", "BT", "IS_VARLEN"],
|
||||
# )
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def merge_16x16_to_64x64_inverse_kernel(
|
||||
A,
|
||||
Ad,
|
||||
Ai,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
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
|
||||
|
||||
A += (bos * H + i_h) * 64
|
||||
Ad += (bos * H + i_h) * 16
|
||||
Ai += (bos * H + i_h) * 64
|
||||
|
||||
p_A_21 = tl.make_block_ptr(
|
||||
A, (T, 64), (H * 64, 1), (i_t * 64 + 16, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_A_32 = tl.make_block_ptr(
|
||||
A, (T, 64), (H * 64, 1), (i_t * 64 + 32, 16), (16, 16), (1, 0)
|
||||
)
|
||||
p_A_31 = tl.make_block_ptr(
|
||||
A, (T, 64), (H * 64, 1), (i_t * 64 + 32, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_A_43 = tl.make_block_ptr(
|
||||
A, (T, 64), (H * 64, 1), (i_t * 64 + 48, 32), (16, 16), (1, 0)
|
||||
)
|
||||
p_A_42 = tl.make_block_ptr(
|
||||
A, (T, 64), (H * 64, 1), (i_t * 64 + 48, 16), (16, 16), (1, 0)
|
||||
)
|
||||
p_A_41 = tl.make_block_ptr(
|
||||
A, (T, 64), (H * 64, 1), (i_t * 64 + 48, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ad_11 = tl.make_block_ptr(
|
||||
Ad, (T, 16), (H * 16, 1), (i_t * 64, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ad_22 = tl.make_block_ptr(
|
||||
Ad, (T, 16), (H * 16, 1), (i_t * 64 + 16, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ad_33 = tl.make_block_ptr(
|
||||
Ad, (T, 16), (H * 16, 1), (i_t * 64 + 32, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ad_44 = tl.make_block_ptr(
|
||||
Ad, (T, 16), (H * 16, 1), (i_t * 64 + 48, 0), (16, 16), (1, 0)
|
||||
)
|
||||
|
||||
A_21 = tl.load(p_A_21, boundary_check=(0, 1)).to(tl.float32)
|
||||
A_32 = tl.load(p_A_32, boundary_check=(0, 1)).to(tl.float32)
|
||||
A_31 = tl.load(p_A_31, boundary_check=(0, 1)).to(tl.float32)
|
||||
A_43 = tl.load(p_A_43, boundary_check=(0, 1)).to(tl.float32)
|
||||
A_42 = tl.load(p_A_42, boundary_check=(0, 1)).to(tl.float32)
|
||||
A_41 = tl.load(p_A_41, boundary_check=(0, 1)).to(tl.float32)
|
||||
|
||||
Ai_11 = tl.load(p_Ad_11, boundary_check=(0, 1)).to(tl.float32)
|
||||
Ai_22 = tl.load(p_Ad_22, boundary_check=(0, 1)).to(tl.float32)
|
||||
Ai_33 = tl.load(p_Ad_33, boundary_check=(0, 1)).to(tl.float32)
|
||||
Ai_44 = tl.load(p_Ad_44, boundary_check=(0, 1)).to(tl.float32)
|
||||
|
||||
Ai_21 = -tl.dot(
|
||||
tl.dot(Ai_22, A_21, input_precision="ieee"), Ai_11, input_precision="ieee"
|
||||
)
|
||||
Ai_32 = -tl.dot(
|
||||
tl.dot(Ai_33, A_32, input_precision="ieee"), Ai_22, input_precision="ieee"
|
||||
)
|
||||
Ai_43 = -tl.dot(
|
||||
tl.dot(Ai_44, A_43, input_precision="ieee"), Ai_33, input_precision="ieee"
|
||||
)
|
||||
|
||||
Ai_31 = -tl.dot(
|
||||
Ai_33,
|
||||
tl.dot(A_31, Ai_11, input_precision="ieee")
|
||||
+ tl.dot(A_32, Ai_21, input_precision="ieee"),
|
||||
input_precision="ieee",
|
||||
)
|
||||
Ai_42 = -tl.dot(
|
||||
Ai_44,
|
||||
tl.dot(A_42, Ai_22, input_precision="ieee")
|
||||
+ tl.dot(A_43, Ai_32, input_precision="ieee"),
|
||||
input_precision="ieee",
|
||||
)
|
||||
Ai_41 = -tl.dot(
|
||||
Ai_44,
|
||||
tl.dot(A_41, Ai_11, input_precision="ieee")
|
||||
+ tl.dot(A_42, Ai_21, input_precision="ieee")
|
||||
+ tl.dot(A_43, Ai_31, input_precision="ieee"),
|
||||
input_precision="ieee",
|
||||
)
|
||||
|
||||
p_Ai_11 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_22 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 16, 16), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_33 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 32, 32), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_44 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 48, 48), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_21 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 16, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_31 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 32, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_32 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 32, 16), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_41 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 48, 0), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_42 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 48, 16), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_43 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 48, 32), (16, 16), (1, 0)
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_11,
|
||||
Ai_11.to(p_Ai_11.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_22,
|
||||
Ai_22.to(p_Ai_22.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_33,
|
||||
Ai_33.to(p_Ai_33.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_44,
|
||||
Ai_44.to(p_Ai_44.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_21,
|
||||
Ai_21.to(p_Ai_21.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_31,
|
||||
Ai_31.to(p_Ai_31.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_32,
|
||||
Ai_32.to(p_Ai_32.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_41,
|
||||
Ai_41.to(p_Ai_41.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_42,
|
||||
Ai_42.to(p_Ai_42.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_43,
|
||||
Ai_43.to(p_Ai_43.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
|
||||
fill_zeros = tl.zeros((16, 16), dtype=tl.float32)
|
||||
p_Ai_12 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64, 16), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_13 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64, 32), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_14 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64, 48), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_23 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 16, 32), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_24 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 16, 48), (16, 16), (1, 0)
|
||||
)
|
||||
p_Ai_34 = tl.make_block_ptr(
|
||||
Ai, (T, 64), (H * 64, 1), (i_t * 64 + 32, 48), (16, 16), (1, 0)
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_12,
|
||||
fill_zeros.to(p_Ai_12.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_13,
|
||||
fill_zeros.to(p_Ai_13.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_14,
|
||||
fill_zeros.to(p_Ai_14.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_23,
|
||||
fill_zeros.to(p_Ai_23.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_24,
|
||||
fill_zeros.to(p_Ai_24.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
tl.store(
|
||||
p_Ai_34,
|
||||
fill_zeros.to(p_Ai_34.dtype.element_ty, fp_downcast_rounding="rtne"),
|
||||
boundary_check=(0, 1),
|
||||
)
|
||||
|
||||
|
||||
@input_guard
|
||||
def solve_tril(
|
||||
A: torch.Tensor,
|
||||
cu_seqlens: Optional[torch.Tensor] = None,
|
||||
output_dtype: torch.dtype = torch.float,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Compute the inverse of the lower triangular matrix
|
||||
A should be strictly lower triangular, i.e., A.triu() == 0.
|
||||
|
||||
Args:
|
||||
A (torch.Tensor):
|
||||
[B, T, H, K]
|
||||
cu_seqlens (torch.Tensor):
|
||||
The cumulative sequence lengths of the input tensor.
|
||||
Default: None.
|
||||
output_dtype (torch.dtype):
|
||||
The dtype of the output tensor. Default: `torch.float`
|
||||
|
||||
Returns:
|
||||
(I + A)^-1 with the same shape as A
|
||||
"""
|
||||
assert A.shape[-1] in [16, 32, 64]
|
||||
|
||||
B, T, H, BT = A.shape
|
||||
Ad = torch.empty(
|
||||
B, T, H, 16, device=A.device, dtype=torch.float if BT != 16 else output_dtype
|
||||
)
|
||||
|
||||
chunk_indices = (
|
||||
prepare_chunk_indices(cu_seqlens, 16) if cu_seqlens is not None else None
|
||||
)
|
||||
NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, 16)
|
||||
solve_tril_16x16_kernel[NT, B * H](
|
||||
A=A,
|
||||
Ad=Ad,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
H=H,
|
||||
BT=BT,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
num_warps=1,
|
||||
num_stages=4,
|
||||
)
|
||||
if BT == 16:
|
||||
return Ad
|
||||
|
||||
Ai = torch.empty(B, T, H, BT, device=A.device, dtype=output_dtype)
|
||||
merge_fn = (
|
||||
merge_16x16_to_32x32_inverse_kernel
|
||||
if BT == 32
|
||||
else merge_16x16_to_64x64_inverse_kernel
|
||||
)
|
||||
chunk_indices = (
|
||||
prepare_chunk_indices(cu_seqlens, BT) if cu_seqlens is not None else None
|
||||
)
|
||||
NT = len(chunk_indices) if cu_seqlens is not None else triton.cdiv(T, BT)
|
||||
merge_fn[NT, B * H](
|
||||
A=A,
|
||||
Ad=Ad,
|
||||
Ai=Ai,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
H=H,
|
||||
BT=BT,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
num_warps=4,
|
||||
num_stages=3,
|
||||
)
|
||||
return Ai
|
||||
@@ -0,0 +1,339 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/utils.py
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
import contextlib
|
||||
import functools
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from enum import Enum
|
||||
from functools import lru_cache
|
||||
from typing import Any, Callable, Dict, Literal, Optional, Tuple
|
||||
|
||||
import torch
|
||||
import triton
|
||||
from packaging import version
|
||||
|
||||
from sglang.srt.utils.common import torch_release
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
COMPILER_MODE = os.getenv("FLA_COMPILER_MODE") == "1"
|
||||
FLA_CI_ENV = os.getenv("FLA_CI_ENV") == "1"
|
||||
FLA_CACHE_RESULTS = os.getenv("FLA_CACHE_RESULTS", "1") == "1"
|
||||
|
||||
|
||||
SUPPORTS_AUTOTUNE_CACHE = (
|
||||
"cache_results" in inspect.signature(triton.autotune).parameters
|
||||
)
|
||||
|
||||
autotune_cache_kwargs = (
|
||||
{"cache_results": FLA_CACHE_RESULTS} if SUPPORTS_AUTOTUNE_CACHE else {}
|
||||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def check_environments():
|
||||
"""
|
||||
Checks the current operating system, Triton version, and Python version,
|
||||
issuing warnings if they don't meet recommendations.
|
||||
This function's body only runs once due to lru_cache.
|
||||
"""
|
||||
# Check Operating System
|
||||
if sys.platform == "win32":
|
||||
logger.warning(
|
||||
"Detected Windows operating system. Triton does not have an official Windows release, "
|
||||
"thus FLA will not be adapted for Windows, and any potential errors will not be fixed. "
|
||||
"Please consider using a Linux environment for compatibility."
|
||||
)
|
||||
|
||||
triton_version = version.parse(triton.__version__)
|
||||
required_triton_version = version.parse("3.2.0")
|
||||
|
||||
if triton_version < required_triton_version:
|
||||
logger.warning(
|
||||
f"Current Triton version {triton_version} is below the recommended 3.2.0 version. "
|
||||
"Errors may occur and these issues will not be fixed. "
|
||||
"Please consider upgrading Triton."
|
||||
)
|
||||
|
||||
# Check Python version
|
||||
py_version = version.parse(f"{sys.version_info.major}.{sys.version_info.minor}")
|
||||
required_py_version = version.parse("3.11")
|
||||
|
||||
if py_version < required_py_version:
|
||||
logger.warning(
|
||||
f"Current Python version {py_version} is below the recommended 3.11 version. "
|
||||
"It is recommended to upgrade to Python 3.11 or higher for the best experience."
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def get_abs_err(x, y):
|
||||
return (x.detach() - y.detach()).flatten().abs().max().item()
|
||||
|
||||
|
||||
def get_err_ratio(x, y):
|
||||
err = (x.detach() - y.detach()).flatten().square().mean().sqrt().item()
|
||||
base = (x.detach()).flatten().square().mean().sqrt().item()
|
||||
return err / (base + 1e-8)
|
||||
|
||||
|
||||
def assert_close(prefix, ref, tri, ratio, warning=False, err_atol=1e-6):
|
||||
abs_atol = get_abs_err(ref, tri)
|
||||
msg = f"{prefix} diff: {abs_atol:.6f} ratio: {get_err_ratio(ref, tri):.6f}"
|
||||
logger.info(msg)
|
||||
error_rate = get_err_ratio(ref, tri)
|
||||
if abs_atol <= err_atol:
|
||||
return
|
||||
if warning or (FLA_CI_ENV and (error_rate < 0.01 or abs_atol <= 0.3)):
|
||||
if error_rate > ratio:
|
||||
import warnings
|
||||
|
||||
warnings.warn(msg)
|
||||
else:
|
||||
assert error_rate < ratio, msg
|
||||
|
||||
|
||||
SUPPRESS_LEVEL = int(os.getenv("GDN_RECOMPUTE_SUPPRESS_LEVEL", "0"))
|
||||
|
||||
|
||||
def tensor_cache(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
|
||||
"""
|
||||
A decorator that caches the most recent results of a function with tensor inputs.
|
||||
This decorator will store the output of the decorated function for the most recent set of input tensors.
|
||||
The cache is limited to a fixed size (default is 4). When the cache is full, the oldest entry will be removed.
|
||||
Args:
|
||||
fn (Callable[..., torch.Tensor]):
|
||||
The function to be decorated. It should take tensor inputs and return tensor outputs.
|
||||
Returns:
|
||||
Callable[..., torch.Tensor]:
|
||||
A wrapped version of the input function with single-entry caching.
|
||||
"""
|
||||
|
||||
cache_entries: Tuple[Optional[Tuple], Optional[Dict], Any] = []
|
||||
cache_size = 4
|
||||
|
||||
@functools.wraps(fn)
|
||||
def wrapper(*args: Any, **kwargs: Any) -> Any:
|
||||
nonlocal cache_entries, cache_size
|
||||
for i, entry in enumerate(cache_entries):
|
||||
last_args, last_kwargs, last_result = entry
|
||||
if len(args) == len(last_args) and len(kwargs) == len(last_kwargs):
|
||||
if all(a is b for a, b in zip(args, last_args)) and all(
|
||||
k in last_kwargs and v is last_kwargs[k] for k, v in kwargs.items()
|
||||
):
|
||||
cache_entries = (
|
||||
cache_entries[:i]
|
||||
+ cache_entries[i + 1 :]
|
||||
+ [(args, kwargs, last_result)]
|
||||
)
|
||||
return last_result
|
||||
|
||||
result = fn(*args, **kwargs)
|
||||
|
||||
if len(cache_entries) >= cache_size:
|
||||
cache_entries = cache_entries[1:]
|
||||
cache_entries.append((args, kwargs, result))
|
||||
return result
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def input_guard(fn: Callable[..., torch.Tensor]) -> Callable[..., torch.Tensor]:
|
||||
"""
|
||||
A decorator to make sure all input tensors are contiguous and set the device based on input tensors.
|
||||
"""
|
||||
|
||||
@functools.wraps(fn)
|
||||
def wrapper(*args, **kwargs):
|
||||
contiguous_args = (
|
||||
i if not isinstance(i, torch.Tensor) else i.contiguous() for i in args
|
||||
)
|
||||
contiguous_kwargs = {
|
||||
k: (v if not isinstance(v, torch.Tensor) else v.contiguous())
|
||||
for k, v in kwargs.items()
|
||||
}
|
||||
|
||||
tensor = None
|
||||
for arg in args:
|
||||
if isinstance(arg, torch.Tensor):
|
||||
tensor = arg
|
||||
break
|
||||
if tensor is None:
|
||||
for value in kwargs.values():
|
||||
if isinstance(value, torch.Tensor):
|
||||
tensor = value
|
||||
break
|
||||
|
||||
if tensor is not None:
|
||||
ctx = custom_device_ctx(tensor.device.index)
|
||||
else:
|
||||
ctx = contextlib.nullcontext()
|
||||
|
||||
with ctx:
|
||||
return fn(*contiguous_args, **contiguous_kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
contiguous = input_guard
|
||||
|
||||
|
||||
def require_version(version, hint):
|
||||
"""
|
||||
Perform a runtime check of the dependency versions, using the exact same syntax used by pip.
|
||||
"""
|
||||
|
||||
def decorator(fn):
|
||||
@functools.wraps(fn)
|
||||
def wrapper(ctx, *args, **kwargs):
|
||||
from transformers.utils.versions import require_version
|
||||
|
||||
require_version(version, hint)
|
||||
return fn(
|
||||
ctx,
|
||||
*(
|
||||
i if not isinstance(i, torch.Tensor) else i.contiguous()
|
||||
for i in args
|
||||
),
|
||||
**{
|
||||
k: (v if not isinstance(v, torch.Tensor) else v.contiguous())
|
||||
for k, v in kwargs.items()
|
||||
},
|
||||
)
|
||||
|
||||
return wrapper
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def checkpoint(fn):
|
||||
def wrapper(*args, **kwargs):
|
||||
return torch.utils.checkpoint.checkpoint(fn, *args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
|
||||
def _cpu_device_warning():
|
||||
import warnings
|
||||
|
||||
warnings.warn(
|
||||
("Triton is not supported on current platform, roll back to CPU."), stacklevel=1
|
||||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def get_multiprocessor_count(tensor_idx: int = 0) -> int:
|
||||
try:
|
||||
return triton.runtime.driver.active.utils.get_device_properties(tensor_idx)[
|
||||
"multiprocessor_count"
|
||||
]
|
||||
except BaseException:
|
||||
_cpu_device_warning()
|
||||
return -1
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def get_available_device() -> str:
|
||||
try:
|
||||
return triton.runtime.driver.active.get_current_target().backend
|
||||
except BaseException:
|
||||
_cpu_device_warning()
|
||||
return "cpu"
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def _check_platform() -> Literal["nvidia", "amd", "intel", "musa"]:
|
||||
device = get_available_device()
|
||||
if device == "cuda":
|
||||
return "nvidia"
|
||||
elif device == "hip":
|
||||
return "amd"
|
||||
elif device == "xpu":
|
||||
return "intel"
|
||||
else:
|
||||
return device
|
||||
|
||||
|
||||
# For AMD GPUs, the triton backend is 'hip', while for Nvidia GPUs, the triton backend is 'cuda'.
|
||||
# However, the torch backend is 'cuda' for both Nvidia and AMD GPUs.
|
||||
# Therefore, we need to check the triton backend to determine the actual GPU vendor.
|
||||
device = get_available_device() if get_available_device() != "hip" else "cuda"
|
||||
device_torch_lib = getattr(torch, device)
|
||||
device_platform = _check_platform()
|
||||
|
||||
is_amd = device_platform == "amd"
|
||||
is_intel = device_platform == "intel"
|
||||
is_nvidia = device_platform == "nvidia"
|
||||
is_intel_alchemist = is_intel and "Intel(R) Arc(TM) A" in torch.xpu.get_device_name(0)
|
||||
is_nvidia_hopper = is_nvidia and (
|
||||
"NVIDIA H" in torch.cuda.get_device_name(0)
|
||||
or torch.cuda.get_device_capability()[0] >= 9
|
||||
)
|
||||
use_cuda_graph = is_nvidia and os.environ.get("FLA_USE_CUDA_GRAPH", "0") == "1"
|
||||
|
||||
# Nvidia Ampere or newer, haven't check AMD and intel yet.
|
||||
is_tf32_supported = is_nvidia and torch.cuda.get_device_capability(0)[0] >= 8
|
||||
is_gather_supported = hasattr(triton.language, "gather")
|
||||
|
||||
|
||||
def get_all_max_shared_mem():
|
||||
try:
|
||||
return [
|
||||
triton.runtime.driver.active.utils.get_device_properties(i)[
|
||||
"max_shared_mem"
|
||||
]
|
||||
for i in range(device_torch_lib.device_count())
|
||||
]
|
||||
except BaseException:
|
||||
_cpu_device_warning()
|
||||
return [-1]
|
||||
|
||||
|
||||
class Backend(Enum):
|
||||
ADA = 101376 # RTX 4090
|
||||
AMPERE = 166912 # A100
|
||||
HOPPER = 232448 # H100
|
||||
DEFAULT = 102400 # Default
|
||||
|
||||
@classmethod
|
||||
def get_shared_memory(cls, arch: str) -> int:
|
||||
try:
|
||||
return cls[arch.upper()].value
|
||||
except KeyError:
|
||||
return cls.DEFAULT.value
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def check_shared_mem(arch: str = "none", tensor_idx: int = 0) -> bool:
|
||||
try:
|
||||
device_shared_mem_list = get_all_max_shared_mem()
|
||||
max_shared_memory = device_shared_mem_list[tensor_idx]
|
||||
return max_shared_memory >= Backend.get_shared_memory(arch)
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
if torch_release >= (2, 4):
|
||||
device = "cuda" if device == "cpu" else device
|
||||
autocast_custom_fwd = functools.partial(torch.amp.custom_fwd, device_type=device)
|
||||
autocast_custom_bwd = functools.partial(torch.amp.custom_bwd, device_type=device)
|
||||
|
||||
def custom_device_ctx(index: int):
|
||||
return device_torch_lib.device(index)
|
||||
|
||||
else:
|
||||
assert (
|
||||
device == "cuda"
|
||||
), "Only cuda device is supported for PyTorch version < 2.4.0."
|
||||
autocast_custom_fwd = device_torch_lib.amp.custom_fwd
|
||||
autocast_custom_bwd = device_torch_lib.amp.custom_bwd
|
||||
|
||||
def custom_device_ctx(index: int):
|
||||
return torch.cuda.device(index)
|
||||
|
||||
|
||||
device_platform = get_available_device()
|
||||
@@ -0,0 +1,156 @@
|
||||
# Adapt from https://github.com/fla-org/flash-linear-attention/blob/main/fla/ops/gated_delta_rule/wy_fast.py
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
|
||||
|
||||
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
|
||||
|
||||
|
||||
# @triton.autotune(
|
||||
# configs=[
|
||||
# triton.Config({}, num_warps=num_warps, num_stages=num_stages)
|
||||
# for num_warps in [2, 4, 8]
|
||||
# for num_stages in [2, 3, 4]
|
||||
# ],
|
||||
# key=["H", "K", "V", "BT", "BK", "BV", "IS_VARLEN"],
|
||||
# )
|
||||
@triton.jit(do_not_specialize=["T"])
|
||||
def recompute_w_u_fwd_kernel(
|
||||
k,
|
||||
v,
|
||||
beta,
|
||||
w,
|
||||
u,
|
||||
A,
|
||||
g,
|
||||
cu_seqlens,
|
||||
chunk_indices,
|
||||
T,
|
||||
H: tl.constexpr,
|
||||
Hg: tl.constexpr,
|
||||
K: tl.constexpr,
|
||||
V: tl.constexpr,
|
||||
BT: tl.constexpr,
|
||||
BK: tl.constexpr,
|
||||
BV: tl.constexpr,
|
||||
IS_VARLEN: tl.constexpr,
|
||||
):
|
||||
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
|
||||
p_beta = tl.make_block_ptr(
|
||||
beta + bos * H + i_h, (T,), (H,), (i_t * BT,), (BT,), (0,)
|
||||
)
|
||||
p_g = tl.make_block_ptr(g + (bos * H + i_h), (T,), (H,), (i_t * BT,), (BT,), (0,))
|
||||
p_A = tl.make_block_ptr(
|
||||
A + (bos * H + i_h) * BT, (T, BT), (H * BT, 1), (i_t * BT, 0), (BT, BT), (1, 0)
|
||||
)
|
||||
b_beta = tl.load(p_beta, boundary_check=(0,))
|
||||
b_A = tl.load(p_A, boundary_check=(0, 1))
|
||||
b_g = tl.exp(tl.load(p_g, boundary_check=(0,)))
|
||||
|
||||
for i_v in range(tl.cdiv(V, BV)):
|
||||
p_v = tl.make_block_ptr(
|
||||
v + (bos * H + i_h) * V,
|
||||
(T, V),
|
||||
(H * V, 1),
|
||||
(i_t * BT, i_v * BV),
|
||||
(BT, BV),
|
||||
(1, 0),
|
||||
)
|
||||
p_u = tl.make_block_ptr(
|
||||
u + (bos * H + i_h) * V,
|
||||
(T, V),
|
||||
(H * V, 1),
|
||||
(i_t * BT, i_v * BV),
|
||||
(BT, BV),
|
||||
(1, 0),
|
||||
)
|
||||
b_v = tl.load(p_v, boundary_check=(0, 1))
|
||||
b_vb = (b_v * b_beta[:, None]).to(b_v.dtype)
|
||||
b_u = tl.dot(b_A, b_vb, allow_tf32=False)
|
||||
tl.store(p_u, b_u.to(p_u.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
for i_k in range(tl.cdiv(K, BK)):
|
||||
p_k = tl.make_block_ptr(
|
||||
k + (bos * Hg + i_h // (H // Hg)) * K,
|
||||
(T, K),
|
||||
(Hg * K, 1),
|
||||
(i_t * BT, i_k * BK),
|
||||
(BT, BK),
|
||||
(1, 0),
|
||||
)
|
||||
p_w = tl.make_block_ptr(
|
||||
w + (bos * H + i_h) * K,
|
||||
(T, K),
|
||||
(H * K, 1),
|
||||
(i_t * BT, i_k * BK),
|
||||
(BT, BK),
|
||||
(1, 0),
|
||||
)
|
||||
b_k = tl.load(p_k, boundary_check=(0, 1))
|
||||
b_kb = (b_k * b_beta[:, None] * b_g[:, None]).to(b_k.dtype)
|
||||
b_w = tl.dot(b_A, b_kb)
|
||||
tl.store(p_w, b_w.to(p_w.dtype.element_ty), boundary_check=(0, 1))
|
||||
|
||||
|
||||
def recompute_w_u_fwd(
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
beta: torch.Tensor,
|
||||
g_cumsum: torch.Tensor,
|
||||
A: torch.Tensor,
|
||||
cu_seqlens: Optional[torch.LongTensor],
|
||||
chunk_indices: torch.LongTensor | None = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
B, T, Hg, K, V = *k.shape, v.shape[-1]
|
||||
H = v.shape[-2]
|
||||
BT = A.shape[-1]
|
||||
|
||||
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)
|
||||
BK = 64
|
||||
BV = 64
|
||||
u = torch.empty_like(v)
|
||||
w = k.new_empty(B, T, H, K)
|
||||
recompute_w_u_fwd_kernel[(NT, B * H)](
|
||||
k=k,
|
||||
v=v,
|
||||
beta=beta,
|
||||
w=w,
|
||||
u=u,
|
||||
A=A,
|
||||
g=g_cumsum,
|
||||
cu_seqlens=cu_seqlens,
|
||||
chunk_indices=chunk_indices,
|
||||
T=T,
|
||||
H=H,
|
||||
Hg=Hg,
|
||||
K=K,
|
||||
V=V,
|
||||
BT=BT,
|
||||
BK=BK,
|
||||
BV=BV,
|
||||
IS_VARLEN=cu_seqlens is not None,
|
||||
num_warps=4,
|
||||
num_stages=3,
|
||||
)
|
||||
return w, u
|
||||
|
||||
|
||||
fwd_recompute_w_u = recompute_w_u_fwd
|
||||
Reference in New Issue
Block a user