From 53b83783073dfb3e25fd02b5b5073a1ad0d797af Mon Sep 17 00:00:00 2001 From: Lianmin Zheng Date: Sun, 31 May 2026 21:27:03 -0700 Subject: [PATCH] Fix weights_checker checksum for 0-dim tensors and multi-GPU (#26863) --- python/sglang/srt/entrypoints/http_server.py | 14 ++++-- python/sglang/srt/layers/multimodal.py | 2 +- python/sglang/srt/managers/io_struct.py | 2 +- python/sglang/srt/managers/mm_utils.py | 4 +- .../scheduler_components/weight_updater.py | 45 ++++++++++++++++++- .../srt/managers/tokenizer_control_mixin.py | 17 +++++-- python/sglang/srt/utils/weight_checker.py | 10 ++++- .../unit/utils/test_weight_checker.py | 5 ++- 8 files changed, 85 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/entrypoints/http_server.py b/python/sglang/srt/entrypoints/http_server.py index 9dc6c0cde..c51dbcd2f 100644 --- a/python/sglang/srt/entrypoints/http_server.py +++ b/python/sglang/srt/entrypoints/http_server.py @@ -1300,15 +1300,21 @@ async def resume_memory_occupation( return _create_error_response(e) -@app.post("/weights_checker") +@app.api_route("/weights_checker", methods=["GET", "POST"]) @auth_level(AuthLevel.ADMIN_OPTIONAL) -async def check_weights(obj: CheckWeightsReqInput, request: Request): - success, message, ranks = await _global_state.tokenizer_manager.check_weights( - obj, request +async def check_weights( + obj: Optional[CheckWeightsReqInput] = None, request: Request = None +): + if obj is None: + obj = CheckWeightsReqInput() + success, message, ranks, per_engine_checksum = ( + await _global_state.tokenizer_manager.check_weights(obj, request) ) body = {"success": success, "message": message} if ranks is not None: body["ranks"] = ranks + if per_engine_checksum is not None: + body["per_engine_checksum"] = per_engine_checksum return ORJSONResponse(body, status_code=200 if success else HTTPStatus.BAD_REQUEST) diff --git a/python/sglang/srt/layers/multimodal.py b/python/sglang/srt/layers/multimodal.py index 738c65830..434a939cc 100644 --- a/python/sglang/srt/layers/multimodal.py +++ b/python/sglang/srt/layers/multimodal.py @@ -119,7 +119,7 @@ def add_tree_reduce_u64_kernel(in_ptr, out_ptr, n_elems, CHUNK: tl.constexpr): def _as_uint32_words(t: torch.Tensor) -> torch.Tensor: assert t.is_cuda, "Use .cuda() first" - tb = t.contiguous().view(torch.uint8) + tb = t.contiguous().reshape(-1).view(torch.uint8) nbytes = tb.numel() pad = (4 - (nbytes & 3)) & 3 if pad: diff --git a/python/sglang/srt/managers/io_struct.py b/python/sglang/srt/managers/io_struct.py index ee718658f..d9d102202 100644 --- a/python/sglang/srt/managers/io_struct.py +++ b/python/sglang/srt/managers/io_struct.py @@ -1650,7 +1650,7 @@ class ResumeMemoryOccupationReqOutput(BaseReq): @dataclass class CheckWeightsReqInput(BaseReq): - action: str + action: str = "checksum" @dataclass diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index 49cb194eb..1810d9cd0 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -1227,7 +1227,7 @@ def tensor_hash(tensor_list) -> int: hasher = hashlib.sha256() for t in tensors: t = t.detach().contiguous() - hasher.update(memoryview(t.view(torch.uint8).numpy())) + hasher.update(memoryview(t.reshape(-1).view(torch.uint8).numpy())) hash_bytes = hasher.digest()[:8] return int.from_bytes(hash_bytes, byteorder="big", signed=False) @@ -1236,7 +1236,7 @@ def tensor_hash(tensor_list) -> int: return gpu_tensor_hash(tensor.cuda()) tensor = tensor.detach().contiguous() hasher = hashlib.sha256() - hasher.update(memoryview(tensor.view(torch.uint8).numpy())) + hasher.update(memoryview(tensor.reshape(-1).view(torch.uint8).numpy())) hash_bytes = hasher.digest()[:8] return int.from_bytes(hash_bytes, byteorder="big", signed=False) diff --git a/python/sglang/srt/managers/scheduler_components/weight_updater.py b/python/sglang/srt/managers/scheduler_components/weight_updater.py index ba7868468..59f6597b3 100644 --- a/python/sglang/srt/managers/scheduler_components/weight_updater.py +++ b/python/sglang/srt/managers/scheduler_components/weight_updater.py @@ -1,9 +1,10 @@ from __future__ import annotations +import hashlib import logging import traceback from dataclasses import dataclass, field -from typing import Any, Callable, Tuple +from typing import Any, Callable, Dict, Tuple import torch @@ -39,6 +40,33 @@ from sglang.srt.managers.io_struct import ( logger = logging.getLogger(__name__) +def _get_draft_model_runner(draft_worker): + # EAGLEWorker (v1): draft_model_runner property -> self.model_runner + runner = getattr(draft_worker, "draft_model_runner", None) + if runner is not None: + return runner + # EAGLEWorkerV2: _draft_worker.draft_runner + inner = getattr(draft_worker, "_draft_worker", None) + if inner is not None: + runner = getattr(inner, "draft_runner", None) + if runner is not None: + return runner + return None + + +def _merge_checksum_payloads(target: Dict, draft: Dict) -> Dict: + merged_checksums = dict(target["checksums"]) + for name, chk in draft["checksums"].items(): + merged_checksums[f"draft.{name}"] = chk + h = hashlib.sha256() + for name in sorted(merged_checksums): + h.update(name.encode()) + h.update(merged_checksums[name].encode()) + target["checksums"] = merged_checksums + target["per_gpu_checksum"] = h.hexdigest() + return target + + @dataclass(kw_only=True, slots=True) class SchedulerWeightUpdaterManager: tp_worker: Any @@ -185,6 +213,21 @@ class SchedulerWeightUpdaterManager: def check_weights(self, recv_req: CheckWeightsReqInput): try: payload = self.tp_worker.model_runner.check_weights(action=recv_req.action) + + if self.draft_worker is not None: + draft_runner = _get_draft_model_runner(self.draft_worker) + if draft_runner is not None: + draft_payload = draft_runner.check_weights(action=recv_req.action) + if payload is not None and draft_payload is not None: + payload = _merge_checksum_payloads(payload, draft_payload) + + tp_size = torch.distributed.get_world_size(group=self.tp_cpu_group) + if tp_size > 1 and payload is not None: + all_payloads = [None] * tp_size + torch.distributed.all_gather_object( + all_payloads, payload, group=self.tp_cpu_group + ) + payload = all_payloads return CheckWeightsReqOutput( success=True, message="Success.", payload=payload ) diff --git a/python/sglang/srt/managers/tokenizer_control_mixin.py b/python/sglang/srt/managers/tokenizer_control_mixin.py index b26a77eec..c9939a1fc 100644 --- a/python/sglang/srt/managers/tokenizer_control_mixin.py +++ b/python/sglang/srt/managers/tokenizer_control_mixin.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import hashlib import logging import time import uuid @@ -757,14 +758,24 @@ class TokenizerControlMixin: self: TokenizerManager, obj: CheckWeightsReqInput, request: Optional[fastapi.Request] = None, - ) -> Tuple[bool, str, Optional[List[Dict]]]: + ) -> Tuple[bool, str, Optional[List[Dict]], Optional[str]]: self.auto_create_handle_loop() results = await self.check_weights_communicator(obj) success, message = FanOutCommunicator.merge_results(results) ranks: Optional[List[Dict]] = None + per_engine_checksum: Optional[str] = None if any(r.payload is not None for r in results): - ranks = [r.payload for r in results] - return success, message, ranks + ranks = [] + for r in results: + if isinstance(r.payload, list): + ranks.extend(r.payload) + else: + ranks.append(r.payload) + h = hashlib.sha256() + for rank in ranks: + h.update(rank["per_gpu_checksum"].encode()) + per_engine_checksum = h.hexdigest() + return success, message, ranks, per_engine_checksum async def slow_down( self: TokenizerManager, diff --git a/python/sglang/srt/utils/weight_checker.py b/python/sglang/srt/utils/weight_checker.py index 55566973a..8e6459a37 100644 --- a/python/sglang/srt/utils/weight_checker.py +++ b/python/sglang/srt/utils/weight_checker.py @@ -1,3 +1,4 @@ +import hashlib import logging import time from typing import Dict, Iterable, Optional, Set, Tuple @@ -32,6 +33,7 @@ class ParallelismInfo(_StrictBaseModel): class ChecksumInfo(_StrictBaseModel): checksums: Dict[str, str] + per_gpu_checksum: str parallelism_info: ParallelismInfo @@ -118,6 +120,12 @@ class WeightChecker: if should_compare } + h = hashlib.sha256() + for name in sorted(checksums): + h.update(name.encode()) + h.update(checksums[name].encode()) + overall = h.hexdigest() + torch.cuda.synchronize() elapsed = time.perf_counter() - start logger.info( @@ -126,6 +134,7 @@ class WeightChecker: info = ChecksumInfo( checksums=checksums, + per_gpu_checksum=overall, parallelism_info=self._parallelism_info(), ) return info.model_dump() @@ -144,7 +153,6 @@ class WeightChecker: ) def _model_state(self): - # TODO: support EAGLE etc (e.g. yield from both main model and draft model) yield from self._model_runner.model.named_parameters() yield from self._model_runner.model.named_buffers() diff --git a/test/registered/unit/utils/test_weight_checker.py b/test/registered/unit/utils/test_weight_checker.py index a253e6155..e54da8ab1 100644 --- a/test/registered/unit/utils/test_weight_checker.py +++ b/test/registered/unit/utils/test_weight_checker.py @@ -491,6 +491,7 @@ class TestHandle(_WeightCheckerTestBase): out = self.checker.handle("checksum") self.assertIsInstance(out, dict) self.assertIn("checksums", out) + self.assertIn("per_gpu_checksum", out) self.assertIn("parallelism_info", out) def test_unknown_action_raises(self): @@ -582,7 +583,9 @@ class TestComputeChecksum(_ChecksumTestBase): def test_returns_dict_with_expected_top_level_keys(self): out = self.checker._compute_checksum() - self.assertEqual(set(out.keys()), {"checksums", "parallelism_info"}) + self.assertEqual( + set(out.keys()), {"checksums", "per_gpu_checksum", "parallelism_info"} + ) def test_skips_non_persistent_buffers(self): out = self.checker._compute_checksum()