[tiny] misc cleanups across configs, attention, jit_kernel (#24350)

This commit is contained in:
Liangsheng Yin
2026-05-04 03:17:14 -07:00
committed by GitHub
parent 952b3caf18
commit 84f3b44916
8 changed files with 62 additions and 22 deletions
+2 -2
View File
@@ -90,7 +90,7 @@ KERNEL_PATH = _resolve_kernel_path()
DEFAULT_INCLUDE = [str(KERNEL_PATH / "include")]
DEFAULT_CFLAGS = ["-std=c++20", "-O3"]
DEFAULT_LDFLAGS = []
CPP_TEMPLATE_TYPE: TypeAlias = Union[int, float, bool, torch.dtype]
CPP_TEMPLATE_TYPE: TypeAlias = Union[int, float, str, bool, torch.dtype]
class CPPArgList(list[str]):
@@ -119,7 +119,7 @@ def make_cpp_args(*args: CPP_TEMPLATE_TYPE) -> CPPArgList:
def _convert(arg: CPP_TEMPLATE_TYPE) -> str:
if isinstance(arg, bool):
return "true" if arg else "false"
if isinstance(arg, (int, float)):
if isinstance(arg, (int, str, float)):
return str(arg)
if isinstance(arg, torch.dtype):
return CPP_DTYPE_MAP[arg]
+23 -16
View File
@@ -82,29 +82,35 @@ class ModelImpl(str, Enum):
MINDSPORE = "mindspore"
def is_deepseek_nsa(config) -> bool:
architectures = (
def _hf_arch(config) -> Optional[str]:
"""First architecture from a HF config dict or PretrainedConfig (or None)."""
archs = (
config.get("architectures")
if isinstance(config, dict)
else getattr(config, "architectures", None)
)
index_topk = (
config.get("index_topk")
if isinstance(config, dict)
else getattr(config, "index_topk", None)
)
return archs[0] if archs else None
def _hf_attr(config, name):
"""Read an arbitrary field from a HF config dict or PretrainedConfig."""
if isinstance(config, dict):
return config.get(name)
return getattr(config, name, None)
def is_deepseek_nsa(config) -> bool:
return (
architectures is not None
and architectures[0]
in [
_hf_arch(config)
in (
"DeepseekV3ForCausalLM",
"DeepseekV32ForCausalLM",
"DeepseekV3ForCausalLMNextN",
"MistralLarge3ForCausalLM",
"PixtralForConditionalGeneration",
"GlmMoeDsaForCausalLM",
]
and index_topk is not None
)
and _hf_attr(config, "index_topk") is not None
)
@@ -968,10 +974,11 @@ class ModelConfig:
return "fp8" # Default fallback
def _get_sliding_window_size(self) -> Optional[int]:
sliding_window_size = getattr(self.hf_text_config, "sliding_window_size", None)
if sliding_window_size is None:
sliding_window_size = getattr(self.hf_text_config, "sliding_window", None)
return sliding_window_size
for key in ("sliding_window_size", "sliding_window", "window_size"):
value = getattr(self.hf_text_config, key, None)
if value is not None:
return value
return None
def _validate_quantize_and_serve_config(self):
"""Validate quantize_and_serve configuration."""
@@ -58,6 +58,15 @@ class AttentionBackend(ABC):
"""Get the fill value for padded seq lens. Typically, it is 0 or 1."""
raise NotImplementedError()
def on_after_cuda_graph_warmup(self):
"""Hook between cuda graph warmup pass and the actual capture.
Override to undo state that warmup mutated or eagerly advanced
(e.g. dirty metadata buffers, raw->full upgrades) before capture
freezes the kernel pointers.
"""
pass
def get_verify_buffers_to_fill_after_draft(self):
"""
Return buffers of verify attention kernels that needs to be filled after draft.
@@ -130,6 +139,7 @@ class AttentionBackend(ABC):
layer: RadixAttention,
forward_batch: ForwardBatch,
save_kv_cache: bool = True,
**kwargs,
):
"""Run a forward for decode."""
raise NotImplementedError()
@@ -142,6 +152,7 @@ class AttentionBackend(ABC):
layer: RadixAttention,
forward_batch: ForwardBatch,
save_kv_cache: bool = True,
**kwargs,
):
"""Run a forward for extend."""
raise NotImplementedError()
@@ -327,6 +327,8 @@ class SetKAndS:
@classmethod
def triton(cls, pool, buf, loc, index_k, index_k_scale):
loc = loc.to(torch.int64)
_set_k_and_s_triton(
buf=buf,
loc=loc,
@@ -969,6 +969,7 @@ class NativeSparseAttnBackend(
spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor],
out_cache_loc: Optional[torch.Tensor] = None,
actual_forward_mode: Optional[ForwardMode] = None,
):
"""Initialize forward metadata for replaying CUDA graph."""
assert seq_lens_cpu is not None
@@ -40,6 +40,9 @@ def grouped_gemm_nt_f8f8bf16_masked(
_sanity_check_input(lhs)
_sanity_check_input(rhs)
lhs = _ensure_cuda(lhs)
rhs = _ensure_cuda(rhs)
with compile_utils.deep_gemm_execution_hook(
expected_m, n, k, num_groups, kernel_type
):
@@ -65,6 +68,15 @@ def grouped_gemm_nt_f8f8bf16_masked(
)
def _ensure_cuda(
pair: Tuple[torch.Tensor, torch.Tensor],
) -> Tuple[torch.Tensor, torch.Tensor]:
return (
pair[0].cuda() if not pair[0].is_cuda else pair[0],
pair[1].cuda() if not pair[1].is_cuda else pair[1],
)
def grouped_gemm_nt_f8f8bf16_contig(
lhs: Tuple[torch.Tensor, torch.Tensor],
rhs: Tuple[torch.Tensor, torch.Tensor],
@@ -75,6 +87,9 @@ def grouped_gemm_nt_f8f8bf16_contig(
num_groups, n, _ = rhs[0].shape
kernel_type = compile_utils.DeepGemmKernelType.GROUPED_GEMM_NT_F8F8BF16_CONTIG
if m == 0:
return
_sanity_check_input(lhs)
_sanity_check_input(rhs)
+6 -2
View File
@@ -268,7 +268,9 @@ class ReplicatedLinear(LinearBase):
param.dtype == loaded_weight.dtype
), "init para dtype and loaded weight dtype should be the same"
assert param.size() == loaded_weight.size()
assert (
param.size() == loaded_weight.size()
), f"{param.shape=} {param.dtype=} {loaded_weight.shape=} {loaded_weight.dtype=}"
param.data.copy_(loaded_weight)
def forward(self, x: torch.Tensor) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
@@ -428,7 +430,9 @@ class ColumnParallelLinear(LinearBase):
if len(loaded_weight.shape) == 0:
loaded_weight = loaded_weight.reshape(1)
assert param_data.shape == loaded_weight.shape
assert (
param_data.shape == loaded_weight.shape
), f"param_data.shape={param_data.shape} != loaded_weight.shape={loaded_weight.shape}"
param_data.copy_(loaded_weight)
def weight_loader_v2(self, param: Parameter, loaded_weight: torch.Tensor):
@@ -105,8 +105,8 @@ class ForwardMode(IntEnum):
# Used in dLLM
DLLM_EXTEND = auto()
def is_prefill(self):
return self.is_extend()
def is_prefill(self, include_draft_extend_v2: bool = False):
return self.is_extend(include_draft_extend_v2=include_draft_extend_v2)
def is_extend(self, include_draft_extend_v2: bool = False):
return (