diff --git a/run_train.sh b/run_train.sh index a8f820fc109..413e5d7a68c 100755 --- a/run_train.sh +++ b/run_train.sh @@ -33,6 +33,7 @@ export LOG_RANK=${LOG_RANK:-0} export OMP_NUM_THREADS=${OMP_NUM_THREADS:-1} export NCCL_P2P_DISABLE=${NCCL_P2P_DISABLE:-1} export NCCL_SHM_LOCALITY=${NCCL_SHM_LOCALITY:-1} +export NCCL_BUFFSIZE=${NCCL_BUFFSIZE:-2097152} MODULE=${MODULE:-"llama3"} CONFIG=${CONFIG:-"llama3_debugmodel"} COMM_MODE=${COMM_MODE:-""} diff --git a/torchtitan/components/metrics.py b/torchtitan/components/metrics.py index 08c9a165d3d..8b775501569 100644 --- a/torchtitan/components/metrics.py +++ b/torchtitan/components/metrics.py @@ -41,9 +41,7 @@ def __init__(self, device: str = f"{device_type}:0"): self.device = torch.device(device) # device object self.device_name = device_module.get_device_name(self.device) self.device_index = device_module.current_device() - self.device_capacity = device_module.get_device_properties( - self.device - ).total_memory + self.device_capacity = device_module.get_device_properties(self.device).total_memory self.device_capacity_gib = self._to_gib(self.device_capacity) device_module.reset_peak_memory_stats() @@ -73,9 +71,7 @@ def get_peak_stats(self): num_ooms = device_info.get("num_ooms", -1) if num_retries > 0: - logger.warning( - f"{num_retries} {device_type.upper()} memory allocation retries." - ) + logger.warning(f"{num_retries} {device_type.upper()} memory allocation retries.") if num_ooms > 0: logger.warning(f"{num_ooms} {device_type.upper()} OOM errors thrown.") @@ -108,9 +104,7 @@ def log(self, metrics: dict[str, Any], step: int) -> None: pass def write_report(self, data: Any, step: int, name: str, output_type: str) -> None: - raise NotImplementedError( - f"{type(self).__name__} does not support write_report" - ) + raise NotImplementedError(f"{type(self).__name__} does not support write_report") def close(self) -> None: pass @@ -168,10 +162,7 @@ def __init__( logger.info("WandB logging enabled") def log(self, metrics: dict[str, Any], step: int) -> None: - wandb_metrics = { - (k if self.tag is None else f"{self.tag}/{k}"): v - for k, v in metrics.items() - } + wandb_metrics = {(k if self.tag is None else f"{self.tag}/{k}"): v for k, v in metrics.items()} self.wandb.log(wandb_metrics, step=step) def close(self) -> None: @@ -193,9 +184,7 @@ def __init__( raise ValueError("REPORTERV2_HOST must be set to use ReporterV2 logging") training_id = os.getenv("REPORTERV2_TRAINING_ID") if not training_id: - raise ValueError( - "REPORTERV2_TRAINING_ID must be set to use ReporterV2 logging" - ) + raise ValueError("REPORTERV2_TRAINING_ID must be set to use ReporterV2 logging") from reporterv2 import ReporterV2 reporter_config = dict(config_dict or {}) @@ -204,11 +193,7 @@ def __init__( metrics_config = {} self.save_freq = int(metrics_config.get("save_freq", 1)) model_spec = reporter_config.get("model_spec", {}) - model_name = ( - model_spec.get("name", "torchtitan") - if isinstance(model_spec, dict) - else "torchtitan" - ) + model_name = model_spec.get("name", "torchtitan") if isinstance(model_spec, dict) else "torchtitan" reporter_config["training_id"] = training_id reporter_config["reporterv2_host"] = host reporter_config.setdefault("trainer", model_name) @@ -231,10 +216,6 @@ def log(self, metrics: dict[str, Any], step: int) -> None: self.reporter.save_metrics() def write_report(self, data: Any, step: int, name: str, output_type: str) -> None: - if output_type == "scalar": - self.reporter.buffer_metrics(step=step, epoch=step, metrics={name: data}) - self.reporter.save_metrics() - return self.reporter.write_report(data, step=step, name=name, output_type=output_type) def close(self) -> None: @@ -273,9 +254,7 @@ def close(self) -> None: logger_instance.close() -def ensure_pp_loss_visible( - *, parallel_dims: ParallelDims, pp_schedule: str, color: Color | NoColor -) -> None: +def ensure_pp_loss_visible(*, parallel_dims: ParallelDims, pp_schedule: str, color: Color | NoColor) -> None: """ Ensures that the loss is visible on the console for pipeline-parallel training. @@ -428,9 +407,7 @@ def __init__( # used for colorful printing self.color = utils.NoColor() if config.disable_color_printing else utils.Color() - self.gpu_peak_flops = utils.get_peak_flops( - self.device_memory_monitor.device_name - ) + self.gpu_peak_flops = utils.get_peak_flops(self.device_memory_monitor.device_name) self.ntokens_since_last_log = 0 self.data_loading_times = [] self.time_last_log = time.perf_counter() @@ -469,23 +446,15 @@ def _build_metric_logger( ) # Check if any logging backend is enabled - has_logging_enabled = ( - config.enable_tensorboard - or config.enable_wandb - or config.enable_reporterv2 - ) + has_logging_enabled = config.enable_tensorboard or config.enable_wandb or config.enable_reporterv2 # Determine if this rank should log should_log = has_logging_enabled if (not config.save_for_all_ranks) and should_log: - metrics_rank = _get_metrics_rank( - parallel_dims=parallel_dims, pp_schedule=pp_schedule - ) + metrics_rank = _get_metrics_rank(parallel_dims=parallel_dims, pp_schedule=pp_schedule) should_log = torch.distributed.get_rank() == metrics_rank - logger.debug( - f"Logging decision: has_logging_enabled={has_logging_enabled}, should_log={should_log}" - ) + logger.debug(f"Logging decision: has_logging_enabled={has_logging_enabled}, should_log={should_log}") if not should_log: logger.debug("Returning BaseLogger due to should_log=False") @@ -505,9 +474,7 @@ def _build_metric_logger( ) if config.save_for_all_ranks: - base_log_dir = os.path.join( - base_log_dir, f"rank_{torch.distributed.get_rank()}" - ) + base_log_dir = os.path.join(base_log_dir, f"rank_{torch.distributed.get_rank()}") # Create logger container logger_container = LoggerContainer() @@ -516,9 +483,7 @@ def _build_metric_logger( if config.enable_wandb: logger.debug("Attempting to create WandB logger") try: - wandb_logger = WandBLogger( - base_log_dir, config_dict=config_dict, tag=tag - ) + wandb_logger = WandBLogger(base_log_dir, config_dict=config_dict, tag=tag) logger_container.add_logger(wandb_logger) except Exception as e: if "No module named 'wandb'" in str(e): @@ -577,9 +542,7 @@ def log( time_delta = time.perf_counter() - self.time_last_log # tokens per second per device, abbreviated as tps - tps = self.ntokens_since_last_log / ( - time_delta * self.parallel_dims.non_data_parallel_size - ) + tps = self.ntokens_since_last_log / (time_delta * self.parallel_dims.non_data_parallel_size) # model FLOPS utilization # For its definition and calculation, please refer to the PaLM paper: # https://arxiv.org/abs/2204.02311 @@ -639,17 +602,13 @@ def log( self.time_last_log = time.perf_counter() self.device_memory_monitor.reset_peak_stats() - def log_validation( - self, loss: float, step: int, extra_metrics: dict[str, Any] | None = None - ): + def log_validation(self, loss: float, step: int, extra_metrics: dict[str, Any] | None = None): time_delta = time.perf_counter() - self.time_last_log device_mem_stats = self.device_memory_monitor.get_peak_stats() # tokens per second per device, abbreviated as tps - tps = self.ntokens_since_last_log / ( - time_delta * self.parallel_dims.non_data_parallel_size - ) + tps = self.ntokens_since_last_log / (time_delta * self.parallel_dims.non_data_parallel_size) metrics = { "validation_metrics/loss": loss, diff --git a/torchtitan/experiments/__init__.py b/torchtitan/experiments/__init__.py index 530a90ba401..12d7298cb97 100644 --- a/torchtitan/experiments/__init__.py +++ b/torchtitan/experiments/__init__.py @@ -13,6 +13,7 @@ "autoparallel.llama3", "autoparallel.local_map_deepseek_v3", "path", + "rldriving", "worldmodel", "torchft.llama3", "rl", diff --git a/torchtitan/experiments/path/__init__.py b/torchtitan/experiments/path/__init__.py index e5e80b8d2aa..2e41cd717f6 100644 --- a/torchtitan/experiments/path/__init__.py +++ b/torchtitan/experiments/path/__init__.py @@ -3,16 +3,3 @@ # # This source code is licensed under the BSD-style license found in the # LICENSE file in the root directory of this source tree. - -from .config_registry import model_registry -from .model import parallelize_path, PathModel -from .trainer import PathTrainer -from .validate import PathValidator - -__all__ = [ - "PathModel", - "PathTrainer", - "PathValidator", - "model_registry", - "parallelize_path", -] diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 694994841de..12f83d0c24b 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -6,21 +6,11 @@ from __future__ import annotations -import math import os from functools import partial +from typing import Literal + from xx.datasets.constants import BASE_DIR_GT, DEFAULT_TEST_5K_LIST_TAGGED, DEFAULT_TRAIN_LIST -from xx.ml_tools.constants.model import ( - frame_constants_from_fps, - FRAME_TYPE, - INPUT_FRAMES_NAMES, - ModelInputs, - N_FRAMES, - SUPERCOMBO_FPS, - TEMPORAL_INPUTS, -) -from xx.training.path.config import DatasetConfig as XXPathDatasetConfig -from xx.training.path.hydra_configs import DRIVING_HEADS, META_HEADS, POSE_HEADS, TEMPORAL_META_HEADS from torchtitan.components.lr_scheduler import LRSchedulersContainer from torchtitan.components.metrics import MetricsProcessor @@ -28,28 +18,19 @@ from torchtitan.components.tokenizer import NoOpTokenizer from torchtitan.config import CompileConfig, DebugConfig, ParallelismConfig, TrainingConfig from torchtitan.distributed.activation_checkpoint import FullAC -from torchtitan.models.common import Embedding, LayerNorm, Linear -from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model_spec import ModelSpec from .dataset import PathDataLoader from .loss import PathLoss -from .model import ( - Hydra, - LinearEncoder, - parallelize_path, - PathHead, - PathMLP, - PathModel, - PathSelfAttention, - PathTransformer, - PathTransformerBlock, - PointSummarizer, - Policy, - ScaleLayer, - TemporalPolicy, - TemporalSummarizer, - Vision, +from .model import parallelize_path +from .model_config import model_config as _model_config +from .model_constants import ( + frame_constants_from_fps, + FRAME_TYPE, + ModelInputs, + N_FRAMES, + SUPERCOMBO_FPS, + TEMPORAL_INPUTS, ) from .onnx_checkpoint import PathOnnxCheckpointManager from .trainer import PathTrainer @@ -170,75 +151,10 @@ def _path(flavor: str) -> PathTrainer.Config: ) -def _model_config(flavor: str) -> PathModel.Config: - vision_features = 512 - n_frames_input = N_FRAMES - input_frame_names = INPUT_FRAMES_NAMES - input_frame_type = FRAME_TYPE - frame_constants = frame_constants_from_fps(n_frames=n_frames_input, frame_type=input_frame_type) - in_channels = sum(frame_constants["frame_shapes"][name][0] for name in input_frame_names) - block_size = len(frame_constants["history_idxs"]) - desire_window_len = frame_constants["desire_window_len"] - history_idxs = tuple(int(x) for x in frame_constants["history_idxs"]) - desire_window_starts = tuple(idx - history_idxs[0] for idx in history_idxs) - dim = vision_features - - return PathModel.Config( - n_frames_input=n_frames_input, - input_frame_names=tuple(input_frame_names), - frame_type=input_frame_type, - vision=Vision.Config( - flavor=flavor, - input_frame_names=tuple(input_frame_names), - in_channels=in_channels, - vision_features=vision_features, - pretrained=True, - drop_path_rate=0.2, - mean=255 / 2, - std=255 / 4, - ), - point_policy=Policy.Config( - summarizer=PointSummarizer.Config( - mlp1=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - mlp2=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - ), - hydra=_hydra(_heads(META_HEADS + POSE_HEADS), in_features=dim, mlp_mult=2), - ), - temporal_policy=TemporalPolicy.Config( - temporal_summarizer=TemporalSummarizer.Config( - mlp1=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - mlp2=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - desire_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.DESIRE][0] * desire_window_len, dim), - desire_window_len=desire_window_len, - desire_window_starts=desire_window_starts, - traffic_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.TRAFFIC][0], dim), - action_t_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.ACTION_T][0], dim), - transformer=PathTransformer.Config( - layers=[ - PathTransformerBlock.Config( - attention=_attention(dim=dim, n_head=8, dropout=0.1), - mlp=_mlp(dim, mlp_mult=2, bias=True, dropout=0.1), - ) - for _ in range(4) - ] - ), - pos_embedding=Embedding.Config( - num_embeddings=block_size, - embedding_dim=dim, - ), - block_size=block_size, - dense_training_outputs=True, - ), - temporal_hydra=_hydra(_heads(DRIVING_HEADS + TEMPORAL_META_HEADS), in_features=dim, mlp_mult=2), - history_idxs=history_idxs, - ), - ) - - def _dataloader_config( *, dataset: str, - split: str, + split: Literal["train", "val"], fps: int, plan_only: bool, limit: int | None, @@ -247,32 +163,17 @@ def _dataloader_config( skip: int, val_skip: int, ) -> PathDataLoader.Config: - base = XXPathDatasetConfig( + return PathDataLoader.Config( + dataset=dataset, + split=split, + deterministic_fidxs=deterministic_fidxs, fps=fps, + pipeline_dir=pipeline_dir, plan_only=plan_only, limit=limit, - pipeline_dir=pipeline_dir, skip=skip, val_skip=val_skip, ) - return PathDataLoader.Config( - dataset=dataset, - split=split, - shuffle_size=_si_int(base.shuffle_size), - deterministic_fidxs=deterministic_fidxs, - min_mixing=base.min_mixing, - num_writers=base.num_writers, - num_readers=base.num_readers, - fps=base.fps, - pipeline_dir=base.pipeline_dir, - plan_only=base.plan_only, - limit=base.limit, - n_frames=base.n_frames, - rgb=base.rgb, - unvision=base.unvision, - skip=base.skip, - val_skip=base.val_skip, - ) def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnxCheckpointManager.Config: @@ -315,12 +216,6 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx ) -def _si_int(value: str | int) -> int: - suffixes = {"k": 1_000, "m": 1_000_000, "g": 1_000_000_000} - value = str(value).strip().lower() - return int(float(value[:-1]) * suffixes[value[-1]]) if value[-1] in suffixes else int(value) - - def _optimizer_config() -> OptimizersContainer.Config: common = {"lr": 1e-3, "betas": (0.9, 0.95), "eps": 1e-8} no_decay = r"(point_policy\.hydra|temporal_policy\.temporal_hydra)\.(final_layer|scale_layer)" @@ -341,70 +236,6 @@ def _optimizer_config() -> OptimizersContainer.Config: ) -def _heads(heads) -> tuple[PathHead, ...]: - return tuple(PathHead(head.name, head.output_size, head.mlp, head.scale) for head in heads) - - -def _hidden_dim(dim: int, mlp_mult: float, multiple_of: int = 256) -> int: - hidden = int(dim * mlp_mult) - return multiple_of * math.ceil(hidden / multiple_of) - - -def _mlp(dim: int, *, mlp_mult: float, bias: bool, dropout: float) -> PathMLP.Config: - hidden = _hidden_dim(dim, mlp_mult) - return PathMLP.Config( - norm=LayerNorm.Config(normalized_shape=dim), - c_fc=Linear.Config(in_features=dim, out_features=hidden, bias=bias), - c_proj=Linear.Config(in_features=hidden, out_features=dim, bias=bias), - act="gelu_tanh", - dropout=dropout, - ) - - -def _encoder(in_features: int, dim: int) -> LinearEncoder.Config: - return LinearEncoder.Config( - in_layer=Linear.Config( - in_features=in_features, - out_features=dim, - bias=True, - ), - out_layer=Linear.Config(in_features=dim, out_features=dim, bias=False), - ) - - -def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Config: - head_dim = dim // n_head - return PathSelfAttention.Config( - norm=LayerNorm.Config(normalized_shape=dim), - q_norm=LayerNorm.Config(normalized_shape=head_dim), - k_norm=LayerNorm.Config(normalized_shape=head_dim), - c_attn=Linear.Config(in_features=dim, out_features=3 * dim, bias=True), - c_proj=Linear.Config(in_features=dim, out_features=dim, bias=True), - inner_attention=ScaledDotProductAttention.Config(), - n_head=n_head, - head_dim=head_dim, - dropout=dropout, - ) - - -def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> Hydra.Config: - return Hydra.Config( - heads=heads, - head_mlps={ - head.name: _mlp(in_features, mlp_mult=mlp_mult, bias=False, dropout=0.0) for head in heads if head.mlp - }, - final_layers={ - head.name: Linear.Config( - in_features=in_features, - out_features=head.output_size, - bias=True, - ) - for head in heads - }, - scale_layers={head.name: ScaleLayer.Config(n_features=head.output_size) for head in heads if head.scale}, - ) - - convnext_atto = partial(_path, "convnext_atto") convnext_femto = partial(_path, "convnext_femto") convnext_pico = partial(_path, "convnext_pico") diff --git a/torchtitan/experiments/path/dataset.py b/torchtitan/experiments/path/dataset.py index a8dafc600ea..f904b2f449f 100644 --- a/torchtitan/experiments/path/dataset.py +++ b/torchtitan/experiments/path/dataset.py @@ -9,32 +9,54 @@ import os from collections.abc import Iterator from dataclasses import dataclass -from typing import Any +from typing import Any, Literal import torch from torchtitan.components.dataloader import BaseDataLoader from torchtitan.components.tokenizer import BaseTokenizer +from .model_constants import FRAME_TYPE, N_FRAMES, SUPERCOMBO_FPS, VisionFrameType + class PathDataLoader(BaseDataLoader): @dataclass(kw_only=True, slots=True) class Config(BaseDataLoader.Config): - split: str - shuffle_size: int - min_mixing: float - num_writers: int - num_readers: int - fps: int - pipeline_dir: str - plan_only: bool - limit: int | None - deterministic_fidxs: bool - n_frames: int - rgb: bool - unvision: bool - skip: int - val_skip: int + dataset: str + split: Literal["train", "val"] + pipeline_dir: str | None = None + shuffle_size: int = 8_000 + min_mixing: float = 0.5 + num_writers: int = 6 + num_readers: int = 1 + fps: int = SUPERCOMBO_FPS + plan_only: bool = False + limit: int | None = 2_500_000 + deterministic_fidxs: bool = False + n_frames: int = N_FRAMES + rgb: bool = FRAME_TYPE is VisionFrameType.RGB + unvision: bool = False + skip: int = 1 + val_skip: int = 1 + + def _build_dataset( + self, + config: Config, + *, + val: bool, + ) -> Any: + if config.pipeline_dir is None: + raise ValueError("pipeline_dir is required for internal PATH datasets") + + from xx.training.path.dataloader import get_dataset + + return get_dataset( + config, + val=val, + local_rank=self.local_rank, + global_rank=self.dp_rank, + global_world_size=self.dp_world_size, + ) def __init__( self, @@ -50,11 +72,7 @@ def __init__( **kwargs: Any, ) -> None: del tokenizer, seq_len, snapshot_every_n_steps, kwargs - from xx.training.lib.dataloader import DataLoader - from xx.training.path.config import DatasetConfig as XXPathDatasetConfig - from xx.training.path.dataloader import get_dataset - - from gigashuffle import DataloaderConfig + from gigashuffle import DataloaderConfig, MultiprocessShuffledDataloader self.config = config self.local_batch_size = local_batch_size @@ -68,25 +86,10 @@ def __init__( val_shuffle_size = local_batch_size * validation_steps * self.local_world_size shuffle_size = val_shuffle_size if val else config.shuffle_size - xx_config = XXPathDatasetConfig( - bs=local_batch_size, - shuffle_size=str(config.shuffle_size), - val_shuffle_size=str(val_shuffle_size), - min_mixing=config.min_mixing, - num_writers=config.num_writers, - num_readers=config.num_readers, - fps=config.fps, - pipeline_dir=config.pipeline_dir, - plan_only=config.plan_only, - limit=config.limit, - deterministic_fidxs=config.deterministic_fidxs, - n_frames=config.n_frames, - rgb=config.rgb, - unvision=config.unvision, - skip=config.skip, - val_skip=config.val_skip, + dataset = self._build_dataset( + config, + val=val, ) - dataset = get_dataset(config.dataset, xx_config, val, self.local_rank, dp_rank, dp_world_size) self.dataset = dataset loader_config = DataloaderConfig( bs=local_batch_size, @@ -101,7 +104,7 @@ def __init__( global_world_size=dp_world_size, queue_name=f"{run_id}-{config.split}-node{node_rank}", ) - self.loader = DataLoader(dataset, loader_config) + self.loader = MultiprocessShuffledDataloader(dataset, loader_config) self._iterator: Any | None = None def __iter__( diff --git a/torchtitan/experiments/path/hydra_configs.py b/torchtitan/experiments/path/hydra_configs.py new file mode 100644 index 00000000000..ad0a0bacc8b --- /dev/null +++ b/torchtitan/experiments/path/hydra_configs.py @@ -0,0 +1,41 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from .model import PathHead +from .model_constants import ( + ACTION_LEN, + DESIRE_LEN, + LEAD_TRAJECTORY_DIM, + LL_SIZE, + META_LEN, + PLAN_SIZE, + POSE_OUTSIZE, + RE_SIZE, +) + +PLAN_HEAD_SIZE = 2 * PLAN_SIZE + +DRIVING_HEADS = [ + PathHead(name="plan", output_size=PLAN_HEAD_SIZE, mlp=True, scale=True), + PathHead(name="lead", output_size=3 * (2 * LEAD_TRAJECTORY_DIM), mlp=True, scale=True), + PathHead(name="lead_prob", output_size=3, mlp=True, scale=False), + PathHead(name="action", output_size=ACTION_LEN * 2, mlp=True, scale=True), +] +TEMPORAL_META_HEADS = [ + PathHead(name="desire_state", output_size=DESIRE_LEN, mlp=True, scale=False), +] +META_HEADS = [ + PathHead(name="lane_lines", output_size=4 * (2 * LL_SIZE), mlp=False, scale=False), + PathHead(name="lane_lines_prob", output_size=8, mlp=False, scale=False), + PathHead(name="road_edges", output_size=2 * (2 * RE_SIZE), mlp=False, scale=False), + PathHead(name="meta", output_size=META_LEN, mlp=False, scale=False), + PathHead(name="desire_pred", output_size=DESIRE_LEN * 4, mlp=False, scale=False), + PathHead(name="road_transform", output_size=POSE_OUTSIZE * 2, mlp=False, scale=False), + PathHead(name="wide_from_device_euler", output_size=3 * 2, mlp=False, scale=False), +] +POSE_HEADS = [ + PathHead(name="pose", output_size=POSE_OUTSIZE * 2, mlp=True, scale=True), +] diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 53e2cd9224c..c01bf26b536 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -7,7 +7,6 @@ from __future__ import annotations from dataclasses import dataclass -from xx.ml_tools.constants.model import frame_constants_from_fps, ModelInputs, TEMPORAL_INPUTS import torch import torch.nn as nn @@ -33,6 +32,7 @@ from torchtitan.tools.logging import logger from . import convnext +from .model_constants import frame_constants_from_fps, ModelInputs, TEMPORAL_INPUTS, VisionFrameType @dataclass(frozen=True) @@ -426,7 +426,7 @@ class PathModel(BaseModel): class Config(BaseModel.Config): n_frames_input: int input_frame_names: tuple[str, ...] - frame_type: str + frame_type: VisionFrameType vision: Vision.Config point_policy: Policy.Config temporal_policy: TemporalPolicy.Config diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py new file mode 100644 index 00000000000..72eded0a7e9 --- /dev/null +++ b/torchtitan/experiments/path/model_config.py @@ -0,0 +1,161 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +import math + +from torchtitan.models.common import Embedding, LayerNorm, Linear +from torchtitan.models.common.attention import ScaledDotProductAttention + +from .hydra_configs import DRIVING_HEADS, META_HEADS, POSE_HEADS, TEMPORAL_META_HEADS +from .model import ( + Hydra, + LinearEncoder, + PathHead, + PathMLP, + PathModel, + PathSelfAttention, + PathTransformer, + PathTransformerBlock, + PointSummarizer, + Policy, + ScaleLayer, + TemporalPolicy, + TemporalSummarizer, + Vision, +) +from .model_constants import ( + frame_constants_from_fps, + FRAME_TYPE, + INPUT_FRAMES_NAMES, + ModelInputs, + N_FRAMES, + TEMPORAL_INPUTS, +) + + +POINT_HEADS = tuple(META_HEADS + POSE_HEADS) +TEMPORAL_HEADS = tuple(DRIVING_HEADS + TEMPORAL_META_HEADS) + + +def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: + vision_features = 512 + frame_constants = frame_constants_from_fps(n_frames=N_FRAMES, frame_type=FRAME_TYPE) + input_frame_names = tuple(INPUT_FRAMES_NAMES) + in_channels = sum(frame_constants["frame_shapes"][name][0] for name in input_frame_names) + + return PathModel.Config( + n_frames_input=N_FRAMES, + input_frame_names=input_frame_names, + frame_type=FRAME_TYPE, + vision=Vision.Config( + flavor=flavor, + input_frame_names=input_frame_names, + in_channels=in_channels, + vision_features=vision_features, + pretrained=True, + drop_path_rate=0.2, + mean=255 / 2, + std=255 / 4, + ), + point_policy=Policy.Config( + summarizer=PointSummarizer.Config( + mlp1=_mlp(vision_features, mlp_mult=2, bias=False, dropout=0.0), + mlp2=_mlp(vision_features, mlp_mult=2, bias=False, dropout=0.0), + ), + hydra=_hydra(POINT_HEADS, in_features=vision_features, mlp_mult=2), + ), + temporal_policy=temporal_policy_config(), + ) + + +def temporal_policy_config( + *, + heads: tuple[PathHead, ...] = TEMPORAL_HEADS, + dropout: float = 0.1, + dense_training_outputs: bool = True, +) -> TemporalPolicy.Config: + vision_features = 512 + frame_constants = frame_constants_from_fps() + history_idxs = tuple(int(index) for index in frame_constants["history_idxs"]) + desire_window_len = frame_constants["desire_window_len"] + desire_window_starts = tuple(index - history_idxs[0] for index in history_idxs) + return TemporalPolicy.Config( + temporal_summarizer=TemporalSummarizer.Config( + mlp1=_mlp(vision_features, mlp_mult=2, bias=False, dropout=0.0), + mlp2=_mlp(vision_features, mlp_mult=2, bias=False, dropout=0.0), + desire_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.DESIRE][0] * desire_window_len, vision_features), + desire_window_len=desire_window_len, + desire_window_starts=desire_window_starts, + traffic_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.TRAFFIC][0], vision_features), + action_t_encoder=_encoder(TEMPORAL_INPUTS[ModelInputs.ACTION_T][0], vision_features), + transformer=PathTransformer.Config( + layers=[ + PathTransformerBlock.Config( + attention=_attention(dim=vision_features, n_head=8, dropout=dropout), + mlp=_mlp(vision_features, mlp_mult=2, bias=True, dropout=dropout), + ) + for _ in range(4) + ] + ), + pos_embedding=Embedding.Config( + num_embeddings=len(history_idxs), + embedding_dim=vision_features, + ), + block_size=len(history_idxs), + dense_training_outputs=dense_training_outputs, + ), + temporal_hydra=_hydra(heads, in_features=vision_features, mlp_mult=2), + history_idxs=history_idxs, + ) + + +def _mlp(dim: int, *, mlp_mult: float, bias: bool, dropout: float) -> PathMLP.Config: + hidden = 256 * math.ceil(int(dim * mlp_mult) / 256) + return PathMLP.Config( + norm=LayerNorm.Config(normalized_shape=dim), + c_fc=Linear.Config(in_features=dim, out_features=hidden, bias=bias), + c_proj=Linear.Config(in_features=hidden, out_features=dim, bias=bias), + act="gelu_tanh", + dropout=dropout, + ) + + +def _encoder(in_features: int, dim: int) -> LinearEncoder.Config: + return LinearEncoder.Config( + in_layer=Linear.Config(in_features=in_features, out_features=dim, bias=True), + out_layer=Linear.Config(in_features=dim, out_features=dim, bias=False), + ) + + +def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Config: + head_dim = dim // n_head + return PathSelfAttention.Config( + norm=LayerNorm.Config(normalized_shape=dim), + q_norm=LayerNorm.Config(normalized_shape=head_dim), + k_norm=LayerNorm.Config(normalized_shape=head_dim), + c_attn=Linear.Config(in_features=dim, out_features=3 * dim, bias=True), + c_proj=Linear.Config(in_features=dim, out_features=dim, bias=True), + inner_attention=ScaledDotProductAttention.Config(), + n_head=n_head, + head_dim=head_dim, + dropout=dropout, + ) + + +def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> Hydra.Config: + return Hydra.Config( + heads=heads, + head_mlps={ + head.name: _mlp(in_features, mlp_mult=mlp_mult, bias=False, dropout=0.0) for head in heads if head.mlp + }, + final_layers={ + head.name: Linear.Config(in_features=in_features, out_features=head.output_size, bias=True) + for head in heads + }, + scale_layers={head.name: ScaleLayer.Config(n_features=head.output_size) for head in heads if head.scale}, + ) diff --git a/torchtitan/experiments/path/model_constants.py b/torchtitan/experiments/path/model_constants.py new file mode 100644 index 00000000000..4dbec10c3ad --- /dev/null +++ b/torchtitan/experiments/path/model_constants.py @@ -0,0 +1,161 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""Path model input, output, and timing contract. + +This is the canonical implementation for both TorchTitan PATH and the legacy +``xx.ml_tools.constants.model`` import path. Keep it free of xx and openpilot +dependencies: these values determine checkpoint shapes and the ONNX interface. +""" + +from __future__ import annotations + +from enum import Enum +from typing import Any, TypedDict + +import numpy as np + + +# These camera/model dimensions are part of the model contract. Keeping them +# here avoids importing openpilot hardware and transformation modules. +CAMERA_FPS = 20 +MEDMODEL_INPUT_SIZE = (512, 256) +BIGMODEL_INPUT_SIZE = (1024, 512) +SBIGMODEL_INPUT_SIZE = (512, 256) + + +def index_function(idx: int, max_val: float = 192) -> float: + return (max_val / 1024) * (idx**2) + + +IDX_N = 33 +T_IDXS = np.array([index_function(idx, max_val=10.0) for idx in range(IDX_N)], dtype=np.float64) +X_IDXS = np.array([index_function(idx, max_val=192.0) for idx in range(IDX_N)], dtype=np.float64) +D_IDXS = np.array([index_function(idx, max_val=512.0) for idx in range(IDX_N)], dtype=np.float64) + +W, H = MEDMODEL_INPUT_SIZE +BIG_W, BIG_H = BIGMODEL_INPUT_SIZE +SBIG_W, SBIG_H = SBIGMODEL_INPUT_SIZE + +POSE_OUTSIZE = 6 +META_LEN = 55 +DESIRE_LEN = 8 +ACTION_LEN = 2 + +LEAD_PRED_DIM = 4 +LEAD_TRAJECTORY_DIM = 6 * 4 +PLAN_WIDTH = 15 +PLAN_SIZE = PLAN_WIDTH * IDX_N +LATERAL_PLANNER_SOLUTION_WIDTH = 4 +LATERAL_PLANNER_SOLUTION_SIZE = LATERAL_PLANNER_SOLUTION_WIDTH * IDX_N +LL_SIZE = 2 * IDX_N +RE_SIZE = 2 * IDX_N + +LEAD_MHP_N = 2 +PLAN_MHP_N = 5 + + +class VisionFrameType(Enum): + YUV = "YUV" + RGB = "RGB" + + +class ModelInputs: + IMG = "img" + BIG_IMG = "big_img" + FEATURES = "features" + DESIRE = "desire_pulse" + TRAFFIC = "traffic_convention" + ACTION_T = "action_t" + ANTI_CHEATING_SAMPLES = "anti_cheating_samples" + FEATURE_QUEUE = "feature_queue" + + +SUPERCOMBO_FPS = 5 +N_FRAMES = 2 +FRAME_TYPE = VisionFrameType.YUV +INPUT_FRAMES_NAMES = [ModelInputs.IMG, ModelInputs.BIG_IMG] + +INPUT_ALIASES = { + getattr(ModelInputs, name): [getattr(ModelInputs, name)] for name in dir(ModelInputs) if not name.startswith("__") +} +INPUT_ALIASES[ModelInputs.IMG] += ["input_imgs", "map"] +INPUT_ALIASES[ModelInputs.BIG_IMG] += ["big_input_imgs"] +INPUT_ALIASES[ModelInputs.FEATURES] += ["vision_features"] +INPUT_ALIASES[ModelInputs.DESIRE] += ["desire_pulse"] +INPUT_ALIASES[ModelInputs.TRAFFIC] += ["traffic_convention"] +INPUT_ALIASES[ModelInputs.ACTION_T] += ["action_t"] +INPUT_ALIASES[ModelInputs.FEATURE_QUEUE] += ["features_buffer"] + +ALIASES_TO_CANONICAL_NAMES = {alias: name for name, aliases in INPUT_ALIASES.items() for alias in aliases} + +VISION_INPUTS_RGB = { + ModelInputs.IMG: (3, H, W), + ModelInputs.BIG_IMG: (3, SBIG_H, SBIG_W), +} + +VISION_INPUTS_YUV = { + ModelInputs.IMG: (6, H // 2, W // 2), + ModelInputs.BIG_IMG: (6, SBIG_H // 2, SBIG_W // 2), +} + +VISION_INPUTS = { + VisionFrameType.RGB: VISION_INPUTS_RGB, + VisionFrameType.YUV: VISION_INPUTS_YUV, +} + +TEMPORAL_INPUTS = { + ModelInputs.FEATURES: (512,), + ModelInputs.DESIRE: (DESIRE_LEN,), + ModelInputs.TRAFFIC: (2,), + ModelInputs.ACTION_T: (ACTION_LEN,), +} + + +def to_canonical(data: dict[str, Any]) -> dict[str, Any]: + return {ALIASES_TO_CANONICAL_NAMES[name]: value for name, value in data.items()} + + +def to_aliases(data: dict[str, Any], aliases: list[str]) -> dict[str, Any]: + canonical_names_to_aliases = to_canonical(dict(zip(aliases, aliases))) + return {canonical_names_to_aliases[name]: value for name, value in data.items()} + + +class FrameConstants(TypedDict): + fps: int + frame_skip: int + n_frames: int + frame_shapes: dict[str, tuple[int, int, int]] + dt: float + desire_window_len: int + temporal_len: int + history_len: int + history_idxs: np.ndarray + + +def frame_constants_from_fps( + fps: int = SUPERCOMBO_FPS, + n_frames: int = N_FRAMES, + frame_type: VisionFrameType = FRAME_TYPE, +) -> FrameConstants: + desire_window_len_seconds = 5 + history_fps = 4 + history_len_seconds = 2 + history_skip = fps // history_fps + history_idxs = -np.arange(1, history_len_seconds * fps, history_skip).astype(np.int32)[::-1] + desire_window_len = fps * desire_window_len_seconds + temporal_len = desire_window_len + int(history_idxs[-1] - history_idxs[0]) + return { + "fps": fps, + "frame_skip": CAMERA_FPS // fps, + "n_frames": n_frames, + "frame_shapes": {name: (shape[0] * n_frames, *shape[1:]) for name, shape in VISION_INPUTS[frame_type].items()}, + "dt": 1 / fps, + "desire_window_len": desire_window_len, + "temporal_len": temporal_len, + "history_len": fps * history_len_seconds, + "history_idxs": history_idxs, + } diff --git a/torchtitan/experiments/path/onnx_checkpoint.py b/torchtitan/experiments/path/onnx_checkpoint.py index d750f9f7849..e5ee56e1743 100644 --- a/torchtitan/experiments/path/onnx_checkpoint.py +++ b/torchtitan/experiments/path/onnx_checkpoint.py @@ -8,7 +8,6 @@ from dataclasses import dataclass, field -from xx.ml_tools.constants.model import ModelInputs from xx.training.lib.onnx_helpers import add_onnx_metadata, patch_depthwise_convs import onnx @@ -19,6 +18,7 @@ from torchtitan.components.onnx_checkpoint import _ONNX_DTYPE_MAP, OnnxCheckpointManager from .model import PathModel +from .model_constants import ModelInputs class _VisionOnnxModel(nn.Module): diff --git a/torchtitan/experiments/rldriving/__init__.py b/torchtitan/experiments/rldriving/__init__.py new file mode 100644 index 00000000000..2e41cd717f6 --- /dev/null +++ b/torchtitan/experiments/rldriving/__init__.py @@ -0,0 +1,5 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. diff --git a/torchtitan/experiments/rldriving/back_port_off_policy_model.py b/torchtitan/experiments/rldriving/back_port_off_policy_model.py new file mode 100644 index 00000000000..e16a6b11dec --- /dev/null +++ b/torchtitan/experiments/rldriving/back_port_off_policy_model.py @@ -0,0 +1,78 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +import json +import re +import sys +from pathlib import Path + +import torch.distributed.checkpoint as dcp +from safetensors.torch import load_file + +from torchtitan.experiments.path.config_registry import _model_config + + +def back_port_off_policy_model(input_path: Path, output_path: Path) -> None: + hparams = json.loads((input_path / "hparams.json").read_text(encoding="utf-8")) + state_dict = load_file(input_path / "model_state_dict.safetensors") + + model_config = _model_config(hparams["model"]["timm_backbone"]) + temporal_hparams = hparams["model"]["temporal_policy"] + hidden_dim = int(temporal_hparams["n_embd"] * temporal_hparams["mlp_mult"]) + for layer in model_config.temporal_policy.temporal_summarizer.transformer.layers: + layer.mlp.c_fc.out_features = hidden_dim + layer.mlp.c_proj.in_features = hidden_dim + + temporal_policy = {} + point_policy = {} + vision = {} + for name, value in state_dict.items(): + if name == "policy.temporal_summarizer.transformer.mask": + continue + if name.startswith("policy."): + name = name.removeprefix("policy.") + name = name.replace("_desire_encode.", "desire_encoder.net.") + name = name.replace("_traffic_encode.", "traffic_encoder.net.") + name = name.replace("_action_t_encode.", "action_t_encoder.net.") + name = re.sub(r"transformer\.(\d+)\.attn\.", r"transformer.layers.\1.attention.", name) + name = re.sub(r"transformer\.(\d+)\.mlp\.", r"transformer.layers.\1.mlp.", name) + temporal_policy[name.replace(".layer_norm.", ".norm.")] = value + elif name.startswith("point_policy."): + name = name.removeprefix("point_policy.").replace(".layer_norm.", ".norm.") + point_policy[name] = value + elif name.startswith("vision."): + name = name.removeprefix("vision.").replace("_en.", "encoder.", 1) + vision[name] = value + + output_hparams = hparams | { + "torchtitan": { + "model": { + "temporal_policy": model_config.temporal_policy.to_dict(), + "point_policy": model_config.point_policy.to_dict(), + "vision": model_config.vision.to_dict() + | { + "act_layer_name": hparams["model"]["timm_kwargs"]["act_layer_name"], + "norm_layer_name": hparams["model"]["timm_kwargs"]["norm_layer_name"], + "norm_eps": 1e-3, + "norm_momentum": 1e-2, + }, + }, + } + } + dcp.save( + { + "temporal_policy": temporal_policy, + "point_policy": point_policy, + "vision": vision, + }, + checkpoint_id=output_path, + no_dist=True, + ) + (output_path / "hparams.json").write_text(json.dumps(output_hparams, indent=2) + "\n", encoding="utf-8") + + +if __name__ == "__main__": + back_port_off_policy_model(Path(sys.argv[1]), Path(sys.argv[2])) diff --git a/torchtitan/experiments/rldriving/config_registry.py b/torchtitan/experiments/rldriving/config_registry.py new file mode 100644 index 00000000000..0b114c6b7ee --- /dev/null +++ b/torchtitan/experiments/rldriving/config_registry.py @@ -0,0 +1,191 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +import os +from typing import cast + +from xx.datasets.constants import BASE_DIR_GT, DEFAULT_TRAIN_LIST +from xx.ml_tools.constants.model import SUPERCOMBO_FPS + +from torchtitan.components.metrics import MetricsProcessor +from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig +from torchtitan.components.tokenizer import NoOpTokenizer +from torchtitan.config import CompileConfig, DebugConfig, ParallelismConfig, TrainingConfig +from torchtitan.protocols.model_spec import ModelSpec + +from .dataset import RLDrivingDataLoader +from .loss import RLDrivingLoss +from .model import actor_config, critic_config, parallelize_rldriving, RLDrivingModel +from .onnx_checkpoint import RLDrivingOnnxCheckpointManager +from .trainer import RLDrivingLRSchedulersConfig, RLDrivingTrainer +from .validate import RLDrivingValidator + + +def model_registry() -> ModelSpec: + actor = actor_config() + critic = critic_config(actor) + return ModelSpec( + name="rldriving", + flavor="default", + model=RLDrivingModel.Config(actor=actor, critic=critic), + parallelize_fn=parallelize_rldriving, + pipelining_fn=None, + post_optimizer_build_fn=None, + state_dict_adapter=None, + ) + + +def rldriving() -> RLDrivingTrainer.Config: + fps = SUPERCOMBO_FPS + num_epochs = 201 + steps_per_epoch = 64 + model_spec = model_registry() + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) + world_size = int(os.environ.get("WORLD_SIZE", str(local_world_size))) + num_nodes = int(os.environ.get("GROUP_WORLD_SIZE", str(world_size // local_world_size))) + reporterv2_host = os.getenv("REPORTERV2_HOST") + reporterv2_training_id = os.getenv("REPORTERV2_TRAINING_ID") + checkpoint_base_folder = f"{reporterv2_host.rstrip('/')}/checkpoint" if reporterv2_host else "" + actor_optim = {"lr": 4e-5, "betas": (0.9, 0.999), "eps": 1e-8} + critic_optim = {"lr": 2e-4, "betas": (0.9, 0.999), "eps": 1e-8} + frequent_report_steps = [(epoch + 1) * steps_per_epoch for epoch in range(0, num_epochs, num_epochs // 10)] + sparse_report_steps = [(epoch + 1) * steps_per_epoch for epoch in range(0, num_epochs, num_epochs // 2)] + reports = dict.fromkeys( + ( + "analyse_lat.no_noise", + "analyse_lat.realistic_noise", + "analyse_long", + ), + frequent_report_steps, + ) | dict.fromkeys( + ( + "analyse_unintended_lead_following", + "analyse_speed_convergence", + "analyse_platform_oscillation", + "analyse_nurec", + ), + sparse_report_steps, + ) + return RLDrivingTrainer.Config( + model_spec=model_spec, + loss=RLDrivingLoss.Config( + action_noise=(0.25, 0.25), + gamma=0.95, + fps=fps, + smooth_lat_cost=0.1, + smooth_long_cost=0.1, + curv_cost=100.0, + ), + warm_start_checkpoint=os.getenv( + "RLDRIVING_WARM_START_CHECKPOINT", + "44b83fa5-2a33-7ee7-40f1-e86e3c24ad36/56320", + ), + tokenizer=NoOpTokenizer.Config(), + dataloader=RLDrivingDataLoader.Config( + dataset=DEFAULT_TRAIN_LIST, + training_id=reporterv2_training_id or "", + pipeline_dir=BASE_DIR_GT, + epochs=num_epochs, + steps_per_epoch=steps_per_epoch, + fps=fps, + ), + optimizer=OptimizersContainer.Config( + implementation="fused", + param_groups=[ + ParamGroupConfig( + pattern=r"^actor\.temporal_hydra\.(final_layer|scale_layer)\.", + optimizer_name="AdamW", + optimizer_kwargs={**actor_optim, "weight_decay": 0.0}, + ), + ParamGroupConfig( + pattern=r"^actor\.", + optimizer_name="AdamW", + optimizer_kwargs={**actor_optim, "weight_decay": 3e-2}, + ), + ParamGroupConfig( + pattern=r"^critic\.(critic1|critic2)\.q_hydra\.(final_layer|scale_layer)\.", + optimizer_name="AdamW", + optimizer_kwargs={**critic_optim, "weight_decay": 0.0}, + ), + ParamGroupConfig( + pattern=r"^critic\.", + optimizer_name="AdamW", + optimizer_kwargs={**critic_optim, "weight_decay": 3e-2}, + ), + ], + ), + lr_scheduler=RLDrivingLRSchedulersConfig( + steps_per_epoch=steps_per_epoch, + num_epochs=num_epochs, + ), + training=TrainingConfig( + local_batch_size=32, + global_batch_size=-1, + seq_len=1, + max_norm=1.0, + steps=num_epochs * steps_per_epoch, + dtype="float32", + mixed_precision_param="bfloat16", + mixed_precision_reduce="float32", + ), + parallelism=ParallelismConfig( + data_parallel_replicate_degree=num_nodes, + data_parallel_shard_degree=local_world_size, + tensor_parallel_degree=1, + context_parallel_degree=1, + pipeline_parallel_degree=1, + expert_parallel_degree=1, + enable_sequence_parallel=False, + ), + checkpoint=_checkpoint_config( + cast(RLDrivingModel.Config, model_spec.model), + base_folder=checkpoint_base_folder, + folder=reporterv2_training_id or "checkpoint", + interval=steps_per_epoch, + ), + steps_per_epoch=steps_per_epoch, + train_step_barrier_timeout_seconds=60 * 60, + ema_tau=128.0, + fps=fps, + activation_checkpoint=None, + compile=CompileConfig(enable=True, components=["model"]), + metrics=MetricsProcessor.Config( + log_freq=16, + enable_reporterv2=True, + save_freq=steps_per_epoch, + ), + validator=RLDrivingValidator.Config( + enable=True, + freq=steps_per_epoch, + fps=fps, + reports=reports, + miniray={"priority": 3}, + ), + debug=DebugConfig(seed=0), + ) + + +def _checkpoint_config( + model: RLDrivingModel.Config, + *, + base_folder: str, + folder: str, + interval: int, +) -> RLDrivingOnnxCheckpointManager.Config: + input_shapes = RLDrivingModel.input_shapes(model) + return RLDrivingOnnxCheckpointManager.Config( + keep_latest_k=0, + enable=True, + checkpoint_base_folder=base_folder, + export_onnx=True, + folder=folder, + interval=interval, + input_names=list(input_shapes), + input_shapes=[list(shape) for shape in input_shapes.values()], + input_dtypes=["float32"] * len(input_shapes), + ) diff --git a/torchtitan/experiments/rldriving/dataset.py b/torchtitan/experiments/rldriving/dataset.py new file mode 100644 index 00000000000..dc782ff06ea --- /dev/null +++ b/torchtitan/experiments/rldriving/dataset.py @@ -0,0 +1,158 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +import os +from collections.abc import Iterator +from dataclasses import dataclass, field +from typing import Any, Literal + +from xx.common.basedir import XX_BASEDIR +from xx.datasets.constants import BASE_DIR_GT +from xx.training.lib.dataloader import DataLoader +from xx.training.rldriving.config import DatasetConfig +from xx.training.rldriving.dataloader import get_dataset, RolloutContext + +import torch + +from torchtitan.components.dataloader import BaseDataLoader +from torchtitan.components.tokenizer import BaseTokenizer + + +class RLDrivingDataLoader(BaseDataLoader): + @dataclass(kw_only=True, slots=True) + class Config(BaseDataLoader.Config): + dataset: str + fps: int + training_id: str = "" + shuffle_size: int = 50_000 + min_mixing: float = 0.9 + num_writers: int = 1 + num_readers: int = 1 + limit: int | None = 500_000 + + codedir: str | None = XX_BASEDIR + pipeline_dir: str | None = BASE_DIR_GT + queue_priority: int = 5 + max_queue_size: int = 256 + max_fq_size: int = 8192 + + train_skip: int = 1 + epochs: int = 0 + steps_per_epoch: int = 1 + save_cache: bool = False + load_caches: list[str] = field(default_factory=list) + + zero_desire: bool = False + photo_noise_model: Literal["NONE", "VISION"] = "VISION" + pre_worldmodel_warmup_seconds: int = 7 + min_simulation_seconds: int = 6 + max_simulation_seconds: int = 7 + worldmodel_future_size_seconds: int = 1 + worldmodel_context_size_seconds: int = 2 + + def __init__( + self, + config: Config, + *, + dp_world_size: int, + dp_rank: int, + tokenizer: BaseTokenizer, + seq_len: int, + local_batch_size: int, + snapshot_every_n_steps: int | None = 1, + validation_steps: int = 1, + **kwargs: Any, + ) -> None: + del tokenizer, seq_len, snapshot_every_n_steps, validation_steps, kwargs + from gigashuffle import DataloaderConfig + + local_rank = int(os.environ.get("LOCAL_RANK", dp_rank)) + local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", dp_world_size)) + node_rank = int(os.environ.get("GROUP_RANK", dp_rank // local_world_size)) + + xx_config = DatasetConfig( + dataset=config.dataset_path or config.dataset, + training_id=config.training_id, + bs=local_batch_size, + nproc_per_node=local_world_size, + nnodes=dp_world_size // local_world_size, + node_rank=node_rank, + shuffle_size=str(config.shuffle_size), + min_mixing=config.min_mixing, + num_writers=config.num_writers, + num_readers=config.num_readers, + limit=config.limit, + codedir=config.codedir, + pipeline_dir=config.pipeline_dir, + queue_priority=config.queue_priority, + max_queue_size=config.max_queue_size, + max_fq_size=config.max_fq_size, + train_skip=config.train_skip, + epochs=config.epochs, + steps_per_epoch=config.steps_per_epoch, + save_cache=config.save_cache, + load_caches=list(config.load_caches), + fps=config.fps, + zero_desire=config.zero_desire, + photo_noise_model=config.photo_noise_model, + pre_worldmodel_warmup_seconds=config.pre_worldmodel_warmup_seconds, + min_simulation_seconds=config.min_simulation_seconds, + max_simulation_seconds=config.max_simulation_seconds, + worldmodel_future_size_seconds=config.worldmodel_future_size_seconds, + worldmodel_context_size_seconds=config.worldmodel_context_size_seconds, + ) + self.dataset = get_dataset(xx_config, local_rank=local_rank) + self._loader_config = DataloaderConfig( + bs=local_batch_size, + shuffle_size=config.shuffle_size, + min_mixing=config.min_mixing, + num_writers=config.num_writers, + num_readers=config.num_readers, + fill_once=False, + local_rank=local_rank, + global_rank=dp_rank, + local_world_size=local_world_size, + global_world_size=dp_world_size, + queue_name=f"{config.training_id or 'rldriving'}-train-node{node_rank}", + ) + self.loader: Any | None = None + self._iterator: Any | None = None + + # pyrefly: ignore [bad-override] + def __iter__( + self, + ) -> Iterator[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor], dict[str, torch.Tensor]]]: + if self.loader is None: + self.loader = DataLoader(self.dataset, self._loader_config) + iterator: Any = iter(self.loader) + self._iterator = iterator + try: + for inputs, targets, metadata in iterator: + yield inputs, targets, metadata + finally: + iterator.close() + if self._iterator is iterator: + self._iterator = None + + def attach_training_context(self, context: RolloutContext) -> None: + self.dataset.context = context + if self.loader is not None: + self.loader.attach_training_context(context) + + def close(self) -> None: + if self._iterator is not None: + self._iterator.close() + self._iterator = None + if self.loader is not None: + self.loader._shutdown_workers() + + def state_dict(self) -> dict[str, int]: + return {} + + def load_state_dict(self, state_dict: dict[str, int]) -> None: + return diff --git a/torchtitan/experiments/rldriving/loss.py b/torchtitan/experiments/rldriving/loss.py new file mode 100644 index 00000000000..91e2811be78 --- /dev/null +++ b/torchtitan/experiments/rldriving/loss.py @@ -0,0 +1,238 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass +from typing import cast + +import torch +import torch.nn as nn +import torch.nn.functional as F + +from torchtitan.components.loss import BaseLoss +from torchtitan.config import CompileConfig +from torchtitan.tools.logging import logger + + +# B: batch, A: action components. +ActorOutputs = dict[str, torch.Tensor] +ModelInputs = dict[str, torch.Tensor] +Targets = dict[str, torch.Tensor] +LossResult = tuple[torch.Tensor, dict[str, torch.Tensor]] + +ACTION_OUTPUT = "action" + + +def _sample_fixed_noise_policy( + action_pred_BA: torch.Tensor, + action_noise_A: torch.Tensor, +) -> torch.Tensor: + action_mean_BA = action_pred_BA[:, :2] + return action_mean_BA + torch.randn_like(action_mean_BA) * action_noise_A + + +def _critic_loss( + *, + next_actor_outputs: ActorOutputs, + targets: Targets, + online_critic: nn.Module, + target_critic: nn.Module, + current_inputs: ModelInputs, + next_inputs: ModelInputs, + action_noise_A: torch.Tensor, + gamma: float, +) -> LossResult: + action_reward_B = targets["action_reward"] + rollout_action_BA = action_reward_B[:, 0:2] + reward_B = action_reward_B[:, 2] + done_B = action_reward_B[:, 3] + + q1_rollout_B, q2_rollout_B = online_critic( + inputs=current_inputs, + action=rollout_action_BA, + ) + + with torch.no_grad(): + next_action_BA = _sample_fixed_noise_policy( + next_actor_outputs[ACTION_OUTPUT], + action_noise_A, + ) + q1_target_B, q2_target_B = target_critic( + inputs=next_inputs, + action=next_action_BA, + ) + bootstrap_B = torch.minimum(q1_target_B, q2_target_B) + target_B = reward_B + gamma * (1.0 - done_B) * bootstrap_B + q_target_abs_gap_B = torch.abs(q1_target_B - q2_target_B) + q_rollout_abs_gap_B = torch.abs(q1_rollout_B - q2_rollout_B) + q_target_clip_correction_B = gamma * (1.0 - done_B) * 0.5 * q_target_abs_gap_B + + critic_loss_B = 0.5 * ( + F.mse_loss(q1_rollout_B, target_B, reduction="none") + F.mse_loss(q2_rollout_B, target_B, reduction="none") + ) + metrics = { + "critic_loss": critic_loss_B.detach(), + "q1_rollout": q1_rollout_B.detach(), + "q2_rollout": q2_rollout_B.detach(), + "reward": reward_B.detach(), + "done": done_B.detach(), + "q_rollout_abs_gap": q_rollout_abs_gap_B.detach(), + "q_target_abs_gap": q_target_abs_gap_B.detach(), + "q_target_clip_correction": q_target_clip_correction_B.detach(), + } + return critic_loss_B, metrics + + +def _actor_loss( + *, + actor_outputs: ActorOutputs, + next_actor_outputs: ActorOutputs, + online_critic: nn.Module, + current_inputs: ModelInputs, + targets: Targets, + action_noise_A: torch.Tensor, + fps: float, + smooth_lat_cost: float, + smooth_long_cost: float, + curv_cost: float, + action_bound: float, + action_bound_loss_weight: float, +) -> LossResult: + action_pred_BA = actor_outputs[ACTION_OUTPUT] + sampled_action_BA = _sample_fixed_noise_policy(action_pred_BA, action_noise_A) + q1_new_B, q2_new_B = online_critic( + inputs=current_inputs, + action=sampled_action_BA, + ) + actor_pi_B = -torch.minimum(q1_new_B, q2_new_B) + actor_q_abs_gap_B = torch.abs(q1_new_B - q2_new_B) + + not_done_B = 1.0 - targets["action_reward"][:, 3] + curvature_B = action_pred_BA[:, 0] / targets["speed"].squeeze(-1).square() + curvature_loss_B = not_done_B * curv_cost * curvature_B.square() + + command_jerk_BA = (next_actor_outputs[ACTION_OUTPUT][:, :2] - action_pred_BA[:, :2]).abs() * fps + smooth_lat_B = not_done_B * smooth_lat_cost * command_jerk_BA[:, 0].square() + smooth_long_B = not_done_B * smooth_long_cost * command_jerk_BA[:, 1].square() + smooth_B = smooth_lat_B + smooth_long_B + actor_loss_B = actor_pi_B + curvature_loss_B + smooth_B + + action_abs_BA = torch.abs(action_pred_BA[..., :2]) + action_bound_excess_BA = torch.clamp(action_abs_BA - action_bound, min=0.0) + action_bound_loss_B = action_bound_excess_BA.square().mean(dim=-1) + loss_B = actor_loss_B + action_bound_loss_weight * action_bound_loss_B + + metrics = { + "loss": loss_B.detach(), + "actor_loss": actor_loss_B.detach(), + "actor_pi": actor_pi_B.detach(), + "actor_q_abs_gap": actor_q_abs_gap_B.detach(), + "actor_curv": (not_done_B * curvature_B).detach(), + "actor_curv_loss": curvature_loss_B.detach(), + "actor_cmd_lat_jerk": (not_done_B * command_jerk_BA[:, 0]).detach(), + "actor_cmd_long_jerk": (not_done_B * command_jerk_BA[:, 1]).detach(), + "actor_smooth_lat_loss": smooth_lat_B.detach(), + "actor_smooth_long_loss": smooth_long_B.detach(), + "actor_smooth_loss": smooth_B.detach(), + "actor_action_bound": action_bound_loss_B.detach(), + "actor_action_bound_max_abs": action_abs_BA.max(dim=-1).values.detach(), + "actor_action_bound_max_excess": action_bound_excess_BA.max(dim=-1).values.detach(), + } + return loss_B, metrics + + +class RLDrivingLoss(BaseLoss): + @dataclass(kw_only=True, slots=True) + class Config(BaseLoss.Config): + action_noise: tuple[float, float] + gamma: float + fps: float + smooth_lat_cost: float = 0.0 + smooth_long_cost: float = 0.0 + curv_cost: float = 0.0 + action_bound: float = 10.0 + action_bound_loss_weight: float = 1.0 + + def __init__( + self, + config: Config, + *, + compile_config: CompileConfig | None = None, + ) -> None: + self.action_noise_A = torch.tensor(config.action_noise) + self.gamma = config.gamma + self.fps = config.fps + self.smooth_lat_cost = config.smooth_lat_cost + self.smooth_long_cost = config.smooth_long_cost + self.curv_cost = config.curv_cost + self.action_bound = config.action_bound + self.action_bound_loss_weight = config.action_bound_loss_weight + + self.critic_fn = _critic_loss + self.actor_fn = _actor_loss + if compile_config is not None and compile_config.enable and "loss" in compile_config.components: + logger.info("Compiling the rldriving loss functions with torch.compile") + self.critic_fn = torch.compile( + self.critic_fn, + backend=compile_config.backend, + ) + self.actor_fn = torch.compile( + self.actor_fn, + backend=compile_config.backend, + ) + + self.fn = cast(Callable[..., torch.Tensor], self.actor_fn) + + def to(self, device: torch.device) -> RLDrivingLoss: + self.action_noise_A = self.action_noise_A.to(device) + return self + + def critic_loss( + self, + *, + next_actor_outputs: ActorOutputs, + targets: Targets, + online_critic: nn.Module, + target_critic: nn.Module, + current_inputs: ModelInputs, + next_inputs: ModelInputs, + ) -> LossResult: + return self.critic_fn( + next_actor_outputs=next_actor_outputs, + targets=targets, + online_critic=online_critic, + target_critic=target_critic, + current_inputs=current_inputs, + next_inputs=next_inputs, + action_noise_A=self.action_noise_A, + gamma=self.gamma, + ) + + def actor_loss( + self, + *, + actor_outputs: ActorOutputs, + next_actor_outputs: ActorOutputs, + online_critic: nn.Module, + current_inputs: ModelInputs, + targets: Targets, + ) -> LossResult: + return self.actor_fn( + actor_outputs=actor_outputs, + next_actor_outputs=next_actor_outputs, + online_critic=online_critic, + current_inputs=current_inputs, + targets=targets, + action_noise_A=self.action_noise_A, + fps=self.fps, + smooth_lat_cost=self.smooth_lat_cost, + smooth_long_cost=self.smooth_long_cost, + curv_cost=self.curv_cost, + action_bound=self.action_bound, + action_bound_loss_weight=self.action_bound_loss_weight, + ) diff --git a/torchtitan/experiments/rldriving/model.py b/torchtitan/experiments/rldriving/model.py new file mode 100644 index 00000000000..e8fc395ae95 --- /dev/null +++ b/torchtitan/experiments/rldriving/model.py @@ -0,0 +1,326 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +import copy +import math +from dataclasses import dataclass +from typing import Any, cast + +from xx.ml_tools.constants.model import ACTION_LEN, ModelInputs + +import torch +import torch.nn as nn +from torch.distributed.checkpoint.state_dict import get_model_state_dict, set_model_state_dict, StateDictOptions +from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy +from torch.utils.flop_counter import FlopCounterMode + +from torchtitan.config import CompileConfig, ParallelismConfig, TORCH_DTYPE_MAP, TrainingConfig +from torchtitan.distributed import ParallelDims +from torchtitan.distributed.activation_checkpoint import ActivationCheckpointingConfig +from torchtitan.distributed.fsdp import enable_fsdp_symm_mem, get_fsdp_reshard_after_forward_policy +from torchtitan.experiments.path.model import ( + Hydra, + LinearEncoder, + PathHead, + PathMLP, + TemporalPolicy, + TemporalSummarizer, +) +from torchtitan.experiments.path.model_config import TEMPORAL_HEADS, temporal_policy_config +from torchtitan.models.common import LayerNorm, Linear +from torchtitan.protocols.model import BaseModel +from torchtitan.protocols.module import Module +from torchtitan.tools.logging import logger + + +# B: batch, T: temporal steps, D: model width, A: action components. +ACTION_HEAD_NAME = "action" +Q_HEAD_NAME = "q" + +TemporalInputs = dict[str, torch.Tensor] +ActorOutputs = dict[str, torch.Tensor] + + +def _policy_forward(policy: TemporalPolicy, inputs: TemporalInputs) -> ActorOutputs: + outputs = policy( + inputs[ModelInputs.FEATURES], + inputs[ModelInputs.DESIRE], + inputs[ModelInputs.TRAFFIC], + inputs[ModelInputs.ACTION_T], + ) + return {ACTION_HEAD_NAME: outputs[ACTION_HEAD_NAME]} + + +class Critic(Module): + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + temporal_summarizer: TemporalSummarizer.Config + history_idxs: tuple[int, ...] + action_encoder: LinearEncoder.Config + post_action_mlp1: PathMLP.Config + post_action_mlp2: PathMLP.Config + q_hydra: Hydra.Config + + def __init__(self, config: Config): + super().__init__() + self.temporal_summarizer = config.temporal_summarizer.build() + self.history_idxs = config.history_idxs + self.action_encoder = config.action_encoder.build() + self.post_action_mlp1 = config.post_action_mlp1.build() + self.post_action_mlp2 = config.post_action_mlp2.build() + self.q_hydra = config.q_hydra.build() + + def forward(self, inputs: TemporalInputs, action: torch.Tensor) -> torch.Tensor: + features_BTD = inputs[ModelInputs.FEATURES] + dtype = features_BTD.dtype + critic_features_BD = self.temporal_summarizer( + features_BTD[:, self.history_idxs], + inputs[ModelInputs.DESIRE].to(dtype), + inputs[ModelInputs.TRAFFIC][:, -1].to(dtype), + inputs[ModelInputs.ACTION_T][:, -1].to(dtype), + ) + critic_features_BD = critic_features_BD + self.post_action_mlp1(critic_features_BD) + critic_features_BD = critic_features_BD + self.action_encoder(action.to(dtype)) + critic_features_BD = critic_features_BD + self.post_action_mlp2(critic_features_BD) + q_B1 = self.q_hydra(critic_features_BD)[Q_HEAD_NAME] + return q_B1.squeeze(-1).clone() + + +def actor_config() -> TemporalPolicy.Config: + action_heads = tuple(head for head in TEMPORAL_HEADS if head.name == ACTION_HEAD_NAME) + return temporal_policy_config(heads=action_heads, dropout=0.0, dense_training_outputs=False) + + +def critic_config(actor: TemporalPolicy.Config) -> Critic.Config: + dim = actor.temporal_summarizer.pos_embedding.embedding_dim + hidden = 256 * math.ceil(2 * dim / 256) + post_action_mlp = PathMLP.Config( + norm=LayerNorm.Config(normalized_shape=dim), + c_fc=Linear.Config(in_features=dim, out_features=hidden, bias=False), + c_proj=Linear.Config(in_features=hidden, out_features=dim, bias=False), + act="gelu_tanh", + dropout=0.0, + ) + return Critic.Config( + temporal_summarizer=copy.deepcopy(actor.temporal_summarizer), + history_idxs=actor.history_idxs, + action_encoder=LinearEncoder.Config( + in_layer=Linear.Config(in_features=ACTION_LEN, out_features=dim, bias=True), + out_layer=Linear.Config(in_features=dim, out_features=dim, bias=False), + ), + post_action_mlp1=post_action_mlp, + post_action_mlp2=copy.deepcopy(post_action_mlp), + q_hydra=Hydra.Config( + heads=(PathHead(name=Q_HEAD_NAME, output_size=1, mlp=False, scale=False),), + head_mlps={}, + final_layers={Q_HEAD_NAME: Linear.Config(in_features=dim, out_features=1, bias=True)}, + scale_layers={}, + ), + ) + + +class TwinCritic(nn.Module): + def __init__(self, config: Critic.Config): + super().__init__() + self.critic1 = config.build() + self.critic2 = config.build() + + def forward( + self, + inputs: TemporalInputs, + action: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + return self.critic1(inputs, action), self.critic2(inputs, action) + + +class RLDrivingModel(BaseModel): + @dataclass(kw_only=True, slots=True) + class Config(BaseModel.Config): + actor: TemporalPolicy.Config + critic: Critic.Config + + def update_from_config(self, *, config, **kwargs) -> None: + parallelism = config.parallelism + if parallelism.spmd_backend == "full_dtensor": + raise ValueError("rldriving does not support full DTensor") + unsupported = { + "tensor parallel": parallelism.tensor_parallel_degree, + "context parallel": parallelism.context_parallel_degree, + "pipeline parallel": parallelism.pipeline_parallel_degree, + "expert parallel": parallelism.expert_parallel_degree, + } + for name, degree in unsupported.items(): + if degree > 1: + raise ValueError(f"rldriving does not support {name}") + if config.activation_checkpoint is not None: + raise ValueError("rldriving does not support activation checkpointing") + + def get_nparams_and_flops(self, model: Module, seq_len: int) -> tuple[int, int]: + rldriving_model = cast(RLDrivingModel, model) + nparams = sum(parameter.numel() for parameter in rldriving_model.parameters()) + device = next(rldriving_model.parameters()).device + inputs = { + name: torch.zeros(shape, dtype=torch.float32, device=device) + for name, shape in RLDrivingModel.input_shapes(self).items() + } + action_dim = self.critic.action_encoder.in_layer.in_features + action_BA = torch.zeros((1, action_dim), dtype=torch.float32, device=device) + with torch.no_grad(), FlopCounterMode(display=False) as counter: + rldriving_model(inputs) + rldriving_model.critic(inputs, action_BA) + # MFU convention estimates backward as twice the counted forward work. + return nparams, 3 * counter.get_total_flops() + + def __init__(self, config: Config): + super().__init__() + self.config = config + self.actor = config.actor.build() + self.critic = TwinCritic(config.critic) + self.target_actor = config.actor.build() + self.target_critic = TwinCritic(config.critic) + + self.target_actor.requires_grad_(False).eval() + self.target_critic.requires_grad_(False).eval() + + @staticmethod + def input_shapes( + config: RLDrivingModel.Config, + batch_size: int = 1, + ) -> dict[str, tuple[int, ...]]: + summarizer = config.actor.temporal_summarizer + temporal_len = max(summarizer.desire_window_starts) + summarizer.desire_window_len + desire_dim = summarizer.desire_encoder.in_layer.in_features // summarizer.desire_window_len + return { + ModelInputs.FEATURES: ( + batch_size, + temporal_len, + summarizer.pos_embedding.embedding_dim, + ), + ModelInputs.DESIRE: (batch_size, temporal_len, desire_dim), + ModelInputs.TRAFFIC: ( + batch_size, + temporal_len, + summarizer.traffic_encoder.in_layer.in_features, + ), + ModelInputs.ACTION_T: ( + batch_size, + temporal_len, + summarizer.action_t_encoder.in_layer.in_features, + ), + } + + def verify_module_protocol(self) -> None: + # Path modules contain parameterless torch.nn activations and dropout. + pass + + def init_states(self, *, buffer_device: torch.device | None = None) -> None: + super().init_states(buffer_device=buffer_device) + self.sync_targets() + + @torch.no_grad() + def sync_targets(self) -> None: + _copy_model_state(self.actor, self.target_actor) + _copy_model_state(self.critic, self.target_critic) + + @torch.no_grad() + def warm_start_critics_from_actor(self) -> None: + for destination in ( + self.critic.critic1.temporal_summarizer, + self.critic.critic2.temporal_summarizer, + ): + _copy_model_state(self.actor.temporal_summarizer, destination) + self.sync_targets() + + def train(self, mode: bool = True) -> RLDrivingModel: + super().train(mode) + self.target_actor.eval() + self.target_critic.eval() + return self + + def forward(self, inputs: TemporalInputs) -> ActorOutputs: + return _policy_forward(self.actor, inputs) + + def target_forward(self, inputs: TemporalInputs) -> ActorOutputs: + return _policy_forward(self.target_actor, inputs) + + +def _copy_model_state(source: nn.Module, destination: nn.Module) -> None: + options = StateDictOptions(full_state_dict=True) + source_state = get_model_state_dict(source, options=options) + set_model_state_dict(destination, source_state, options=options) + + +def parallelize_rldriving( + model: RLDrivingModel, + *, + parallel_dims: ParallelDims, + training: TrainingConfig, + parallelism: ParallelismConfig, + compile_config: CompileConfig, + ac_config: ActivationCheckpointingConfig, + dump_folder: str, +) -> RLDrivingModel: + if compile_config.enable and "model" in compile_config.components: + torch._dynamo.config.capture_scalar_outputs = True + model.actor.compile(backend=compile_config.backend) + model.critic.compile(backend=compile_config.backend) + model.target_actor.compile(backend=compile_config.backend) + model.target_critic.compile(backend=compile_config.backend) + logger.info("Compiling rldriving model components with torch.compile") + + names = ["dp_replicate", "fsdp"] if parallel_dims.dp_replicate_enabled else ["fsdp"] + mp_policy = MixedPrecisionPolicy( + param_dtype=TORCH_DTYPE_MAP[training.mixed_precision_param], + reduce_dtype=TORCH_DTYPE_MAP[training.mixed_precision_reduce], + cast_forward_inputs=True, + ) + fsdp_config: dict[str, Any] = {"mesh": parallel_dims.get_mesh(names), "mp_policy": mp_policy} + if training.enable_cpu_offload: + fsdp_config["offload_policy"] = CPUOffloadPolicy() + reshard_after_forward = get_fsdp_reshard_after_forward_policy( + parallelism.fsdp_reshard_after_forward, + parallel_dims.pp_enabled, + ) + + def shard(module: nn.Module, reshard: bool = reshard_after_forward) -> None: + fully_shard(module, **fsdp_config, reshard_after_forward=reshard) + + for temporal_summarizer in ( + model.actor.temporal_summarizer, + model.critic.critic1.temporal_summarizer, + model.critic.critic2.temporal_summarizer, + model.target_actor.temporal_summarizer, + model.target_critic.critic1.temporal_summarizer, + model.target_critic.critic2.temporal_summarizer, + ): + temporal_summarizer.transformer.apply_fsdp(shard, reshard_after_forward) + + shard(model.actor) + shard(model.target_actor) + for critic in ( + model.critic.critic1, + model.critic.critic2, + model.target_critic.critic1, + model.target_critic.critic2, + ): + shard(critic) + shard(model.critic) + shard(model.target_critic) + fully_shard(model, **fsdp_config) + + if parallelism.enable_fsdp_symm_mem: + enable_fsdp_symm_mem(model) + + logger.info( + "Applied HSDP to the rldriving model" + if parallel_dims.dp_replicate_enabled + else "Applied FSDP to the rldriving model" + ) + if training.enable_cpu_offload: + logger.info("Applied CPU Offloading to the rldriving model") + return model diff --git a/torchtitan/experiments/rldriving/onnx_checkpoint.py b/torchtitan/experiments/rldriving/onnx_checkpoint.py new file mode 100644 index 00000000000..c475d87436d --- /dev/null +++ b/torchtitan/experiments/rldriving/onnx_checkpoint.py @@ -0,0 +1,50 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +from dataclasses import dataclass +from typing import cast + +import torch +import torch.nn as nn + +from torchtitan.components.checkpoint import OPTIMIZER +from torchtitan.components.onnx_checkpoint import OnnxCheckpointManager + +from .model import ACTION_HEAD_NAME, RLDrivingModel + + +class _TargetActorOnnxModel(nn.Module): + def __init__(self, model: RLDrivingModel) -> None: + super().__init__() + self.model = model + + def forward(self, inputs: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: + return self.model.target_forward(inputs) + + +class RLDrivingOnnxCheckpointManager(OnnxCheckpointManager): + @dataclass(kw_only=True, slots=True) + class Config(OnnxCheckpointManager.Config): + checkpoint_base_folder: str = "" + + def __init__(self, config: Config, **kwargs) -> None: + if config.checkpoint_base_folder: + kwargs["base_folder"] = config.checkpoint_base_folder + super().__init__(config, **kwargs) + if self.enable: + self.states.pop(OPTIMIZER) + + def _export_onnx(self, model: nn.Module, path: str) -> None: + model = cast(RLDrivingModel, model) + inputs = dict(zip(self.input_names, self._build_onnx_inputs())) + self._export_one( + _TargetActorOnnxModel(model).eval(), + inputs, + path, + output_names=[ACTION_HEAD_NAME], + ) diff --git a/torchtitan/experiments/rldriving/supercombo.py b/torchtitan/experiments/rldriving/supercombo.py new file mode 100644 index 00000000000..5b3fe018190 --- /dev/null +++ b/torchtitan/experiments/rldriving/supercombo.py @@ -0,0 +1,79 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from types import MethodType +from xx.ml_tools.constants.model import ModelInputs + +import torch + +from torchtitan.experiments.path.model import PathSelfAttention +from torchtitan.experiments.path.model_config import model_config +from .model import actor_config + + +VISION_OUTPUT_ORDER = tuple( + "lane_lines lane_lines_prob road_edges meta desire_pred pose wide_from_device_euler road_transform".split() +) +OFF_POLICY_OUTPUT_ORDER = ("plan", "lead", "lead_prob", "desire_state") +ON_POLICY_OUTPUT_ORDER = ("action",) +OUTPUT_ORDER = (*VISION_OUTPUT_ORDER, *OFF_POLICY_OUTPUT_ORDER, *ON_POLICY_OUTPUT_ORDER, "hidden_state") + + +def _naive_attention(self: PathSelfAttention, x: torch.Tensor) -> torch.Tensor: + b, t, _ = x.shape + qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) + q, k, v = (value.squeeze(0) for value in qkv.permute(2, 0, 3, 1, 4).split(1)) + q, k = self.q_norm(q), self.k_norm(k) + scores = (q @ k.transpose(-2, -1)) * self.head_dim**-0.5 + x = (scores.masked_fill(~self._supercombo_mask, float("-inf")).softmax(-1) @ v).transpose(1, 2) + return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) + + +# some micro optimizations can be made but not worth it for now +# desire is never -1 in runtime, se we can drop the unknown_desire_embedding code path +# there are some Unsqueeze -> Gather that can be bipassed (traffic_convention, action_t) +class Supercombo(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + config = model_config() + config.temporal_policy.temporal_summarizer.dense_training_outputs = False + self.vision = config.vision.build() + self.point_policy = config.point_policy.build() + self.off_policy = config.temporal_policy.build() + self.on_policy = actor_config().build() + output_size = self.vision.config.vision_features + sum( + hydra.final_layer[name].out_features + for hydra, names in ( + (self.point_policy.hydra, VISION_OUTPUT_ORDER), + (self.off_policy.temporal_hydra, OFF_POLICY_OUTPUT_ORDER), + (self.on_policy.temporal_hydra, ON_POLICY_OUTPUT_ORDER), + ) + for name in names + ) + self.register_buffer("pad", torch.zeros(1, -output_size % 4), persistent=False) + for policy in (self.off_policy, self.on_policy): + for layer in policy.temporal_summarizer.transformer.layers: + attention = layer.attention + mask = torch.ones( + 1, 1, policy.temporal_summarizer.block_size, policy.temporal_summarizer.block_size, dtype=torch.bool + ) + attention.register_buffer("_supercombo_mask", mask.tril(), persistent=False) + attention.forward = MethodType(_naive_attention, attention) + + def forward(self, inputs: dict[str, torch.Tensor]) -> torch.Tensor: + current = self.vision(inputs) + features = torch.cat((inputs["features_buffer"], current[:, None]), dim=1) + outputs = self.point_policy(current) + for policy, names in ((self.off_policy, OFF_POLICY_OUTPUT_ORDER), (self.on_policy, ON_POLICY_OUTPUT_ORDER)): + policy_outputs = policy( + features, + inputs[ModelInputs.DESIRE], + inputs[ModelInputs.TRAFFIC][:, None], + inputs[ModelInputs.ACTION_T][:, None], + ) + outputs.update({name: policy_outputs[name] for name in names}) + outputs["hidden_state"] = current.detach() + return torch.cat([outputs[name] for name in OUTPUT_ORDER] + [self.pad], dim=1) diff --git a/torchtitan/experiments/rldriving/trainer.py b/torchtitan/experiments/rldriving/trainer.py new file mode 100644 index 00000000000..fe7ab16a7d4 --- /dev/null +++ b/torchtitan/experiments/rldriving/trainer.py @@ -0,0 +1,416 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +import os +import time +from collections.abc import Iterable, Iterator +from dataclasses import dataclass +from datetime import timedelta +from functools import cache +from typing import Any, cast, Literal + +from xx.common.helpers import parse_info +from xx.ml_tools.constants.model import TEMPORAL_INPUTS +from xx.training.lib.checkpoint import Checkpoint +from xx.training.rldriving.dataloader import RolloutContext + +import torch +import torch.distributed as dist +import torch.distributed.checkpoint as dcp +import torch.nn as nn +from torch.distributed.checkpoint._fsspec_filesystem import FsspecReader +from torch.distributed.elastic.multiprocessing.errors import record +from torch.optim.lr_scheduler import LambdaLR + +from torchtitan.components.dataloader import DataloaderExhaustedError +from torchtitan.components.lr_scheduler import LRSchedulersContainer +from torchtitan.components.optimizer import OptimizersContainer +from torchtitan.components.unique_counter import StringUniqueCounter +from torchtitan.distributed import utils as dist_utils +from torchtitan.observability import structured_logger as sl +from torchtitan.tools.logging import logger +from torchtitan.trainer import Trainer + +from .dataset import RLDrivingDataLoader +from .loss import RLDrivingLoss +from .model import RLDrivingModel +from .onnx_checkpoint import RLDrivingOnnxCheckpointManager +from .validate import RLDrivingValidator + + +Batch = tuple[ + dict[str, torch.Tensor], + dict[str, torch.Tensor], + dict[str, torch.Tensor], +] +PreparedBatch = tuple[ + dict[str, torch.Tensor], + dict[str, torch.Tensor], + dict[str, torch.Tensor], + dict[str, torch.Tensor], +] + + +_get_path_checkpoint = cache(Checkpoint) + + +@dataclass(kw_only=True, slots=True) +class RLDrivingLRSchedulersConfig(LRSchedulersContainer.Config): + steps_per_epoch: int + num_epochs: int + actor_delay_epochs: float = 15.0 + actor_warmup_fraction: float = 0.1 + cooldown_fraction: float = 0.4 + min_lr_factor: float = 0.05 + critic_switch_epoch: float = 15.0 + critic_second_lr: float = 4e-5 + + # pyrefly: ignore [bad-override] + def build(self, *, optimizers, training_steps): + return RLDrivingLRSchedulers(self, optimizers=optimizers) + + +class RLDrivingLRSchedulers(LRSchedulersContainer): + Config = RLDrivingLRSchedulersConfig + + def __init__( + self, + config: Config, + *, + optimizers: OptimizersContainer, + ) -> None: + self.config = config + self.optimizer_container = optimizers + self.optimizer = next(iter(optimizers)) + self.schedulers = [ + LambdaLR( + self.optimizer, + [self._lr_lambda(group) for group in self.optimizer.param_groups], + ) + ] + + def _lr_lambda(self, group): + phase = group["param_names"][0].split(".", 1)[0] + base_lr = float(group["lr"]) + + def lr_lambda(current_step: int) -> float: + config = self.config + epoch = current_step / config.steps_per_epoch + max_epoch = config.num_epochs - 1.0 + cooldown_start = max_epoch * (1.0 - config.cooldown_fraction) + if phase == "actor": + if epoch < config.actor_delay_epochs: + return 0.0 + warmup_end = config.actor_delay_epochs + max_epoch * config.actor_warmup_fraction + if epoch < warmup_end: + return (epoch - config.actor_delay_epochs) / (max_epoch * config.actor_warmup_fraction) + if epoch < cooldown_start: + return 1.0 + progress = min(1.0, (epoch - cooldown_start) / (max_epoch - cooldown_start)) + return 1.0 + progress * (config.min_lr_factor - 1.0) + + if epoch < config.critic_switch_epoch: + return 1.0 + lr = config.critic_second_lr + if epoch >= cooldown_start: + progress = min(1.0, (epoch - cooldown_start) / (max_epoch - cooldown_start)) + lr *= 1.0 + progress * (config.min_lr_factor - 1.0) + return lr / base_lr + + return lr_lambda + + def step_phase(self, phase: Literal["actor", "critic"]) -> None: + all_param_groups = self.optimizer.param_groups + self.optimizer.param_groups = [ + group for group in all_param_groups if group["param_names"][0].startswith(f"{phase}.") + ] + self.optimizer_container.step() + self.optimizer.param_groups = all_param_groups + + +class RLDrivingTrainer(Trainer): + @dataclass(kw_only=True, slots=True) + class Config(Trainer.Config): + loss: RLDrivingLoss.Config # pyrefly: ignore [bad-override] + dataloader: RLDrivingDataLoader.Config # pyrefly: ignore [bad-override] + checkpoint: RLDrivingOnnxCheckpointManager.Config # pyrefly: ignore [bad-override] + lr_scheduler: RLDrivingLRSchedulers.Config # pyrefly: ignore [bad-override] + validator: RLDrivingValidator.Config # pyrefly: ignore [bad-override] + warm_start_checkpoint: str + steps_per_epoch: int + train_step_barrier_timeout_seconds: int + ema_tau: float + fps: int + + def __post_init__(self) -> None: + Trainer.Config.__post_init__(self) + if self.codedir: + self.dataloader.codedir = self.codedir + self.validator.miniray = {**self.validator.miniray, "codedir": self.codedir} + if self.steps_per_epoch != self.dataloader.steps_per_epoch: + raise ValueError("trainer and dataloader steps_per_epoch must match") + if self.ema_tau < 1.0: + raise ValueError("ema_tau must be at least 1") + + config: Config # pyrefly: ignore [bad-override] + loss_fn: RLDrivingLoss # pyrefly: ignore [bad-override] + dataloader: RLDrivingDataLoader # pyrefly: ignore [bad-override] + lr_schedulers: RLDrivingLRSchedulers # pyrefly: ignore [bad-override] + validator: RLDrivingValidator # pyrefly: ignore [bad-override] + + def __init__(self, config: Config): + super().__init__(config) + if self.gradient_accumulation_steps != 1: + raise ValueError("rldriving does not support gradient accumulation") + self.train_step_barrier_group: dist.ProcessGroup | None = None + if dist.get_world_size() > 1: + self.train_step_barrier_group = dist.new_group( + backend="gloo", timeout=timedelta(seconds=config.train_step_barrier_timeout_seconds) + ) + training_id = os.getenv("REPORTERV2_TRAINING_ID") or "local" + self.unique_segment_counter = StringUniqueCounter(f"unique_ids:{training_id}:rldriving:train") + self.loss_fn.to(self.device) + self.model = cast(RLDrivingModel, self.model_parts[0]) + dcp.load( + {"temporal_policy": self.model.actor}, + storage_reader=FsspecReader(_get_path_checkpoint(config.warm_start_checkpoint).url_or_file()), + ) + self.model.warm_start_critics_from_actor() + + # pyrefly: ignore [bad-override] + def batch_generator(self, data_iterable: Iterable[Batch]) -> Iterator[Batch]: + data_iterator = iter(data_iterable) + while True: + data_load_start = time.perf_counter() + try: + batch = next(data_iterator) + except StopIteration as ex: + raise DataloaderExhaustedError() from ex + batch_size = next(iter(batch[0].values())).shape[0] + self.metrics_processor.ntokens_since_last_log += batch_size + self.metrics_processor.data_loading_times.append(time.perf_counter() - data_load_start) + yield batch + + def prepare_batch(self, batch: Batch) -> PreparedBatch: + inputs, targets, metadata = batch + inputs = {name: value.to(self.device) for name, value in inputs.items()} + targets = {name: value.to(self.device) for name, value in targets.items()} + metadata = {name: value.to(self.device) for name, value in metadata.items()} + current_inputs = {name: inputs[name].float() for name in TEMPORAL_INPUTS} + next_inputs = { + name: torch.cat((inputs[name][:, 1:], inputs[f"next_{name}"]), dim=1).float() for name in current_inputs + } + return current_inputs, next_inputs, targets, metadata + + # pyrefly: ignore [bad-override] + def train_step(self, data_iterator: Iterator[Batch]) -> None: + steps_per_epoch = self.config.steps_per_epoch + rollout_epoch = ((self.step - 1) // steps_per_epoch) * steps_per_epoch + 1 + self.dataloader.attach_training_context(RolloutContext(epoch=rollout_epoch)) + batch = next(data_iterator) + info = batch[0].get("info") + if info is not None: + self.unique_segment_counter.update(parse_info(value)["name"] for value in info.cpu().numpy()) + current_inputs, next_inputs, targets, metadata = self.prepare_batch(batch) + batch_size = next(iter(current_inputs.values())).shape[0] + self.ntokens_seen += batch_size + local_samples = torch.tensor(batch_size, dtype=torch.float32, device=self.device) + + lr_metrics = self.lr_schedulers.get_metrics() + metric_sums: dict[str, torch.Tensor] = {} + self.optimizers.zero_grad() + if self.train_step_barrier_group is not None: + dist.barrier(group=self.train_step_barrier_group) + with self.train_context(): + actor_outputs = self.model(current_inputs) + next_actor_outputs = self.model(next_inputs) + actor_loss_B, actor_metrics = self.loss_fn.actor_loss( + actor_outputs=actor_outputs, + next_actor_outputs=next_actor_outputs, + online_critic=self.model.critic, + current_inputs=current_inputs, + targets=targets, + ) + actor_loss = actor_loss_B.sum() / local_samples + actor_loss.backward() + actor_loss = actor_loss.detach() + self._accumulate_metrics(metric_sums, actor_metrics) + del actor_outputs, next_actor_outputs, actor_loss_B, actor_metrics + actor_grad_norm = self._clip_phase_grad_norm(self.model.actor) + self.checkpointer.maybe_wait_for_staging() + self.lr_schedulers.step_phase("actor") + + self.optimizers.zero_grad() + with self.train_context(): + with torch.no_grad(): + next_actor_outputs = self.model.target_forward(next_inputs) + critic_loss_B, critic_metrics = self.loss_fn.critic_loss( + next_actor_outputs=next_actor_outputs, + targets=targets, + online_critic=self.model.critic, + target_critic=self.model.target_critic, + current_inputs=current_inputs, + next_inputs=next_inputs, + ) + critic_loss = critic_loss_B.sum() / local_samples + critic_loss.backward() + critic_loss = critic_loss.detach() + self._accumulate_metrics(metric_sums, critic_metrics) + del next_actor_outputs, critic_loss_B, critic_metrics + critic_grad_norm = self._clip_phase_grad_norm(self.model.critic) + self.lr_schedulers.step_phase("critic") + self.optimizers.zero_grad() + + with torch.no_grad(): + decay = 1.0 - 1.0 / self.config.ema_tau + for online, target in ( + (self.model.actor, self.model.target_actor), + (self.model.critic, self.model.target_critic), + ): + for online_param, target_param in zip(online.parameters(), target.parameters()): + target_param.mul_(decay).add_(online_param, alpha=1.0 - decay) + for online_buffer, target_buffer in zip(online.buffers(), target.buffers()): + target_buffer.copy_(online_buffer) + self.lr_schedulers.step() + + if self.step == 0 or not self.metrics_processor.should_log(self.step): + return + + loss = actor_loss + critic_loss + loss_mesh = self.parallel_dims.get_optional_mesh("loss") + if loss_mesh is not None: + global_samples = float(dist_utils.dist_sum(local_samples, loss_mesh)) + global_avg_loss = dist_utils.dist_sum(loss * local_samples, loss_mesh) / global_samples + global_max_loss = dist_utils.dist_max(loss, loss_mesh) + global_samples_seen = dist_utils.dist_sum( + torch.tensor(self.ntokens_seen, dtype=torch.int64, device=self.device), + loss_mesh, + ) + metric_averages = { + name: dist_utils.dist_sum(value, loss_mesh) / global_samples for name, value in metric_sums.items() + } + else: + global_avg_loss = global_max_loss = float(loss.item()) + global_samples_seen = self.ntokens_seen + metric_averages = {name: float(value.item()) / batch_size for name, value in metric_sums.items()} + + batch_mesh = self.parallel_dims.get_optional_mesh("batch") + unique_segments_seen = ( + self.unique_segment_counter.global_count(batch_mesh.get_group()) + if batch_mesh is not None + else self.unique_segment_counter.local_count() + ) + + metadata_averages = {} + for name, value in metadata.items(): + value = value.float() + finite = torch.isfinite(value) + value_sum = torch.where(finite, value, 0.0).sum() + value_count = finite.sum() + if loss_mesh is not None: + total = dist_utils.dist_sum(value_sum, loss_mesh) + count = dist_utils.dist_sum(value_count, loss_mesh) + else: + total = float(value_sum.item()) + count = float(value_count.item()) + metadata_averages[name] = total / count if count else float("nan") + + self.metrics_processor.log( + self.step, + global_avg_loss, + global_max_loss, + float(torch.maximum(actor_grad_norm, critic_grad_norm).item()), + extra_metrics={ + "n_samples_seen": global_samples_seen, + "actor_grad_norm": float(actor_grad_norm.item()), + "critic_grad_norm": float(critic_grad_norm.item()), + "dataset/unique_segments_seen": unique_segments_seen, + **lr_metrics, + **{f"rldriving/{name}": value for name, value in metric_averages.items()}, + **{f"sim/{name}": value for name, value in metadata_averages.items()}, + }, + ) + + def _clip_phase_grad_norm(self, module: nn.Module) -> torch.Tensor: + return dist_utils.clip_grad_norm_( + module.parameters(), + self.config.training.max_norm, + foreach=True, + ) + + @staticmethod + def _accumulate_metrics( + sums: dict[str, torch.Tensor], + metrics: dict[str, torch.Tensor], + ) -> None: + for name, value in metrics.items(): + sums[name] = sums.get(name, torch.zeros((), device=value.device)) + value.float().sum() + + @record + def train(self) -> None: + config = self.config + sl.log_trace_instant("training_start") + loaded = self.checkpointer.load(step=config.checkpoint.load_step) + if not loaded: + self.checkpointer.save(0) + self.set_runtime_seed() + loaded_step = self.step + logger.info(f"Training starts at step {self.step + 1}") + + with config.profiler.build( + global_step=self.step, + base_folder=config.dump_folder, + ) as profiler: + data_iterator = self.batch_generator(self.dataloader) + while self.should_continue_training(): + self.step += 1 + sl.set_step(self.step, relative_step=self.step - loaded_step) + with sl.log_trace_span("step"): + self.gc_handler.run(self.step) + try: + self.train_step(data_iterator) + except DataloaderExhaustedError: + logger.warning("Ran out of data; last step was canceled.") + break + self.checkpointer.save( + self.step, + last_step=(self.step == config.training.steps), + ) + if config.validator.enable and self.validator.should_validate(self.step): + self.validator.validate(self.model_parts, self.step) + profiler.step() + if self.step - loaded_step == 1: + dist_utils.set_pg_timeouts( + timeout=timedelta(seconds=config.comm.train_timeout_seconds), + parallel_dims=self.parallel_dims, + ) + + if torch.distributed.get_rank() == 0: + logger.info("Sleeping 2 seconds for other ranks to complete") + time.sleep(2) + logger.info("Training completed") + + def close(self) -> None: + self.dataloader.close() + if self.config.validator.enable: + self.validator.close() + super().close() + if self.train_step_barrier_group is not None: + dist.destroy_process_group(self.train_step_barrier_group) + self.train_step_barrier_group = None + + def state_dict(self) -> dict[str, Any]: + return { + **super().state_dict(), + "unique_segment_counter": self.unique_segment_counter.state_dict(), + } + + def load_state_dict(self, state_dict: dict[str, Any]) -> None: + super().load_state_dict(state_dict) + if "unique_segment_counter" in state_dict: + self.unique_segment_counter.load_state_dict(state_dict["unique_segment_counter"]) diff --git a/torchtitan/experiments/rldriving/validate.py b/torchtitan/experiments/rldriving/validate.py new file mode 100644 index 00000000000..f55cf6b2f78 --- /dev/null +++ b/torchtitan/experiments/rldriving/validate.py @@ -0,0 +1,90 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from __future__ import annotations + +import os +from dataclasses import dataclass, field +from typing import Any +from xx.release_tests.lib.base_report import ReportFormat +from xx.training.lib.checkpoint import wait_for_checkpoint +from xx.training.rldriving.test import MODEL_REPORTS, SCALAR_REPORTS + +import torch.distributed as dist +import torch.nn as nn + +from torchtitan.components.metrics import MetricsProcessor +from torchtitan.components.report_runner import ReportRunner, ReportSpec +from torchtitan.components.validate import BaseValidator + + +class RLDrivingValidator(BaseValidator): + @dataclass(kw_only=True, slots=True) + class Config(BaseValidator.Config): + enable: bool + fps: int + reports: dict[str, list[int]] = field(default_factory=dict) + miniray: dict[str, Any] = field(default_factory=dict) + + config: Config # pyrefly: ignore [bad-override] + + def __init__( + self, + config: Config, + *, + metrics_processor: MetricsProcessor, + **kwargs: Any, + ) -> None: + del kwargs + super().__init__(config=config) + self.metrics_processor = metrics_processor + self.training_id = os.getenv("REPORTERV2_TRAINING_ID") or "local" + self.miniray = { + **config.miniray, + "job_group": f"rldriving_validation_{self.training_id}", + } + self.report_runner = ReportRunner(metrics_processor=metrics_processor, enabled=dist.get_rank() == 0) + + def validate(self, model_parts: list[nn.Module], step: int) -> None: + del model_parts + self._submit_reports(step) + + def close(self) -> None: + self.report_runner.close() + + def _submit_reports(self, step: int) -> None: + current_checkpoint = f"{self.training_id}/{step}" + + def _run_report(TestCls: type, test_config: Any, include_scalars: bool) -> tuple[Any, ...]: + wait_for_checkpoint(current_checkpoint) + test = TestCls(test_config) + html = test.run_report() + return (html, test.scalars) if include_scalars else (html,) + + report_specs = {} + for report_name, (TestCls, ReportConfigCls) in MODEL_REPORTS.items(): + include_scalars = report_name in SCALAR_REPORTS + report_config = ReportConfigCls( + rollout={ + "agent": { + "supercombo": current_checkpoint, + "model_trained_fps": self.config.fps, + } + }, + report_name=f"rldriving_{report_name}", + save_tmp=False, + format=ReportFormat.HTML, + miniray=self.miniray, + ) + report_specs[report_name] = ReportSpec( + output_names=(report_name, f"{report_name}_scalar") if include_scalars else (report_name,), + output_types=("html", "scalar") if include_scalars else ("html",), + steps=self.config.reports.get(report_name, []), + func=_run_report, + arguments=[TestCls, report_config, include_scalars], + ) + + self.report_runner.submit_due(step=step, report_specs=report_specs)