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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 14 additions & 4 deletions python/sglang/kernels/ops/attention/dsv4/fp4_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -647,14 +647,24 @@ def index_k_rope_pack(
@triton.jit
def _e2m1_decode(code):
# code: uint 0..15 -> e2m1 value. exp = bits 2..1, mantissa = bit 0, sign = bit 3.
# |v| = 0, 0.5, 1, 1.5, 2, 3, 4, 6.
e = (code >> 1) & 3
m = (code & 1).to(tl.float32)
sub = m * 0.5
nor = (1.0 + m * 0.5) * tl.exp2((e - 1).to(tl.float32))
v = tl.where(e == 0, sub, nor)
pow2 = tl.where(e == 3, 4.0, tl.where(e == 2, 2.0, 1.0))
v = tl.where(e == 0, m * 0.5, (1.0 + m * 0.5) * pow2)
return tl.where((code >> 3) == 1, -v, v)


@triton.jit
def _e8m0_decode(code):
# code: uint 0..255 -> 2^(code - 127), written directly as the FP32 exponent
# field. The scorer only sees codes 1..254: the quantizer clamps to that range
# and masked loads use 127. Code 0 gives 0 and 255 gives inf here, unlike
# _ue8m0_to_fp32 in dequant_k_cache.py, which follows torch.float8_e8m0fnu
# (2^-127 and NaN).
return (code.to(tl.int32) << 23).to(tl.float32, bitcast=True)


@triton.jit
def _fp4_index_logits_tile(
q_ptr,
Expand Down Expand Up @@ -698,7 +708,7 @@ def _fp4_index_logits_tile(
mask=valid[:, None],
other=127,
)
scale = tl.exp2(exps.to(tl.float32) - 127.0)
scale = _e8m0_decode(exps)
k_low = (low * scale).to(tl.bfloat16) # [BLOCK_L, HALF_D] elements 2i
k_high = (high * scale).to(tl.bfloat16) # elements 2i+1

Expand Down
26 changes: 26 additions & 0 deletions test/registered/kernels/ops/attention/test_fp4_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@

import pytest
import torch
import triton
import triton.language as tl

from sglang.kernels.ops.attention.deepseek_v4_rope import (
apply_rotary_emb_triton,
Expand All @@ -17,6 +19,8 @@
fused_q_indexer_rope_hadamard_fp4_quant,
)
from sglang.kernels.ops.attention.dsv4.fp4_indexer import (
_e2m1_decode,
_e8m0_decode,
finish_paged_indexer_topk,
fp4_index_logits_decode,
fp4_index_logits_paged,
Expand Down Expand Up @@ -398,6 +402,28 @@ def _reference_logits(q, weights, slots, lens, table):
return logits.masked_fill(position[None, :] >= lens[:, None], -torch.inf)


@triton.jit
def _decode_all_codes_kernel(e2m1_ptr, e8m0_ptr):
code = tl.arange(0, 256).to(tl.uint8)
tl.store(e2m1_ptr + tl.arange(0, 16), _e2m1_decode(tl.arange(0, 16).to(tl.uint8)))
tl.store(e8m0_ptr + tl.arange(0, 256), _e8m0_decode(code))


def test_fp4_scorer_decodes_every_code_exactly():
e2m1 = torch.empty(16, device=get_device(), dtype=torch.float32)
e8m0 = torch.empty(256, device=get_device(), dtype=torch.float32)
_decode_all_codes_kernel[(1,)](e2m1, e8m0)
levels = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0]
# Value equality: the sign of zero (E2M1 code 8) does not reach the scores.
# E8M0 code 0 is never stored and decodes to 0; code 255 overflows to inf.
expected_e2m1 = torch.tensor(levels + [-v for v in levels], dtype=torch.float32)
expected_e8m0 = torch.tensor(
[0.0] + [2.0 ** (e - 127) for e in range(1, 256)], dtype=torch.float32
)
assert torch.equal(e2m1.cpu(), expected_e2m1)
assert torch.equal(e8m0.cpu(), expected_e8m0)


@pytest.mark.parametrize("heads", [32, 64])
@pytest.mark.parametrize("rows", [1, 6, 13])
@pytest.mark.parametrize("width", [0, 1, 63, 64, 65, 129])
Expand Down
Loading