Narrow component dependencies to injected fields instead of ModelRunner (#31166)

This commit is contained in:
fzyzcjy
2026-07-14 16:03:07 +08:00
committed by GitHub
parent 6999007a13
commit 54f99a21d5
5 changed files with 95 additions and 44 deletions
@@ -20,6 +20,7 @@ from unittest.mock import patch
import torch
from torch import nn
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.layers.quantization.fp8_utils import (
quant_weight_ue8m0,
transform_scale_ue8m0,
@@ -129,14 +130,18 @@ class _FakeModelRunner:
dp_size: int = 1,
pp_rank: int = 0,
pp_size: int = 1,
attn_dp_size: int | None = None,
):
self.model = model
self.tp_rank = tp_rank
self.tp_size = tp_size
self.dp_rank = dp_rank
self.dp_size = dp_size
self.pp_rank = pp_rank
self.pp_size = pp_size
self.ps = ParallelState.trivial(
tp_rank=tp_rank,
tp_size=tp_size,
dp_rank=dp_rank,
dp_size=dp_size,
attn_dp_size=attn_dp_size if attn_dp_size is not None else dp_size,
pp_rank=pp_rank,
pp_size=pp_size,
)
# ---------------------------------------------------------------------------
@@ -490,7 +495,8 @@ class _WeightCheckerTestBase(CustomTestCase):
def setUp(self):
torch.manual_seed(0)
self.model = _TinyModel().cuda()
self.checker = WeightChecker(model_runner=_FakeModelRunner(self.model))
runner = _FakeModelRunner(self.model)
self.checker = WeightChecker(get_model=lambda: runner.model, ps=runner.ps)
class TestSnapshot(_WeightCheckerTestBase):
@@ -697,7 +703,9 @@ class _ChecksumTestBase(CustomTestCase):
pp_rank=0,
pp_size=1,
)
self.checker = WeightChecker(model_runner=self.runner)
self.checker = WeightChecker(
get_model=lambda: self.runner.model, ps=self.runner.ps
)
class TestComputeChecksum(_ChecksumTestBase):