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
89 changes: 89 additions & 0 deletions dingo/core/nn/compile_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,89 @@
"""
The neural spline flow launches tens of thousands of tiny CUDA kernels per step
(one per spline/indexing op in every coupling transform), so on a modern GPU the
step is bound by kernel-launch overhead rather than by arithmetic. ``torch.compile``
fuses these kernels and removes most of that overhead.
"""

import contextlib
import os
import tempfile

import torch


@contextlib.contextmanager
def eager_mode():
"""Run compiled networks eagerly inside the block, without (re)compiling.

The test epoch runs the network in eval mode. This is a different graph than
what you would see at training. Therefore, it would trigger a
second full compilation. Instead, we do eager evaluation (not using the
fused kernels) to avoid the test loop taking a long time. We could also
compile a "test time" graph, but it does not amortize as well as training time.
"""
with torch.compiler.set_stance("force_eager"):
yield


def is_compiled(network: torch.nn.Module) -> bool:
"""True if ``network`` is (or wraps) a ``torch.compile``-d module."""
while network is not None:
if hasattr(network, "_orig_mod"): # torch.compile OptimizedModule
return True
network = getattr(network, "module", None) # DDP
return False


def compile_network(
network: torch.nn.Module, rank: int = None, cache_dir: str = None
) -> torch.nn.Module:
"""Return ``torch.compile(network)``, set up for (optionally) DDP training.

Parameters
----------
network : torch.nn.Module
The network to compile. Under DDP pass the DDP-wrapped network, so that
the gradient all-reduce can still overlap with the backward pass.
rank : int, optional
DDP rank. When given, each rank is pointed at its own on-disk
Inductor/Triton cache: the ranks compile concurrently and would otherwise
race on the shared cache files.
cache_dir : str, optional
Base directory for the on-disk Inductor/Triton cache (default: the system
temp dir). It must be **node-local**: Triton shared objects written to a
network filesystem can be unloadable from another process, which hangs
the run.
"""
if rank is not None or cache_dir is not None:
base = cache_dir or tempfile.gettempdir()
name = "dingo_inductor" if rank is None else f"dingo_inductor_rank{rank}"
cache = os.path.join(base, name)
os.environ["TORCHINDUCTOR_CACHE_DIR"] = cache
os.environ["TRITON_CACHE_DIR"] = os.path.join(cache, "triton")
return torch.compile(network)


def reset_graphs_if_requires_grad_changes(
network: torch.nn.Module, name_contains: str, requires_grad: bool
) -> bool:
"""Discard the compiled graphs of ``network`` if setting ``requires_grad`` on the
parameters whose name contains ``name_contains`` would change their state.

Dynamo does not guard on ``requires_grad`` of parameters: a graph traced while
a layer was frozen is reused after the layer is unfrozen, and its backward
never produces gradients for that layer (no error, the loss looks normal).
Call this *before* flipping the flags at a stage boundary; the next forward
then re-traces with the new set of trainable parameters. Returns True if the
graphs were reset. No-op for uncompiled networks.
"""
if not is_compiled(network):
return False
params = [p for n, p in network.named_parameters() if name_contains in n]
if not params:
return False
currently_trainable = any(p.requires_grad for p in params)
if currently_trainable == bool(requires_grad):
return False
torch.compiler.reset()
return True
21 changes: 12 additions & 9 deletions dingo/core/posterior_models/base_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,6 @@
import torch.distributed as dist
from threadpoolctl import threadpool_limits
from torch.amp import autocast
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import Dataset

try:
Expand All @@ -40,8 +39,10 @@ def __new__(cls, device="cuda", **kwargs):

import dingo.core.utils as utils
import dingo.core.utils.trainutils
from dingo.core.nn.compile_utils import eager_mode
from dingo.core.utils.backward_compatibility import update_model_config
from dingo.core.utils.misc import get_version
from dingo.core.utils.torchutils import get_ddp_module, unwrap_network
from dingo.core.utils.trainutils import EarlyStopping, RuntimeLimits


Expand Down Expand Up @@ -259,11 +260,9 @@ def save_model(
saved, e.g. optimizer state dict

"""
# Strip the DDP wrapper so the checkpoint can be loaded on any number of GPUs.
if isinstance(self.network, DDP):
model_state_dict = self.network.module.state_dict()
else:
model_state_dict = self.network.state_dict()
# Strip the DDP and torch.compile wrappers so the checkpoint can be loaded
# on any number of GPUs, with or without compilation.
model_state_dict = unwrap_network(self.network).state_dict()

model_dict = {
"model_kwargs": self.model_kwargs,
Expand Down Expand Up @@ -648,7 +647,7 @@ def train_epoch(
if scaler is None:
scaler = _build_grad_scaler(pm.device)

is_ddp = isinstance(pm.network, DDP)
ddp_module = get_ddp_module(pm.network)

for batch_idx, data in enumerate(dataloader):
loss_info.update_timer("Dataloader")
Expand All @@ -663,7 +662,9 @@ def train_epoch(
# Under DDP, gradients only need to be all-reduced on the final backward
# pass of an accumulation window; skip the synchronization otherwise.
sync_ctx = (
pm.network.no_sync() if is_ddp and not is_step_batch else nullcontext()
ddp_module.no_sync()
if ddp_module is not None and not is_step_batch
else nullcontext()
)

# Gradients are summed over the accumulated mini-batches, so divide each
Expand Down Expand Up @@ -725,7 +726,9 @@ def test_epoch(
float
Average loss over the test set.
"""
with torch.no_grad():
# eager_mode: evaluating a compiled network in eval mode would trigger another
# full compilation, which a short test epoch never amortizes.
with torch.no_grad(), eager_mode():
pm.network.eval()

if pm.rank is None:
Expand Down
58 changes: 58 additions & 0 deletions dingo/core/utils/torchutils.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,26 @@ def fix_random_seeds(_):
pass


def set_float32_matmul_precision(local_settings: dict) -> None:
"""Apply ``local.float32_matmul_precision`` (``highest`` | ``high`` | ``medium``).

``high`` lets float32 matrix multiplications use TensorFloat-32 on Ampere and
newer GPUs (10-bit mantissa inputs, fp32 accumulation), roughly doubling the
speed of the network's linear layers; ``medium`` additionally allows bfloat16
inputs. PyTorch's default, ``highest``, keeps full fp32 and is left unchanged
when the setting is absent.
"""
precision = local_settings.get("float32_matmul_precision")
if precision is None:
return
if precision not in ("highest", "high", "medium"):
raise ValueError(
f"float32_matmul_precision must be 'highest', 'high' or 'medium', "
f"got {precision!r}."
)
torch.set_float32_matmul_precision(precision)


def get_cuda_info() -> dict[str, Any]:
"""Get information about the CUDA devices available in the system."""
if not torch.cuda.is_available():
Expand Down Expand Up @@ -129,6 +149,34 @@ def replace_BatchNorm_with_SyncBatchNorm(network: nn.Module) -> nn.Module:
return nn.SyncBatchNorm.convert_sync_batchnorm(network)


def get_ddp_module(network: nn.Module) -> Optional[DDP]:
"""Return the DDP wrapper inside *network* (looking through a ``torch.compile``
wrapper), or ``None`` if the network is not DDP-wrapped."""
while True:
if isinstance(network, DDP):
return network
if hasattr(network, "_orig_mod"): # torch.compile OptimizedModule
network = network._orig_mod
else:
return None


def unwrap_network(network: nn.Module) -> nn.Module:
"""Strip ``torch.compile`` and DDP wrappers, returning the bare network.

Used to save checkpoints whose state-dict keys carry no wrapper prefixes, so
they load on any number of GPUs with or without compilation."""
# we need a while loop here because the network can be wrapped twice
# once by the DDP and once by torch.compile
while True:
if isinstance(network, DDP):
network = network.module
elif hasattr(network, "_orig_mod"): # torch.compile OptimizedModule
network = network._orig_mod
else:
return network


def print_number_of_model_parameters(network: nn.Module) -> None:
"""
Print the number of fixed and learnable parameters of *network*.
Expand Down Expand Up @@ -329,6 +377,7 @@ def build_train_and_test_loaders(
num_workers: int,
world_size: Optional[int] = None,
rank: Optional[int] = None,
drop_last: bool = False,
) -> Tuple[DataLoader, DataLoader, Optional[DistributedSampler]]:
"""
Split the dataset into train and test sets, and build corresponding DataLoaders.
Expand All @@ -350,6 +399,11 @@ def build_train_and_test_loaders(
Total number of DDP processes (GPUs).
rank : int, optional
Rank of the current DDP process.
drop_last : bool
Drop the last, smaller batch of each training epoch, so that a compiled
network (specialized to the batch shape) is not recompiled for it. The test
loader always keeps its last batch: the test epoch runs eagerly, and dropping
it could leave a small per-rank test split with no batches at all.

Returns
-------
Expand Down Expand Up @@ -380,6 +434,7 @@ def build_train_and_test_loaders(
num_workers=num_workers,
worker_init_fn=fix_random_seeds,
persistent_workers=persistent_workers,
drop_last=drop_last,
)
test_loader = DataLoader(
test_dataset,
Expand All @@ -389,6 +444,7 @@ def build_train_and_test_loaders(
num_workers=num_workers,
worker_init_fn=fix_random_seeds,
persistent_workers=persistent_workers,
drop_last=False,
)
else:
train_sampler = None
Expand All @@ -400,6 +456,7 @@ def build_train_and_test_loaders(
num_workers=num_workers,
worker_init_fn=fix_random_seeds,
persistent_workers=persistent_workers,
drop_last=drop_last,
)
test_loader = DataLoader(
test_dataset,
Expand All @@ -409,6 +466,7 @@ def build_train_and_test_loaders(
num_workers=num_workers,
worker_init_fn=fix_random_seeds,
persistent_workers=persistent_workers,
drop_last=False,
)

return train_loader, test_loader, train_sampler
Expand Down
Loading
Loading