[FlashInfer v0.6.10] [RL] [DSv32] [GLM-5] Add --dsa-topk-backend and integrate FlashInfer and pytorch topk (#22851)
This commit is contained in:
@@ -465,6 +465,8 @@ class Envs:
|
||||
|
||||
# DSA Backend (canonical names; fall back to SGLANG_NSA_* with deprecation warning)
|
||||
SGLANG_DSA_FUSE_TOPK = EnvBoolWithAlias(True, deprecated_name="SGLANG_NSA_FUSE_TOPK")
|
||||
SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC = EnvBool(False)
|
||||
SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK = EnvStr(None)
|
||||
SGLANG_DSA_ENABLE_MTP_PRECOMPUTE_METADATA = EnvBoolWithAlias(
|
||||
True, deprecated_name="SGLANG_NSA_ENABLE_MTP_PRECOMPUTE_METADATA"
|
||||
)
|
||||
|
||||
@@ -0,0 +1,271 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from enum import Enum, IntEnum, auto
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
_FLASHINFER_TIE_BREAK_VALUES = {
|
||||
"small": 1,
|
||||
"large": 2,
|
||||
}
|
||||
|
||||
|
||||
class TopkTransformMethod(IntEnum):
|
||||
# Transform topk indices to indices to the page table (page_size = 1)
|
||||
PAGED = auto()
|
||||
# Transform topk indices to indices to ragged kv (non-paged)
|
||||
RAGGED = auto()
|
||||
|
||||
|
||||
class DSATopKBackend(Enum):
|
||||
SGL_KERNEL = "sgl-kernel"
|
||||
TORCH = "torch"
|
||||
FLASHINFER = "flashinfer"
|
||||
|
||||
def is_sgl_kernel(self) -> bool:
|
||||
return self == DSATopKBackend.SGL_KERNEL
|
||||
|
||||
def is_torch(self) -> bool:
|
||||
return self == DSATopKBackend.TORCH
|
||||
|
||||
def is_flashinfer(self) -> bool:
|
||||
return self == DSATopKBackend.FLASHINFER
|
||||
|
||||
def topk_func(
|
||||
self,
|
||||
score: torch.Tensor,
|
||||
lengths: torch.Tensor,
|
||||
topk: int,
|
||||
row_starts: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if self.is_sgl_kernel():
|
||||
from sgl_kernel import fast_topk_v2
|
||||
|
||||
return fast_topk_v2(score, lengths, topk, row_starts=row_starts)
|
||||
if self.is_torch():
|
||||
return _topk_unfused(
|
||||
score,
|
||||
lengths,
|
||||
topk,
|
||||
row_starts=row_starts,
|
||||
topk_op=torch.topk,
|
||||
topk_op_kwargs={"dim": -1},
|
||||
)
|
||||
if self.is_flashinfer():
|
||||
import flashinfer
|
||||
|
||||
return _topk_unfused(
|
||||
score,
|
||||
lengths,
|
||||
topk,
|
||||
row_starts=row_starts,
|
||||
topk_op=flashinfer.top_k,
|
||||
topk_op_kwargs={
|
||||
"sorted": False,
|
||||
"deterministic": envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(),
|
||||
"tie_break": _flashinfer_tie_break_value(),
|
||||
"dsa_graph_safe": True,
|
||||
},
|
||||
)
|
||||
raise RuntimeError(f"Unsupported {self = }.")
|
||||
|
||||
def topk_transform(
|
||||
self,
|
||||
logits: torch.Tensor,
|
||||
lengths: torch.Tensor,
|
||||
topk: int,
|
||||
topk_transform_method: TopkTransformMethod,
|
||||
attn_metadata,
|
||||
cu_seqlens_q_topk: Optional[torch.Tensor] = None,
|
||||
topk_indices_offset: Optional[torch.Tensor] = None,
|
||||
row_starts: Optional[torch.Tensor] = None,
|
||||
batch_idx_list: Optional[List[int]] = None,
|
||||
force_unfused_topk: bool = False,
|
||||
) -> torch.Tensor:
|
||||
if not envs.SGLANG_DSA_FUSE_TOPK.get() or force_unfused_topk:
|
||||
return self.topk_func(logits, lengths, topk, row_starts=row_starts)
|
||||
|
||||
if self.is_sgl_kernel():
|
||||
from sgl_kernel import (
|
||||
fast_topk_transform_fused,
|
||||
fast_topk_transform_ragged_fused,
|
||||
)
|
||||
|
||||
if topk_transform_method == TopkTransformMethod.PAGED:
|
||||
page_table_size_1 = (
|
||||
attn_metadata.page_table_1[batch_idx_list]
|
||||
if batch_idx_list is not None
|
||||
else attn_metadata.page_table_1
|
||||
)
|
||||
return fast_topk_transform_fused(
|
||||
score=logits,
|
||||
lengths=lengths,
|
||||
page_table_size_1=page_table_size_1,
|
||||
cu_seqlens_q=cu_seqlens_q_topk,
|
||||
topk=topk,
|
||||
row_starts=row_starts,
|
||||
)
|
||||
if topk_transform_method == TopkTransformMethod.RAGGED:
|
||||
if topk_indices_offset is None:
|
||||
raise RuntimeError(
|
||||
"RAGGED topk_transform requires topk_indices_offset; "
|
||||
"expected extend-without-speculative metadata."
|
||||
)
|
||||
return fast_topk_transform_ragged_fused(
|
||||
score=logits,
|
||||
lengths=lengths,
|
||||
topk_indices_offset=topk_indices_offset,
|
||||
topk=topk,
|
||||
row_starts=row_starts,
|
||||
)
|
||||
raise RuntimeError(f"Unsupported {topk_transform_method = }.")
|
||||
|
||||
if self.is_flashinfer():
|
||||
import flashinfer
|
||||
|
||||
if topk_transform_method == TopkTransformMethod.PAGED:
|
||||
row_to_batch, local_row_starts = _build_flashinfer_paged_args(
|
||||
attn_metadata=attn_metadata,
|
||||
row_starts=row_starts,
|
||||
cu_seqlens_q_topk=cu_seqlens_q_topk,
|
||||
batch_idx_list=batch_idx_list,
|
||||
device=logits.device,
|
||||
num_rows=logits.shape[0],
|
||||
)
|
||||
return flashinfer.top_k_page_table_transform(
|
||||
logits.contiguous(),
|
||||
attn_metadata.page_table_1.contiguous(),
|
||||
lengths.contiguous(),
|
||||
topk,
|
||||
row_to_batch=row_to_batch,
|
||||
deterministic=envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(),
|
||||
tie_break=_flashinfer_tie_break_value(),
|
||||
dsa_graph_safe=True,
|
||||
row_starts=local_row_starts,
|
||||
)
|
||||
if topk_transform_method == TopkTransformMethod.RAGGED:
|
||||
if topk_indices_offset is None:
|
||||
raise RuntimeError(
|
||||
"RAGGED topk_transform requires topk_indices_offset; "
|
||||
"expected extend-without-speculative metadata."
|
||||
)
|
||||
return flashinfer.top_k_ragged_transform(
|
||||
logits.contiguous(),
|
||||
topk_indices_offset.contiguous(),
|
||||
lengths.contiguous(),
|
||||
topk,
|
||||
deterministic=envs.SGLANG_DSA_TOPK_FLASHINFER_DETERMINISTIC.get(),
|
||||
tie_break=_flashinfer_tie_break_value(),
|
||||
dsa_graph_safe=True,
|
||||
row_starts=row_starts,
|
||||
)
|
||||
raise RuntimeError(f"Unsupported {topk_transform_method = }.")
|
||||
|
||||
raise RuntimeError(f"Unsupported {self = } for SGLANG_DSA_FUSE_TOPK.")
|
||||
|
||||
|
||||
def _topk_unfused(
|
||||
score: torch.Tensor,
|
||||
lengths: torch.Tensor,
|
||||
topk: int,
|
||||
row_starts: Optional[torch.Tensor] = None,
|
||||
topk_op: Callable[..., Tuple[torch.Tensor, torch.Tensor]] = torch.topk,
|
||||
topk_op_kwargs: Optional[Dict[str, object]] = None,
|
||||
) -> torch.Tensor:
|
||||
batch_size, max_score_len = score.shape
|
||||
topk_indices = score.new_full((batch_size, topk), -1, dtype=torch.int32)
|
||||
if batch_size == 0 or topk == 0 or max_score_len == 0:
|
||||
return topk_indices
|
||||
|
||||
if row_starts is None:
|
||||
row_starts = torch.zeros_like(lengths, dtype=torch.int32, device=score.device)
|
||||
else:
|
||||
row_starts = row_starts.to(dtype=torch.int32, device=score.device)
|
||||
lengths = lengths.to(dtype=torch.int32, device=score.device)
|
||||
|
||||
col_indices = torch.arange(max_score_len, dtype=torch.int32, device=score.device)
|
||||
col_indices = col_indices.unsqueeze(0)
|
||||
row_starts_unsqueezed = row_starts.unsqueeze(1)
|
||||
row_ends_unsqueezed = (row_starts + lengths).unsqueeze(1)
|
||||
valid_mask = (col_indices >= row_starts_unsqueezed) & (
|
||||
col_indices < row_ends_unsqueezed
|
||||
)
|
||||
|
||||
masked_logits = score.masked_fill(~valid_mask, float("-inf"))
|
||||
valid_topk = min(topk, max_score_len)
|
||||
topk_kwargs = topk_op_kwargs or {}
|
||||
topk_scores, topk_col_indices = topk_op(masked_logits, valid_topk, **topk_kwargs)
|
||||
topk_local_indices = topk_col_indices.to(torch.int32) - row_starts_unsqueezed
|
||||
topk_local_indices = topk_local_indices.masked_fill(
|
||||
topk_scores == float("-inf"), -1
|
||||
)
|
||||
topk_indices[:, :valid_topk] = topk_local_indices
|
||||
|
||||
return topk_indices
|
||||
|
||||
|
||||
def _build_flashinfer_paged_args(
|
||||
attn_metadata,
|
||||
row_starts: Optional[torch.Tensor],
|
||||
cu_seqlens_q_topk: Optional[torch.Tensor],
|
||||
batch_idx_list: Optional[List[int]],
|
||||
device: torch.device,
|
||||
num_rows: int,
|
||||
) -> Tuple[Optional[torch.Tensor], Optional[torch.Tensor]]:
|
||||
row_to_batch = (
|
||||
torch.as_tensor(batch_idx_list, dtype=torch.int32, device=device)
|
||||
if batch_idx_list is not None
|
||||
else None
|
||||
)
|
||||
|
||||
if (
|
||||
row_to_batch is not None
|
||||
and cu_seqlens_q_topk is not None
|
||||
and row_to_batch.shape[0] != num_rows
|
||||
):
|
||||
q_lens = (cu_seqlens_q_topk[1:] - cu_seqlens_q_topk[:-1]).to(
|
||||
dtype=torch.int32, device=device
|
||||
)
|
||||
row_to_batch = torch.repeat_interleave(row_to_batch, q_lens)
|
||||
|
||||
if row_to_batch is None and cu_seqlens_q_topk is not None:
|
||||
# Decode-like case (one query row per batch) does not need an explicit mapping.
|
||||
# Avoid dynamic tensor construction in this branch to keep CUDA graph capture safe.
|
||||
num_batches = cu_seqlens_q_topk.shape[0] - 1
|
||||
if not (row_starts is None and num_rows == num_batches):
|
||||
q_lens = (cu_seqlens_q_topk[1:] - cu_seqlens_q_topk[:-1]).to(
|
||||
dtype=torch.int32, device=device
|
||||
)
|
||||
row_to_batch = torch.repeat_interleave(
|
||||
torch.arange(q_lens.shape[0], dtype=torch.int32, device=device),
|
||||
q_lens,
|
||||
)
|
||||
|
||||
if row_starts is not None and row_to_batch is None:
|
||||
raise RuntimeError(
|
||||
"PAGED topk_transform with row_starts requires cu_seqlens_q metadata."
|
||||
)
|
||||
|
||||
local_row_starts = row_starts
|
||||
if local_row_starts is not None and row_to_batch is not None:
|
||||
local_row_starts = (
|
||||
local_row_starts - attn_metadata.cu_seqlens_k[:-1][row_to_batch]
|
||||
)
|
||||
|
||||
return row_to_batch, local_row_starts
|
||||
|
||||
|
||||
def _flashinfer_tie_break_value() -> int:
|
||||
mode = envs.SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK.get()
|
||||
if mode is None:
|
||||
return 0
|
||||
mode = mode.lower()
|
||||
if mode not in _FLASHINFER_TIE_BREAK_VALUES:
|
||||
raise RuntimeError(
|
||||
"SGLANG_DSA_TOPK_FLASHINFER_TIE_BREAK must be one of "
|
||||
f"{tuple(_FLASHINFER_TIE_BREAK_VALUES)} or unset, got {mode!r}."
|
||||
)
|
||||
return _FLASHINFER_TIE_BREAK_VALUES[mode]
|
||||
@@ -2,8 +2,15 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum, auto
|
||||
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Tuple, TypeAlias
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Dict,
|
||||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Tuple,
|
||||
TypeAlias,
|
||||
)
|
||||
|
||||
import torch
|
||||
|
||||
@@ -19,6 +26,10 @@ from sglang.srt.layers.attention.dsa.dsa_backend_mtp_precompute import (
|
||||
compute_cu_seqlens,
|
||||
)
|
||||
from sglang.srt.layers.attention.dsa.dsa_indexer import BaseIndexerMetadata
|
||||
from sglang.srt.layers.attention.dsa.dsa_topk_backend import (
|
||||
DSATopKBackend,
|
||||
TopkTransformMethod,
|
||||
)
|
||||
from sglang.srt.layers.attention.dsa.quant_k_cache import quantize_k_cache
|
||||
from sglang.srt.layers.attention.dsa.transform_index import (
|
||||
transform_index_page_table_decode,
|
||||
@@ -161,13 +172,6 @@ class DSAMetadata:
|
||||
token_to_batch_idx: Optional[torch.Tensor] = None
|
||||
|
||||
|
||||
class TopkTransformMethod(IntEnum):
|
||||
# Transform topk indices to indices to the page table (page_size = 1)
|
||||
PAGED = auto()
|
||||
# Transform topk indices to indices to ragged kv (non-paged)
|
||||
RAGGED = auto()
|
||||
|
||||
|
||||
@torch.compile
|
||||
def _compiled_cat(tensors: list[torch.Tensor], dim: int = -1) -> torch.Tensor:
|
||||
return torch.cat(tensors, dim=dim)
|
||||
@@ -193,6 +197,7 @@ def _cat(tensors: list[torch.Tensor], dim: int = -1) -> torch.Tensor:
|
||||
class DSAIndexerMetadata(BaseIndexerMetadata):
|
||||
attn_metadata: DSAMetadata
|
||||
topk_transform_method: TopkTransformMethod
|
||||
topk_backend: DSATopKBackend = DSATopKBackend.SGL_KERNEL
|
||||
paged_mqa_schedule_metadata: Optional[torch.Tensor] = None
|
||||
force_unfused_topk: bool = False
|
||||
|
||||
@@ -231,17 +236,11 @@ class DSAIndexerMetadata(BaseIndexerMetadata):
|
||||
logits: torch.Tensor,
|
||||
topk: int,
|
||||
ks: Optional[torch.Tensor] = None,
|
||||
cu_seqlens_q: torch.Tensor = None,
|
||||
ke_offset: torch.Tensor = None,
|
||||
batch_idx_list: List[int] = None,
|
||||
cu_seqlens_q: Optional[torch.Tensor] = None,
|
||||
ke_offset: Optional[torch.Tensor] = None,
|
||||
batch_idx_list: Optional[List[int]] = None,
|
||||
topk_indices_offset_override: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
from sgl_kernel import (
|
||||
fast_topk_transform_fused,
|
||||
fast_topk_transform_ragged_fused,
|
||||
fast_topk_v2,
|
||||
)
|
||||
|
||||
if topk_indices_offset_override is not None:
|
||||
cu_topk_indices_offset = topk_indices_offset_override
|
||||
cu_seqlens_q_topk = None
|
||||
@@ -259,38 +258,18 @@ class DSAIndexerMetadata(BaseIndexerMetadata):
|
||||
seq_lens_topk = ke_offset
|
||||
else:
|
||||
seq_lens_topk = self.get_seqlens_expanded()
|
||||
if batch_idx_list is not None:
|
||||
page_table_size_1 = self.attn_metadata.page_table_1[batch_idx_list]
|
||||
else:
|
||||
page_table_size_1 = self.attn_metadata.page_table_1
|
||||
|
||||
if not envs.SGLANG_DSA_FUSE_TOPK.get() or self.force_unfused_topk:
|
||||
return fast_topk_v2(logits, seq_lens_topk, topk, row_starts=ks)
|
||||
elif self.topk_transform_method == TopkTransformMethod.PAGED:
|
||||
# NOTE(dark): if fused, we return a transformed page table directly
|
||||
return fast_topk_transform_fused(
|
||||
score=logits,
|
||||
lengths=seq_lens_topk,
|
||||
page_table_size_1=page_table_size_1,
|
||||
cu_seqlens_q=cu_seqlens_q_topk,
|
||||
topk=topk,
|
||||
row_starts=ks,
|
||||
)
|
||||
elif self.topk_transform_method == TopkTransformMethod.RAGGED:
|
||||
if cu_topk_indices_offset is None:
|
||||
raise RuntimeError(
|
||||
"RAGGED topk_transform requires topk_indices_offset; "
|
||||
"expected extend-without-speculative metadata."
|
||||
)
|
||||
return fast_topk_transform_ragged_fused(
|
||||
score=logits,
|
||||
lengths=seq_lens_topk,
|
||||
topk_indices_offset=cu_topk_indices_offset,
|
||||
topk=topk,
|
||||
row_starts=ks,
|
||||
)
|
||||
else:
|
||||
assert False, f"Unsupported {self.topk_transform_method = }"
|
||||
return self.topk_backend.topk_transform(
|
||||
logits=logits,
|
||||
lengths=seq_lens_topk,
|
||||
topk=topk,
|
||||
topk_transform_method=self.topk_transform_method,
|
||||
attn_metadata=self.attn_metadata,
|
||||
cu_seqlens_q_topk=cu_seqlens_q_topk,
|
||||
topk_indices_offset=cu_topk_indices_offset,
|
||||
row_starts=ks,
|
||||
batch_idx_list=batch_idx_list,
|
||||
force_unfused_topk=self.force_unfused_topk,
|
||||
)
|
||||
|
||||
|
||||
_DSA_IMPL_T: TypeAlias = Literal[
|
||||
@@ -343,6 +322,9 @@ class DeepseekSparseAttnBackend(
|
||||
model_runner.server_args.dsa_prefill_backend
|
||||
)
|
||||
self.dsa_decode_impl: _DSA_IMPL_T = model_runner.server_args.dsa_decode_backend
|
||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
|
||||
model_runner.server_args.dsa_topk_backend
|
||||
)
|
||||
if self.num_q_heads <= 64:
|
||||
self.flashmla_kv_num_q_heads = 64
|
||||
elif self.num_q_heads <= 128:
|
||||
@@ -398,6 +380,16 @@ class DeepseekSparseAttnBackend(
|
||||
else:
|
||||
self.workspace_buffer = None
|
||||
|
||||
def _get_fused_topk_page_table(self, topk_indices: torch.Tensor) -> torch.Tensor:
|
||||
if (
|
||||
self.dsa_topk_backend.is_sgl_kernel()
|
||||
or self.dsa_topk_backend.is_flashinfer()
|
||||
):
|
||||
return topk_indices
|
||||
raise RuntimeError(
|
||||
f"Unsupported {self.dsa_topk_backend = } for SGLANG_DSA_FUSE_TOPK."
|
||||
)
|
||||
|
||||
def get_device_int32_arange(self, l: int) -> torch.Tensor:
|
||||
if l > len(self._arange_buf):
|
||||
next_pow_of_2 = 1 << (l - 1).bit_length()
|
||||
@@ -1431,7 +1423,7 @@ class DeepseekSparseAttnBackend(
|
||||
forward_batch.forward_mode
|
||||
)
|
||||
if envs.SGLANG_DSA_FUSE_TOPK.get():
|
||||
page_table_1 = topk_indices
|
||||
page_table_1 = self._get_fused_topk_page_table(topk_indices)
|
||||
else:
|
||||
if topk_transform_method == TopkTransformMethod.RAGGED:
|
||||
topk_indices_offset = metadata.topk_indices_offset
|
||||
@@ -1619,7 +1611,7 @@ class DeepseekSparseAttnBackend(
|
||||
layer.layer_id,
|
||||
)
|
||||
elif envs.SGLANG_DSA_FUSE_TOPK.get():
|
||||
page_table_1 = topk_indices
|
||||
page_table_1 = self._get_fused_topk_page_table(topk_indices)
|
||||
else:
|
||||
page_table_1 = transform_index_page_table_decode(
|
||||
page_table=metadata.page_table_1,
|
||||
@@ -2126,7 +2118,7 @@ class DeepseekSparseAttnBackend(
|
||||
topk_indices = self._pad_topk_indices(topk_indices, q.shape[0])
|
||||
|
||||
if envs.SGLANG_DSA_FUSE_TOPK.get():
|
||||
page_table_1 = topk_indices
|
||||
page_table_1 = self._get_fused_topk_page_table(topk_indices)
|
||||
elif is_prefill:
|
||||
page_table_1 = transform_index_page_table_prefill(
|
||||
page_table=metadata.page_table_1,
|
||||
@@ -2287,6 +2279,7 @@ class DeepseekSparseAttnBackend(
|
||||
topk_transform_method=self.get_topk_transform_method(
|
||||
forward_batch.forward_mode
|
||||
),
|
||||
topk_backend=self.dsa_topk_backend,
|
||||
paged_mqa_schedule_metadata=self.forward_metadata.paged_mqa_schedule_metadata,
|
||||
force_unfused_topk=force_unfused,
|
||||
)
|
||||
|
||||
@@ -266,6 +266,8 @@ DSA_CHOICES = [
|
||||
]
|
||||
NSA_CHOICES = DSA_CHOICES # deprecated alias
|
||||
|
||||
DSA_TOPK_BACKEND_CHOICES = ["sgl-kernel", "torch", "flashinfer"]
|
||||
|
||||
MAMBA_SCHEDULER_STRATEGY_CHOICES = ["auto", "no_buffer", "extra_buffer"]
|
||||
|
||||
MAMBA_BACKEND_CHOICES = ["triton", "flashinfer"]
|
||||
@@ -557,6 +559,7 @@ class ServerArgs:
|
||||
dsa_decode_backend: Optional[str] = (
|
||||
None # auto-detect based on hardware/kv_cache_dtype
|
||||
)
|
||||
dsa_topk_backend: str = "sgl-kernel"
|
||||
disable_flashinfer_autotune: bool = False
|
||||
mamba_backend: str = "triton"
|
||||
|
||||
@@ -5494,6 +5497,15 @@ class ServerArgs:
|
||||
choices=DSA_CHOICES,
|
||||
help="[Deprecated] Use --dsa-decode-backend instead.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dsa-topk-backend",
|
||||
dest="dsa_topk_backend",
|
||||
default=ServerArgs.dsa_topk_backend,
|
||||
type=str,
|
||||
choices=DSA_TOPK_BACKEND_CHOICES,
|
||||
help="DSA indexer top-k backend. Options: 'sgl-kernel', 'torch', 'flashinfer'. "
|
||||
"The 'torch' backend currently requires SGLANG_DSA_FUSE_TOPK=false.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fp8-gemm-backend",
|
||||
type=str,
|
||||
|
||||
Reference in New Issue
Block a user