[AMD] fix dsv4 indexer dtype dispatch on gfx950 (#29479)
Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
@@ -22,6 +22,7 @@ from sglang.srt.layers.attention.dsv4.metadata import (
|
||||
PagedIndexerMetadata,
|
||||
)
|
||||
from sglang.srt.layers.linear import ReplicatedLinear
|
||||
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context import (
|
||||
is_in_breakable_cuda_graph,
|
||||
@@ -44,12 +45,8 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
|
||||
if is_hip():
|
||||
FP8_DTYPE = torch.float8_e4m3fnuz
|
||||
FP8_MAX = torch.finfo(FP8_DTYPE).max
|
||||
else:
|
||||
FP8_DTYPE = torch.float8_e4m3fn
|
||||
FP8_MAX = torch.finfo(FP8_DTYPE).max
|
||||
FP8_DTYPE = torch.float8_e4m3fnuz if is_fp8_fnuz() else torch.float8_e4m3fn
|
||||
|
||||
|
||||
IndexerQuery: TypeAlias = Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]
|
||||
|
||||
|
||||
@@ -58,12 +58,15 @@ import triton
|
||||
import triton.language as tl
|
||||
from aiter.ops.triton.utils.device_info import get_num_sms
|
||||
|
||||
from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz
|
||||
|
||||
LOG2E = 1.4426950408889634 # log2(e); folded into qk_scale so softmax can use exp2.
|
||||
_MAX_KV_SPLITS = 64 # Hard cap on kv_splits (see _kv_splits_heuristic).
|
||||
|
||||
# FP8 KV cache (1xGROUP_SIZE block-scale quantization).
|
||||
#
|
||||
# Storage: unified_kv[total_pages, D] in e4m3fnuz + kv_scales[total_pages,
|
||||
# Storage: unified_kv[total_pages, D] in e4m3fnuz (gfx942 only) / e4m3fn +
|
||||
# kv_scales[total_pages,
|
||||
# D // GROUP_SIZE] in fp32. Per-slot, D is split into NUM_GROUPS chunks of
|
||||
# GROUP_SIZE elements; each chunk shares one fp32 scale.
|
||||
# Dequant in-kernel: kv_bf16 = kv_fp8.to(fp32) * scale[d // GROUP_SIZE], cast
|
||||
@@ -73,7 +76,9 @@ _MAX_KV_SPLITS = 64 # Hard cap on kv_splits (see _kv_splits_heuristic).
|
||||
# per slot, 4 bytes each → +6.25% storage on top of the fp8 pool (vs the
|
||||
# halving from bf16→fp8 = 2× saving — net ~46% read bandwidth reduction).
|
||||
_FP8_GROUP_SIZE = 64
|
||||
_FP8_DTYPE = torch.float8_e4m3fnuz
|
||||
# Match the KV-cache write format: fnuz only on AMD gfx942; Other platform e.g.
|
||||
# gfx950 store e4m3fn. See quant_k_cache.py / indexer.py FP8_DTYPE.
|
||||
_FP8_DTYPE = torch.float8_e4m3fnuz if is_fp8_fnuz() else torch.float8_e4m3fn
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=1)
|
||||
|
||||
Reference in New Issue
Block a user