"""
Trainer implementations and lifecycle orchestration for pimm.
The default trainer builds runtime components, runs hooks around train, epoch,
and step boundaries, moves batches to the selected parallel device, and records
checkpointable resume state after each optimization step.
Author: Xiaoyang Wu (xiaoyang.wu.cs@gmail.com)
Please cite our work if the code is helpful to you.
"""
import contextlib
import os
import sys
import weakref
from functools import partial
from typing import Any
import torch
import torch.nn as nn
import torch.utils.data
from packaging import version
if sys.version_info >= (3, 10):
from collections.abc import Iterator
else:
from collections import Iterator
from tensorboardX import SummaryWriter
from torchdata.stateful_dataloader import StatefulDataLoader
import pimm.utils.comm as comm
from pimm.datasets import (
build_dataset,
collate_fn,
inseg_collate_fn,
point_collate_fn,
StatefulRandomSampler,
set_dataloader_epoch,
)
from pimm.distributed import (
create_parallel_context,
move_batch_to_device,
prepare_model,
unwrap_model,
)
from pimm.models import build_model
from pimm.observability import structured_logger as sl
from pimm.observability.structured_logger.structured_logging import (
_structured_logger_disabled,
)
from pimm.utils.events import EventStorage, ExceptionWriter, WandbSummaryWriter
from pimm.utils.logger import get_root_logger
from pimm.utils.optimizer import build_optimizer
from pimm.utils.registry import Registry
from pimm.utils.scheduler import build_scheduler
from ._train_utils import worker_init_fn
from .hooks import HookBase, build_hooks
from ._train_utils import TrainState
TRAINERS = Registry("trainers")
AMP_DTYPE = dict(
bfloat16=torch.bfloat16,
)
_END_OF_DATALOADER = object()
class TrainerBase:
"""Abstract hook-driven training lifecycle.
Defines the generic train loop -- ``before_train`` -> per-epoch
(``before_epoch`` -> per-step ``before_step``/``run_step``/``after_step``) ->
``after_epoch`` -> ``after_train`` -- where each phase fans out to the
registered :class:`HookBase` instances (checkpointing, logging, evaluation,
schedulers). Holds the shared mutable state hooks read and write (``model``,
loaders, ``optimizer``, ``scheduler``, ``scaler``, ``epoch``/``global_step``
counters, ``comm_info``, ``storage``, ``writer``). Subclasses must implement
:meth:`run_step`; this base class is not registered and is not selected
directly via config.
"""
def __init__(self) -> None:
"""Initialize shared lifecycle counters and hook-visible state."""
self.hooks = []
self.cfg: Any = None
self.logger: Any = None
self.model: nn.Module | None = None
self.train_loader: Any = None
self.val_loader: Any = None
self.test_loader: Any = None
self.optimizer: Any = None
self.scheduler: Any = None
self.scaler: Any = None
self.epoch = 0
self.start_epoch = 0
self.start_iter = 0 # First dataloader position to consume on resume.
self.max_epoch = 0
self.max_iter = 0
self.global_step = 0
self.samples_seen = 0
self.best_metric_value = -torch.inf
self.train_state = TrainState()
self.comm_info = dict()
self.data_iterator: Iterator = enumerate([])
self.storage: EventStorage | None = None
self.writer: SummaryWriter | None = None
self._trace_hooks = False
self._trace_batch_stats_every = 0
def register_hooks(self, hooks) -> None:
"""Build hooks and attach this trainer through weak references."""
hooks = build_hooks(hooks)
for h in hooks:
assert isinstance(h, HookBase)
# To avoid circular reference, hooks and trainer cannot own each other.
# This normally does not matter, but will cause memory leak if the
# involved objects contain __del__:
# See http://engineering.hearsaysocial.com/2013/06/16/circular-references-in-python/
h.trainer = weakref.proxy(self)
self.hooks.extend(hooks)
def _call_hooks(self, method_name, *args) -> None:
"""Call one lifecycle method on every hook, optionally tracing each one."""
for hook in self.hooks:
callback = getattr(hook, method_name)
if self._trace_hooks:
hook_name = type(hook).__name__
with sl.log_trace_span(f"hook.{hook_name}.{method_name}"):
callback(*args)
else:
callback(*args)
def train(self):
"""Run the generic train/epoch/step lifecycle."""
with EventStorage() as self.storage:
# Hooks bracket the whole run, each epoch, and each step.
self.before_train()
if self._training_already_complete():
self._finish_completed_resume()
return
for self.epoch in range(self.start_epoch, self.max_epoch):
self.before_epoch()
for (
self.comm_info["iter"],
self.comm_info["input_dict"],
) in self.data_iterator:
self.before_step()
self.run_step()
self.after_step()
self.after_epoch()
self.after_train()
def before_train(self):
"""Apply global numeric settings and call before-train hooks."""
if self.cfg.matmul_precision is not None:
torch.set_float32_matmul_precision(self.cfg.matmul_precision)
with sl.log_trace_span("hooks.before_train"):
self._call_hooks("before_train")
def before_epoch(self):
"""Call hooks before the current epoch starts."""
with sl.log_trace_span("hooks.before_epoch"):
self._call_hooks("before_epoch")
def before_step(self):
"""Call hooks before consuming the current batch."""
self._flush_writer_step()
with sl.log_trace_span("hooks.before_step"):
self._call_hooks("before_step")
def run_step(self):
"""Run one optimization step for the current batch."""
raise NotImplementedError
def after_step(self):
"""Call hooks after the current optimization step."""
with sl.log_trace_span("hooks.after_step"):
self._call_hooks("after_step")
def after_epoch(self):
"""Call epoch-end hooks and reset per-epoch event histories."""
with sl.log_trace_span("hooks.after_epoch"):
self._call_hooks("after_epoch")
self._flush_writer_step()
self.storage.reset_histories()
def _flush_writer_step(self):
"""Commit a writer row after all metrics for its optimizer step exist."""
if not comm.is_main_process():
return
writer = getattr(self, "writer", None)
flush_step = getattr(writer, "flush_step", None)
if flush_step is not None:
flush_step()
def after_train(self):
"""Synchronize workers, call final hooks, and close the writer."""
# Sync GPU before running train hooks
with sl.log_trace_span("training_synchronize"):
comm.synchronize()
with sl.log_trace_span("hooks.after_train"):
self._call_hooks("after_train")
self._close_writer()
def _training_already_complete(self):
"""Return whether restored progress is already at the run horizon."""
return int(getattr(self, "start_epoch", 0)) >= int(
getattr(self, "max_epoch", 0)
)
def _finish_completed_resume(self):
"""Exit a resumed, already-complete run without final checkpoint churn."""
if hasattr(self, "logger"):
self.logger.info(
"Training already complete: "
f"start_epoch={self.start_epoch}, max_epoch={self.max_epoch}. "
"Exiting without running more steps."
)
sl.log_trace_instant("training_already_complete")
with sl.log_trace_span("training_synchronize"):
comm.synchronize()
self._close_writer()
def _close_writer(self):
"""Close the event writer on the main process when one exists."""
if comm.is_main_process():
writer = getattr(self, "writer", None)
if writer is not None:
writer.close()
[docs]
@TRAINERS.register_module("DefaultTrainer")
class Trainer(TrainerBase):
"""Default single-dataset supervised/SSL trainer.
Builds everything from one ``cfg`` (model, train/val/test loaders, optimizer,
scheduler, AMP scaler, hooks, writer, and the parallel context) and runs the
standard AMP forward/backward optimization step in :meth:`run_step`. The
train loop is resume-aware: it restores ``start_epoch``/``start_iter`` and the
global step from checkpoint hooks, re-aligns metric writers and the dataloader
cursor, and treats ``cfg.epoch`` as the absolute horizon so a run can be
extended without changing schedules. Registered as ``DefaultTrainer`` --
select via ``train = dict(type="DefaultTrainer")`` (the default).
Args:
cfg: Fully-resolved run config providing ``save_path``, ``epoch``,
``resume``, the ``model``/``data``/``optimizer``/``scheduler``/``hooks``
sub-configs, and the distributed/AMP settings.
"""
# Set when the train dataset is an IterableDataset (streaming). Affects loader
# construction (no sampler) and epoch-length accounting (see _iters_per_epoch).
_train_is_iterable = False
_iter_per_epoch_value = None
def __init__(self, cfg):
"""Build model, data loaders, optimizer, scheduler, hooks, and writer."""
super().__init__()
self.cfg = cfg
# deal with structured logging
structured_cfg = cfg.get("structured_logging", {})
tracing_enabled = not _structured_logger_disabled()
self._trace_hooks = tracing_enabled and bool(
structured_cfg.get("trace_hooks", False)
)
self._trace_batch_stats_every = (
max(0, int(structured_cfg.get("batch_stats_every", 1) or 0))
if tracing_enabled
else 0
)
self.epoch = 0
self.start_epoch = 0
self.start_iter = 0
self.global_step = 0
self.samples_seen = 0
self.train_state = TrainState()
with sl.log_trace_span("build.parallel_context"):
self.parallel_context = create_parallel_context(cfg)
# When resuming, use cfg.epoch as the absolute horizon so we can
# extend training beyond the previous end without changing schedules.
self.max_epoch = cfg.epoch
self.best_metric_value = -torch.inf
self.logger = get_root_logger(
log_file=os.path.join(cfg.save_path, "train.log"),
file_mode="a" if cfg.resume else "w",
)
self.logger.info("=> Loading config ...")
self.logger.info(f"Save path: {cfg.save_path}")
self.logger.info(f"Config:\n{cfg.pretty_text}")
self.logger.info("=> Building model ...")
with sl.log_trace_span("build.model"):
self.model = self.build_model()
self.logger.info("=> Building train dataset & dataloader ...")
with sl.log_trace_span("build.train_loader"):
self.train_loader = self.build_train_loader()
self.logger.info("=> Building val dataset & dataloader ...")
with sl.log_trace_span("build.evaluation_loaders"):
self.val_loader = self.build_val_loader()
self.test_loader = self.build_test_loader()
self.logger.info("=> Building optimize, scheduler, scaler(amp) ...")
with sl.log_trace_span("build.optimization"):
self.optimizer = self.build_optimizer()
self.scheduler = self.build_scheduler()
self.scaler = self.build_scaler()
self.logger.info("=> Building hooks ...")
with sl.log_trace_span("build.hooks"):
self.register_hooks(self.cfg.hooks)
self.logger.info("=> Running config modifiers ...")
with sl.log_trace_span("hooks.modify_config"):
self._call_hooks("modify_config", self.cfg)
self.logger.info("=> Building writer ...")
with sl.log_trace_span("build.writer"):
self.writer = self.build_writer()
[docs]
def train(self):
"""Run training from the configured or restored epoch/iteration."""
anomaly_context = torch.autograd.detect_anomaly() if self.cfg.detect_anomaly else contextlib.nullcontext()
with EventStorage() as self.storage, ExceptionWriter(), anomaly_context:
# Checkpoint hooks can restore start_epoch, start_iter, and counters.
self.before_train()
# Keep metric writers aligned with the absolute optimization step.
iter_per_epoch = self._iters_per_epoch()
resumed_iter = self.global_step or (
self.start_epoch * iter_per_epoch + self.start_iter
)
if resumed_iter > 0:
self.storage.iter = resumed_iter
self._align_writer_step(resumed_iter)
self.logger.info(f"Resuming from iteration {resumed_iter}")
if self._training_already_complete():
if resumed_iter > 0 and iter_per_epoch > 0:
last_completed_index = resumed_iter - 1
sl.set_step(
resumed_iter,
relative_step=0,
epoch=last_completed_index // iter_per_epoch,
iteration=last_completed_index % iter_per_epoch,
)
self._finish_completed_resume()
return
self.logger.info(">>>>>>>>>>>>>>>> Start Training >>>>>>>>>>>>>>>>")
sl.log_trace_instant("training_start")
process_step = 0
for self.epoch in range(self.start_epoch, self.max_epoch):
resume_mid_epoch = (
self.epoch == self.start_epoch and self.start_iter > 0
)
set_dataloader_epoch(
self.train_loader,
self.epoch,
reset_position=not resume_mid_epoch,
)
self.comm_info["epoch"] = self.epoch
self.comm_info["iter_per_epoch"] = iter_per_epoch
self.model.train()
start_iter = self.start_iter if resume_mid_epoch else 0
if start_iter > 0:
self.logger.info(
f"Resuming epoch {self.epoch} from dataloader position {start_iter}"
)
self.data_iterator = iter(self.train_loader)
# Epoch-boundary hooks belong to the first intended step of
# this epoch, not the final step context left by the previous
# epoch.
sl.set_step(
self.epoch * iter_per_epoch + start_iter + 1,
relative_step=process_step + 1,
epoch=self.epoch,
iteration=start_iter,
)
self.before_epoch()
for iteration in range(start_iter, iter_per_epoch):
absolute_step = self.epoch * iter_per_epoch + iteration + 1
self.comm_info["iter"] = iteration
if iteration != start_iter:
sl.set_step(
absolute_step,
relative_step=process_step + 1,
epoch=self.epoch,
iteration=iteration,
)
with sl.log_trace_span("step"):
with sl.log_trace_span("data_fetch"):
input_dict = next(self.data_iterator, _END_OF_DATALOADER)
if input_dict is _END_OF_DATALOADER:
sl.log_trace_instant("dataloader_exhausted")
break
process_step += 1
self.comm_info["input_dict"] = input_dict
self._log_trace_batch_stats(absolute_step)
self.before_step()
self.run_step()
# Capture the next resume point before checkpoint hooks run.
self._record_step_state()
self.after_step()
# Iterable streams have no natural length. The bounded range
# above keeps every rank at the configured epoch length.
if self._train_is_iterable and iteration + 1 >= iter_per_epoch:
break
self.start_iter = 0
self.after_epoch()
self.after_train()
sl.log_trace_instant("training_end")
def _log_trace_batch_stats(self, absolute_step):
"""Record cheap rank-local batch shape scalars without synchronizing CUDA."""
frequency = self._trace_batch_stats_every
if frequency <= 0 or absolute_step % frequency != 0:
return
input_dict = self.comm_info.get("input_dict")
if not isinstance(input_dict, dict):
return
scalars = {}
offset = input_dict.get("offset")
if offset is not None:
try:
scalars["batch.local_samples"] = len(offset)
except TypeError:
pass
coord = input_dict.get("coord")
if hasattr(coord, "shape") and len(coord.shape) > 0:
scalars["batch.local_points"] = int(coord.shape[0])
if scalars:
sl.log_trace_scalar(scalars)
def _align_writer_step(self, global_step):
"""Align writer-internal step counters after checkpoint resume."""
if self.writer is None:
return
if hasattr(self.writer, "step"):
self.writer.step = max(int(getattr(self.writer, "step", 0)), int(global_step))
self.cfg.log_step_offset = int(getattr(self.cfg, "log_step_offset", 0) or 0)
def _record_step_state(self):
"""Update checkpointable counters for the next batch to consume."""
iter_per_epoch = int(self.comm_info.get("iter_per_epoch", self._iters_per_epoch()))
iter_in_epoch = int(self.comm_info.get("iter", 0)) + 1
self.global_step = self.epoch * iter_per_epoch + iter_in_epoch
input_dict = self.comm_info.get("input_dict", {})
local_batch = self.cfg.batch_size_per_gpu
# Point datasets use offset length as the actual per-rank sample count.
if isinstance(input_dict, dict) and "offset" in input_dict:
try:
local_batch = len(input_dict["offset"])
except TypeError:
local_batch = self.cfg.batch_size_per_gpu
self.samples_seen += int(local_batch) * comm.get_world_size()
self.train_state = TrainState.from_trainer(self)
[docs]
def run_step(self):
"""Move one batch to device, run forward/backward, and update LR."""
if version.parse(torch.__version__) >= version.parse("2.4"):
auto_cast = partial(
torch.amp.autocast,
device_type=self.parallel_context.device.type,
)
else:
# deprecated warning
auto_cast = torch.cuda.amp.autocast
with sl.log_trace_span("batch_to_device"):
input_dict = move_batch_to_device(
self.comm_info["input_dict"],
self.parallel_context.device,
)
# Store the device-resident batch so hooks and checkpoint state agree.
self.comm_info["input_dict"] = input_dict
with sl.log_trace_span("forward"):
with auto_cast(
enabled=self.cfg.enable_amp,
dtype=AMP_DTYPE[self.cfg.amp_dtype],
):
output_dict = self.model(input_dict)
loss = output_dict["loss"]
# Log average points per sample for throughput analysis
if "offset" in input_dict:
output_dict["avg_pts"] = input_dict["coord"].shape[0] / len(
input_dict["offset"]
)
self.optimizer.zero_grad()
if self.cfg.enable_amp:
with sl.log_trace_span("backward"):
self.scaler.scale(loss).backward()
self.scaler.unscale_(self.optimizer)
if self.cfg.clip_grad is not None:
torch.nn.utils.clip_grad_norm_(
self.model.parameters(), self.cfg.clip_grad
)
with sl.log_trace_span("optimizer"):
self.scaler.step(self.optimizer)
# When enable amp, optimizer.step call are skipped if the loss scaling factor is too large.
# Fix torch warning scheduler step before optimizer step.
scaler = self.scaler.get_scale()
self.scaler.update()
if scaler <= self.scaler.get_scale():
self.scheduler.step()
else:
with sl.log_trace_span("backward"):
loss.backward()
if self.cfg.clip_grad is not None:
torch.nn.utils.clip_grad_norm_(
self.model.parameters(), self.cfg.clip_grad
)
with sl.log_trace_span("optimizer"):
self.optimizer.step()
self.scheduler.step()
if self.cfg.empty_cache and self.parallel_context.device.type == "cuda":
with sl.log_trace_span("cuda_empty_cache"):
torch.cuda.empty_cache()
self.comm_info["model_output_dict"] = output_dict
[docs]
def after_epoch(self):
"""Run epoch-end hooks, clear histories, and optionally empty CUDA cache."""
with sl.log_trace_span("hooks.after_epoch"):
self._call_hooks("after_epoch")
self._flush_writer_step()
self.storage.reset_histories()
if self.cfg.empty_cache_per_epoch:
with sl.log_trace_span("cuda_empty_cache"):
torch.cuda.empty_cache()
[docs]
def build_model(self):
"""Construct the model and wrap it for the configured parallel strategy."""
model = build_model(self.cfg.model)
n_parameters = sum(p.numel() for p in model.parameters() if p.requires_grad)
# logger.info(f"Model: \n{self.model}")
self.logger.info(f"Num params: {n_parameters:,}")
model = prepare_model(model, self.cfg, self.parallel_context)
self.logger.info(
f"Parallel strategy: {self.parallel_context.strategy}, device: {self.parallel_context.device}"
)
return model
[docs]
def build_writer(self):
"""Create a main-rank summary writer for TensorBoard or W&B."""
if self.cfg.get("use_wandb", False):
wandb_kwargs = dict(
project=self.cfg.get("wandb_project", "pimm"),
name=self.cfg.get(
"wandb_run_name", os.path.basename(self.cfg.save_path)
),
config=self.cfg,
step_offset=self.cfg.get("log_step_offset", 0),
)
for cfg_key, wandb_key in (
("wandb_group", "group"),
("wandb_job_type", "job_type"),
("wandb_run_id", "id"),
("wandb_resume", "resume"),
):
value = self.cfg.get(cfg_key, None)
if value is not None:
wandb_kwargs[wandb_key] = value
# Surface the original `pimm submit/launch` command in the run Notes
# (top of the Overview tab) so it is visible without digging into the
# Config. wandb's own "Command" panel still shows the long auto-
# captured train.py argv, which we cannot override.
launch_command = self.cfg.get("launch_command", None)
if launch_command:
wandb_kwargs["notes"] = launch_command
writer = (
WandbSummaryWriter(**wandb_kwargs) if comm.is_main_process() else None
)
self.logger.info(
f"Weights & Biases writer initialized with project: {self.cfg.get('wandb_project', 'pimm')}"
)
else:
writer = (
SummaryWriter(self.cfg.save_path) if comm.is_main_process() else None
)
self.logger.info(f"Tensorboard writer logging dir: {self.cfg.save_path}")
return writer
[docs]
def build_train_loader(self):
"""Build the stateful training loader used for mid-epoch resume."""
train_data = build_dataset(self.cfg.data.train)
return self._build_stateful_train_loader(
train_data, partial(collate_fn, mix_prob=self.cfg.mix_prob)
)
def _build_stateful_train_loader(self, train_data, collate_fn, **loader_kwargs):
"""Build a ``StatefulDataLoader`` for a map-style OR an iterable train
dataset. Extra ``loader_kwargs`` (e.g. ``in_order``) pass through to the
``StatefulDataLoader``.
Map-style datasets get a rank-aware ``StatefulRandomSampler`` (shuffle +
checkpointable position). ``IterableDataset``s get no sampler -- they own
shuffling and DDP/worker sharding internally (see the dataset's
``__iter__``); ``drop_last`` plus the ``iter_per_epoch`` cap in
:meth:`train` keeps every rank's step count uniform for DDP lockstep.
Works for any ``IterableDataset``, not just the parquet one.
"""
init_fn = (
partial(
worker_init_fn,
num_workers=self.cfg.num_worker_per_gpu,
rank=comm.get_rank(),
seed=self.cfg.seed,
)
if self.cfg.seed is not None
else None
)
common = dict(
batch_size=self.cfg.batch_size_per_gpu,
num_workers=self.cfg.num_worker_per_gpu,
collate_fn=collate_fn,
pin_memory=True,
worker_init_fn=init_fn,
persistent_workers=(self.cfg.num_worker_per_gpu > 0),
snapshot_every_n_steps=self.cfg.get("dataloader_snapshot_every_n_steps", 1),
)
common.update(loader_kwargs)
if isinstance(train_data, torch.utils.data.IterableDataset):
self._train_is_iterable = True
return StatefulDataLoader(train_data, drop_last=True, **common)
self._train_is_iterable = False
drop_last = len(train_data) > self.cfg.batch_size
sampler = StatefulRandomSampler(
train_data,
shuffle=True,
seed=self.cfg.seed if self.cfg.seed is not None else 0,
num_replicas=comm.get_world_size(),
rank=comm.get_rank(),
drop_last=drop_last,
)
return StatefulDataLoader(
train_data, sampler=sampler, drop_last=drop_last, **common
)
def _iters_per_epoch(self):
"""Optimizer steps per epoch on THIS rank (cached).
Map-style: ``len(train_loader)``. Iterable: ``cfg.iters_per_epoch`` when
set, else derived from the dataset's ``num_samples()`` floored per rank
(``num_samples // world_size // batch_size_per_gpu``) so every rank runs
the same count -- required for DDP lockstep. Raises if an iterable train
dataset provides neither.
"""
if self._iter_per_epoch_value is not None:
return self._iter_per_epoch_value
if not self._train_is_iterable:
val = len(self.train_loader)
else:
cfg_ipe = self.cfg.get("iters_per_epoch", None)
if cfg_ipe:
val = int(cfg_ipe)
else:
dataset = getattr(self.train_loader, "dataset", None)
num_samples = getattr(dataset, "num_samples", None)
if not callable(num_samples):
raise ValueError(
"Iterable train dataset requires `iters_per_epoch` in the "
"config, or a `num_samples()` method on the dataset."
)
world = max(1, comm.get_world_size())
per_rank = int(num_samples()) // world
val = per_rank // self.cfg.batch_size_per_gpu
if val <= 0:
raise ValueError(
f"Derived iters_per_epoch={val} <= 0 (num_samples per "
f"rank={per_rank}, batch_per_gpu={self.cfg.batch_size_per_gpu})."
)
self.logger.info(f"Iterable train loader: iters_per_epoch={val}")
self._iter_per_epoch_value = val
return val
[docs]
def build_val_loader(self):
"""Build the optional validation loader."""
val_loader = None
if self.cfg.evaluate:
val_data = build_dataset(self.cfg.data.val)
if comm.get_world_size() > 1:
val_sampler = torch.utils.data.distributed.DistributedSampler(val_data)
else:
val_sampler = None
val_loader = torch.utils.data.DataLoader(
val_data,
batch_size=self.cfg.batch_size_val_per_gpu,
shuffle=False,
num_workers=0,
pin_memory=True,
sampler=val_sampler,
collate_fn=collate_fn,
# in_order=self.cfg.deterministic,
)
return val_loader
[docs]
def build_test_loader(self):
"""Build the optional test loader used by evaluation hooks."""
test_loader = None
if self.cfg.evaluate and hasattr(self.cfg.data, "test"):
test_data = build_dataset(self.cfg.data.test)
if comm.get_world_size() > 1:
test_sampler = torch.utils.data.distributed.DistributedSampler(
test_data
)
else:
test_sampler = None
test_loader = torch.utils.data.DataLoader(
test_data,
batch_size=self.cfg.batch_size_val_per_gpu,
shuffle=False,
num_workers=0,
pin_memory=True,
sampler=test_sampler,
collate_fn=collate_fn,
# in_order=self.cfg.deterministic,
)
return test_loader
[docs]
def build_optimizer(self):
"""Build the optimizer from config."""
return build_optimizer(self.cfg.optimizer, self.model, self.cfg.param_dicts)
[docs]
def build_scheduler(self):
"""Build a scheduler sized to all optimizer steps in training."""
assert hasattr(self, "optimizer")
assert hasattr(self, "train_loader")
self.cfg.scheduler.total_steps = self._iters_per_epoch() * self.cfg.epoch
return build_scheduler(self.cfg.scheduler, self.optimizer)
[docs]
def build_scaler(self):
"""Build an AMP gradient scaler when mixed precision is enabled."""
if not self.cfg.enable_amp:
return None
# Use standard grad scaler for DDP
if version.parse(torch.__version__) >= version.parse("2.4"):
grad_scaler = partial(torch.amp.GradScaler, device="cuda")
else:
# deprecated warning
grad_scaler = torch.cuda.amp.GradScaler
scaler = grad_scaler()
return scaler
[docs]
@TRAINERS.register_module("GRPOTrainer")
class GRPOTrainer(Trainer):
"""Reinforcement-learning trainer implementing GRPO at the trainer level.
Replaces the single supervised optimization step with a rollout-based loop:
each batch is sampled into a group of trajectories, then the policy is updated
``policy_updates_per_rollout`` times over the cached rollout (optionally split
into trajectory microbatches to bound memory). The scheduler is sized to
``len(train_loader) * epoch * policy_updates_per_rollout`` accordingly, and
scalar GRPO metrics are reduced across ranks with key-specific ops
(min/max/mean). Inherits all model/loader/optimizer construction from
:class:`Trainer`. Registered as ``GRPOTrainer`` -- select via
``train = dict(type="GRPOTrainer")``.
"""
def _sync_grpo_scalar_metrics(self, output_dict):
"""Reduce scalar GRPO metrics across ranks with key-specific ops."""
if comm.get_world_size() < 2:
return output_dict
synced = {}
world_size = comm.get_world_size()
for key, value in output_dict.items():
if not torch.is_tensor(value) or value.numel() != 1:
synced[key] = value
continue
metric = value.detach().float()
if metric.device.type == "cpu":
metric = metric.cuda()
if "_min" in key:
reduce_op = torch.distributed.ReduceOp.MIN
elif "_max" in key or "abs_max" in key:
reduce_op = torch.distributed.ReduceOp.MAX
else:
reduce_op = torch.distributed.ReduceOp.SUM
torch.distributed.all_reduce(metric, op=reduce_op)
if reduce_op == torch.distributed.ReduceOp.SUM:
metric /= world_size
synced[key] = metric
return synced
def _policy_updates_per_rollout(self):
"""Return the number of policy updates to run per sampled rollout."""
train_cfg = getattr(self.cfg, "train", {})
if hasattr(train_cfg, "get"):
return max(1, int(train_cfg.get("policy_updates_per_rollout", 1)))
return max(1, int(getattr(train_cfg, "policy_updates_per_rollout", 1)))
def _trajectory_microbatch_size(self):
"""Return trajectory microbatch size, or zero to disable splitting."""
train_cfg = getattr(self.cfg, "train", {})
if hasattr(train_cfg, "get"):
value = train_cfg.get("trajectory_microbatch_size", 0)
else:
value = getattr(train_cfg, "trajectory_microbatch_size", 0)
return max(0, int(value or 0))
[docs]
def build_scheduler(self):
"""Build a scheduler sized to rollout count times policy updates."""
assert hasattr(self, "optimizer")
assert hasattr(self, "train_loader")
self.cfg.scheduler.total_steps = (
len(self.train_loader) * self.cfg.epoch * self._policy_updates_per_rollout()
)
return build_scheduler(self.cfg.scheduler, self.optimizer)
def _optimizer_update(self, loss):
"""Apply one optimizer/scheduler update for a GRPO loss."""
self.optimizer.zero_grad()
if self.cfg.enable_amp:
self.scaler.scale(loss).backward()
self.scaler.unscale_(self.optimizer)
if self.cfg.clip_grad is not None:
torch.nn.utils.clip_grad_norm_(
self.model.parameters(), self.cfg.clip_grad
)
self.scaler.step(self.optimizer)
scaler = self.scaler.get_scale()
self.scaler.update()
if scaler <= self.scaler.get_scale():
self.scheduler.step()
else:
loss.backward()
if self.cfg.clip_grad is not None:
torch.nn.utils.clip_grad_norm_(
self.model.parameters(), self.cfg.clip_grad
)
self.optimizer.step()
self.scheduler.step()
def _rollout_metric_sums(self, trajectories):
"""Aggregate scalar rollout metrics over trajectories."""
metric_sums = {}
for traj in trajectories:
for key, value in traj.metrics.items():
metric_sums[key] = metric_sums.get(key, 0.0) + float(value)
return metric_sums
def _slice_grpo_event(self, event, start, end):
"""Return an event view containing a trajectory slice."""
trajectories = event["trajectories"][start:end]
sliced = dict(event)
sliced["trajectories"] = trajectories
if "advantages" in sliced:
sliced["advantages"] = sliced["advantages"][start:end]
if "step_advantages" in sliced and sliced["step_advantages"] is not None:
sliced["step_advantages"] = sliced["step_advantages"][start:end]
return sliced
def _iter_grpo_trajectory_microbatches(self, rollout_batch, microbatch_size):
"""Yield rollout-batch chunks split by trajectory count."""
if microbatch_size <= 0:
yield rollout_batch
return
for event in rollout_batch["events"]:
trajectories = event["trajectories"]
for start in range(0, len(trajectories), microbatch_size):
end = min(start + microbatch_size, len(trajectories))
sliced_event = self._slice_grpo_event(event, start, end)
sliced_trajectories = sliced_event["trajectories"]
chunk = dict(rollout_batch)
chunk["events"] = [sliced_event]
chunk["metric_count"] = len(sliced_trajectories)
chunk["metric_sums"] = self._rollout_metric_sums(sliced_trajectories)
if "reward_stds" in chunk:
chunk["reward_stds"] = [
sliced_event.get("reward_std", event.get("reward_std"))
]
if "rloo_score_means" in chunk and "rloo_score_mean" in sliced_event:
chunk["rloo_score_means"] = [sliced_event["rloo_score_mean"]]
if "rloo_score_stds" in chunk and "rloo_score_std" in sliced_event:
chunk["rloo_score_stds"] = [sliced_event["rloo_score_std"]]
yield chunk
def _combine_weighted_grpo_outputs(self, weighted_outputs):
"""Combine microbatch outputs using their metric-count weights."""
if not weighted_outputs:
raise RuntimeError("GRPOTrainer received no microbatch outputs")
combined = {}
first_output = weighted_outputs[0][0]
tensor_keys = [
key for key, value in first_output.items() if torch.is_tensor(value)
]
for key in tensor_keys:
values = [
(output[key], weight)
for output, weight in weighted_outputs
if key in output
]
if not values:
continue
if "_min" in key:
combined[key] = torch.stack(
[value.detach() for value, _ in values]
).min()
elif "_max" in key or "abs_max" in key:
combined[key] = torch.stack(
[value.detach() for value, _ in values]
).max()
else:
combined[key] = torch.stack(
[value.detach() * float(weight) for value, weight in values]
).sum()
for key, value in first_output.items():
if key not in combined and not torch.is_tensor(value):
combined[key] = value
return combined
def _optimizer_update_grpo_microbatched(
self,
model_impl,
rollout_batch,
*,
update_index,
policy_updates,
microbatch_size,
auto_cast,
):
"""Backpropagate GRPO loss over trajectory microbatches."""
chunks = list(
self._iter_grpo_trajectory_microbatches(rollout_batch, microbatch_size)
)
total_count = sum(max(0, int(chunk.get("metric_count", 0))) for chunk in chunks)
total_count = max(total_count, 1)
weighted_outputs = []
self.optimizer.zero_grad()
for chunk in chunks:
weight = max(0, int(chunk.get("metric_count", 0))) / total_count
with auto_cast(
enabled=self.cfg.enable_amp, dtype=AMP_DTYPE[self.cfg.amp_dtype]
):
output_dict = model_impl.grpo_loss_from_batch(
chunk,
update_index=update_index,
policy_updates_per_rollout=policy_updates,
)
loss = output_dict["loss"] * float(weight)
detached = {
key: value.detach() if torch.is_tensor(value) else value
for key, value in output_dict.items()
}
weighted_outputs.append((detached, weight))
if self.cfg.enable_amp:
self.scaler.scale(loss).backward()
else:
loss.backward()
if self.cfg.enable_amp:
self.scaler.unscale_(self.optimizer)
if self.cfg.clip_grad is not None:
torch.nn.utils.clip_grad_norm_(
self.model.parameters(), self.cfg.clip_grad
)
self.scaler.step(self.optimizer)
scaler = self.scaler.get_scale()
self.scaler.update()
if scaler <= self.scaler.get_scale():
self.scheduler.step()
else:
if self.cfg.clip_grad is not None:
torch.nn.utils.clip_grad_norm_(
self.model.parameters(), self.cfg.clip_grad
)
self.optimizer.step()
self.scheduler.step()
combined = self._combine_weighted_grpo_outputs(weighted_outputs)
if hasattr(model_impl, "_rollout_metric_tensors"):
combined.update(model_impl._rollout_metric_tensors(rollout_batch))
return combined
def _combine_grpo_outputs(self, outputs):
"""Summarize first, last, and mean metrics across policy updates."""
if not outputs:
raise RuntimeError("GRPOTrainer received no update outputs")
first = outputs[0]
last = outputs[-1]
combined = {}
for key, value in last.items():
combined[key] = value.detach() if torch.is_tensor(value) else value
if "loss" in first:
combined["loss"] = torch.stack(
[output["loss"].detach() for output in outputs]
).mean()
tracked = [
"rl_pg_loss",
"rl_kl_loss",
"rl_kl",
"rl_ratio_geom",
"rl_log_ratio",
"rl_log_ratio_min",
"rl_log_ratio_max",
"rl_stop_log_ratio",
"rl_stop_log_ratio_min",
"rl_stop_log_ratio_max",
"rl_class_log_ratio",
"rl_class_log_ratio_min",
"rl_class_log_ratio_max",
"rl_kernel_log_ratio",
"rl_kernel_log_ratio_min",
"rl_kernel_log_ratio_max",
"rl_kernel_dim_log_ratio",
"rl_kernel_dim_log_ratio_min",
"rl_kernel_dim_log_ratio_max",
"rl_clip_frac",
"rl_logprob",
"rl_advantage_abs_max",
]
for key in tracked:
if key not in first or key not in last:
continue
first_value = (
first[key].detach() if torch.is_tensor(first[key]) else first[key]
)
last_value = last[key].detach() if torch.is_tensor(last[key]) else last[key]
combined[f"{key}_update0"] = first_value
combined[f"{key}_last"] = last_value
if torch.is_tensor(first[key]):
combined[f"{key}_mean_update"] = torch.stack(
[output[key].detach() for output in outputs]
).mean()
return combined
[docs]
def run_step(self):
"""Sample rollouts, run policy updates, and publish GRPO metrics."""
if version.parse(torch.__version__) >= version.parse("2.4"):
auto_cast = partial(
torch.amp.autocast,
device_type=self.parallel_context.device.type,
)
else:
auto_cast = torch.cuda.amp.autocast
with sl.log_trace_span("batch_to_device"):
input_dict = move_batch_to_device(
self.comm_info["input_dict"],
self.parallel_context.device,
)
# The model rollout and loss APIs expect all tensors on one device.
self.comm_info["input_dict"] = input_dict
model_impl = unwrap_model(self.model)
if not hasattr(model_impl, "sample_grpo_batch"):
raise RuntimeError("GRPOTrainer requires model.sample_grpo_batch")
if not hasattr(model_impl, "grpo_loss_from_batch"):
raise RuntimeError("GRPOTrainer requires model.grpo_loss_from_batch")
policy_updates = self._policy_updates_per_rollout()
microbatch_size = self._trajectory_microbatch_size()
with sl.log_trace_span("grpo_rollout"):
with auto_cast(
enabled=self.cfg.enable_amp,
dtype=AMP_DTYPE[self.cfg.amp_dtype],
):
rollout_batch = model_impl.sample_grpo_batch(input_dict)
update_outputs = []
for update_index in range(policy_updates):
with sl.log_trace_span("grpo_policy_update"):
if microbatch_size > 0:
output_dict = self._optimizer_update_grpo_microbatched(
model_impl,
rollout_batch,
update_index=update_index,
policy_updates=policy_updates,
microbatch_size=microbatch_size,
auto_cast=auto_cast,
)
else:
with auto_cast(
enabled=self.cfg.enable_amp,
dtype=AMP_DTYPE[self.cfg.amp_dtype],
device_type=self.parallel_context.device.type,
):
output_dict = model_impl.grpo_loss_from_batch(
rollout_batch,
update_index=update_index,
policy_updates_per_rollout=policy_updates,
)
loss = output_dict["loss"]
self._optimizer_update(loss)
update_outputs.append(output_dict)
if self.cfg.empty_cache and self.parallel_context.device.type == "cuda":
with sl.log_trace_span("cuda_empty_cache"):
torch.cuda.empty_cache()
combined = self._combine_grpo_outputs(update_outputs)
if "offset" in input_dict:
avg_pts = input_dict["coord"].shape[0] / len(input_dict["offset"])
combined["avg_pts"] = torch.as_tensor(
avg_pts, device=input_dict["coord"].device, dtype=torch.float32
)
with sl.log_trace_span("grpo_metric_sync"):
combined = self._sync_grpo_scalar_metrics(combined)
self.comm_info["model_output_dict"] = combined
[docs]
@TRAINERS.register_module("MultiDatasetTrainer")
class MultiDatasetTrainer(Trainer):
"""Trainer that draws mixed batches from several datasets.
Identical to :class:`Trainer` except :meth:`build_train_loader` swaps the
standard loader for ``MultiDatasetDataloader``, which samples across the
configured datasets (honoring per-dataset ratios and ``mix_prob``) and defines
the epoch length. Registered as ``MultiDatasetTrainer`` -- select via
``train = dict(type="MultiDatasetTrainer")``.
"""
[docs]
def build_train_loader(self):
"""Build a multi-dataset train loader and expose its epoch length."""
from pointcept.datasets import MultiDatasetDataloader
train_data = build_dataset(self.cfg.data.train)
train_loader = MultiDatasetDataloader(
train_data,
self.cfg.batch_size_per_gpu,
self.cfg.num_worker_per_gpu,
self.cfg.mix_prob,
self.cfg.seed,
)
self.comm_info["iter_per_epoch"] = len(train_loader)
return train_loader
[docs]
@TRAINERS.register_module("InsegTrainer")
class InsegTrainer(Trainer):
"""Trainer for instance segmentation with instance-aware collation.
Identical to :class:`Trainer` except the train and val loaders use
``inseg_collate_fn`` (which preserves variable per-sample instance/query
targets and applies ``mix_prob`` only at train time) over a stateful,
resume-able sampler/loader. Use with the instance-segmentation losses
(e.g. ``FastInstanceSegmentationLoss``). Registered as ``InsegTrainer`` --
select via ``train = dict(type="InsegTrainer")``.
"""
[docs]
def build_train_loader(self):
"""Build the stateful instance-segmentation training loader."""
train_data = build_dataset(self.cfg.data.train)
return self._build_stateful_train_loader(
train_data,
partial(inseg_collate_fn, mix_prob=self.cfg.mix_prob),
in_order=self.cfg.deterministic,
)
[docs]
def build_val_loader(self):
"""Build the optional instance-segmentation validation loader."""
val_loader = None
if self.cfg.evaluate:
val_data = build_dataset(self.cfg.data.val)
if comm.get_world_size() > 1:
val_sampler = torch.utils.data.distributed.DistributedSampler(val_data)
else:
val_sampler = None
val_loader = torch.utils.data.DataLoader(
val_data,
batch_size=self.cfg.batch_size_val_per_gpu,
shuffle=False,
num_workers=self.cfg.num_worker_per_gpu,
pin_memory=True,
sampler=val_sampler,
# Use inseg_collate_fn for validation as well
collate_fn=partial(inseg_collate_fn, mix_prob=0),
# in_order=self.cfg.deterministic,
)
return val_loader
[docs]
@TRAINERS.register_module("ImageClassTrainer")
class ImageClassTrainer(Trainer):
"""Trainer for dense 2D image batches (e.g. rasterized ring images).
Identical to :class:`Trainer` except the train loader uses ``default_collate``
to stack per-event ``image`` into ``(B, C, H, W)`` and scalar labels/momenta
into ``(B, 1)``, instead of the point-cloud collate that concatenates
variable-length clouds along a single axis. Registered as
``ImageClassTrainer`` -- select via ``train = dict(type="ImageClassTrainer")``.
"""
[docs]
def build_train_loader(self):
"""Build the stateful image-classification training loader."""
train_data = build_dataset(self.cfg.data.train)
return self._build_stateful_train_loader(
train_data, torch.utils.data.default_collate
)
[docs]
def build_val_loader(self):
"""Build the optional image-classification validation loader."""
val_loader = None
if self.cfg.evaluate:
val_data = build_dataset(self.cfg.data.val)
if comm.get_world_size() > 1:
val_sampler = torch.utils.data.distributed.DistributedSampler(val_data)
else:
val_sampler = None
val_loader = torch.utils.data.DataLoader(
val_data,
batch_size=self.cfg.batch_size_val_per_gpu,
shuffle=False,
num_workers=self.cfg.num_worker_per_gpu,
pin_memory=True,
sampler=val_sampler,
collate_fn=torch.utils.data.default_collate,
)
return val_loader