diff --git a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py index decfcdd56..42af1fc0f 100644 --- a/python/sglang/srt/model_executor/runner/flashinfer_autotune.py +++ b/python/sglang/srt/model_executor/runner/flashinfer_autotune.py @@ -17,6 +17,7 @@ import contextlib import datetime import functools import hashlib +import json import logging from pathlib import Path from typing import TYPE_CHECKING, Callable, Optional @@ -34,6 +35,7 @@ from sglang.srt.runtime_context import ( from sglang.srt.utils import empty_context, log_info_on_rank0 if TYPE_CHECKING: + from sglang.srt.distributed.parallel_state import GroupCoordinator from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.runner.base_runner import BaseRunner @@ -170,14 +172,90 @@ def flashinfer_autotune_cache_path(model_runner: ModelRunner) -> Path: ) +def _autotune_tactic_sync_group( + tp_group: GroupCoordinator, +) -> Optional[torch.distributed.ProcessGroup]: + """CPU group over the ranks that must agree on the tuned tactics. + + Per-rank timing noise alone makes each rank's ``argmin`` pick a different + tactic for the same shape. FlashInfer all-reduces the timings over this + group so every rank minimizes over the same numbers. TP is the scope: those + ranks run the same dummy forward, and PP stages are already separate groups. + """ + if tp_group.world_size <= 1: + return None + # The CPU group keeps the reduction of these scalars off the profiled stream. + return tp_group.cpu_group + + +@contextlib.contextmanager +def _autotune_process_group(group: Optional[torch.distributed.ProcessGroup]): + """Set FlashInfer's timing-reduction group, restoring the previous one after.""" + from flashinfer.autotuner import ( + get_autotune_process_group, + set_autotune_process_group, + ) + + previous = get_autotune_process_group() + set_autotune_process_group(group) + try: + yield + finally: + set_autotune_process_group(previous) + + +def _autotune_cache_digest(cache_path: Path, env: dict[str, str]) -> str: + """Hash of what this rank would load from ``cache_path`` ("" for nothing). + + Includes the environment: ``load_configs`` ignores the whole file when its + ``_metadata`` stamp disagrees with the environment reading it, so equal + tactics alone do not mean two ranks load the same thing. + """ + if not cache_path.is_file(): + return "" + try: + configs = json.loads(cache_path.read_text()) + except (OSError, ValueError): + return "" + if not isinstance(configs, dict): + return "" + payload = {"file": configs, "env": env} + return hashlib.sha256(json.dumps(payload, sort_keys=True).encode()).hexdigest() + + +def _drop_diverged_autotune_cache( + cache_path: Path, group: torch.distributed.ProcessGroup, env: dict[str, str] +) -> None: + """Enter tuning with the same cache on every rank, or with none at all. + + A cache hit skips a profile, so caches that disagree desync the reduction. + """ + digests: list[str] = [""] * torch.distributed.get_world_size(group) + torch.distributed.all_gather_object( + digests, _autotune_cache_digest(cache_path, env), group=group + ) + if len(set(digests)) == 1: + return + log_info_on_rank0( + logger, + "FlashInfer autotune: per-rank caches disagree, discarding them and " + "tuning from scratch so all ranks agree on the tactics.", + ) + cache_path.unlink(missing_ok=True) + + @contextlib.contextmanager def flashinfer_autotune_context(model_runner: ModelRunner, *, run_lm_head: bool): - from flashinfer.autotuner import autotune + # The gate below decides on the same inputs load_configs does. + from flashinfer.autotuner import _collect_metadata, autotune mr = model_runner cache_path = flashinfer_autotune_cache_path(mr) + sync_group = _autotune_tactic_sync_group(mr.tp_group) if envs.SGLANG_FLASHINFER_AUTOTUNE_CACHE.get(): autotune_cache = cache_path + if sync_group is not None: + _drop_diverged_autotune_cache(cache_path, sync_group, _collect_metadata()) logger.info("Running FlashInfer autotune with cache: %s", autotune_cache) else: timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S") @@ -197,7 +275,7 @@ def flashinfer_autotune_context(model_runner: ModelRunner, *, run_lm_head: bool) from sglang.srt.layers.logits_processor import autotune_dummy_run_mode skip_ops = get_flashinfer_autotune_skip_ops(mr) - with autotune( + with _autotune_process_group(sync_group), autotune( True, cache=str(autotune_cache), skip_ops=skip_ops, @@ -321,6 +399,11 @@ def maybe_flashinfer_autotune_extend( try: run_flashinfer_autotune_forward(mr, forward_fn, run_lm_head=False) except torch.OutOfMemoryError: + if _autotune_tactic_sync_group(mr.tp_group) is not None: + # Tuning is collective: this rank has stopped reducing while its + # peers wait on the next tactic, so skipping the pass would hang + # them. Fail instead of degrading alone. + raise # The pass is an optimization; without headroom for the extend-shaped # forward, fall back to untuned extend buckets instead of failing. log_info_on_rank0( diff --git a/test/registered/unit/model_executor/runner/test_flashinfer_autotune_sync.py b/test/registered/unit/model_executor/runner/test_flashinfer_autotune_sync.py new file mode 100644 index 000000000..c2a783e91 --- /dev/null +++ b/test/registered/unit/model_executor/runner/test_flashinfer_autotune_sync.py @@ -0,0 +1,164 @@ +"""FlashInfer autotune must reach the same tactics on every TP rank. + +Without a cross-rank reduction each rank's ``argmin`` follows local timing noise +(measured: 20/20 tuned MoE shapes diverged across 4 ranks on gpt-oss-120b). The +reduction holds only if ranks also enter tuning with the same cache, so these +cover that gate and the digest it decides on. +""" + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=50, suite="base-a-test-cpu") + +import json +import multiprocessing +import os +import tempfile +import traceback +import unittest +from pathlib import Path +from types import SimpleNamespace + +import torch.distributed as dist + +from sglang.srt.model_executor.runner.flashinfer_autotune import ( + _autotune_cache_digest, + _autotune_tactic_sync_group, + _drop_diverged_autotune_cache, +) +from sglang.test.test_utils import CustomTestCase, find_available_port + +ENV = {"flashinfer_version": "0.6.17", "gpu": "NVIDIA GB300"} + + +def _gate_worker(rank, world_size, master_port, cache_path, writer): + """Run the entry gate on one rank; report whether the cache survived.""" + try: + os.environ.update( + RANK=str(rank), + WORLD_SIZE=str(world_size), + MASTER_ADDR="localhost", + MASTER_PORT=str(master_port), + ) + dist.init_process_group("gloo", rank=rank, world_size=world_size) + _drop_diverged_autotune_cache(Path(cache_path), dist.group.WORLD, ENV) + writer.send(("ok", Path(cache_path).is_file())) + except Exception as e: # noqa: BLE001 + traceback.print_exc() + writer.send(("error", f"{e}")) + finally: + writer.close() + if dist.is_initialized(): + dist.destroy_process_group() + + +class TestAutotuneTacticSyncGroup(CustomTestCase): + def test_single_rank_has_nobody_to_agree_with(self): + # A 1-rank group would add a collective per tactic for no agreement. + tp_group = SimpleNamespace(world_size=1, cpu_group=object()) + self.assertIsNone(_autotune_tactic_sync_group(tp_group)) + + +class TestAutotuneCacheDigest(CustomTestCase): + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + self.addCleanup(self.tmp.cleanup) + self.dir = Path(self.tmp.name) + + def _write(self, name: str, configs) -> Path: + path = self.dir / name + path.write_text(json.dumps(configs)) + return path + + def _digest(self, path: Path, env=ENV) -> str: + return _autotune_cache_digest(path, env) + + def test_unusable_caches_read_as_empty(self): + # Files yielding no loadable entries must digest alike, whichever way + # they are unusable; a non-dict also has to not raise. + self.assertEqual(self._digest(self.dir / "absent.json"), "") + corrupt = self.dir / "corrupt.json" + corrupt.write_text("{not json") + self.assertEqual(self._digest(corrupt), "") + self.assertEqual(self._digest(self._write("null.json", None)), "") + self.assertEqual(self._digest(self._write("list.json", [])), "") + + def test_metadata_stamp_decides_whether_entries_load(self): + # Equal tactics, different stamps: one rank loads them, the other + # ignores the file. + rank0 = self._write("rank0.json", {"_metadata": {"cublas": "12.8"}, "op": 7}) + rank1 = self._write("rank1.json", {"_metadata": {"cublas": "12.9"}, "op": 7}) + self.assertNotEqual(self._digest(rank0), self._digest(rank1)) + + def test_environment_is_part_of_the_load_decision(self): + # Same file, drifted environment on one rank: that rank loads nothing. + cache = self._write("rank.json", {"_metadata": {"cublas": "12.8"}, "op": 7}) + self.assertNotEqual( + self._digest(cache), self._digest(cache, {**ENV, "gpu": "NVIDIA B200"}) + ) + + def test_key_order_does_not_matter(self): + # Pins sort_keys: the same tactics must digest alike in any order. + rank0 = self._write("rank0.json", {"a": 1, "b": 2}) + rank1 = self._write("rank1.json", {"b": 2, "a": 1}) + self.assertEqual(self._digest(rank0), self._digest(rank1)) + + +class TestDropDivergedAutotuneCache(CustomTestCase): + """The gate itself, over a real gloo group and real files.""" + + def setUp(self): + self.tmp = tempfile.TemporaryDirectory() + self.addCleanup(self.tmp.cleanup) + self.dir = Path(self.tmp.name) + + def _run_gate(self, per_rank_configs) -> list: + world_size = len(per_rank_configs) + port = find_available_port(23456) + ctx = multiprocessing.get_context("spawn") + procs, readers = [], [] + for rank, configs in enumerate(per_rank_configs): + path = self.dir / f"rank{rank}.json" + path.write_text(json.dumps(configs)) + reader, writer = ctx.Pipe(duplex=False) + proc = ctx.Process( + target=_gate_worker, + args=(rank, world_size, port, str(path), writer), + ) + proc.start() + writer.close() + procs.append(proc) + readers.append(reader) + results = [r.recv() for r in readers] + for proc in procs: + proc.join(timeout=120) + for status, value in results: + self.assertEqual(status, "ok", msg=value) + return [value for _, value in results] + + def test_matching_caches_are_kept(self): + entries = {"_metadata": {"cublas": "12.8"}, "op": 7} + self.assertEqual(self._run_gate([entries, entries]), [True, True]) + + def test_diverged_caches_are_dropped_on_every_rank(self): + # A rank that kept its cache would skip profiles its peer still runs. + meta = {"_metadata": {"cublas": "12.8"}} + self.assertEqual( + self._run_gate([{**meta, "op": 7}, {**meta, "op": 8}]), [False, False] + ) + + def test_caches_diverging_only_in_metadata_are_dropped(self): + # Same desync, reached through the stamp instead of the tactics. + self.assertEqual( + self._run_gate( + [ + {"_metadata": {"cublas": "12.8"}, "op": 7}, + {"_metadata": {"cublas": "12.9"}, "op": 7}, + ] + ), + [False, False], + ) + + +if __name__ == "__main__": + unittest.main()