[AMD] Avoid the FP8 wo_a path when the weight is BF16 (#38453)
Signed-off-by: bingxche <bingxche@amd.com>
This commit is contained in:
@@ -3,7 +3,9 @@
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from sglang.srt.arg_groups.overrides import (
|
from sglang.srt.arg_groups.overrides import (
|
||||||
@@ -81,6 +83,59 @@ def _rocm_fp8_wo_a_supported() -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _probe_wo_a_weight_dtype(model_config: Any, download_dir: str | None) -> str | None:
|
||||||
|
"""Read one indexed wo_a dtype without downloading a weight shard."""
|
||||||
|
try:
|
||||||
|
from huggingface_hub import (
|
||||||
|
parse_local_safetensors_file_metadata,
|
||||||
|
parse_safetensors_file_metadata,
|
||||||
|
)
|
||||||
|
from transformers.utils.hub import cached_file
|
||||||
|
|
||||||
|
model_path = model_config.model_path
|
||||||
|
revision = (
|
||||||
|
getattr(model_config.hf_config, "_commit_hash", None)
|
||||||
|
or model_config.revision
|
||||||
|
)
|
||||||
|
index_path = cached_file(
|
||||||
|
model_path,
|
||||||
|
"model.safetensors.index.json",
|
||||||
|
revision=revision,
|
||||||
|
cache_dir=download_dir,
|
||||||
|
)
|
||||||
|
with open(index_path) as f:
|
||||||
|
weight_map = json.load(f).get("weight_map", {})
|
||||||
|
name = next((key for key in weight_map if key.endswith(".wo_a.weight")), None)
|
||||||
|
if name is None:
|
||||||
|
return None
|
||||||
|
|
||||||
|
shard = weight_map[name]
|
||||||
|
local_shard = os.path.join(os.path.dirname(index_path), shard)
|
||||||
|
metadata = (
|
||||||
|
parse_local_safetensors_file_metadata(local_shard)
|
||||||
|
if os.path.isfile(local_shard)
|
||||||
|
else parse_safetensors_file_metadata(model_path, shard, revision=revision)
|
||||||
|
)
|
||||||
|
return getattr(metadata.tensors.get(name), "dtype", None)
|
||||||
|
except Exception:
|
||||||
|
logger.debug("Unable to inspect the checkpoint wo_a dtype", exc_info=True)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _configure_rocm_fp8_wo_a_gemm(model_config: Any, download_dir: str | None) -> None:
|
||||||
|
flag = envs.SGLANG_OPT_FP8_WO_A_GEMM
|
||||||
|
if not _rocm_fp8_wo_a_supported():
|
||||||
|
flag.set(False)
|
||||||
|
return
|
||||||
|
if flag.is_set():
|
||||||
|
return
|
||||||
|
|
||||||
|
dtype = _probe_wo_a_weight_dtype(model_config, download_dir)
|
||||||
|
if dtype is not None and dtype != "F8_E4M3":
|
||||||
|
flag.set(False)
|
||||||
|
logger.info("Disabled ROCm fp8 wo_a GEMM for checkpoint dtype %s", dtype)
|
||||||
|
|
||||||
|
|
||||||
def handle_model_specific_adjustments(server_args: Any):
|
def handle_model_specific_adjustments(server_args: Any):
|
||||||
|
|
||||||
cfg = resolving_view(server_args)
|
cfg = resolving_view(server_args)
|
||||||
@@ -387,11 +442,7 @@ def handle_model_specific_adjustments(server_args: Any):
|
|||||||
envs.SGLANG_OPT_USE_TILELANG_INDEXER.set(True)
|
envs.SGLANG_OPT_USE_TILELANG_INDEXER.set(True)
|
||||||
elif get_platform().is_hip:
|
elif get_platform().is_hip:
|
||||||
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
|
envs.SGLANG_OPT_DEEPGEMM_HC_PRENORM.set(False)
|
||||||
# The fp8 wo_a GEMM is DeepGEMM-based on CUDA. ROCm has an aiter
|
_configure_rocm_fp8_wo_a_gemm(model_config, cfg.download_dir)
|
||||||
# e8m0 block-scale equivalent, but only on gfx950 -- everywhere else
|
|
||||||
# keeps the bf16 absorb GEMM.
|
|
||||||
if not _rocm_fp8_wo_a_supported():
|
|
||||||
envs.SGLANG_OPT_FP8_WO_A_GEMM.set(False)
|
|
||||||
envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.set(False)
|
envs.SGLANG_OPT_USE_JIT_INDEXER_METADATA.set(False)
|
||||||
envs.SGLANG_OPT_USE_TOPK_V2.set(True)
|
envs.SGLANG_OPT_USE_TOPK_V2.set(True)
|
||||||
envs.SGLANG_OPT_USE_AITER_INDEXER.set(True)
|
envs.SGLANG_OPT_USE_AITER_INDEXER.set(True)
|
||||||
|
|||||||
Reference in New Issue
Block a user