[AMD] enable CUDA graph for NSA backend and fix NSA FP8 fused RMSNorm group quant (#16841)
Co-authored-by: wufann <715544327@qq.com>
This commit is contained in:
@@ -4,6 +4,8 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
|
from sglang.srt.utils import is_hip
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
|
||||||
|
|
||||||
@@ -347,12 +349,17 @@ def _set_k_and_s_triton(
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"index_k_scale must be 1D or 2D, got shape {index_k_scale.shape}"
|
f"index_k_scale must be 1D or 2D, got shape {index_k_scale.shape}"
|
||||||
)
|
)
|
||||||
|
if is_hip():
|
||||||
assert buf_numel_per_page == 64 * (128 + 4)
|
assert buf_numel_per_page == 1 * (128 + 4)
|
||||||
|
else:
|
||||||
|
assert buf_numel_per_page == 64 * (128 + 4)
|
||||||
assert num_tokens_to_write == num_tokens_to_write_ == num_tokens_to_write__
|
assert num_tokens_to_write == num_tokens_to_write_ == num_tokens_to_write__
|
||||||
assert index_head_dim == 128
|
assert index_head_dim == 128
|
||||||
assert scale_dim == 1
|
assert scale_dim == 1
|
||||||
assert page_size == 64
|
if is_hip():
|
||||||
|
assert page_size == 1
|
||||||
|
else:
|
||||||
|
assert page_size == 64
|
||||||
|
|
||||||
assert buf.dtype == torch.uint8
|
assert buf.dtype == torch.uint8
|
||||||
assert loc.dtype == torch.int64, f"{loc.dtype=}" # can be int32
|
assert loc.dtype == torch.int64, f"{loc.dtype=}" # can be int32
|
||||||
|
|||||||
@@ -12,14 +12,16 @@ from sglang.srt.layers.utils import MultiPlatformOp
|
|||||||
from sglang.srt.utils import add_prefix, ceil_align, is_cuda, is_hip, is_npu
|
from sglang.srt.utils import add_prefix, ceil_align, is_cuda, is_hip, is_npu
|
||||||
|
|
||||||
global _use_multi_stream
|
global _use_multi_stream
|
||||||
|
_is_cuda = is_cuda()
|
||||||
if is_cuda():
|
_is_hip = is_hip()
|
||||||
|
_is_npu = is_npu()
|
||||||
|
if _is_cuda:
|
||||||
try:
|
try:
|
||||||
import deep_gemm
|
import deep_gemm
|
||||||
except ImportError as e:
|
except ImportError as e:
|
||||||
deep_gemm = e
|
deep_gemm = e
|
||||||
|
|
||||||
if is_npu():
|
if _is_npu:
|
||||||
import custom_ops # noqa: F401
|
import custom_ops # noqa: F401
|
||||||
import torch_npu
|
import torch_npu
|
||||||
from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream
|
from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream
|
||||||
@@ -42,7 +44,8 @@ from sglang.srt.server_args import get_global_server_args
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool
|
||||||
|
|
||||||
DUAL_STREAM_TOKEN_THRESHOLD = 1024 if is_cuda() else 0
|
|
||||||
|
DUAL_STREAM_TOKEN_THRESHOLD = 1024 if _is_cuda else 0
|
||||||
|
|
||||||
|
|
||||||
class BaseIndexerMetadata(ABC):
|
class BaseIndexerMetadata(ABC):
|
||||||
@@ -59,6 +62,13 @@ class BaseIndexerMetadata(ABC):
|
|||||||
The page size of the table is 64.
|
The page size of the table is 64.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_page_table_1(self) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
Return: (batch_size, num_blocks) int32, page table.
|
||||||
|
The page size of the table is 1.
|
||||||
|
"""
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def get_seqlens_expanded(self) -> torch.Tensor:
|
def get_seqlens_expanded(self) -> torch.Tensor:
|
||||||
"""
|
"""
|
||||||
@@ -101,7 +111,11 @@ class BaseIndexerMetadata(ABC):
|
|||||||
|
|
||||||
def rotate_activation(x: torch.Tensor) -> torch.Tensor:
|
def rotate_activation(x: torch.Tensor) -> torch.Tensor:
|
||||||
assert x.dtype == torch.bfloat16
|
assert x.dtype == torch.bfloat16
|
||||||
from sgl_kernel import hadamard_transform
|
# from sgl_kernel import hadamard_transform
|
||||||
|
if _is_hip:
|
||||||
|
from fast_hadamard_transform import hadamard_transform
|
||||||
|
else:
|
||||||
|
from sgl_kernel import hadamard_transform
|
||||||
|
|
||||||
hidden_size = x.size(-1)
|
hidden_size = x.size(-1)
|
||||||
assert (
|
assert (
|
||||||
@@ -145,7 +159,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
else:
|
else:
|
||||||
self.cp_size = None
|
self.cp_size = None
|
||||||
self.cp_rank = None
|
self.cp_rank = None
|
||||||
if is_cuda():
|
if _is_cuda:
|
||||||
self.sm_count = deep_gemm.get_num_sms()
|
self.sm_count = deep_gemm.get_num_sms()
|
||||||
self.half_device_sm_count = ceil_align(self.sm_count // 2, 8)
|
self.half_device_sm_count = ceil_align(self.sm_count // 2, 8)
|
||||||
pp_size = get_global_server_args().pp_size
|
pp_size = get_global_server_args().pp_size
|
||||||
@@ -205,13 +219,13 @@ class Indexer(MultiPlatformOp):
|
|||||||
else:
|
else:
|
||||||
yield
|
yield
|
||||||
|
|
||||||
@torch.compile(dynamic=True)
|
@torch.compile(dynamic=True) if not _is_hip else lambda f: f
|
||||||
def _project_and_scale_head_gates(self, x: torch.Tensor):
|
def _project_and_scale_head_gates(self, x: torch.Tensor):
|
||||||
weights, _ = self.weights_proj(x.float())
|
weights, _ = self.weights_proj(x.float())
|
||||||
weights = weights * self.n_heads**-0.5
|
weights = weights * self.n_heads**-0.5
|
||||||
return weights
|
return weights
|
||||||
|
|
||||||
@torch.compile(dynamic=True)
|
@torch.compile(dynamic=True) if not _is_hip else lambda f: f
|
||||||
def _get_logits_head_gate(self, x: torch.Tensor, q_scale: torch.Tensor):
|
def _get_logits_head_gate(self, x: torch.Tensor, q_scale: torch.Tensor):
|
||||||
weights, _ = self.weights_proj(x.float())
|
weights, _ = self.weights_proj(x.float())
|
||||||
weights = weights * self.n_heads**-0.5
|
weights = weights * self.n_heads**-0.5
|
||||||
@@ -323,10 +337,13 @@ class Indexer(MultiPlatformOp):
|
|||||||
|
|
||||||
page_size = forward_batch.token_to_kv_pool.page_size
|
page_size = forward_batch.token_to_kv_pool.page_size
|
||||||
# NOTE(dark): blocksize = 64 is hardcoded in deep_gemm
|
# NOTE(dark): blocksize = 64 is hardcoded in deep_gemm
|
||||||
assert page_size == 64, "only support page size 64"
|
if _is_hip:
|
||||||
|
assert page_size == 1, "only support page size 1"
|
||||||
# NOTE(dark): this support extend/decode/decode+graph
|
block_tables = metadata.get_page_table_1()
|
||||||
block_tables = metadata.get_page_table_64()
|
else:
|
||||||
|
assert page_size == 64, "only support page size 64"
|
||||||
|
# NOTE(dark): this support extend/decode/decode+graph
|
||||||
|
block_tables = metadata.get_page_table_64()
|
||||||
|
|
||||||
max_seq_len = block_tables.shape[1] * page_size
|
max_seq_len = block_tables.shape[1] * page_size
|
||||||
kv_cache_fp8 = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer(
|
kv_cache_fp8 = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer(
|
||||||
@@ -344,32 +361,64 @@ class Indexer(MultiPlatformOp):
|
|||||||
# Reuse pre-computed schedule metadata if available (from init_forward_metadata),
|
# Reuse pre-computed schedule metadata if available (from init_forward_metadata),
|
||||||
# otherwise fall back to computing it here.
|
# otherwise fall back to computing it here.
|
||||||
schedule_metadata = getattr(metadata, "paged_mqa_schedule_metadata", None)
|
schedule_metadata = getattr(metadata, "paged_mqa_schedule_metadata", None)
|
||||||
if schedule_metadata is None:
|
if _is_cuda:
|
||||||
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
|
if schedule_metadata is None:
|
||||||
seqlens_32, blocksize, self.sm_count
|
schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata(
|
||||||
)
|
seqlens_32, blocksize, self.sm_count
|
||||||
|
)
|
||||||
|
|
||||||
assert len(q_fp8.shape) == 3
|
assert len(q_fp8.shape) == 3
|
||||||
q_fp8 = q_fp8.unsqueeze(1) # the next_n dim is 1 now
|
q_fp8 = q_fp8.unsqueeze(1) # the next_n dim is 1 now
|
||||||
assert len(kv_cache_fp8.shape) == 2
|
assert len(kv_cache_fp8.shape) == 2
|
||||||
block_kv = 64
|
block_kv = 1 if _is_hip else 64
|
||||||
num_heads_kv = 1
|
num_heads_kv = 1
|
||||||
head_dim_with_sf = 132
|
head_dim_with_sf = 132
|
||||||
kv_cache_fp8 = kv_cache_fp8.view(
|
if _is_hip:
|
||||||
kv_cache_fp8.shape[0], block_kv, num_heads_kv, head_dim_with_sf
|
kv_cache_fp8 = kv_cache_fp8.view(
|
||||||
)
|
-1, block_kv, num_heads_kv, head_dim_with_sf
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
kv_cache_fp8 = kv_cache_fp8.view(
|
||||||
|
kv_cache_fp8.shape[0], block_kv, num_heads_kv, head_dim_with_sf
|
||||||
|
)
|
||||||
assert len(weights.shape) == 3
|
assert len(weights.shape) == 3
|
||||||
weights = weights.squeeze(2)
|
weights = weights.squeeze(2)
|
||||||
logits = deep_gemm.fp8_paged_mqa_logits(
|
|
||||||
q_fp8,
|
if _is_hip:
|
||||||
kv_cache_fp8,
|
from aiter.ops.triton.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits
|
||||||
weights,
|
|
||||||
seqlens_32,
|
batch_size, next_n, heads, _ = q_fp8.shape
|
||||||
block_tables,
|
logits = torch.full(
|
||||||
schedule_metadata,
|
(batch_size * next_n, max_seq_len),
|
||||||
max_seq_len,
|
float("-inf"),
|
||||||
clean_logits=False,
|
device=q_fp8.device,
|
||||||
)
|
dtype=torch.float32,
|
||||||
|
)
|
||||||
|
deepgemm_fp8_paged_mqa_logits(
|
||||||
|
q_fp8,
|
||||||
|
kv_cache_fp8,
|
||||||
|
weights,
|
||||||
|
logits,
|
||||||
|
seqlens_32,
|
||||||
|
block_tables,
|
||||||
|
max_seq_len,
|
||||||
|
Preshuffle=False,
|
||||||
|
KVBlockSize=block_kv,
|
||||||
|
ChunkK=128,
|
||||||
|
TotalCuCount=256,
|
||||||
|
WavePerEU=5,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logits = deep_gemm.fp8_paged_mqa_logits(
|
||||||
|
q_fp8,
|
||||||
|
kv_cache_fp8,
|
||||||
|
weights,
|
||||||
|
seqlens_32,
|
||||||
|
block_tables,
|
||||||
|
schedule_metadata,
|
||||||
|
max_seq_len,
|
||||||
|
clean_logits=False,
|
||||||
|
)
|
||||||
|
|
||||||
# NOTE(dark): logits should be cleaned in topk_transform
|
# NOTE(dark): logits should be cleaned in topk_transform
|
||||||
topk_result = metadata.topk_transform(logits, self.index_topk)
|
topk_result = metadata.topk_transform(logits, self.index_topk)
|
||||||
@@ -408,13 +457,20 @@ class Indexer(MultiPlatformOp):
|
|||||||
assert forward_batch.forward_mode.is_extend_without_speculative()
|
assert forward_batch.forward_mode.is_extend_without_speculative()
|
||||||
|
|
||||||
page_size = forward_batch.token_to_kv_pool.page_size
|
page_size = forward_batch.token_to_kv_pool.page_size
|
||||||
assert page_size == 64, "only support page size 64"
|
if _is_hip:
|
||||||
|
assert page_size == 1, "only support page size 1"
|
||||||
|
else:
|
||||||
|
assert page_size == 64, "only support page size 64"
|
||||||
|
|
||||||
assert len(weights.shape) == 3
|
assert len(weights.shape) == 3
|
||||||
weights = weights.squeeze(-1)
|
weights = weights.squeeze(-1)
|
||||||
k_fp8_list = []
|
k_fp8_list = []
|
||||||
k_scale_list = []
|
k_scale_list = []
|
||||||
|
|
||||||
block_tables = metadata.get_page_table_64()
|
if _is_hip:
|
||||||
|
block_tables = metadata.get_page_table_1()
|
||||||
|
else:
|
||||||
|
block_tables = metadata.get_page_table_64()
|
||||||
|
|
||||||
assert (
|
assert (
|
||||||
forward_batch.seq_lens_cpu is not None
|
forward_batch.seq_lens_cpu is not None
|
||||||
@@ -459,14 +515,22 @@ class Indexer(MultiPlatformOp):
|
|||||||
if not need_chunk:
|
if not need_chunk:
|
||||||
assert q_fp8[:q_offset].shape[0] != 0
|
assert q_fp8[:q_offset].shape[0] != 0
|
||||||
with self._with_real_sm_count():
|
with self._with_real_sm_count():
|
||||||
logits = deep_gemm.fp8_mqa_logits(
|
if _is_hip:
|
||||||
q_fp8[:q_offset],
|
from aiter.ops.triton.fp8_mqa_logits import fp8_mqa_logits
|
||||||
kv_fp8,
|
|
||||||
weights[:q_offset],
|
kv, scale = kv_fp8
|
||||||
ks,
|
logits = fp8_mqa_logits(
|
||||||
ke,
|
q_fp8[:q_offset], kv, scale, weights[:q_offset], ks, ke
|
||||||
clean_logits=False,
|
)
|
||||||
)
|
else:
|
||||||
|
logits = deep_gemm.fp8_mqa_logits(
|
||||||
|
q_fp8[:q_offset],
|
||||||
|
kv_fp8,
|
||||||
|
weights[:q_offset],
|
||||||
|
ks,
|
||||||
|
ke,
|
||||||
|
clean_logits=False,
|
||||||
|
)
|
||||||
assert logits.shape[0] == len(seq_lens_expanded)
|
assert logits.shape[0] == len(seq_lens_expanded)
|
||||||
assert logits.shape[1] == k_offset
|
assert logits.shape[1] == k_offset
|
||||||
|
|
||||||
@@ -496,14 +560,27 @@ class Indexer(MultiPlatformOp):
|
|||||||
end = min(start + max_rows, q_offset)
|
end = min(start + max_rows, q_offset)
|
||||||
|
|
||||||
with self._with_real_sm_count():
|
with self._with_real_sm_count():
|
||||||
logits_chunk = deep_gemm.fp8_mqa_logits(
|
if _is_hip:
|
||||||
q_fp8[start:end],
|
from aiter.ops.triton.fp8_mqa_logits import fp8_mqa_logits
|
||||||
kv_fp8,
|
|
||||||
weights[start:end],
|
kv, scale = kv_fp8
|
||||||
ks[start:end],
|
logits = fp8_mqa_logits(
|
||||||
ke[start:end],
|
q_fp8[start:end],
|
||||||
clean_logits=False,
|
kv_fp8,
|
||||||
)
|
scale,
|
||||||
|
weights[start:end],
|
||||||
|
ks[start:end],
|
||||||
|
ke[start:end],
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logits_chunk = deep_gemm.fp8_mqa_logits(
|
||||||
|
q_fp8[start:end],
|
||||||
|
kv_fp8,
|
||||||
|
weights[start:end],
|
||||||
|
ks[start:end],
|
||||||
|
ke[start:end],
|
||||||
|
clean_logits=False,
|
||||||
|
)
|
||||||
|
|
||||||
lengths_chunk = seq_lens_expanded[start:end]
|
lengths_chunk = seq_lens_expanded[start:end]
|
||||||
|
|
||||||
@@ -548,6 +625,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
return_indices: bool = True,
|
return_indices: bool = True,
|
||||||
) -> Optional[torch.Tensor]:
|
) -> Optional[torch.Tensor]:
|
||||||
assert forward_batch.forward_mode.is_extend_without_speculative()
|
assert forward_batch.forward_mode.is_extend_without_speculative()
|
||||||
|
x_meta = x[0] if isinstance(x, tuple) else x
|
||||||
|
|
||||||
# Fast path: only compute and store k cache, skip all q and weights ops
|
# Fast path: only compute and store k cache, skip all q and weights ops
|
||||||
key = self._get_k_bf16(x, positions, enable_dual_stream)
|
key = self._get_k_bf16(x, positions, enable_dual_stream)
|
||||||
@@ -573,7 +651,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
seq_lens_expanded.shape[0],
|
seq_lens_expanded.shape[0],
|
||||||
self.index_topk,
|
self.index_topk,
|
||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
device=x.device,
|
device=x_meta.device,
|
||||||
)
|
)
|
||||||
return metadata.topk_transform(dummy_logits, self.index_topk)
|
return metadata.topk_transform(dummy_logits, self.index_topk)
|
||||||
|
|
||||||
@@ -734,7 +812,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
topk: int,
|
topk: int,
|
||||||
layer_id: int,
|
layer_id: int,
|
||||||
) -> Optional[torch.Tensor]:
|
) -> Optional[torch.Tensor]:
|
||||||
if not is_npu():
|
if not _is_npu:
|
||||||
from sglang.srt.layers.attention.nsa.tilelang_kernel import fp8_index
|
from sglang.srt.layers.attention.nsa.tilelang_kernel import fp8_index
|
||||||
|
|
||||||
page_size = forward_batch.token_to_kv_pool.page_size
|
page_size = forward_batch.token_to_kv_pool.page_size
|
||||||
@@ -818,14 +896,18 @@ class Indexer(MultiPlatformOp):
|
|||||||
layer_id: int,
|
layer_id: int,
|
||||||
return_indices: bool = True,
|
return_indices: bool = True,
|
||||||
) -> Optional[torch.Tensor]:
|
) -> Optional[torch.Tensor]:
|
||||||
if is_hip():
|
if _is_hip:
|
||||||
from sglang.srt.layers.attention.nsa.tilelang_kernel import act_quant
|
from sglang.srt.layers.attention.nsa.tilelang_kernel import act_quant
|
||||||
elif not is_npu():
|
elif not _is_npu:
|
||||||
from sglang.srt.layers.attention.nsa.triton_kernel import act_quant
|
from sglang.srt.layers.attention.nsa.triton_kernel import act_quant
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
assert isinstance(forward_batch.token_to_kv_pool, NSATokenToKVPool)
|
assert isinstance(forward_batch.token_to_kv_pool, NSATokenToKVPool)
|
||||||
|
|
||||||
|
# When upstream uses fused FP8 RMSNorm+quant, activations may be passed as
|
||||||
|
# a tuple like (x_fp8, x_scale[, y]). Use `x_meta` for shape/device queries.
|
||||||
|
x_meta = x[0] if isinstance(x, tuple) else x
|
||||||
|
|
||||||
metadata = forward_batch.attn_backend.get_indexer_metadata(
|
metadata = forward_batch.attn_backend.get_indexer_metadata(
|
||||||
layer_id, forward_batch
|
layer_id, forward_batch
|
||||||
)
|
)
|
||||||
@@ -891,7 +973,38 @@ class Indexer(MultiPlatformOp):
|
|||||||
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt)
|
||||||
k_fp8, k_scale = act_quant(key, self.block_size, self.scale_fmt)
|
k_fp8, k_scale = act_quant(key, self.block_size, self.scale_fmt)
|
||||||
|
|
||||||
weights = self._get_logits_head_gate(x, q_scale)
|
# `_get_logits_head_gate` expects a Tensor. For tuple activations, dequantize
|
||||||
|
# to a float tensor here (callsite), keeping `_get_logits_head_gate` backend-agnostic.
|
||||||
|
if isinstance(x, tuple):
|
||||||
|
assert len(x) in (
|
||||||
|
2,
|
||||||
|
3,
|
||||||
|
), "For tuple input, only (x, x_s) or (x, x_s, y) formats are accepted"
|
||||||
|
x_q, x_s = x[0], x[1]
|
||||||
|
if (
|
||||||
|
x_s is not None
|
||||||
|
and x_q.dim() == 2
|
||||||
|
and x_s.dim() == 2
|
||||||
|
and x_q.shape[0] == x_s.shape[0]
|
||||||
|
):
|
||||||
|
m, n = x_q.shape
|
||||||
|
ng = x_s.shape[1]
|
||||||
|
if ng > 0 and n % ng == 0:
|
||||||
|
group = n // ng
|
||||||
|
x_for_gate = (
|
||||||
|
x_q.to(torch.float32)
|
||||||
|
.view(m, ng, group)
|
||||||
|
.mul_(x_s.to(torch.float32).unsqueeze(-1))
|
||||||
|
.view(m, n)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
x_for_gate = x_q.to(torch.float32)
|
||||||
|
else:
|
||||||
|
x_for_gate = x_q.to(torch.float32)
|
||||||
|
else:
|
||||||
|
x_for_gate = x
|
||||||
|
|
||||||
|
weights = self._get_logits_head_gate(x_for_gate, q_scale)
|
||||||
|
|
||||||
# k_fp8: (seq_len, head_dim) fp8_e4m3fn
|
# k_fp8: (seq_len, head_dim) fp8_e4m3fn
|
||||||
# k_buffer: (num_total_tokens + page_size, head_dim) fp8_e4m3fn
|
# k_buffer: (num_total_tokens + page_size, head_dim) fp8_e4m3fn
|
||||||
@@ -906,7 +1019,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
index_k_scale=k_scale,
|
index_k_scale=k_scale,
|
||||||
)
|
)
|
||||||
|
|
||||||
if is_cuda():
|
if _is_cuda or _is_hip:
|
||||||
assert forward_batch.seq_lens_cpu is not None
|
assert forward_batch.seq_lens_cpu is not None
|
||||||
if len(forward_batch.seq_lens_cpu) == 0:
|
if len(forward_batch.seq_lens_cpu) == 0:
|
||||||
# this seems b/c max-pad, no worries?
|
# this seems b/c max-pad, no worries?
|
||||||
@@ -915,7 +1028,10 @@ class Indexer(MultiPlatformOp):
|
|||||||
# "HACK: seq_lens empty but x not empty, hackily return all-invalid topk_result"
|
# "HACK: seq_lens empty but x not empty, hackily return all-invalid topk_result"
|
||||||
# )
|
# )
|
||||||
return torch.full(
|
return torch.full(
|
||||||
(x.shape[0], self.index_topk), -1, dtype=torch.int, device="cuda"
|
(x_meta.shape[0], self.index_topk),
|
||||||
|
-1,
|
||||||
|
dtype=torch.int,
|
||||||
|
device=x_meta.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -147,6 +147,8 @@ def fp8_index_kernel(h: int, d: int, clear_accum=True):
|
|||||||
T.copy(k_s[i_b, i1_n * blk_n1 + i2_n * blk_n2], k_s_frag)
|
T.copy(k_s[i_b, i1_n * blk_n1 + i2_n * blk_n2], k_s_frag)
|
||||||
|
|
||||||
logits = T.alloc_fragment((blk_n2, h), FP32)
|
logits = T.alloc_fragment((blk_n2, h), FP32)
|
||||||
|
if not clear_accum:
|
||||||
|
T.fill(logits, 0)
|
||||||
T.gemm(
|
T.gemm(
|
||||||
k_smem,
|
k_smem,
|
||||||
q_smem,
|
q_smem,
|
||||||
|
|||||||
@@ -174,6 +174,9 @@ class NSAIndexerMetadata(BaseIndexerMetadata):
|
|||||||
def get_page_table_64(self) -> torch.Tensor:
|
def get_page_table_64(self) -> torch.Tensor:
|
||||||
return self.attn_metadata.real_page_table
|
return self.attn_metadata.real_page_table
|
||||||
|
|
||||||
|
def get_page_table_1(self) -> torch.Tensor:
|
||||||
|
return self.attn_metadata.page_table_1
|
||||||
|
|
||||||
def get_seqlens_expanded(self) -> torch.Tensor:
|
def get_seqlens_expanded(self) -> torch.Tensor:
|
||||||
return self.attn_metadata.nsa_seqlens_expanded
|
return self.attn_metadata.nsa_seqlens_expanded
|
||||||
|
|
||||||
|
|||||||
@@ -52,7 +52,7 @@ from sglang.srt.mem_cache.utils import (
|
|||||||
set_mla_kv_buffer_triton,
|
set_mla_kv_buffer_triton,
|
||||||
set_mla_kv_scale_buffer_triton,
|
set_mla_kv_scale_buffer_triton,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import is_cuda, is_npu, next_power_of_2
|
from sglang.srt.utils import is_cuda, is_hip, is_npu, next_power_of_2
|
||||||
from sglang.srt.utils.custom_op import register_custom_op
|
from sglang.srt.utils.custom_op import register_custom_op
|
||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||||
|
|
||||||
@@ -68,6 +68,7 @@ logger = logging.getLogger(__name__)
|
|||||||
GB = 1024 * 1024 * 1024
|
GB = 1024 * 1024 * 1024
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
|
_is_hip = is_hip()
|
||||||
|
|
||||||
|
|
||||||
def get_tensor_size_bytes(t: Union[torch.Tensor, List[torch.Tensor]]):
|
def get_tensor_size_bytes(t: Union[torch.Tensor, List[torch.Tensor]]):
|
||||||
@@ -1724,7 +1725,10 @@ class NSATokenToKVPool(MLATokenToKVPool):
|
|||||||
# num head == 1 and head dim == 128 for index_k in NSA
|
# num head == 1 and head dim == 128 for index_k in NSA
|
||||||
assert index_head_dim == 128
|
assert index_head_dim == 128
|
||||||
|
|
||||||
assert self.page_size == 64
|
if _is_hip:
|
||||||
|
assert self.page_size == 1
|
||||||
|
else:
|
||||||
|
assert self.page_size == 64
|
||||||
with (
|
with (
|
||||||
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
torch.cuda.use_mem_pool(self.custom_mem_pool)
|
||||||
if self.custom_mem_pool
|
if self.custom_mem_pool
|
||||||
|
|||||||
@@ -1523,10 +1523,34 @@ class DeepseekV2AttentionMLA(nn.Module):
|
|||||||
# NSA Indexer: cache quantized keys, auto-skip topk for sequences <= nsa_index_topk
|
# NSA Indexer: cache quantized keys, auto-skip topk for sequences <= nsa_index_topk
|
||||||
|
|
||||||
if self.use_nsa:
|
if self.use_nsa:
|
||||||
q_lora = self.q_a_layernorm(q)
|
# NSA requires unquantized q_lora for the indexer. When q_b_proj is FP8
|
||||||
q = self.q_b_proj(q_lora)[0].view(
|
# on gfx95, we can still use fused RMSNorm+FP8 quant, but MUST request
|
||||||
-1, self.num_local_heads, self.qk_head_dim
|
# the unquantized output for q_lora; otherwise q_lora becomes the (fp8,scale)
|
||||||
)
|
# tuple.
|
||||||
|
if (
|
||||||
|
_use_aiter_gfx95
|
||||||
|
and self.q_b_proj.weight.dtype == torch.float8_e4m3fn
|
||||||
|
):
|
||||||
|
q_quanted, q_lora, _, _ = fused_rms_fp8_group_quant(
|
||||||
|
q,
|
||||||
|
self.q_a_layernorm.weight,
|
||||||
|
self.q_a_layernorm.variance_epsilon,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
group_size=128,
|
||||||
|
dtype_quant=torch.float8_e4m3fn,
|
||||||
|
res1=None,
|
||||||
|
output_unquantized_inp1=True,
|
||||||
|
)
|
||||||
|
q = self.q_b_proj(q_quanted)[0].view(
|
||||||
|
-1, self.num_local_heads, self.qk_head_dim
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
q_lora = self.q_a_layernorm(q)
|
||||||
|
q = self.q_b_proj(q_lora)[0].view(
|
||||||
|
-1, self.num_local_heads, self.qk_head_dim
|
||||||
|
)
|
||||||
_ = self.indexer(
|
_ = self.indexer(
|
||||||
x=hidden_states,
|
x=hidden_states,
|
||||||
q_lora=q_lora,
|
q_lora=q_lora,
|
||||||
@@ -1703,23 +1727,38 @@ class DeepseekV2AttentionMLA(nn.Module):
|
|||||||
self.kv_a_layernorm.variance_epsilon,
|
self.kv_a_layernorm.variance_epsilon,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
q_lora = None
|
||||||
if (
|
if (
|
||||||
_use_aiter_gfx95
|
_use_aiter_gfx95
|
||||||
and self.q_b_proj.weight.dtype == torch.float8_e4m3fn
|
and self.q_b_proj.weight.dtype == torch.float8_e4m3fn
|
||||||
):
|
):
|
||||||
|
if self.use_nsa:
|
||||||
q, _, k_nope, _ = fused_rms_fp8_group_quant(
|
q_quanted, q_lora, k_nope, _ = fused_rms_fp8_group_quant(
|
||||||
q,
|
q,
|
||||||
self.q_a_layernorm.weight,
|
self.q_a_layernorm.weight,
|
||||||
self.q_a_layernorm.variance_epsilon,
|
self.q_a_layernorm.variance_epsilon,
|
||||||
k_nope,
|
k_nope,
|
||||||
self.kv_a_layernorm.weight,
|
self.kv_a_layernorm.weight,
|
||||||
self.kv_a_layernorm.variance_epsilon,
|
self.kv_a_layernorm.variance_epsilon,
|
||||||
group_size=128,
|
group_size=128,
|
||||||
dtype_quant=torch.float8_e4m3fn,
|
dtype_quant=torch.float8_e4m3fn,
|
||||||
res1=None,
|
res1=None,
|
||||||
output_unquantized_inp1=False,
|
output_unquantized_inp1=True,
|
||||||
)
|
)
|
||||||
|
q = q_quanted
|
||||||
|
else:
|
||||||
|
q, _, k_nope, _ = fused_rms_fp8_group_quant(
|
||||||
|
q,
|
||||||
|
self.q_a_layernorm.weight,
|
||||||
|
self.q_a_layernorm.variance_epsilon,
|
||||||
|
k_nope,
|
||||||
|
self.kv_a_layernorm.weight,
|
||||||
|
self.kv_a_layernorm.variance_epsilon,
|
||||||
|
group_size=128,
|
||||||
|
dtype_quant=torch.float8_e4m3fn,
|
||||||
|
res1=None,
|
||||||
|
output_unquantized_inp1=False,
|
||||||
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
q = self.q_a_layernorm(q)
|
q = self.q_a_layernorm(q)
|
||||||
@@ -1727,7 +1766,8 @@ class DeepseekV2AttentionMLA(nn.Module):
|
|||||||
|
|
||||||
# q_lora needed by indexer
|
# q_lora needed by indexer
|
||||||
if self.use_nsa:
|
if self.use_nsa:
|
||||||
q_lora = q
|
if q_lora is None:
|
||||||
|
q_lora = q
|
||||||
|
|
||||||
# overlap q_b_proj and indexer during decode
|
# overlap q_b_proj and indexer during decode
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -1090,7 +1090,7 @@ class ServerArgs:
|
|||||||
self.attention_backend = "nsa"
|
self.attention_backend = "nsa"
|
||||||
logger.info("Use nsa attention backend for DeepSeek with DSA.")
|
logger.info("Use nsa attention backend for DeepSeek with DSA.")
|
||||||
|
|
||||||
if not is_npu(): # CUDA GPU
|
if not is_npu(): # CUDA or ROCm GPU
|
||||||
if self.enable_nsa_prefill_context_parallel:
|
if self.enable_nsa_prefill_context_parallel:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"Context parallel feature is still under experiment. It has only been verified on Hopper platform."
|
f"Context parallel feature is still under experiment. It has only been verified on Hopper platform."
|
||||||
@@ -1126,8 +1126,15 @@ class ServerArgs:
|
|||||||
f"attn_tp_size={self.tp_size}, attention weights will be sharded across {self.tp_size} ranks."
|
f"attn_tp_size={self.tp_size}, attention weights will be sharded across {self.tp_size} ranks."
|
||||||
)
|
)
|
||||||
|
|
||||||
self.page_size = 64
|
if is_hip():
|
||||||
logger.warning("Setting page size to 64 for DeepSeek DSA.")
|
self.page_size = 1
|
||||||
|
logger.warning(
|
||||||
|
"Setting page size to 1 for DeepSeek DSA on ROCm."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# For CUDA GPU
|
||||||
|
self.page_size = 64
|
||||||
|
logger.warning("Setting page size to 64 for DeepSeek DSA.")
|
||||||
|
|
||||||
# For Hopper, we support both bf16 and fp8 kv cache; for Blackwell, we support fp8 only currently
|
# For Hopper, we support both bf16 and fp8 kv cache; for Blackwell, we support fp8 only currently
|
||||||
import torch
|
import torch
|
||||||
|
|||||||
Reference in New Issue
Block a user