[core/loader] Add presharded load format (#24256)
Co-authored-by: Shu Wang <shuwanguc@google.com>
This commit is contained in:
@@ -0,0 +1,960 @@
|
||||
"""Unit tests for PreshardedModelLoader's pure helpers.
|
||||
|
||||
These tests exercise the deterministic pieces of the presharding algorithm
|
||||
(tensor hashing, plan construction, file naming, dedup, file-size cap, and
|
||||
per-rank workload balance) without needing a GPU or distributed setup.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import tempfile
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.model_loader.loader import PreshardedModelLoader
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
class TestPreshardedHashTensor(unittest.TestCase):
|
||||
def test_same_content_same_hash(self):
|
||||
a = torch.arange(100, dtype=torch.float32).reshape(10, 10)
|
||||
b = torch.arange(100, dtype=torch.float32).reshape(10, 10)
|
||||
self.assertEqual(
|
||||
PreshardedModelLoader._hash_tensor(a),
|
||||
PreshardedModelLoader._hash_tensor(b),
|
||||
)
|
||||
|
||||
def test_different_content_different_hash(self):
|
||||
a = torch.arange(100, dtype=torch.float32)
|
||||
b = a.clone()
|
||||
b[0] = 999
|
||||
self.assertNotEqual(
|
||||
PreshardedModelLoader._hash_tensor(a),
|
||||
PreshardedModelLoader._hash_tensor(b),
|
||||
)
|
||||
|
||||
def test_dtype_changes_hash(self):
|
||||
a = torch.zeros(8, dtype=torch.float32)
|
||||
b = torch.zeros(16, dtype=torch.float16) # same byte content (zeros)
|
||||
self.assertNotEqual(
|
||||
PreshardedModelLoader._hash_tensor(a),
|
||||
PreshardedModelLoader._hash_tensor(b),
|
||||
)
|
||||
|
||||
def test_shape_changes_hash(self):
|
||||
a = torch.zeros(16, dtype=torch.float32)
|
||||
b = torch.zeros((4, 4), dtype=torch.float32)
|
||||
self.assertNotEqual(
|
||||
PreshardedModelLoader._hash_tensor(a),
|
||||
PreshardedModelLoader._hash_tensor(b),
|
||||
)
|
||||
|
||||
def test_empty_tensor_is_hashable(self):
|
||||
a = torch.empty(0, dtype=torch.float32)
|
||||
digest = PreshardedModelLoader._hash_tensor(a)
|
||||
self.assertIsInstance(digest, str)
|
||||
self.assertEqual(
|
||||
len(digest), PreshardedModelLoader._CONTENT_HASH_HEX_LEN
|
||||
) # xxh3-128 hex
|
||||
|
||||
def test_cuda_and_cpu_digests_agree(self):
|
||||
# Streaming D2H + host xxh3 must match a full CPU hash of the same
|
||||
# bytes; otherwise dump/reload verify would false-fail across devices.
|
||||
if not torch.cuda.is_available():
|
||||
self.skipTest("CUDA required")
|
||||
cpu = torch.arange(10_000, dtype=torch.float32)
|
||||
gpu = cpu.cuda()
|
||||
self.assertEqual(
|
||||
PreshardedModelLoader._hash_tensor(cpu),
|
||||
PreshardedModelLoader._hash_tensor(gpu),
|
||||
)
|
||||
|
||||
|
||||
class TestPreshardedFilename(unittest.TestCase):
|
||||
def test_common_filename(self):
|
||||
self.assertEqual(
|
||||
PreshardedModelLoader._make_filename(0, (0, 1, 2, 3), is_common=True),
|
||||
"model-00000-common.safetensor",
|
||||
)
|
||||
|
||||
def test_rank_filename_three_digit_padding(self):
|
||||
self.assertEqual(
|
||||
PreshardedModelLoader._make_filename(5, (1, 3, 5, 7), is_common=False),
|
||||
"model-00005-rank-001,003,005,007.safetensor",
|
||||
)
|
||||
|
||||
def test_file_id_zero_padding(self):
|
||||
self.assertTrue(
|
||||
PreshardedModelLoader._make_filename(42, (0,), is_common=False).startswith(
|
||||
"model-00042-"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class TestBuildDumpPlan(unittest.TestCase):
|
||||
def _write_manifests(self, tmp_dir, manifests):
|
||||
for r, m in manifests.items():
|
||||
with open(os.path.join(tmp_dir, f"manifest_{r:05d}.json"), "w") as f:
|
||||
json.dump(m, f)
|
||||
|
||||
def test_dedup_across_ranks(self):
|
||||
# Same content (same checksum) on all 4 ranks → single file marked common.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
shared = {
|
||||
"checksum": "deadbeef",
|
||||
"size": 1024,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [256],
|
||||
}
|
||||
manifests = {r: {"shared.weight": shared} for r in range(4)}
|
||||
self._write_manifests(tmp, manifests)
|
||||
plan = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=4, tmp_dir=tmp, max_file_bytes=10**12
|
||||
)
|
||||
self.assertEqual(len(plan["files"]), 1)
|
||||
self.assertTrue(plan["files"][0]["is_common"])
|
||||
self.assertIn("common", plan["files"][0]["filename"])
|
||||
# Each rank should still have a read entry pointing at the file.
|
||||
for r in range(4):
|
||||
reads = plan["rank_to_reads"][str(r)]
|
||||
self.assertEqual(len(reads), 1)
|
||||
self.assertEqual(reads[0]["name"], "shared.weight")
|
||||
self.assertEqual(reads[0]["filename"], plan["files"][0]["filename"])
|
||||
self.assertEqual(reads[0]["stored_key"], "deadbeef")
|
||||
self.assertIn("rank_checksums", plan)
|
||||
self.assertEqual(set(plan["rank_checksums"].keys()), {"0", "1", "2", "3"})
|
||||
|
||||
def test_per_rank_unique_tensors(self):
|
||||
# Each rank has its own tensor (different content). 4 distinct files,
|
||||
# filenames should be rank-{rrr}.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
manifests = {
|
||||
r: {
|
||||
"layer.weight": {
|
||||
"checksum": f"hash_{r}",
|
||||
"size": 2048,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [512],
|
||||
}
|
||||
}
|
||||
for r in range(4)
|
||||
}
|
||||
self._write_manifests(tmp, manifests)
|
||||
plan = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=4, tmp_dir=tmp, max_file_bytes=10**12
|
||||
)
|
||||
self.assertEqual(len(plan["files"]), 4)
|
||||
for f in plan["files"]:
|
||||
self.assertFalse(f["is_common"])
|
||||
self.assertEqual(len(f["rank_list"]), 1)
|
||||
# Writer is the only rank in the rank_list.
|
||||
self.assertEqual(f["writer_rank"], f["rank_list"][0])
|
||||
self.assertIn(f"-rank-{f['rank_list'][0]:03d}.safetensor", f["filename"])
|
||||
|
||||
def test_partial_share_has_correct_rank_list(self):
|
||||
# Tensor shared by ranks 1,3,5,7 only.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
shared = {
|
||||
"checksum": "shared_hash",
|
||||
"size": 1024,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [256],
|
||||
}
|
||||
manifests = {
|
||||
r: ({"x": shared} if r in (1, 3, 5, 7) else {}) for r in range(8)
|
||||
}
|
||||
self._write_manifests(tmp, manifests)
|
||||
plan = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=8, tmp_dir=tmp, max_file_bytes=10**12
|
||||
)
|
||||
self.assertEqual(len(plan["files"]), 1)
|
||||
f = plan["files"][0]
|
||||
self.assertFalse(f["is_common"])
|
||||
self.assertEqual(f["rank_list"], [1, 3, 5, 7])
|
||||
self.assertIn(f["writer_rank"], (1, 3, 5, 7))
|
||||
self.assertEqual(f["filename"], "model-00000-rank-001,003,005,007.safetensor")
|
||||
# Only the 4 sharing ranks have read entries.
|
||||
for r in range(8):
|
||||
reads = plan["rank_to_reads"].get(str(r), [])
|
||||
if r in (1, 3, 5, 7):
|
||||
self.assertEqual(len(reads), 1)
|
||||
else:
|
||||
self.assertEqual(len(reads), 0)
|
||||
|
||||
def test_max_file_size_caps_files(self):
|
||||
# Two tensors of 1 MiB each shared by all 2 ranks; cap = 1.5 MiB →
|
||||
# two files.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
t1 = {
|
||||
"checksum": "h1",
|
||||
"size": 1024 * 1024,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [256, 1024],
|
||||
}
|
||||
t2 = dict(t1, checksum="h2")
|
||||
manifests = {0: {"a": t1, "b": t2}, 1: {"a": t1, "b": t2}}
|
||||
self._write_manifests(tmp, manifests)
|
||||
plan = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=2,
|
||||
tmp_dir=tmp,
|
||||
max_file_bytes=int(1.5 * 1024 * 1024),
|
||||
)
|
||||
# Both tensors share the same rank_list (0,1). With balanced packing,
|
||||
# each writer (rank 0 and rank 1) gets one tensor → 2 files.
|
||||
self.assertEqual(len(plan["files"]), 2)
|
||||
for f in plan["files"]:
|
||||
self.assertEqual(len(f["tensors"]), 1)
|
||||
# rank_list is full world ⇒ marked common
|
||||
self.assertTrue(f["is_common"])
|
||||
|
||||
def test_workload_balanced_within_rank_list(self):
|
||||
# 4 same-size tensors all shared by ranks (0,1,2,3) → with balanced
|
||||
# packing each writer rank should get exactly one tensor.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
tensors = {
|
||||
f"t{i}": {
|
||||
"checksum": f"h{i}",
|
||||
"size": 1024,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [256],
|
||||
}
|
||||
for i in range(4)
|
||||
}
|
||||
manifests = {r: dict(tensors) for r in range(4)}
|
||||
self._write_manifests(tmp, manifests)
|
||||
plan = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=4,
|
||||
tmp_dir=tmp,
|
||||
max_file_bytes=10**12,
|
||||
)
|
||||
# 4 tensors, 4 writers, 4 files (one per writer).
|
||||
self.assertEqual(len(plan["files"]), 4)
|
||||
writers = sorted(f["writer_rank"] for f in plan["files"])
|
||||
self.assertEqual(writers, [0, 1, 2, 3])
|
||||
|
||||
def test_round_trip_dump_and_read_back(self):
|
||||
# End-to-end on disk for world_size=1: build manifests for tensors,
|
||||
# construct plan, write safetensors per the plan, read back and
|
||||
# verify both checksums and bit-identity.
|
||||
from safetensors.torch import safe_open, save_file
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
tensors = {
|
||||
"embed.weight": torch.arange(64, dtype=torch.float32).reshape(8, 8),
|
||||
"norm.weight": torch.full((16,), 0.5, dtype=torch.float32),
|
||||
"head.weight": torch.linspace(-1, 1, 32, dtype=torch.float32),
|
||||
}
|
||||
manifest_dir = os.path.join(tmp, "manifests")
|
||||
presharded_dir = os.path.join(tmp, "presharded")
|
||||
os.makedirs(manifest_dir)
|
||||
os.makedirs(presharded_dir)
|
||||
|
||||
checksums = {
|
||||
name: PreshardedModelLoader._hash_tensor(t)
|
||||
for name, t in tensors.items()
|
||||
}
|
||||
manifest = {
|
||||
name: {
|
||||
"checksum": checksums[name],
|
||||
"size": t.numel() * t.element_size(),
|
||||
"dtype": str(t.dtype),
|
||||
"shape": list(t.shape),
|
||||
}
|
||||
for name, t in tensors.items()
|
||||
}
|
||||
with open(os.path.join(manifest_dir, "manifest_00000.json"), "w") as f:
|
||||
json.dump(manifest, f)
|
||||
|
||||
plan = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=1, tmp_dir=manifest_dir, max_file_bytes=10**12
|
||||
)
|
||||
|
||||
# Single rank → all tensors live in one common file.
|
||||
self.assertEqual(len(plan["files"]), 1)
|
||||
f = plan["files"][0]
|
||||
self.assertTrue(f["is_common"])
|
||||
|
||||
# Write the file with stored_keys mapped to tensor content.
|
||||
to_save = {}
|
||||
for t_entry in f["tensors"]:
|
||||
name = t_entry["rank_to_names"]["0"][0]
|
||||
to_save[t_entry["stored_key"]] = tensors[name]
|
||||
save_file(to_save, os.path.join(presharded_dir, f["filename"]))
|
||||
|
||||
# Read back: verify each tensor's checksum and content.
|
||||
with safe_open(
|
||||
os.path.join(presharded_dir, f["filename"]), framework="pt"
|
||||
) as fh:
|
||||
for r in plan["rank_to_reads"]["0"]:
|
||||
loaded = fh.get_tensor(r["stored_key"])
|
||||
self.assertEqual(
|
||||
PreshardedModelLoader._hash_tensor(loaded), r["stored_key"]
|
||||
)
|
||||
torch.testing.assert_close(loaded, tensors[r["name"]])
|
||||
|
||||
def test_dedup_with_multiple_names_per_rank(self):
|
||||
# Same checksum can appear under MULTIPLE param names on the same
|
||||
# rank (e.g., k_scale and v_scale both default to 1.0). Both names
|
||||
# must end up in rank_to_reads.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
shared = {
|
||||
"checksum": "scale_hash",
|
||||
"size": 4,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [],
|
||||
}
|
||||
manifests = {
|
||||
0: {
|
||||
"layers.0.attn.k_scale": shared,
|
||||
"layers.0.attn.v_scale": shared,
|
||||
"layers.1.attn.k_scale": shared,
|
||||
},
|
||||
1: {
|
||||
"layers.0.attn.k_scale": shared,
|
||||
"layers.0.attn.v_scale": shared,
|
||||
"layers.1.attn.k_scale": shared,
|
||||
},
|
||||
}
|
||||
self._write_manifests(tmp, manifests)
|
||||
plan = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=2, tmp_dir=tmp, max_file_bytes=10**12
|
||||
)
|
||||
# All 6 (rank,name) pairs must be readable, even though there is one
|
||||
# underlying tensor stored on disk.
|
||||
self.assertEqual(len(plan["files"]), 1)
|
||||
for r in (0, 1):
|
||||
reads = plan["rank_to_reads"][str(r)]
|
||||
names = sorted(rd["name"] for rd in reads)
|
||||
self.assertEqual(
|
||||
names,
|
||||
[
|
||||
"layers.0.attn.k_scale",
|
||||
"layers.0.attn.v_scale",
|
||||
"layers.1.attn.k_scale",
|
||||
],
|
||||
)
|
||||
# All point at the same stored_key (deduplicated content).
|
||||
self.assertEqual({rd["stored_key"] for rd in reads}, {"scale_hash"})
|
||||
|
||||
def test_collision_size_mismatch_raises(self):
|
||||
# Same checksum but different sizes ⇒ plan builder rejects.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
manifests = {
|
||||
0: {
|
||||
"x": {
|
||||
"checksum": "same",
|
||||
"size": 1024,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [256],
|
||||
}
|
||||
},
|
||||
1: {
|
||||
"x": {
|
||||
"checksum": "same",
|
||||
"size": 2048,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [512],
|
||||
}
|
||||
},
|
||||
}
|
||||
self._write_manifests(tmp, manifests)
|
||||
with self.assertRaises(RuntimeError):
|
||||
PreshardedModelLoader._build_dump_plan(
|
||||
world_size=2, tmp_dir=tmp, max_file_bytes=10**12
|
||||
)
|
||||
|
||||
def test_rank_checksum_deterministic(self):
|
||||
# rank_checksums must be reproducible from the same manifest input
|
||||
# and depend on (name, content-SHA) pairs of every tensor a rank
|
||||
# owns. Permuting the manifest's insertion order must not change
|
||||
# the rank checksum.
|
||||
with tempfile.TemporaryDirectory() as tmp_a, tempfile.TemporaryDirectory() as tmp_b:
|
||||
base_entries = {
|
||||
"alpha.weight": {
|
||||
"checksum": "h_alpha",
|
||||
"size": 16,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [4],
|
||||
},
|
||||
"beta.weight": {
|
||||
"checksum": "h_beta",
|
||||
"size": 16,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [4],
|
||||
},
|
||||
}
|
||||
self._write_manifests(tmp_a, {0: dict(base_entries)})
|
||||
# Insertion-order-permuted copy.
|
||||
permuted = {k: base_entries[k] for k in reversed(list(base_entries))}
|
||||
self._write_manifests(tmp_b, {0: permuted})
|
||||
plan_a = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=1, tmp_dir=tmp_a, max_file_bytes=10**12
|
||||
)
|
||||
plan_b = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=1, tmp_dir=tmp_b, max_file_bytes=10**12
|
||||
)
|
||||
self.assertEqual(plan_a["rank_checksums"], plan_b["rank_checksums"])
|
||||
|
||||
def test_rank_checksum_distinguishes_content(self):
|
||||
# Changing one tensor's content-SHA must change the rank checksum.
|
||||
with tempfile.TemporaryDirectory() as tmp_a, tempfile.TemporaryDirectory() as tmp_b:
|
||||
entries_a = {
|
||||
"x.weight": {
|
||||
"checksum": "ha",
|
||||
"size": 16,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [4],
|
||||
},
|
||||
}
|
||||
entries_b = {
|
||||
"x.weight": {
|
||||
"checksum": "hb", # different content
|
||||
"size": 16,
|
||||
"dtype": "torch.float32",
|
||||
"shape": [4],
|
||||
},
|
||||
}
|
||||
self._write_manifests(tmp_a, {0: entries_a})
|
||||
self._write_manifests(tmp_b, {0: entries_b})
|
||||
plan_a = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=1, tmp_dir=tmp_a, max_file_bytes=10**12
|
||||
)
|
||||
plan_b = PreshardedModelLoader._build_dump_plan(
|
||||
world_size=1, tmp_dir=tmp_b, max_file_bytes=10**12
|
||||
)
|
||||
self.assertNotEqual(
|
||||
plan_a["rank_checksums"]["0"],
|
||||
plan_b["rank_checksums"]["0"],
|
||||
)
|
||||
|
||||
def test_presharded_ready_sentinel(self):
|
||||
# Loader treats a dir as a valid presharded ckpt only when the
|
||||
# READY sentinel exists. A bare checksum.json (or partial files)
|
||||
# is NOT enough.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
self.assertFalse(PreshardedModelLoader._presharded_ready(tmp))
|
||||
# checksum.json alone is not sufficient.
|
||||
with open(
|
||||
os.path.join(tmp, PreshardedModelLoader.CHECKSUM_FILENAME), "w"
|
||||
) as f:
|
||||
f.write("{}")
|
||||
self.assertFalse(PreshardedModelLoader._presharded_ready(tmp))
|
||||
# READY makes it ready.
|
||||
with open(
|
||||
os.path.join(tmp, PreshardedModelLoader.READY_FILENAME), "w"
|
||||
) as f:
|
||||
f.write("{}")
|
||||
self.assertTrue(PreshardedModelLoader._presharded_ready(tmp))
|
||||
|
||||
def test_separate_presharded_path_overrides_for_target_and_draft(self):
|
||||
# Target and draft get distinct roots via presharded_path vs
|
||||
# draft_presharded_path. Failure mode if this regresses: draft
|
||||
# re-dump collides with / wipes target READY under a shared path.
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
target_root = os.path.join(tmp, "target_cache")
|
||||
draft_root = os.path.join(tmp, "draft_cache")
|
||||
loader._presharded_path_override = target_root
|
||||
loader._draft_presharded_path_override = draft_root
|
||||
cfg = {"tp": 8, "structural_signature": "same_sig"}
|
||||
sub = loader._build_subfolder_name(cfg)
|
||||
target_mc = SimpleNamespace(model_path="/models/dsv3", is_draft_model=False)
|
||||
draft_mc = SimpleNamespace(model_path="/models/dsv3", is_draft_model=True)
|
||||
target_dir = loader._presharded_dir(target_mc, cfg)
|
||||
draft_dir = loader._presharded_dir(draft_mc, cfg)
|
||||
self.assertEqual(target_dir, os.path.join(target_root, sub))
|
||||
self.assertEqual(draft_dir, os.path.join(draft_root, sub))
|
||||
self.assertNotEqual(target_dir, draft_dir)
|
||||
|
||||
# Target-only override: draft must not fall back into the target
|
||||
# root (same model_path is common for DeepSeek MTP).
|
||||
loader._draft_presharded_path_override = None
|
||||
draft_fallback = loader._presharded_dir(draft_mc, cfg)
|
||||
self.assertEqual(
|
||||
draft_fallback,
|
||||
os.path.join(
|
||||
"/models/dsv3",
|
||||
PreshardedModelLoader.DEFAULT_SUBDIR,
|
||||
sub,
|
||||
),
|
||||
)
|
||||
self.assertFalse(draft_fallback.startswith(target_root))
|
||||
|
||||
# No overrides: both use model_path/presharded/<subfolder>.
|
||||
loader._presharded_path_override = None
|
||||
self.assertEqual(
|
||||
loader._presharded_dir(target_mc, cfg),
|
||||
os.path.join(
|
||||
"/models/dsv3",
|
||||
PreshardedModelLoader.DEFAULT_SUBDIR,
|
||||
sub,
|
||||
),
|
||||
)
|
||||
|
||||
def test_apply_shape_mismatch_raises(self):
|
||||
# Reload must not silently copy_ into a wrong layout when process
|
||||
# shapes and dumped tensors disagree (previously only warned).
|
||||
# Use a *larger* dumped tensor so prefix-narrow does not paper over it.
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
loader._verify_on_load = False
|
||||
param = torch.nn.Parameter(torch.zeros(2, 2))
|
||||
state_dict = {"w": param}
|
||||
items = [{"name": "w", "stored_key": "w", "is_extra": False}]
|
||||
cached = {"w": torch.ones(4, 4)}
|
||||
loaded: set = set()
|
||||
with self.assertRaises(ValueError) as ctx:
|
||||
loader._apply_presharded_file(
|
||||
items=items,
|
||||
cached=cached,
|
||||
model=torch.nn.Module(),
|
||||
state_dict=state_dict,
|
||||
target_device=torch.device("cpu"),
|
||||
loaded_param_keys=loaded,
|
||||
verify_hashes=[],
|
||||
)
|
||||
self.assertIn("shape mismatch", str(ctx.exception).lower())
|
||||
|
||||
def test_build_dump_plan_missing_manifest_mentions_shared_fs(self):
|
||||
# Multi-node without a shared dump dir fails at plan build with a
|
||||
# clear message (not a bare FileNotFoundError path).
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
with self.assertRaises(FileNotFoundError) as ctx:
|
||||
PreshardedModelLoader._build_dump_plan(
|
||||
world_size=2, tmp_dir=tmp, max_file_bytes=1024
|
||||
)
|
||||
msg = str(ctx.exception)
|
||||
self.assertIn("Rank 0", msg)
|
||||
self.assertIn("shared", msg.lower())
|
||||
|
||||
def test_ensure_presharded_dir_writable_rejects_readonly(self):
|
||||
# Guard against spending a full source load before discovering a
|
||||
# read-only dump root (HF cache mounts). Mock OSError because root
|
||||
# can often still write to mode-0555 dirs under CAP_DAC_OVERRIDE.
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
with mock.patch.object(
|
||||
loader, "_world_rank_and_size", return_value=(0, 1)
|
||||
), mock.patch.object(loader, "_world_barrier"), mock.patch(
|
||||
"os.makedirs", side_effect=OSError("Read-only file system")
|
||||
):
|
||||
with self.assertRaises(RuntimeError) as ctx:
|
||||
loader._ensure_presharded_dir_writable("/ro/presharded")
|
||||
self.assertIn("not writable", str(ctx.exception).lower())
|
||||
self.assertIn("presharded_path", str(ctx.exception))
|
||||
|
||||
def test_ensure_presharded_dir_writable_ok_rank0(self):
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
leaf = os.path.join(tmp, "TP-8-sig-test")
|
||||
with mock.patch.object(
|
||||
loader, "_world_rank_and_size", return_value=(0, 1)
|
||||
), mock.patch.object(loader, "_world_barrier") as barrier:
|
||||
loader._ensure_presharded_dir_writable(leaf)
|
||||
self.assertTrue(os.path.isdir(leaf))
|
||||
barrier.assert_called_once()
|
||||
|
||||
|
||||
class TestStructuralSignature(unittest.TestCase):
|
||||
"""The structural signature is the long-term fix for the underlying
|
||||
problem `moe_dense_tp_size` was an instance of: instead of hand-
|
||||
enumerating every parallelism knob that might change a rank's tensor
|
||||
shapes, hash the (name, shape, dtype) of every parameter in a
|
||||
meta-device model skeleton built under the live parallel state. Any
|
||||
future sharding knob that changes shapes is automatically caught
|
||||
without touching `_build_subfolder_name`."""
|
||||
|
||||
def test_hash_changes_with_shape_dtype_or_added_tensor(self):
|
||||
# Regression: a future sharding knob that changes shapes/dtypes or
|
||||
# adds per-rank buffers must change the digest, otherwise cache
|
||||
# collisions load wrong weights. One case covers the three axes that
|
||||
# `_hash_structural_signature` is responsible for.
|
||||
base = [("a.weight", (4, 4), "torch.float32")]
|
||||
shape = [("a.weight", (4, 8), "torch.float32")]
|
||||
dtype = [("a.weight", (4, 4), "torch.float16")]
|
||||
extra = [
|
||||
("a.weight", (4, 4), "torch.float32"),
|
||||
("a.extra_buf", (4,), "torch.float32"),
|
||||
]
|
||||
h = PreshardedModelLoader._hash_structural_signature
|
||||
self.assertNotEqual(h(base), h(shape))
|
||||
self.assertNotEqual(h(base), h(dtype))
|
||||
self.assertNotEqual(h(base), h(extra))
|
||||
|
||||
def test_local_signature_sorts_state_dict_order(self):
|
||||
# Production path sorts state_dict items before hashing. If that
|
||||
# sorted(...) is dropped, two identical models whose state_dict
|
||||
# iteration order differs would get different signatures and thrash
|
||||
# the cache. Guard the sort, not just the pure hash helper.
|
||||
import torch.nn as nn
|
||||
|
||||
class OrderedModule(nn.Module):
|
||||
def __init__(self, order):
|
||||
super().__init__()
|
||||
for name in order:
|
||||
self.register_parameter(
|
||||
name, nn.Parameter(torch.zeros(2, 2), requires_grad=False)
|
||||
)
|
||||
|
||||
def _init_stub(model_config, load_config, quant_config):
|
||||
return OrderedModule(model_config.order)
|
||||
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
loader.load_config = SimpleNamespace()
|
||||
with mock.patch(
|
||||
"sglang.srt.model_loader.loader._get_quantization_config",
|
||||
return_value=None,
|
||||
), mock.patch(
|
||||
"sglang.srt.model_loader.loader._initialize_model",
|
||||
side_effect=_init_stub,
|
||||
), mock.patch(
|
||||
"sglang.srt.model_loader.loader.set_default_torch_dtype",
|
||||
return_value=mock.MagicMock(__enter__=mock.Mock(), __exit__=mock.Mock()),
|
||||
):
|
||||
sig_ab = loader._compute_local_structural_signature(
|
||||
SimpleNamespace(
|
||||
quantization=None, dtype=torch.float32, order=["a", "b"]
|
||||
)
|
||||
)
|
||||
sig_ba = loader._compute_local_structural_signature(
|
||||
SimpleNamespace(
|
||||
quantization=None, dtype=torch.float32, order=["b", "a"]
|
||||
)
|
||||
)
|
||||
self.assertIsNotNone(sig_ab)
|
||||
self.assertEqual(sig_ab, sig_ba)
|
||||
|
||||
def test_rank_invariant_signature_aggregates_per_rank_locals(self):
|
||||
# Under PP, ranks build different local digests. The shared cache key
|
||||
# must still agree across ranks: all-gather then hash the ordered
|
||||
# list. Without this, each PP stage would pick a different
|
||||
# presharded subfolder and the multi-rank dump protocol breaks.
|
||||
fake_group = mock.Mock()
|
||||
fake_group.world_size = 2
|
||||
# Simulate two ranks each calling with their own local sig; both
|
||||
# must see the same gathered list and thus the same aggregate.
|
||||
fake_group.all_gather_object.side_effect = lambda local: ["sig-pp0", "sig-pp1"]
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.distributed.get_world_group", return_value=fake_group
|
||||
):
|
||||
agg_from_rank0 = (
|
||||
PreshardedModelLoader._make_rank_invariant_structural_signature(
|
||||
"sig-pp0"
|
||||
)
|
||||
)
|
||||
agg_from_rank1 = (
|
||||
PreshardedModelLoader._make_rank_invariant_structural_signature(
|
||||
"sig-pp1"
|
||||
)
|
||||
)
|
||||
self.assertIsNotNone(agg_from_rank0)
|
||||
self.assertEqual(agg_from_rank0, agg_from_rank1)
|
||||
# Changing any rank's local contribution must change the aggregate.
|
||||
fake_group.all_gather_object.side_effect = lambda local: [
|
||||
"sig-pp0",
|
||||
"sig-pp1-changed",
|
||||
]
|
||||
with mock.patch(
|
||||
"sglang.srt.distributed.get_world_group", return_value=fake_group
|
||||
):
|
||||
agg_changed = (
|
||||
PreshardedModelLoader._make_rank_invariant_structural_signature(
|
||||
"sig-pp0"
|
||||
)
|
||||
)
|
||||
self.assertNotEqual(agg_from_rank0, agg_changed)
|
||||
|
||||
def test_compute_structural_signature_returns_none_on_failure(self):
|
||||
# No self.load_config (and no real model class behind
|
||||
# SimpleNamespace) means _get_quantization_config / _initialize_model
|
||||
# will raise; this must be swallowed and return None, never raise,
|
||||
# since a model class that can't be built on meta device must not
|
||||
# break the overall load.
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
model_config = SimpleNamespace(quantization=None)
|
||||
self.assertIsNone(loader._compute_structural_signature(model_config))
|
||||
|
||||
def test_compute_structural_signature_picks_up_meta_model_shapes(self):
|
||||
# End-to-end on a tiny real nn.Module standing in for the model
|
||||
# class, to prove the meta-device construction + hashing wiring
|
||||
# actually reflects shapes that depend on the live parallel state
|
||||
# (here simulated via a width captured at construction time).
|
||||
import torch.nn as nn
|
||||
|
||||
class FakeModelLoader(PreshardedModelLoader):
|
||||
def __init__(self, width):
|
||||
self._width = width
|
||||
self.load_config = SimpleNamespace()
|
||||
|
||||
def _initialize_model_stub(model_config, load_config, quant_config):
|
||||
return nn.Linear(4, model_config.width, bias=False)
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.model_loader.loader._get_quantization_config",
|
||||
return_value=None,
|
||||
), mock.patch(
|
||||
"sglang.srt.model_loader.loader._initialize_model",
|
||||
side_effect=_initialize_model_stub,
|
||||
), mock.patch(
|
||||
"sglang.srt.model_loader.loader.set_default_torch_dtype",
|
||||
return_value=mock.MagicMock(__enter__=mock.Mock(), __exit__=mock.Mock()),
|
||||
):
|
||||
loader_narrow = FakeModelLoader(width=2)
|
||||
loader_wide = FakeModelLoader(width=8)
|
||||
model_config_narrow = SimpleNamespace(
|
||||
quantization=None, dtype=torch.float32, width=2
|
||||
)
|
||||
model_config_wide = SimpleNamespace(
|
||||
quantization=None, dtype=torch.float32, width=8
|
||||
)
|
||||
sig_narrow = loader_narrow._compute_structural_signature(
|
||||
model_config_narrow
|
||||
)
|
||||
sig_wide = loader_wide._compute_structural_signature(model_config_wide)
|
||||
self.assertIsNotNone(sig_narrow)
|
||||
self.assertIsNotNone(sig_wide)
|
||||
self.assertNotEqual(sig_narrow, sig_wide)
|
||||
|
||||
def test_meta_rope_cache_cleared_even_on_failure(self):
|
||||
# If _initialize_model partially populates _ROPE_DICT with meta-device
|
||||
# entries before raising, _compute_structural_signature must still
|
||||
# clean them up (via finally), otherwise the real model init reuses
|
||||
# the meta module and fails with "Cannot copy out of meta tensor".
|
||||
import torch.nn as nn
|
||||
|
||||
from sglang.srt.layers.rotary_embedding.factory import _ROPE_DICT
|
||||
|
||||
# Plant a fake meta-device rotary module in the global cache.
|
||||
fake_key = ("_test_meta_rope_cleanup_sentinel",)
|
||||
with torch.device("meta"):
|
||||
fake_module = nn.Linear(4, 4)
|
||||
_ROPE_DICT[fake_key] = fake_module
|
||||
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
loader.load_config = SimpleNamespace()
|
||||
|
||||
with mock.patch(
|
||||
"sglang.srt.model_loader.loader._get_quantization_config",
|
||||
return_value=None,
|
||||
), mock.patch(
|
||||
"sglang.srt.model_loader.loader._initialize_model",
|
||||
side_effect=RuntimeError("simulated init failure"),
|
||||
), mock.patch(
|
||||
"sglang.srt.model_loader.loader.set_default_torch_dtype",
|
||||
return_value=mock.MagicMock(__enter__=mock.Mock(), __exit__=mock.Mock()),
|
||||
):
|
||||
result = loader._compute_structural_signature(
|
||||
SimpleNamespace(quantization=None, dtype=torch.float32)
|
||||
)
|
||||
|
||||
self.assertIsNone(result)
|
||||
self.assertNotIn(fake_key, _ROPE_DICT)
|
||||
|
||||
|
||||
class TestShardConfig(unittest.TestCase):
|
||||
"""Bookkeeping guards for the enumerated cache-key fields and the
|
||||
load-time match path. Dropping a field from `_collect_shard_config` or
|
||||
breaking `_shard_config_matches` would silently collide caches; these
|
||||
cases pin the failure modes the PR is meant to prevent."""
|
||||
|
||||
def _base_config(self, **overrides):
|
||||
cfg = {
|
||||
"tp": 8,
|
||||
"dp": 1,
|
||||
"ep": 1,
|
||||
"pp": 1,
|
||||
"moe_dense_tp_size": None,
|
||||
"moe_dp_size": 1,
|
||||
"enable_dp_lm_head": False,
|
||||
"enable_fp32_lm_head": False,
|
||||
"quantization": None,
|
||||
"model_dtype": "torch.bfloat16",
|
||||
"ep_num_redundant_experts": 0,
|
||||
"enable_eplb": False,
|
||||
"init_expert_location": "trivial",
|
||||
"structural_signature": "deadbeef",
|
||||
}
|
||||
cfg.update(overrides)
|
||||
return cfg
|
||||
|
||||
def test_collect_shard_config_includes_required_keys(self):
|
||||
# Registry completeness: dropping a key from the dict literal in
|
||||
# `_collect_shard_config` is the exact failure mode that left
|
||||
# moe_dense_tp_size / LM-head flags out of the cache key before.
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
server_args = SimpleNamespace(
|
||||
moe_dense_tp_size=1,
|
||||
moe_dp_size=2,
|
||||
enable_dp_lm_head=True,
|
||||
enable_fp32_lm_head=True,
|
||||
ep_num_redundant_experts=4,
|
||||
enable_eplb=True,
|
||||
init_expert_location="trivial",
|
||||
)
|
||||
model_config = SimpleNamespace(quantization="fp8", dtype=torch.bfloat16)
|
||||
required = {
|
||||
"tp",
|
||||
"dp",
|
||||
"ep",
|
||||
"pp",
|
||||
"moe_dense_tp_size",
|
||||
"moe_dp_size",
|
||||
"enable_dp_lm_head",
|
||||
"enable_fp32_lm_head",
|
||||
"quantization",
|
||||
"model_dtype",
|
||||
"ep_num_redundant_experts",
|
||||
"enable_eplb",
|
||||
"init_expert_location",
|
||||
"structural_signature",
|
||||
}
|
||||
parallel = SimpleNamespace(tp_size=8, moe_dp_size=2, moe_ep_size=4, pp_size=1)
|
||||
with mock.patch(
|
||||
"sglang.srt.model_loader.loader.get_server_args",
|
||||
return_value=server_args,
|
||||
), mock.patch(
|
||||
"sglang.srt.model_loader.loader.get_parallel",
|
||||
return_value=parallel,
|
||||
), mock.patch.object(
|
||||
loader, "_compute_structural_signature", return_value="sig16"
|
||||
):
|
||||
cfg = loader._collect_shard_config(model_config)
|
||||
self.assertEqual(required, set(cfg.keys()))
|
||||
self.assertEqual(cfg["tp"], 8)
|
||||
self.assertEqual(cfg["dp"], 2)
|
||||
self.assertEqual(cfg["ep"], 4)
|
||||
self.assertEqual(cfg["pp"], 1)
|
||||
self.assertEqual(cfg["moe_dense_tp_size"], 1)
|
||||
self.assertEqual(cfg["moe_dp_size"], 2)
|
||||
self.assertTrue(cfg["enable_dp_lm_head"])
|
||||
self.assertTrue(cfg["enable_fp32_lm_head"])
|
||||
self.assertEqual(cfg["init_expert_location"], "trivial")
|
||||
|
||||
def test_enumerated_fields_change_subfolder_hash(self):
|
||||
# Each enumerated content/shape knob must feed the subfolder name.
|
||||
# A silent drop from `_collect_shard_config` would leave this red.
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
base = self._base_config()
|
||||
base_name = loader._build_subfolder_name(base)
|
||||
field_variants = {
|
||||
"moe_dense_tp_size": 1,
|
||||
"moe_dp_size": 2,
|
||||
"enable_dp_lm_head": True,
|
||||
"enable_fp32_lm_head": True,
|
||||
"ep_num_redundant_experts": 8,
|
||||
"enable_eplb": True,
|
||||
"init_expert_location": "file:map.json:sha1:abcd",
|
||||
"structural_signature": "cafebabe",
|
||||
"quantization": "fp8",
|
||||
}
|
||||
for field, value in field_variants.items():
|
||||
with self.subTest(field=field):
|
||||
other = self._base_config(**{field: value})
|
||||
other_name = loader._build_subfolder_name(other)
|
||||
self.assertNotEqual(
|
||||
base_name,
|
||||
other_name,
|
||||
f"changing {field} must change the cache subfolder name",
|
||||
)
|
||||
|
||||
def test_shard_config_matches_equality_and_missing(self):
|
||||
# Match must require exact stored equality; missing/mismatched
|
||||
# shard_config is a cache miss (never raises) so reload re-dumps.
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
cfg = self._base_config()
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
# No checksum.json → miss.
|
||||
self.assertFalse(loader._shard_config_matches(tmp, cfg))
|
||||
# Matching stored config → hit.
|
||||
with open(
|
||||
os.path.join(tmp, PreshardedModelLoader.CHECKSUM_FILENAME), "w"
|
||||
) as f:
|
||||
json.dump({"shard_config": cfg}, f)
|
||||
self.assertTrue(loader._shard_config_matches(tmp, cfg))
|
||||
# One field differs → miss.
|
||||
mismatched = self._base_config(moe_dense_tp_size=1)
|
||||
self.assertFalse(loader._shard_config_matches(tmp, mismatched))
|
||||
# Plan without shard_config key → miss (upgrade path).
|
||||
with open(
|
||||
os.path.join(tmp, PreshardedModelLoader.CHECKSUM_FILENAME), "w"
|
||||
) as f:
|
||||
json.dump({"version": 1}, f)
|
||||
self.assertFalse(loader._shard_config_matches(tmp, cfg))
|
||||
|
||||
def test_init_expert_location_hashes_file_contents(self):
|
||||
# Overwriting the same path with a different expert map must bust
|
||||
# the cache; keying on the path alone would silently reuse weights.
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
path = os.path.join(tmp, "experts.json")
|
||||
with open(path, "w") as f:
|
||||
json.dump({"logical_count": [[1, 0], [0, 1]]}, f)
|
||||
key_a = PreshardedModelLoader._normalize_init_expert_location(path)
|
||||
with open(path, "w") as f:
|
||||
json.dump({"logical_count": [[0, 1], [1, 0]]}, f)
|
||||
key_b = PreshardedModelLoader._normalize_init_expert_location(path)
|
||||
self.assertNotEqual(key_a, key_b)
|
||||
self.assertTrue(key_a.startswith("file:experts.json:sha1:"))
|
||||
self.assertEqual(
|
||||
PreshardedModelLoader._normalize_init_expert_location("trivial"),
|
||||
"trivial",
|
||||
)
|
||||
|
||||
def test_redump_clears_ready_before_rewrite(self):
|
||||
# Config-mismatch re-dump into an already-ready dir must drop READY
|
||||
# before mutating files, otherwise concurrent readers can observe
|
||||
# READY + partial checksum/safetensors.
|
||||
loader = object.__new__(PreshardedModelLoader)
|
||||
loader._hash_num_threads = 1
|
||||
loader._max_file_bytes = 10**12
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
ready_path = os.path.join(tmp, PreshardedModelLoader.READY_FILENAME)
|
||||
with open(ready_path, "w") as f:
|
||||
f.write("{}")
|
||||
self.assertTrue(os.path.isfile(ready_path))
|
||||
|
||||
# Force rank 0 / world 1 so the method runs the rank-0 prologue
|
||||
# and then fails early on empty state (no need for full dump).
|
||||
with mock.patch.object(
|
||||
PreshardedModelLoader,
|
||||
"_world_rank_and_size",
|
||||
return_value=(0, 1),
|
||||
), mock.patch.object(
|
||||
PreshardedModelLoader, "_world_barrier", return_value=None
|
||||
), mock.patch.object(
|
||||
PreshardedModelLoader,
|
||||
"_build_dump_plan",
|
||||
return_value={
|
||||
"version": PreshardedModelLoader.PLAN_VERSION,
|
||||
"files": [],
|
||||
"rank_to_reads": {"0": []},
|
||||
"rank_checksums": {"0": "0"},
|
||||
},
|
||||
), mock.patch.object(
|
||||
PreshardedModelLoader, "_dump_files_for_rank", return_value=None
|
||||
):
|
||||
loader._dump_state_to_disk(
|
||||
state_dict={},
|
||||
extras={},
|
||||
presharded_dir=tmp,
|
||||
shard_config=self._base_config(),
|
||||
)
|
||||
# Dump rewrote READY at the end; the critical property is that
|
||||
# the prologue unlinked the *previous* READY before rewriting
|
||||
# checksum.json. Assert checksum was written and READY exists
|
||||
# only as the fresh sentinel from this dump.
|
||||
self.assertTrue(os.path.isfile(ready_path))
|
||||
with open(ready_path) as f:
|
||||
sentinel = json.load(f)
|
||||
self.assertIn("plan_version", sentinel)
|
||||
with open(os.path.join(tmp, PreshardedModelLoader.CHECKSUM_FILENAME)) as f:
|
||||
plan = json.load(f)
|
||||
self.assertEqual(plan["shard_config"]["tp"], 8)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user