Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions magi_compiler/_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -457,6 +457,14 @@ def _compilation_context(state: MagiCompileState):
{
"TORCHINDUCTOR_CACHE_DIR": (_inductor_cache_dump_path).as_posix(),
"TRITON_CACHE_DIR": (_triton_cache_dump_path).as_posix(),
# Enable Triton autotune result disk cache. With
# autotune_at_compile_time=False, autotuning runs on the first
# forward pass. This flag persists the selected configs as
# .autotune.json in TRITON_CACHE_DIR (which points to
# persistent AFS). Subsequent cold starts (verify pods) read
# the cached results and skip benchmarking entirely,
# eliminating ~3-5 min first-forward overhead.
"TRITON_CACHE_AUTOTUNING": "1",
},
),
explain_compilation(_debug_dump_path.as_posix()),
Expand Down
20 changes: 11 additions & 9 deletions magi_compiler/magi_backend/piecewise_compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,18 +332,20 @@ def compile(
) -> tuple[Callable | None, CacheHandle | None]:
# Step1: Update compile settings
compilation_counter.num_inductor_compiles += 1
current_config = {
# standalone_compile hardcodes autotune_at_compile_time=True, but
# Triton autotune benchmarks with unbacked SymInt dimensions cause
# CUDA illegal-memory-access errors. Disable compile-time autotune
# so that tuning happens at first runtime invocation instead (same
# kernel quality, avoids the crash, and tunes on real shapes).
"triton.autotune_at_compile_time": False
}

# standalone_compile hardcodes autotune_at_compile_time=True, but
# unbacked SymInt dimensions cause CUDA illegal-memory-access when
# autotune benchmarks run at compile time (confirmed on PT 2.9 / B300).
# Override to False: autotuning defers to first forward pass instead.
#
# The runtime autotune results are persisted via TRITON_CACHE_AUTOTUNING=1
# (set in _api.py) — Triton writes .autotune.json to TRITON_CACHE_DIR
# (persistent AFS). Subsequent cold starts read cached results and
# skip benchmarking, eliminating the ~3-5 min overhead.
current_config: dict[str, Any] = {"triton.autotune_at_compile_time": False}
if inductor_compile_config is not None:
current_config.update(inductor_compile_config)
if isinstance(runtime_shape, int):
# for a specific sequence length, tuning triton kernel parameters can be beneficial
current_config.update(
{
"max_autotune": self.compile_config.enable_inductor_max_autotune,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
# Copyright (c) 2025 SandAI. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Subprocess worker for test_autotune_at_compile_time.py.

Runs a Triton autotune kernel and prints elapsed time. The @triton.autotune
decorator triggers benchmarking of all configs on the first invocation.
When TRITON_CACHE_AUTOTUNING=1 and TRITON_CACHE_DIR is set, benchmark results
are persisted to .autotune.json — subsequent processes read the cache and skip
benchmarking.
"""

import time

import torch
import triton
import triton.language as tl


@triton.autotune(
configs=[
triton.Config({"BLOCK_SIZE": 64}, num_warps=2),
triton.Config({"BLOCK_SIZE": 128}, num_warps=4),
triton.Config({"BLOCK_SIZE": 256}, num_warps=4),
triton.Config({"BLOCK_SIZE": 512}, num_warps=8),
],
key=["n_elements"],
)
@triton.jit
def add_kernel(x_ptr, y_ptr, output_ptr, n_elements, BLOCK_SIZE: tl.constexpr):
pid = tl.program_id(axis=0)
offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
mask = offsets < n_elements
x = tl.load(x_ptr + offsets, mask=mask)
y = tl.load(y_ptr + offsets, mask=mask)
tl.store(output_ptr + offsets, x + y, mask=mask)


def main() -> None:
n = 1024 * 1024
x = torch.randn(n, device="cuda")
y = torch.randn(n, device="cuda")
out = torch.empty_like(x)

t0 = time.time()
grid = lambda meta: (triton.cdiv(n, meta["BLOCK_SIZE"]),)
add_kernel[grid](x, y, out, n)
torch.cuda.synchronize()
elapsed = time.time() - t0
print(f"ELAPSED={elapsed:.4f}")


if __name__ == "__main__":
main()
211 changes: 211 additions & 0 deletions tests/feature_tests/cache/test_autotune_at_compile_time.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,211 @@
# Copyright (c) 2025 SandAI. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Tests for autotune_at_compile_time override and TRITON_CACHE_AUTOTUNING.

Bug: autotune_at_compile_time=False defers autotuning to first forward pass.
Triton has a built-in disk cache (Autotuner.check_disk_cache → .autotune.json)
but it's disabled by default (TRITON_CACHE_AUTOTUNING unset). Each new pod
re-benchmarks all kernel configs: ~3-5 min overhead on cold start.

Fix: Set TRITON_CACHE_AUTOTUNING=1 in _compilation_context() so that
bake warmup autotune results persist to AFS and verify pods skip benchmarking.
"""

import glob
import os
import subprocess
import sys
import tempfile
from pathlib import Path
from unittest.mock import MagicMock, patch

import pytest

from magi_compiler.config import CompileConfig

# ──────────────────────────────────────────────────────────────────────
# Part 1: piecewise_compiler must override autotune_at_compile_time=False
# ──────────────────────────────────────────────────────────────────────


class TestAutotuneOverriddenToFalse:
"""Verify config_patches contain autotune_at_compile_time=False."""

@staticmethod
def _capture_config_patches(runtime_shape=None):
import torch.fx as fx

from magi_compiler.magi_backend.piecewise_compiler import InductorStandaloneAdaptor

captured = {}

def fake_standalone_compile(graph, example_inputs, *, dynamic_shapes, options):
captured.update(options.get("config_patches", {}))
mock_artifact = MagicMock()
mock_artifact.compiled_fn = lambda *a: None
mock_artifact.artifacts = None
return mock_artifact

adaptor = InductorStandaloneAdaptor(CompileConfig())
adaptor.initialize_cache(Path(CompileConfig().cache_root_dir) / "test")

graph = fx.Graph()
graph.output(None)
gm = fx.GraphModule({}, graph)

with patch("torch._inductor.standalone_compile", fake_standalone_compile):
try:
adaptor.compile(gm, [], {}, runtime_shape=runtime_shape, key="test")
except Exception:
pass

return captured

def test_dynamic_shape_autotune_is_false(self):
config = self._capture_config_patches(runtime_shape=None)
assert config.get("triton.autotune_at_compile_time") is False

def test_static_shape_autotune_is_false(self):
config = self._capture_config_patches(runtime_shape=128)
assert config.get("triton.autotune_at_compile_time") is False


# ──────────────────────────────────────────────────────────────────────
# Part 2: _compilation_context must set TRITON_CACHE_AUTOTUNING=1
# ──────────────────────────────────────────────────────────────────────


class TestTritonCacheAutotuningEnvVar:
"""Verify _compilation_context sets TRITON_CACHE_AUTOTUNING=1."""

@staticmethod
def _make_state():
from magi_compiler.magi_backend.magi_compiler_base import MagiCompileState

def _dummy_fn(x):
return x

return MagiCompileState(
obj=_dummy_fn, compile_config=CompileConfig(), model_idx=0, model_tag="test", dynamic_arg_dims={}
)

def test_compile_context_enables_triton_cache_autotuning(self):
from magi_compiler._api import _compilation_context

state = self._make_state()
with _compilation_context(state):
assert os.environ.get("TRITON_CACHE_AUTOTUNING") == "1", (
"TRITON_CACHE_AUTOTUNING must be '1' during compilation " "to persist autotune results for cold start reuse"
)

def test_compile_context_sets_triton_cache_dir(self):
from magi_compiler._api import _compilation_context

state = self._make_state()
with _compilation_context(state):
cache_dir = os.environ.get("TRITON_CACHE_DIR", "")
assert cache_dir, "TRITON_CACHE_DIR must be set"
assert "triton_cache" in cache_dir, "TRITON_CACHE_DIR should point to persistent triton_cache dir"


# ──────────────────────────────────────────────────────────────────────
# Part 3: Reproduce the bug + verify the fix (2-process simulation)
# ──────────────────────────────────────────────────────────────────────

_WORKER_SCRIPT = Path(__file__).parent / "cache_reuse_helper" / "autotune_cache_worker.py"


def _run_worker(cache_dir: str, enable_cache: bool) -> float:
"""Run the worker in a subprocess, return elapsed time for the kernel call."""
env = os.environ.copy()
env["TRITON_CACHE_DIR"] = cache_dir
if enable_cache:
env["TRITON_CACHE_AUTOTUNING"] = "1"
else:
env.pop("TRITON_CACHE_AUTOTUNING", None)

result = subprocess.run([sys.executable, str(_WORKER_SCRIPT)], env=env, capture_output=True, text=True, timeout=120)
if result.returncode != 0:
raise RuntimeError(f"Worker failed (rc={result.returncode}):\n{result.stderr[-500:]}")
for line in result.stdout.strip().split("\n"):
if line.startswith("ELAPSED="):
return float(line.split("=")[1])
raise RuntimeError(f"Worker did not print ELAPSED line: {result.stdout}")


@pytest.mark.skipif(not os.path.exists("/dev/nvidia0"), reason="GPU required")
class TestAutotuneCachePersistence:
"""Reproduce the bug and verify the fix via 2-process simulation.

Bug: without TRITON_CACHE_AUTOTUNING=1, autotune results are NOT
persisted to disk → no .autotune.json → 2nd process re-benchmarks.

Fix: with TRITON_CACHE_AUTOTUNING=1, autotune results ARE persisted
→ .autotune.json exists → 2nd process reads cache, skips benchmark.
"""

def test_bug_reproduced_no_cache_without_env(self):
"""WITHOUT TRITON_CACHE_AUTOTUNING=1:
- Process 1 benchmarks and finishes
- No .autotune.json is written to disk
This proves the bug: autotune results are lost between processes.
"""
with tempfile.TemporaryDirectory() as cache_dir:
_run_worker(cache_dir, enable_cache=False)

jsons = glob.glob(os.path.join(cache_dir, "**/*.autotune.json"), recursive=True)
assert len(jsons) == 0, (
f"Bug not reproduced: expected 0 .autotune.json files "
f"without TRITON_CACHE_AUTOTUNING=1, but found {len(jsons)}: {jsons}"
)

def test_bug_fixed_cache_persisted_with_env(self):
"""WITH TRITON_CACHE_AUTOTUNING=1, .autotune.json is written to disk.

Primary evidence: .autotune.json file existence (deterministic).

Supplementary evidence (3-process controlled experiment):
- Process 1 warms both compilation cache AND autotune cache
- Process 2 (cache=True): same compilation cache + autotune HIT
- Process 3 (cache=False): same compilation cache + autotune MISS

Comparing P2 vs P3 isolates the autotune cache variable: both share
the same Triton .cubin compilation cache and CUDA init overhead.
For a single kernel with 4 configs the gap is ~1.5-2x; real models
with hundreds of kernels see ~3-5 min cumulative savings.
"""
with tempfile.TemporaryDirectory() as cache_dir:
# Process 1: warm both compilation and autotune caches
_run_worker(cache_dir, enable_cache=True)

# ── Primary evidence: .autotune.json must exist ──
jsons = glob.glob(os.path.join(cache_dir, "**/*.autotune.json"), recursive=True)
assert len(jsons) >= 1, (
f"Fix not working: expected ≥1 .autotune.json file " f"with TRITON_CACHE_AUTOTUNING=1, but found {len(jsons)}"
)

# ── Supplementary evidence: timing comparison ──
# Process 2: compilation cache warm + autotune cache HIT
t_cached = _run_worker(cache_dir, enable_cache=True)

# Process 3: compilation cache warm + autotune cache MISS
# (must re-benchmark 4 configs × 3 runs = 12 kernel launches)
t_uncached = _run_worker(cache_dir, enable_cache=False)

assert t_uncached > t_cached, (
f"Autotune cache miss should be slower than cache hit "
f"(same compilation cache): cached={t_cached:.3f}s, "
f"uncached={t_uncached:.3f}s, ratio={t_uncached / t_cached:.1f}x"
)
Loading