Refactor buffer patterns in weight checker (#24538)
This commit is contained in:
@@ -11,6 +11,18 @@ from sglang.srt.layers.quantization.fp8_utils import (
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
_NON_PERSISTENT_BUFFER_PATTERNS = (
|
||||||
|
"cos_sin_cache",
|
||||||
|
"inv_freq",
|
||||||
|
"freqs_cis",
|
||||||
|
"_weight_fp32",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _is_non_persistent_buffer_name(name: str) -> bool:
|
||||||
|
return any(pat in name for pat in _NON_PERSISTENT_BUFFER_PATTERNS)
|
||||||
|
|
||||||
|
|
||||||
class WeightChecker:
|
class WeightChecker:
|
||||||
def __init__(self, model_runner):
|
def __init__(self, model_runner):
|
||||||
self._model_runner = model_runner
|
self._model_runner = model_runner
|
||||||
@@ -38,7 +50,7 @@ class WeightChecker:
|
|||||||
|
|
||||||
def _reset_tensors(self):
|
def _reset_tensors(self):
|
||||||
for name, param in self._model_state():
|
for name, param in self._model_state():
|
||||||
if "cos_sin_cache" in name or "freqs_cis" in name or "_weight_fp32" in name:
|
if _is_non_persistent_buffer_name(name):
|
||||||
continue
|
continue
|
||||||
param.copy_(_random_like(param))
|
param.copy_(_random_like(param))
|
||||||
|
|
||||||
@@ -125,21 +137,12 @@ def _postprocess_tensors(
|
|||||||
|
|
||||||
skip_compare_names = []
|
skip_compare_names = []
|
||||||
|
|
||||||
# Skip non-persistent buffers like cos_sin_cache
|
# Skip non-persistent buffers (registered with persistent=False; recomputed
|
||||||
# These buffers are registered with persistent=False and are not saved in checkpoints
|
# after weight load and not part of the synced payload).
|
||||||
# They should be recomputed after loading weights, so we don't compare them here
|
|
||||||
non_persistent_buffer_patterns = [
|
|
||||||
"cos_sin_cache", # RoPE cache
|
|
||||||
"inv_freq", # RoPE inverse frequency (if it exists as buffer)
|
|
||||||
"_weight_fp32", # FP32 cache of gate weight (e.g. Glm4MoeGate)
|
|
||||||
]
|
|
||||||
|
|
||||||
for name in raw:
|
for name in raw:
|
||||||
for pattern in non_persistent_buffer_patterns:
|
if _is_non_persistent_buffer_name(name):
|
||||||
if pattern in name:
|
skip_compare_names.append(name)
|
||||||
skip_compare_names.append(name)
|
logger.info(f"[check_tensors] Skipping non-persistent buffer: {name}")
|
||||||
logger.info(f"[check_tensors] Skipping non-persistent buffer: {name}")
|
|
||||||
break
|
|
||||||
|
|
||||||
# dequant fp8
|
# dequant fp8
|
||||||
quant_names = [
|
quant_names = [
|
||||||
|
|||||||
Reference in New Issue
Block a user