[diffusion] feat: persist torch.compile inductor/triton cache across restarts (#28205)
This commit is contained in:
@@ -5,6 +5,7 @@ import gc
|
|||||||
import logging
|
import logging
|
||||||
import multiprocessing as mp
|
import multiprocessing as mp
|
||||||
import os
|
import os
|
||||||
|
import tempfile
|
||||||
import time
|
import time
|
||||||
from contextlib import ExitStack
|
from contextlib import ExitStack
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
@@ -145,6 +146,35 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
|||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
return OutputBatch(output={"released": released, "session_id": session_id})
|
return OutputBatch(output={"released": released, "session_id": session_id})
|
||||||
|
|
||||||
|
def _configure_persistent_torch_compile_cache(self) -> None:
|
||||||
|
"""Persist torch.compile's Inductor/Triton cache across restarts"""
|
||||||
|
compile_cache_root = os.path.join(
|
||||||
|
envs.SGLANG_DIFFUSION_CACHE_ROOT, "torch_compile_cache"
|
||||||
|
)
|
||||||
|
tmp_root = tempfile.gettempdir()
|
||||||
|
for env_name, sub in (
|
||||||
|
("TORCHINDUCTOR_CACHE_DIR", "inductor"),
|
||||||
|
("TRITON_CACHE_DIR", "triton"),
|
||||||
|
):
|
||||||
|
current = os.environ.get(env_name)
|
||||||
|
if current and not current.startswith(tmp_root):
|
||||||
|
# Respect an explicit, non-ephemeral user-provided cache dir.
|
||||||
|
continue
|
||||||
|
cache_path = os.path.join(compile_cache_root, sub)
|
||||||
|
try:
|
||||||
|
os.makedirs(cache_path, exist_ok=True)
|
||||||
|
except OSError as e:
|
||||||
|
logger.warning(
|
||||||
|
"Could not create torch.compile cache dir %s: %s", cache_path, e
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
os.environ[env_name] = cache_path
|
||||||
|
logger.info(
|
||||||
|
"torch.compile cache: TORCHINDUCTOR_CACHE_DIR=%s TRITON_CACHE_DIR=%s",
|
||||||
|
os.environ.get("TORCHINDUCTOR_CACHE_DIR"),
|
||||||
|
os.environ.get("TRITON_CACHE_DIR"),
|
||||||
|
)
|
||||||
|
|
||||||
def init_device_and_model(self) -> None:
|
def init_device_and_model(self) -> None:
|
||||||
"""Initialize the device and load the model."""
|
"""Initialize the device and load the model."""
|
||||||
torch.get_device_module().set_device(self.local_rank)
|
torch.get_device_module().set_device(self.local_rank)
|
||||||
@@ -154,6 +184,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
|
|||||||
os.environ["LOCAL_RANK"] = str(self.local_rank)
|
os.environ["LOCAL_RANK"] = str(self.local_rank)
|
||||||
os.environ["RANK"] = str(self.rank)
|
os.environ["RANK"] = str(self.rank)
|
||||||
os.environ["WORLD_SIZE"] = str(self.server_args.num_gpus)
|
os.environ["WORLD_SIZE"] = str(self.server_args.num_gpus)
|
||||||
|
self._configure_persistent_torch_compile_cache()
|
||||||
# initialize the distributed environment
|
# initialize the distributed environment
|
||||||
maybe_init_distributed_environment_and_model_parallel(
|
maybe_init_distributed_environment_and_model_parallel(
|
||||||
tp_size=self.server_args.tp_size,
|
tp_size=self.server_args.tp_size,
|
||||||
|
|||||||
Reference in New Issue
Block a user