1010 lines
42 KiB
Python
1010 lines
42 KiB
Python
"""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.srt.runtime_context import get_context, get_parallel
|
|
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.parallel_state.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.parallel_state.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)
|
|
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",
|
|
}
|
|
# The sizes go through both channels: some entries read the published
|
|
# leaf, others the live property. `get_moe_cp_size` is imported inside
|
|
# `_collect_shard_config`, so it is patched where it is defined.
|
|
override = get_context().override_server_args(
|
|
tp_size=8,
|
|
pp_size=1,
|
|
moe_dp_size=2,
|
|
moe_dense_tp_size=1,
|
|
enable_dp_lm_head=True,
|
|
)
|
|
override.install()
|
|
self.addCleanup(override.restore)
|
|
with (
|
|
get_parallel().override(
|
|
tp_size=8, pp_size=1, moe_dp_size=2, moe_ep_size=4, moe_tp_size=1
|
|
),
|
|
mock.patch(
|
|
"sglang.srt.layers.dp_attention.get_moe_cp_size",
|
|
return_value=2,
|
|
),
|
|
mock.patch(
|
|
"sglang.srt.model_loader.loader.get_exec",
|
|
return_value=SimpleNamespace(
|
|
features=SimpleNamespace(enable_fp32_lm_head=True),
|
|
moe=SimpleNamespace(
|
|
ep_num_redundant_experts=4,
|
|
enable_eplb=True,
|
|
init_expert_location="trivial",
|
|
),
|
|
),
|
|
),
|
|
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()
|