From c1072c59fea3ac88b1881ded49a80df12bcea8f8 Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Sat, 15 Aug 2026 19:54:46 -0700 Subject: [PATCH 1/9] path comma 1M dataset --- .../experiments/path/comma1m_dataset.py | 338 ++++++++++++++++++ .../experiments/path/config_registry.py | 22 ++ torchtitan/experiments/path/dataset.py | 56 +-- 3 files changed, 391 insertions(+), 25 deletions(-) create mode 100644 torchtitan/experiments/path/comma1m_dataset.py diff --git a/torchtitan/experiments/path/comma1m_dataset.py b/torchtitan/experiments/path/comma1m_dataset.py new file mode 100644 index 00000000000..774b4a102ee --- /dev/null +++ b/torchtitan/experiments/path/comma1m_dataset.py @@ -0,0 +1,338 @@ +# 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 random +from collections import defaultdict +from typing import Any + +import cv2 +import fsspec +import numpy as np +import PyNvVideoCodec +import reverse_geocoder + +from openpilot.cereal import log +from openpilot.common.transformations.camera import denormalize, DEVICE_CAMERAS, view_frame_from_device_frame +from openpilot.common.transformations.coordinates import ecef2geodetic +from openpilot.common.transformations.model import ( + BIGMODEL_INPUT_SIZE, + bigmodel_intrinsics, + MEDMODEL_INPUT_SIZE, + medmodel_intrinsics, + SBIGMODEL_INPUT_SIZE, + sbigmodel_intrinsics, +) +from openpilot.common.transformations.orientation import euler_from_rot, rot_from_euler, rot_from_quat +from openpilot.selfdrive.controls.lib.drive_helpers import get_accel_from_plan, get_curvature_from_plan, MIN_SPEED +from openpilot.selfdrive.modeld.constants import ModelConstants, Plan +from openpilot.system.loggerd.config import CAMERA_FPS +from scipy.ndimage import gaussian_filter1d +from torch.utils.data import get_worker_info + + +COMMA1M_REPO_ID = "commaai/comma1M" +T_IDXS = np.asarray(ModelConstants.T_IDXS) +W, H = MEDMODEL_INPUT_SIZE +BIG_W, BIG_H = BIGMODEL_INPUT_SIZE +SBIG_W, SBIG_H = SBIGMODEL_INPUT_SIZE +BIG_MODEL_CORNERS = np.asarray([[0, 0, 1], [BIG_W, 0, 1], [0, BIG_H, 1], [BIG_W, BIG_H, 1]]) +CAMERA_BY_DEVICE_TYPE = { + log.InitData.DeviceType.neo: ("neo", "unknown"), + log.InitData.DeviceType.tici: ("tici", "ar0231"), + log.InitData.DeviceType.tizi: ("tizi", "ar0231"), + log.InitData.DeviceType.mici: ("mici", "os04c10"), +} +ROAD_POINTS = np.asarray( + [[-1.0, 1.22, 100.0], [1.0, 1.22, 100.0], [-1.0, 1.22, 200.0], [1.0, 1.22, 200.0]], + dtype=np.float32, +) +LHT_COUNTRIES = frozenset( + "AG AI AU BB BD BM BN BS BT BW CC CK CX CY DM FJ FK GB GD GG GY HK ID IE IM IN JE JM JP KE KI KN KY LC " + "LK LS MO MS MT MU MV MW MY MZ NA NF NP NR NU NZ PG PK PN SB SC SG SH SR SZ TC TH TK TL TO TT TV TZ UG " + "VC VG VI WS ZA ZM ZW".split() +) + + +class Comma1MDataset: + """Iterable over public comma1M segments.""" + + def __init__( + self, + config: Any, + val: bool, + local_rank: int, + global_rank: int, + global_world_size: int, + ) -> None: + assert config.plan_only + assert not config.rgb + assert not config.unvision + assert config.dataset_path is not None + assert not config.deterministic_fidxs + + self.config = config + self.val = val + self.local_rank = local_rank + self.global_world_size = global_world_size + self.fs, root = fsspec.url_to_fs(config.dataset_path) + self.data_dir = f"{root.rstrip('/')}/data" + segments = sorted(path.rsplit("/", 2)[-2] for path in self.fs.glob(f"{self.data_dir}/*/fcamera.hevc"))[ + : config.limit + ] + segments = [segment for segment in segments if (hash(int(segment, 16)) % 10 == 0) == val] + self.segments = segments[global_rank::global_world_size] + + def __iter__(self): + worker = get_worker_info() + segments = self.segments + if worker is not None: + writer_count = worker.num_workers // self.global_world_size + segments = segments[worker.id % writer_count :: writer_count] + if not segments: + raise ValueError("comma1M split is empty; increase limit or reduce ranks/writers") + + while True: + order = segments.copy() + random.shuffle(order) + for segment in order: + yield _load_segment(self.fs, f"{self.data_dir}/{segment}", self.config, self.val, self.local_rank) + + +def _decode(fs, path: str, index: np.ndarray, wanted: set[int], gpu_id: int) -> dict[int, np.ndarray]: + if not wanted: + return {} + + iframe_indexes = np.flatnonzero(index[:-1, 0] == 2) + next_iframe = np.searchsorted(iframe_indexes, max(wanted), side="right") + read_end = int(iframe_indexes[next_iframe]) if next_iframe < len(iframe_indexes) else len(index) - 1 + payload = fs.cat_file(path, end=int(index[read_end, 1])) + + decoder = PyNvVideoCodec.CreateDecoder( + gpuid=gpu_id, + codec=PyNvVideoCodec.cudaVideoCodec.HEVC, + usedevicememory=False, + outputColorType=PyNvVideoCodec.OutputColorType.RGB, + ) + decoded = {} + frame_idx = 0 + for packet in (payload, b""): + packet_array = np.frombuffer(packet, dtype=np.uint8) + packet_data = PyNvVideoCodec.PacketData() + packet_data.bsl = len(packet) + packet_data.bsl_data = packet_array.ctypes.data + for frame in decoder.Decode(packet_data): + if frame_idx in wanted: + decoded[frame_idx] = np.from_dlpack(frame).copy() + frame_idx += 1 + return decoded + + +def _transform_matrix(from_intrinsics: np.ndarray, to_intrinsics: np.ndarray, eulers: np.ndarray) -> np.ndarray: + after = ROAD_POINTS @ rot_from_euler(eulers) + before_pixels = denormalize(ROAD_POINTS[:, :2] / ROAD_POINTS[:, 2, None], intrinsics=from_intrinsics) + after_pixels = denormalize(after[:, :2] / after[:, 2, None], intrinsics=to_intrinsics) + return cv2.getPerspectiveTransform(before_pixels.astype(np.float32), after_pixels.astype(np.float32)) + + +def _pack_image(rgb: np.ndarray) -> np.ndarray: + height, width = rgb.shape[:2] + yuv = cv2.cvtColor(rgb, cv2.COLOR_RGB2YUV_I420) + y = yuv[:height].reshape(height // 2, 2, width // 2, 2).transpose(3, 1, 0, 2).reshape(4, height // 2, width // 2) + uv = yuv[height:].reshape(2, height // 2, width // 2) + return np.concatenate((y, uv)) + + +def _calibration_augmentations(rpy: np.ndarray, num_samples: int, val: bool) -> tuple[np.ndarray, np.ndarray]: + if val: + multipliers = np.zeros(num_samples) + else: + multipliers = np.asarray( + [float(np.random.uniform()) if np.random.uniform() > 0.8 else 0.0 for _ in range(num_samples)] + ) + augment_from_calib = rot_from_euler(rpy * multipliers[:, None]) + augment_from_device = np.einsum("tij,jk->tik", augment_from_calib, rot_from_euler(rpy).T) + eulers_view = np.einsum( + "ij,tj->ti", view_frame_from_device_frame, euler_from_rot(augment_from_device.swapaxes(1, 2)) + ) + return augment_from_calib, eulers_view + + +def _load_images( + fs, + segment_dir: str, + frame_info: dict[str, np.ndarray], + fidxs: np.ndarray, + temporal_fidxs: np.ndarray, + eulers_view: np.ndarray, + frame_skip: int, + n_frames: int, + local_rank: int, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + big_camera = "ecamera" if "ecamera/index" in frame_info else "fcamera" + device_type = int(frame_info["device_type"].item()) + cameras = DEVICE_CAMERAS[CAMERA_BY_DEVICE_TYPE[device_type]] + + fcamera_intrinsics = cameras.fcam.intrinsics + big_intrinsics = (cameras.ecam if big_camera == "ecamera" else cameras.fcam).intrinsics + med_matrices = [_transform_matrix(fcamera_intrinsics, medmodel_intrinsics, eulers) for eulers in eulers_view] + big_matrices = [_transform_matrix(big_intrinsics, sbigmodel_intrinsics, eulers) for eulers in eulers_view] + + source_width = int(frame_info[f"{big_camera}/width"].item()) + source_height = int(frame_info[f"{big_camera}/height"].item()) + valid = np.ones(len(fidxs), dtype=bool) + for sample_idx, eulers in enumerate(eulers_view): + inverse = np.linalg.inv(_transform_matrix(big_intrinsics, bigmodel_intrinsics, eulers)) + source = BIG_MODEL_CORNERS @ inverse.T + source = source[:, :2] / source[:, 2, None] + valid[sample_idx] = np.all( + (source[:, 0] > 0) & (source[:, 0] < source_width) & (source[:, 1] > 0) & (source[:, 1] < source_height) + ) + + requests: dict[int, list[tuple[int, int, int]]] = defaultdict(list) + for (sample_idx, temporal_idx), temporal_fidx in np.ndenumerate(temporal_fidxs): + if valid[sample_idx]: + for frame_number in range(n_frames): + source_fidx = int(temporal_fidx + (frame_number + 1 - n_frames) * frame_skip) + requests[source_fidx].append((sample_idx, temporal_idx, frame_number)) + + wanted = set(requests) + fcamera_frames = _decode(fs, f"{segment_dir}/fcamera.hevc", frame_info["fcamera/index"], wanted, local_rank) + big_frames = ( + fcamera_frames + if big_camera == "fcamera" + else _decode(fs, f"{segment_dir}/ecamera.hevc", frame_info["ecamera/index"], wanted, local_rank) + ) + images = np.zeros((len(fidxs), temporal_fidxs.shape[1], 6 * n_frames, H // 2, W // 2), dtype=np.uint8) + big_images = np.zeros((len(fidxs), temporal_fidxs.shape[1], 6 * n_frames, SBIG_H // 2, SBIG_W // 2), dtype=np.uint8) + for source_fidx, destinations in requests.items(): + for sample_idx, temporal_idx, frame_number in destinations: + channel_slice = slice(6 * frame_number, 6 * (frame_number + 1)) + image = cv2.warpPerspective( + fcamera_frames[source_fidx], med_matrices[sample_idx], (W, H), borderMode=cv2.BORDER_REPLICATE + ) + big_image = cv2.warpPerspective( + big_frames[source_fidx], big_matrices[sample_idx], (SBIG_W, SBIG_H), borderMode=cv2.BORDER_REPLICATE + ) + images[sample_idx, temporal_idx, channel_slice] = _pack_image(image) + big_images[sample_idx, temporal_idx, channel_slice] = _pack_image(big_image) + return images, big_images, valid + + +def _load_targets( + localizer: dict[str, np.ndarray], + temporal_fidxs: np.ndarray, + augment_from_calib: np.ndarray, + action_t: np.ndarray, +) -> dict[str, np.ndarray]: + frame_t = localizer["frame_t"] + states = localizer["frame_states"].copy() + states[:, 10:13] = gaussian_filter1d(states[:, 10:13], 2, axis=0, radius=5, mode="nearest") + states[:, 19:22] = gaussian_filter1d(states[:, 19:22], 2, axis=0, radius=5, mode="nearest") + calib_from_device = rot_from_euler(localizer["rpy"]).T + plan = np.empty((*temporal_fidxs.shape, len(T_IDXS), 15), dtype=np.float32) + + for (sample_idx, temporal_idx), fidx in np.ndenumerate(temporal_fidxs): + query_t = T_IDXS + frame_t[fidx] + indexes = np.clip(np.searchsorted(frame_t, query_t) - 1, 0, len(frame_t) - 2) + distance = (query_t - frame_t[indexes]) / (frame_t[indexes + 1] - frame_t[indexes]) + future = (states[indexes].T * (1 - distance)).T + (states[indexes + 1].T * distance).T + device_from_ecef = rot_from_quat(future[:, 3:7]).swapaxes(1, 2) + calib_from_ecef = np.einsum("ij,tjk->tik", calib_from_device, device_from_ecef) + + value = np.empty((len(T_IDXS), 15), dtype=np.float64) + value[:, Plan.POSITION] = np.einsum("ij,tj->ti", calib_from_ecef[0], future[:, 0:3] - future[0, 0:3]) + value[:, Plan.VELOCITY] = np.einsum("tij,tj->ti", calib_from_ecef, future[:, 7:10]) + value[:, Plan.ACCELERATION] = np.einsum("ij,tj->ti", calib_from_device, future[:, 19:22]) + value[:, Plan.T_FROM_CURRENT_EULER] = euler_from_rot( + np.einsum("ij,tjk->tik", calib_from_ecef[0], calib_from_ecef.swapaxes(1, 2)) + ) + value[:, Plan.ORIENTATION_RATE] = np.einsum("ij,tj->ti", calib_from_device, future[:, 10:13]) + value = value.astype(np.float32) + + augment = augment_from_calib[sample_idx] + for value_slice in (Plan.POSITION, Plan.VELOCITY, Plan.ACCELERATION, Plan.ORIENTATION_RATE): + value[:, value_slice] = np.einsum("ij,tj->ti", augment, value[:, value_slice]) + value[:, Plan.T_FROM_CURRENT_EULER] = euler_from_rot( + np.einsum("ij,tjk->tik", augment, rot_from_euler(value[:, Plan.T_FROM_CURRENT_EULER])) + ) + plan[sample_idx, temporal_idx] = value + + flat_plan = plan.reshape(-1, len(T_IDXS), 15) + flat_action_t = np.broadcast_to(action_t[:, None], (*plan.shape[:2], 2)).reshape(-1, 2) + action = np.zeros((len(flat_plan), 2), dtype=np.float32) + for idx, (flat_value, times) in enumerate(zip(flat_plan, flat_action_t, strict=True)): + velocity = max(float(flat_value[0, Plan.VELOCITY.start]), MIN_SPEED) + curvature = get_curvature_from_plan( + flat_value[:, Plan.T_FROM_CURRENT_EULER.start + 2], + flat_value[:, Plan.ORIENTATION_RATE.start + 2], + T_IDXS, + velocity, + float(times[0]), + ) + action[idx, 0] = curvature * velocity**2 + action[idx, 1] = get_accel_from_plan( + flat_value[:, Plan.VELOCITY.start], + flat_value[:, Plan.ACCELERATION.start], + T_IDXS, + action_t=float(times[1]), + ) + return {"plan": plan, "action": action.reshape(*plan.shape[:2], 2)} + + +def _load_segment( + fs, segment_dir: str, config: Any, val: bool, local_rank: int +) -> tuple[dict[str, np.ndarray], dict[str, np.ndarray]]: + from safetensors.numpy import load + + frame_info = load(fs.cat_file(f"{segment_dir}/frame_info.safetensors")) + localizer = load(fs.cat_file(f"{segment_dir}/localizer.safetensors")) + frame_skip = CAMERA_FPS // config.fps + history_idxs = -np.arange(1, 2 * config.fps, config.fps // 4, dtype=np.int32)[::-1] + temporal_len = 5 * config.fps + int(history_idxs[-1] - history_idxs[0]) + start = frame_skip * (temporal_len + 1) + int(np.random.randint(0, 5)) + candidates = np.arange(start, len(frame_info["fcamera/t"]) - 200, 40) + sample_count = min(len(candidates), 8) + fidxs = np.random.choice(candidates, sample_count, replace=False) + + temporal_fidxs = np.asarray( + [[fidx + (history_idx + 1) * frame_skip for history_idx in history_idxs] for fidx in fidxs] + ) + augment_from_calib, eulers_view = _calibration_augmentations(localizer["rpy"], len(fidxs), val) + + images, big_images, valid = _load_images( + fs, + segment_dir, + frame_info, + fidxs, + temporal_fidxs, + eulers_view, + frame_skip, + config.n_frames, + local_rank, + ) + action_t = np.random.uniform(0.0, 1.0, size=(len(fidxs), 2)).astype(np.float32) + targets = _load_targets(localizer, temporal_fidxs, augment_from_calib, action_t) + + traffic = np.zeros((len(fidxs), temporal_len, 2), dtype=np.float32) + latitude, longitude, _ = ecef2geodetic(localizer["frame_states"][-1, 0:3]) + country = reverse_geocoder.search((latitude, longitude), mode=1, verbose=False)[0]["cc"] + traffic[:, :, int(country in LHT_COUNTRIES)] = 1 + inputs = { + "img": images, + "big_img": big_images, + "desire_pulse": np.full((len(fidxs), temporal_len, 8), -1.0, dtype=np.float32), + "traffic_convention": traffic, + "action_t": np.broadcast_to(action_t[:, None], (len(fidxs), temporal_len, 2)).copy(), + } + + inputs = {name: np.ascontiguousarray(value[valid]) for name, value in inputs.items()} + targets = { + name: np.ascontiguousarray(value[valid]).reshape(valid.sum(), temporal_fidxs.shape[1], -1) + for name, value in targets.items() + } + return inputs, targets diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 694994841de..29d87da0f76 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -8,6 +8,7 @@ import math import os +from dataclasses import replace from functools import partial from xx.datasets.constants import BASE_DIR_GT, DEFAULT_TEST_5K_LIST_TAGGED, DEFAULT_TRAIN_LIST from xx.ml_tools.constants.model import ( @@ -31,6 +32,7 @@ from torchtitan.models.common import Embedding, LayerNorm, Linear from torchtitan.models.common.attention import ScaledDotProductAttention from torchtitan.protocols.model_spec import ModelSpec +from .comma1m_dataset import COMMA1M_REPO_ID from .dataset import PathDataLoader from .loss import PathLoss @@ -170,6 +172,25 @@ def _path(flavor: str) -> PathTrainer.Config: ) +def _comma1m_path(flavor: str) -> PathTrainer.Config: + config = _path(flavor) + dataloader = replace( + config.dataloader, + dataset=COMMA1M_REPO_ID, + plan_only=True, + limit=None, + ) + validation_dataloader = replace( + config.validator.dataloader, + dataset=COMMA1M_REPO_ID, + plan_only=True, + limit=None, + deterministic_fidxs=False, + ) + validator = replace(config.validator, dataloader=validation_dataloader, reports={}) + return replace(config, dataloader=dataloader, validator=validator) + + def _model_config(flavor: str) -> PathModel.Config: vision_features = 512 n_frames_input = N_FRAMES @@ -406,6 +427,7 @@ def _hydra(heads: tuple[PathHead, ...], *, in_features: int, mlp_mult: float) -> convnext_atto = partial(_path, "convnext_atto") +convnext_atto_comma1m = partial(_comma1m_path, "convnext_atto") convnext_femto = partial(_path, "convnext_femto") convnext_pico = partial(_path, "convnext_pico") convnext_tiny = partial(_path, "convnext_tiny") diff --git a/torchtitan/experiments/path/dataset.py b/torchtitan/experiments/path/dataset.py index a8dafc600ea..041ee290b25 100644 --- a/torchtitan/experiments/path/dataset.py +++ b/torchtitan/experiments/path/dataset.py @@ -50,11 +50,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 +64,35 @@ 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 = get_dataset(config.dataset, xx_config, val, self.local_rank, dp_rank, dp_world_size) + from .comma1m_dataset import COMMA1M_REPO_ID + + if config.dataset == COMMA1M_REPO_ID: + from .comma1m_dataset import Comma1MDataset + + dataset = Comma1MDataset(config, val, self.local_rank, dp_rank, dp_world_size) + else: + from xx.training.path.config import DatasetConfig as XXPathDatasetConfig + from xx.training.path.dataloader import get_dataset as get_xx_dataset + + 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 = get_xx_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 +107,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__( From 1aed333b8987ff9a4b5338fbb49be00f6d373bf3 Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Tue, 25 Aug 2026 00:13:43 -0700 Subject: [PATCH 2/9] works --- .../experiments/path/comma1m_dataset.py | 7 +- .../experiments/path/comma1m_trainer.py | 33 ++++++ .../experiments/path/config_registry.py | 104 +++++++++++++----- torchtitan/experiments/path/dataset.py | 14 ++- torchtitan/experiments/path/loss.py | 27 ++++- 5 files changed, 149 insertions(+), 36 deletions(-) create mode 100644 torchtitan/experiments/path/comma1m_trainer.py diff --git a/torchtitan/experiments/path/comma1m_dataset.py b/torchtitan/experiments/path/comma1m_dataset.py index 774b4a102ee..9e56a6d81d8 100644 --- a/torchtitan/experiments/path/comma1m_dataset.py +++ b/torchtitan/experiments/path/comma1m_dataset.py @@ -34,8 +34,8 @@ from scipy.ndimage import gaussian_filter1d from torch.utils.data import get_worker_info +from .dataset import COMMA1M_REPO_ID as COMMA1M_REPO_ID -COMMA1M_REPO_ID = "commaai/comma1M" T_IDXS = np.asarray(ModelConstants.T_IDXS) W, H = MEDMODEL_INPUT_SIZE BIG_W, BIG_H = BIGMODEL_INPUT_SIZE @@ -71,7 +71,6 @@ def __init__( ) -> None: assert config.plan_only assert not config.rgb - assert not config.unvision assert config.dataset_path is not None assert not config.deterministic_fidxs @@ -177,8 +176,8 @@ def _load_images( device_type = int(frame_info["device_type"].item()) cameras = DEVICE_CAMERAS[CAMERA_BY_DEVICE_TYPE[device_type]] - fcamera_intrinsics = cameras.fcam.intrinsics - big_intrinsics = (cameras.ecam if big_camera == "ecamera" else cameras.fcam).intrinsics + fcamera_intrinsics = cameras.narrow_road.intrinsics + big_intrinsics = (cameras.wide_road if big_camera == "ecamera" else cameras.narrow_road).intrinsics med_matrices = [_transform_matrix(fcamera_intrinsics, medmodel_intrinsics, eulers) for eulers in eulers_view] big_matrices = [_transform_matrix(big_intrinsics, sbigmodel_intrinsics, eulers) for eulers in eulers_view] diff --git a/torchtitan/experiments/path/comma1m_trainer.py b/torchtitan/experiments/path/comma1m_trainer.py new file mode 100644 index 00000000000..84a64d8aac8 --- /dev/null +++ b/torchtitan/experiments/path/comma1m_trainer.py @@ -0,0 +1,33 @@ +# 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 + +import torch + +from torchtitan.observability import structured_logger as sl +from torchtitan.trainer import Trainer + + +class Comma1MPathTrainer(Trainer): + @dataclass(kw_only=True, slots=True) + class Config(Trainer.Config): + pass + + @sl.log_trace_span("post_dataloading_process") + def post_dataloading_process( + self, + input_dict: dict[str, torch.Tensor], + labels: torch.Tensor, + ) -> tuple[dict[str, torch.Tensor], torch.Tensor, dict]: + self.ntokens_seen += labels.shape[0] + return input_dict, labels, {} + + def close(self) -> None: + self.dataloader.close() + super().close() diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 4ea5359890d..7bc4e40a78c 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -7,12 +7,10 @@ from __future__ import annotations import os -from dataclasses import replace from functools import partial -from typing import Literal - -from xx.comma_data.constants import BASE_DIR_GT, DEFAULT_TEST_5K_LIST_TAGGED, DEFAULT_TRAIN_LIST +from typing import TYPE_CHECKING, Literal +from torchtitan.components.checkpoint import CheckpointManager from torchtitan.components.lr_scheduler import LRSchedulersContainer from torchtitan.components.metrics import MetricsProcessor from torchtitan.components.optimizer import OptimizersContainer, ParamGroupConfig @@ -21,9 +19,7 @@ from torchtitan.distributed.activation_checkpoint import FullAC from torchtitan.protocols.model_spec import ModelSpec -from .comma1m_dataset import COMMA1M_REPO_ID -from .dataset import PathDataLoader -from .loss import PathLoss +from .dataset import COMMA1M_REPO_ID, PathDataLoader from .model import parallelize_path from .model_config import model_config as _model_config from .model_constants import ( @@ -34,9 +30,11 @@ SUPERCOMBO_FPS, TEMPORAL_INPUTS, ) -from .onnx_checkpoint import PathOnnxCheckpointManager -from .trainer import PathTrainer -from .validate import PathValidator + +if TYPE_CHECKING: + from .comma1m_trainer import Comma1MPathTrainer + from .onnx_checkpoint import PathOnnxCheckpointManager + from .trainer import PathTrainer def model_registry(flavor: str) -> ModelSpec: @@ -59,6 +57,12 @@ def _dp_degrees() -> tuple[int, int]: def _path(flavor: str) -> PathTrainer.Config: + from xx.comma_data.constants import BASE_DIR_GT, DEFAULT_TEST_5K_LIST_TAGGED, DEFAULT_TRAIN_LIST + + from .loss import PathLoss + from .trainer import PathTrainer + from .validate import PathValidator + steps = 1024 * 55 validation_freq = 1024 reports = { @@ -152,39 +156,76 @@ def _path(flavor: str) -> PathTrainer.Config: ) -def _comma1m_path(flavor: str) -> PathTrainer.Config: - config = _path(flavor) - dataloader = replace( - config.dataloader, - dataset=COMMA1M_REPO_ID, - plan_only=True, - limit=None, - ) - validation_dataloader = replace( - config.validator.dataloader, - dataset=COMMA1M_REPO_ID, - plan_only=True, - limit=None, - deterministic_fidxs=False, +def _comma1m_path(flavor: str) -> Comma1MPathTrainer.Config: + from .comma1m_trainer import Comma1MPathTrainer + from .loss import PathMSELoss + + steps = 1024 * 55 + num_nodes, local_world_size = _dp_degrees() + return Comma1MPathTrainer.Config( + loss=PathMSELoss.Config(), + model_spec=model_registry(flavor), + tokenizer=NoOpTokenizer.Config(), + dataloader=_dataloader_config( + dataset=COMMA1M_REPO_ID, + dataset_path=os.getenv("COMMA1M_DATASET_PATH"), + split="train", + fps=SUPERCOMBO_FPS, + plan_only=True, + limit=None, + deterministic_fidxs=False, + pipeline_dir=None, + skip=1, + val_skip=1, + ), + optimizer=_optimizer_config(), + lr_scheduler=LRSchedulersContainer.Config( + warmup_steps=1024, + total_steps=steps, + decay_ratio=0.1, + decay_type="linear", + min_lr_factor=0.0, + ), + training=TrainingConfig( + local_batch_size=16, + seq_len=1, + steps=steps, + mixed_precision_param="bfloat16", + ), + parallelism=ParallelismConfig( + data_parallel_replicate_degree=num_nodes, + data_parallel_shard_degree=local_world_size, + enable_sequence_parallel=False, + ), + checkpoint=CheckpointManager.Config( + enable=True, + folder="checkpoint", + interval=500, + enable_first_step_checkpoint=True, + ), + activation_checkpoint=FullAC.Config(), + compile=CompileConfig(enable=True, components=["model", "loss"]), + metrics=MetricsProcessor.Config(log_freq=16, enable_wandb=True), + debug=DebugConfig(seed=0), ) - validator = replace(config.validator, dataloader=validation_dataloader, reports={}) - return replace(config, dataloader=dataloader, validator=validator) def _dataloader_config( *, dataset: str, + dataset_path: str | None = None, split: Literal["train", "val"], fps: int, plan_only: bool, limit: int | None, deterministic_fidxs: bool, - pipeline_dir: str, + pipeline_dir: str | None, skip: int, val_skip: int, ) -> PathDataLoader.Config: return PathDataLoader.Config( dataset=dataset, + dataset_path=dataset_path, split=split, deterministic_fidxs=deterministic_fidxs, fps=fps, @@ -196,7 +237,13 @@ def _dataloader_config( ) -def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnxCheckpointManager.Config: +def _checkpoint_config( + folder: str, + base_folder: str, + interval: int, +) -> PathOnnxCheckpointManager.Config: + from .onnx_checkpoint import PathOnnxCheckpointManager + frame_constants = frame_constants_from_fps(n_frames=N_FRAMES, frame_type=FRAME_TYPE) temporal_len = frame_constants["temporal_len"] vision_input_names = [ModelInputs.IMG, ModelInputs.BIG_IMG] @@ -258,6 +305,7 @@ def _optimizer_config() -> OptimizersContainer.Config: convnext_atto = partial(_path, "convnext_atto") convnext_atto_comma1m = partial(_comma1m_path, "convnext_atto") +convnext_xxlarge_comma1m = partial(_comma1m_path, "convnext_xxlarge") convnext_femto = partial(_path, "convnext_femto") convnext_pico = partial(_path, "convnext_pico") convnext_tiny = partial(_path, "convnext_tiny") diff --git a/torchtitan/experiments/path/dataset.py b/torchtitan/experiments/path/dataset.py index 97bee8368a1..841a1d4491c 100644 --- a/torchtitan/experiments/path/dataset.py +++ b/torchtitan/experiments/path/dataset.py @@ -19,6 +19,9 @@ from .model_constants import FRAME_TYPE, N_FRAMES, SUPERCOMBO_FPS, VisionFrameType +COMMA1M_REPO_ID = "commaai/comma1M" + + class PathDataLoader(BaseDataLoader): @dataclass(kw_only=True, slots=True) class Config(BaseDataLoader.Config): @@ -44,8 +47,6 @@ def _build_dataset( *, val: bool, ) -> Any: - from .comma1m_dataset import COMMA1M_REPO_ID - if config.dataset == COMMA1M_REPO_ID: from .comma1m_dataset import Comma1MDataset @@ -115,11 +116,18 @@ def __init__( def __iter__( self, - ) -> Iterator[tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]]: + ) -> Iterator[ + tuple[ + dict[str, torch.Tensor], + torch.Tensor | dict[str, torch.Tensor], + ] + ]: iterator = iter(self.loader) self._iterator = iterator try: for inputs, targets in iterator: + if self.config.dataset == COMMA1M_REPO_ID: + targets = targets["plan"] yield inputs, targets finally: iterator.close() diff --git a/torchtitan/experiments/path/loss.py b/torchtitan/experiments/path/loss.py index 486886d6a19..b5071db0776 100644 --- a/torchtitan/experiments/path/loss.py +++ b/torchtitan/experiments/path/loss.py @@ -7,7 +7,6 @@ from __future__ import annotations from dataclasses import dataclass -from xx.training.lib.driving import DrivingLoss, DrivingMetric import torch @@ -17,12 +16,38 @@ from torchtitan.tools.logging import logger +def path_plan_mse( + pred: dict[str, torch.Tensor], + target: torch.Tensor, +) -> torch.Tensor: + plan = pred["plan"][..., : target.shape[-1]] + return torch.nn.functional.mse_loss(plan.float(), target.float().detach(), reduction="sum") + + +class PathMSELoss(BaseLoss): + @dataclass(kw_only=True, slots=True) + class Config(BaseLoss.Config): + pass + + def __init__( + self, + config: Config, + *, + compile_config: CompileConfig | None = None, + ) -> None: + del config + self.fn = path_plan_mse + self._maybe_compile(compile_config) + + class PathLoss(BaseLoss): @dataclass(kw_only=True, slots=True) class Config(BaseLoss.Config): pass def __init__(self, config: Config, *, compile_config: CompileConfig | None = None): + from xx.training.lib.driving import DrivingLoss, DrivingMetric + del config self.fn = None self.loss_fn = DrivingLoss() From 4c2e19c316ddf24b1f12e9f8af57c7d3f2ffbebf Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Tue, 25 Aug 2026 01:08:28 -0700 Subject: [PATCH 3/9] works --- .../experiments/path/comma1m_dataset.py | 29 ++++++++++++++++- .../experiments/path/comma1m_trainer.py | 11 +++++-- .../experiments/path/config_registry.py | 32 +++++++++++-------- torchtitan/experiments/path/dataset.py | 2 ++ torchtitan/experiments/path/loss.py | 25 ++++++++++++--- 5 files changed, 77 insertions(+), 22 deletions(-) diff --git a/torchtitan/experiments/path/comma1m_dataset.py b/torchtitan/experiments/path/comma1m_dataset.py index 9e56a6d81d8..9bb6428fb9d 100644 --- a/torchtitan/experiments/path/comma1m_dataset.py +++ b/torchtitan/experiments/path/comma1m_dataset.py @@ -35,6 +35,7 @@ from torch.utils.data import get_worker_info from .dataset import COMMA1M_REPO_ID as COMMA1M_REPO_ID +from .model_constants import ModelInputs T_IDXS = np.asarray(ModelConstants.T_IDXS) W, H = MEDMODEL_INPUT_SIZE @@ -146,6 +147,25 @@ def _pack_image(rgb: np.ndarray) -> np.ndarray: return np.concatenate((y, uv)) +def _rgb_from_tensor(tensor: np.ndarray, size: tuple[int, int] = (128, 256)) -> np.ndarray: + height = tensor.shape[2] * 2 + width = tensor.shape[3] * 2 + frames = np.zeros((tensor.shape[0], height * 3 // 2, width), dtype=np.uint8) + frames[:, 0:height:2, 0::2] = tensor[:, 0] + frames[:, 1:height:2, 0::2] = tensor[:, 1] + frames[:, 0:height:2, 1::2] = tensor[:, 2] + frames[:, 1:height:2, 1::2] = tensor[:, 3] + frames[:, height : height + height // 4] = tensor[:, 4].reshape(-1, height // 4, width) + frames[:, height + height // 4 : height + height // 2] = tensor[:, 5].reshape(-1, height // 4, width) + + output_height, output_width = size + output = np.zeros((tensor.shape[0], output_height, output_width, 3), dtype=np.uint8) + for index, frame in enumerate(frames): + rgb = cv2.cvtColor(frame, cv2.COLOR_YUV2RGB_I420) + output[index] = cv2.resize(rgb, (output_width, output_height)) + return output + + def _calibration_augmentations(rpy: np.ndarray, num_samples: int, val: bool) -> tuple[np.ndarray, np.ndarray]: if val: multipliers = np.zeros(num_samples) @@ -324,7 +344,7 @@ def _load_segment( inputs = { "img": images, "big_img": big_images, - "desire_pulse": np.full((len(fidxs), temporal_len, 8), -1.0, dtype=np.float32), + "desire_pulse": np.zeros((len(fidxs), temporal_len, 8), dtype=np.float32), # TODO: get from log "traffic_convention": traffic, "action_t": np.broadcast_to(action_t[:, None], (len(fidxs), temporal_len, 2)).copy(), } @@ -334,4 +354,11 @@ def _load_segment( name: np.ascontiguousarray(value[valid]).reshape(valid.sum(), temporal_fidxs.shape[1], -1) for name, value in targets.items() } + targets["imgs"] = np.concatenate( + [ + _rgb_from_tensor(inputs[ModelInputs.IMG][:, -1, -6:]), + _rgb_from_tensor(inputs[ModelInputs.BIG_IMG][:, -1, -6:]), + ], + axis=3, + ) return inputs, targets diff --git a/torchtitan/experiments/path/comma1m_trainer.py b/torchtitan/experiments/path/comma1m_trainer.py index 84a64d8aac8..c4072723f86 100644 --- a/torchtitan/experiments/path/comma1m_trainer.py +++ b/torchtitan/experiments/path/comma1m_trainer.py @@ -13,6 +13,8 @@ from torchtitan.observability import structured_logger as sl from torchtitan.trainer import Trainer +from .dataset import COMMA1M_IMGS_TARGET + class Comma1MPathTrainer(Trainer): @dataclass(kw_only=True, slots=True) @@ -24,9 +26,14 @@ def post_dataloading_process( self, input_dict: dict[str, torch.Tensor], labels: torch.Tensor, - ) -> tuple[dict[str, torch.Tensor], torch.Tensor, dict]: + ) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor], dict]: + inputs = {name: value for name, value in input_dict.items() if name != COMMA1M_IMGS_TARGET} + targets = { + "plan": labels, + "imgs": input_dict[COMMA1M_IMGS_TARGET], + } self.ntokens_seen += labels.shape[0] - return input_dict, labels, {} + return inputs, targets, {} def close(self) -> None: self.dataloader.close() diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 7bc4e40a78c..6a6a85ab496 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -162,22 +162,26 @@ def _comma1m_path(flavor: str) -> Comma1MPathTrainer.Config: steps = 1024 * 55 num_nodes, local_world_size = _dp_degrees() + dataloader = _dataloader_config( + dataset=COMMA1M_REPO_ID, + dataset_path=os.getenv("COMMA1M_DATASET_PATH"), + split="train", + fps=SUPERCOMBO_FPS, + plan_only=True, + limit=None, + deterministic_fidxs=False, + pipeline_dir=None, + skip=1, + val_skip=1, + ) + dataloader.num_writers = 1 + dataloader.shuffle_size = 32 + dataloader.min_mixing = 0 return Comma1MPathTrainer.Config( loss=PathMSELoss.Config(), model_spec=model_registry(flavor), tokenizer=NoOpTokenizer.Config(), - dataloader=_dataloader_config( - dataset=COMMA1M_REPO_ID, - dataset_path=os.getenv("COMMA1M_DATASET_PATH"), - split="train", - fps=SUPERCOMBO_FPS, - plan_only=True, - limit=None, - deterministic_fidxs=False, - pipeline_dir=None, - skip=1, - val_skip=1, - ), + dataloader=dataloader, optimizer=_optimizer_config(), lr_scheduler=LRSchedulersContainer.Config( warmup_steps=1024, @@ -187,9 +191,9 @@ def _comma1m_path(flavor: str) -> Comma1MPathTrainer.Config: min_lr_factor=0.0, ), training=TrainingConfig( - local_batch_size=16, + local_batch_size=1, seq_len=1, - steps=steps, + steps=1, mixed_precision_param="bfloat16", ), parallelism=ParallelismConfig( diff --git a/torchtitan/experiments/path/dataset.py b/torchtitan/experiments/path/dataset.py index 841a1d4491c..7f74d44d5b2 100644 --- a/torchtitan/experiments/path/dataset.py +++ b/torchtitan/experiments/path/dataset.py @@ -20,6 +20,7 @@ COMMA1M_REPO_ID = "commaai/comma1M" +COMMA1M_IMGS_TARGET = "_comma1m_imgs_target" class PathDataLoader(BaseDataLoader): @@ -127,6 +128,7 @@ def __iter__( try: for inputs, targets in iterator: if self.config.dataset == COMMA1M_REPO_ID: + inputs[COMMA1M_IMGS_TARGET] = targets["imgs"] targets = targets["plan"] yield inputs, targets finally: diff --git a/torchtitan/experiments/path/loss.py b/torchtitan/experiments/path/loss.py index b5071db0776..817fa654adc 100644 --- a/torchtitan/experiments/path/loss.py +++ b/torchtitan/experiments/path/loss.py @@ -16,12 +16,27 @@ from torchtitan.tools.logging import logger -def path_plan_mse( +def path_mse( pred: dict[str, torch.Tensor], - target: torch.Tensor, + targets: dict[str, torch.Tensor], ) -> torch.Tensor: - plan = pred["plan"][..., : target.shape[-1]] - return torch.nn.functional.mse_loss(plan.float(), target.float().detach(), reduction="sum") + target_plan = targets["plan"] + plan = pred["plan"][..., : target_plan.shape[-1]] + plan_mse = torch.nn.functional.mse_loss( + plan.float(), + target_plan.float().detach(), + reduction="sum", + ) + + pred_imgs = (pred["imgs"].float() / 127.5 - 1.0).clamp(-1.0, 1.0) + target_imgs = (targets["imgs"].permute(0, 3, 1, 2).float() / 127.5 - 1.0).clamp(-1.0, 1.0) + imgs_mse = torch.nn.functional.mse_loss( + pred_imgs, + target_imgs.detach(), + reduction="sum", + ) + imgs_mse = imgs_mse * (target_plan.numel() / target_imgs.numel()) + return plan_mse + imgs_mse class PathMSELoss(BaseLoss): @@ -36,7 +51,7 @@ def __init__( compile_config: CompileConfig | None = None, ) -> None: del config - self.fn = path_plan_mse + self.fn = path_mse self._maybe_compile(compile_config) From 8eb3c30c36d1e90f4dc0375f6a147cc46cb938ed Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Tue, 25 Aug 2026 01:15:52 -0700 Subject: [PATCH 4/9] linty --- torchtitan/experiments/path/comma1m_dataset.py | 6 ++++-- torchtitan/experiments/path/config_registry.py | 2 +- torchtitan/experiments/path/dataset.py | 7 +------ 3 files changed, 6 insertions(+), 9 deletions(-) diff --git a/torchtitan/experiments/path/comma1m_dataset.py b/torchtitan/experiments/path/comma1m_dataset.py index 9bb6428fb9d..6206d716ebe 100644 --- a/torchtitan/experiments/path/comma1m_dataset.py +++ b/torchtitan/experiments/path/comma1m_dataset.py @@ -34,9 +34,11 @@ from scipy.ndimage import gaussian_filter1d from torch.utils.data import get_worker_info -from .dataset import COMMA1M_REPO_ID as COMMA1M_REPO_ID +from .dataset import COMMA1M_REPO_ID from .model_constants import ModelInputs +__all__ = ["COMMA1M_REPO_ID"] + T_IDXS = np.asarray(ModelConstants.T_IDXS) W, H = MEDMODEL_INPUT_SIZE BIG_W, BIG_H = BIGMODEL_INPUT_SIZE @@ -344,7 +346,7 @@ def _load_segment( inputs = { "img": images, "big_img": big_images, - "desire_pulse": np.zeros((len(fidxs), temporal_len, 8), dtype=np.float32), # TODO: get from log + "desire_pulse": np.zeros((len(fidxs), temporal_len, 8), dtype=np.float32), # TODO: get from log "traffic_convention": traffic, "action_t": np.broadcast_to(action_t[:, None], (len(fidxs), temporal_len, 2)).copy(), } diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 6a6a85ab496..afcf7e48223 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -8,7 +8,7 @@ import os from functools import partial -from typing import TYPE_CHECKING, Literal +from typing import Literal, TYPE_CHECKING from torchtitan.components.checkpoint import CheckpointManager from torchtitan.components.lr_scheduler import LRSchedulersContainer diff --git a/torchtitan/experiments/path/dataset.py b/torchtitan/experiments/path/dataset.py index 7f74d44d5b2..3b1b8d77c96 100644 --- a/torchtitan/experiments/path/dataset.py +++ b/torchtitan/experiments/path/dataset.py @@ -117,12 +117,7 @@ def __init__( def __iter__( self, - ) -> Iterator[ - tuple[ - dict[str, torch.Tensor], - torch.Tensor | dict[str, torch.Tensor], - ] - ]: + ) -> Iterator[tuple[dict[str, torch.Tensor], torch.Tensor | dict[str, torch.Tensor],]]: iterator = iter(self.loader) self._iterator = iterator try: From 21cc7e0dd55c476fccb8901bdde8f2bc8a4ec78f Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Tue, 25 Aug 2026 13:49:46 -0700 Subject: [PATCH 5/9] works --- .../experiments/path/comma1m_trainer.py | 43 ++++++++++++++++ .../experiments/path/config_registry.py | 28 +++++----- torchtitan/experiments/path/loss.py | 51 +++++++++++++++++-- 3 files changed, 106 insertions(+), 16 deletions(-) diff --git a/torchtitan/experiments/path/comma1m_trainer.py b/torchtitan/experiments/path/comma1m_trainer.py index c4072723f86..3f9397a5976 100644 --- a/torchtitan/experiments/path/comma1m_trainer.py +++ b/torchtitan/experiments/path/comma1m_trainer.py @@ -6,21 +6,64 @@ from __future__ import annotations +from collections.abc import Iterator from dataclasses import dataclass +from typing import Any import torch +from torchtitan.distributed import utils as dist_utils from torchtitan.observability import structured_logger as sl from torchtitan.trainer import Trainer from .dataset import COMMA1M_IMGS_TARGET +from .loss import PathMSELoss class Comma1MPathTrainer(Trainer): + path_loss: PathMSELoss + @dataclass(kw_only=True, slots=True) class Config(Trainer.Config): pass + def __init__(self, config: Config) -> None: + super().__init__(config) + if not isinstance(self.loss_fn, PathMSELoss): + raise TypeError(f"Comma1MPathTrainer requires PathMSELoss, got {type(self.loss_fn).__name__}") + self.path_loss = self.loss_fn + + base_log = self.metrics_processor.log + + def log_with_loss_components( + step: int, + global_avg_loss: float, + global_max_loss: float, + grad_norm: float, + extra_metrics: dict[str, Any] | None = None, + ) -> None: + metrics = dict(extra_metrics or {}) + loss_mesh = self.parallel_dims.get_optional_mesh("loss") + metrics.update( + { + name: dist_utils.dist_sum(value, loss_mesh) + for name, value in self.path_loss.get_component_metrics().items() + } + ) + base_log( + step, + global_avg_loss, + global_max_loss, + grad_norm, + extra_metrics=metrics, + ) + + self.metrics_processor.log = log_with_loss_components + + def train_step(self, data_iterator: Iterator[tuple[dict[str, torch.Tensor], torch.Tensor]]) -> None: + self.path_loss.reset_component_metrics() + super().train_step(data_iterator) + @sl.log_trace_span("post_dataloading_process") def post_dataloading_process( self, diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index afcf7e48223..88c60e6349c 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -160,7 +160,7 @@ def _comma1m_path(flavor: str) -> Comma1MPathTrainer.Config: from .comma1m_trainer import Comma1MPathTrainer from .loss import PathMSELoss - steps = 1024 * 55 + steps = 16 num_nodes, local_world_size = _dp_degrees() dataloader = _dataloader_config( dataset=COMMA1M_REPO_ID, @@ -175,25 +175,25 @@ def _comma1m_path(flavor: str) -> Comma1MPathTrainer.Config: val_skip=1, ) dataloader.num_writers = 1 - dataloader.shuffle_size = 32 + dataloader.shuffle_size = 64 dataloader.min_mixing = 0 return Comma1MPathTrainer.Config( loss=PathMSELoss.Config(), model_spec=model_registry(flavor), tokenizer=NoOpTokenizer.Config(), dataloader=dataloader, - optimizer=_optimizer_config(), + optimizer=_optimizer_config(lr=1e-6), lr_scheduler=LRSchedulersContainer.Config( - warmup_steps=1024, + warmup_steps=0, total_steps=steps, - decay_ratio=0.1, + decay_ratio=0, decay_type="linear", - min_lr_factor=0.0, + min_lr_factor=0, ), training=TrainingConfig( - local_batch_size=1, + local_batch_size=8, seq_len=1, - steps=1, + steps=steps, mixed_precision_param="bfloat16", ), parallelism=ParallelismConfig( @@ -204,12 +204,12 @@ def _comma1m_path(flavor: str) -> Comma1MPathTrainer.Config: checkpoint=CheckpointManager.Config( enable=True, folder="checkpoint", - interval=500, + interval=steps, enable_first_step_checkpoint=True, ), activation_checkpoint=FullAC.Config(), compile=CompileConfig(enable=True, components=["model", "loss"]), - metrics=MetricsProcessor.Config(log_freq=16, enable_wandb=True), + metrics=MetricsProcessor.Config(log_freq=1, enable_wandb=True), debug=DebugConfig(seed=0), ) @@ -287,8 +287,12 @@ def _checkpoint_config( ) -def _optimizer_config() -> OptimizersContainer.Config: - common = {"lr": 1e-3, "betas": (0.9, 0.95), "eps": 1e-8} +def _optimizer_config( + lr: float = 1e-3, + betas: tuple[float, float] = (0.9, 0.95), + eps: float = 1e-8, +) -> OptimizersContainer.Config: + common = {"lr": lr, "betas": betas, "eps": eps} no_decay = r"(point_policy\.hydra|temporal_policy\.temporal_hydra)\.(final_layer|scale_layer)" return OptimizersContainer.Config( implementation="fused_opt_states_bf16", diff --git a/torchtitan/experiments/path/loss.py b/torchtitan/experiments/path/loss.py index 817fa654adc..37c7ed103ca 100644 --- a/torchtitan/experiments/path/loss.py +++ b/torchtitan/experiments/path/loss.py @@ -6,6 +6,7 @@ from __future__ import annotations +from collections.abc import Callable from dataclasses import dataclass import torch @@ -16,10 +17,10 @@ from torchtitan.tools.logging import logger -def path_mse( +def path_mse_components( pred: dict[str, torch.Tensor], targets: dict[str, torch.Tensor], -) -> torch.Tensor: +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: target_plan = targets["plan"] plan = pred["plan"][..., : target_plan.shape[-1]] plan_mse = torch.nn.functional.mse_loss( @@ -36,7 +37,15 @@ def path_mse( reduction="sum", ) imgs_mse = imgs_mse * (target_plan.numel() / target_imgs.numel()) - return plan_mse + imgs_mse + return plan_mse + 100.0 * imgs_mse, plan_mse, imgs_mse + + +def path_mse( + pred: dict[str, torch.Tensor], + targets: dict[str, torch.Tensor], +) -> torch.Tensor: + loss, _, _ = path_mse_components(pred, targets) + return loss class PathMSELoss(BaseLoss): @@ -52,7 +61,41 @@ def __init__( ) -> None: del config self.fn = path_mse - self._maybe_compile(compile_config) + self.component_fn: Callable[ + [dict[str, torch.Tensor], dict[str, torch.Tensor]], + tuple[torch.Tensor, torch.Tensor, torch.Tensor], + ] = path_mse_components + self._component_metrics: dict[str, torch.Tensor] = {} + if compile_config is not None and compile_config.enable and "loss" in compile_config.components: + logger.info("Compiling the loss function with torch.compile") + self.component_fn = torch.compile(self.component_fn, backend=compile_config.backend) + + def reset_component_metrics(self) -> None: + self._component_metrics.clear() + + def get_component_metrics(self) -> dict[str, torch.Tensor]: + return self._component_metrics.copy() + + def __call__( + self, + pred: dict[str, torch.Tensor], + targets: dict[str, torch.Tensor], + global_valid_tokens: float | torch.Tensor | None = None, + ) -> torch.Tensor: + loss, plan_mse, imgs_mse = self.component_fn(pred, targets) + if global_valid_tokens is not None: + loss = loss / global_valid_tokens + plan_mse = plan_mse / global_valid_tokens + imgs_mse = imgs_mse / global_valid_tokens + + for name, value in ( + ("loss_metrics/plan_mse", plan_mse), + ("loss_metrics/imgs_mse", imgs_mse), + ): + value = value.detach() + accumulated = self._component_metrics.get(name) + self._component_metrics[name] = value if accumulated is None else accumulated + value + return loss class PathLoss(BaseLoss): From fd0d36b2c2cf3e694ee3a4b94aadf5f7ca2dd1f3 Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Tue, 25 Aug 2026 14:21:32 -0700 Subject: [PATCH 6/9] fix mfu --- .../experiments/path/comma1m_trainer.py | 31 ++++++++++++++++++- 1 file changed, 30 insertions(+), 1 deletion(-) diff --git a/torchtitan/experiments/path/comma1m_trainer.py b/torchtitan/experiments/path/comma1m_trainer.py index 3f9397a5976..2804b5915f2 100644 --- a/torchtitan/experiments/path/comma1m_trainer.py +++ b/torchtitan/experiments/path/comma1m_trainer.py @@ -6,12 +6,14 @@ from __future__ import annotations -from collections.abc import Iterator +import time +from collections.abc import Iterable, Iterator from dataclasses import dataclass from typing import Any import torch +from torchtitan.components.dataloader import DataloaderExhaustedError from torchtitan.distributed import utils as dist_utils from torchtitan.observability import structured_logger as sl from torchtitan.trainer import Trainer @@ -20,6 +22,16 @@ from .loss import PathMSELoss +def _copy_microbatch( + batch: tuple[dict[str, torch.Tensor], torch.Tensor], +) -> tuple[dict[str, torch.Tensor], torch.Tensor]: + inputs, labels = batch + return ( + {name: value.clone() for name, value in inputs.items()}, + labels.clone(), + ) + + class Comma1MPathTrainer(Trainer): path_loss: PathMSELoss @@ -62,8 +74,25 @@ def log_with_loss_components( def train_step(self, data_iterator: Iterator[tuple[dict[str, torch.Tensor], torch.Tensor]]) -> None: self.path_loss.reset_component_metrics() + if self.gradient_accumulation_steps > 1: + data_iterator = map(_copy_microbatch, data_iterator) super().train_step(data_iterator) + def batch_generator( + self, + data_iterable: Iterable[tuple[dict[str, torch.Tensor], torch.Tensor]], + ) -> Iterator[tuple[dict[str, torch.Tensor], torch.Tensor]]: + data_iterator = iter(data_iterable) + while True: + data_load_start = time.perf_counter() + try: + input_dict, labels = next(data_iterator) + except StopIteration as ex: + raise DataloaderExhaustedError() from ex + self.metrics_processor.ntokens_since_last_log += labels.shape[0] + self.metrics_processor.data_loading_times.append(time.perf_counter() - data_load_start) + yield input_dict, labels + @sl.log_trace_span("post_dataloading_process") def post_dataloading_process( self, From 776394369a04ba68616979d15d23b9276083eea3 Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Tue, 25 Aug 2026 15:19:31 -0700 Subject: [PATCH 7/9] log --- torchtitan/experiments/path/comma1m_dataset.py | 7 +++++++ 1 file changed, 7 insertions(+) diff --git a/torchtitan/experiments/path/comma1m_dataset.py b/torchtitan/experiments/path/comma1m_dataset.py index 6206d716ebe..e03726cd3f0 100644 --- a/torchtitan/experiments/path/comma1m_dataset.py +++ b/torchtitan/experiments/path/comma1m_dataset.py @@ -33,6 +33,7 @@ from openpilot.system.loggerd.config import CAMERA_FPS from scipy.ndimage import gaussian_filter1d from torch.utils.data import get_worker_info +from torchtitan.observability import structured_logger as sl from .dataset import COMMA1M_REPO_ID from .model_constants import ModelInputs @@ -88,6 +89,12 @@ def __init__( ] segments = [segment for segment in segments if (hash(int(segment, 16)) % 10 == 0) == val] self.segments = segments[global_rank::global_world_size] + sl.log_trace_scalar( + { + "comma1m_segments": len(segments), + "local_comma1m_segments": len(self.segments), + } + ) def __iter__(self): worker = get_worker_info() From 91e2f776429e7bc5f71bb2b977cb45a221823543 Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Tue, 25 Aug 2026 15:26:44 -0700 Subject: [PATCH 8/9] lint! --- torchtitan/experiments/path/comma1m_dataset.py | 1 + 1 file changed, 1 insertion(+) diff --git a/torchtitan/experiments/path/comma1m_dataset.py b/torchtitan/experiments/path/comma1m_dataset.py index e03726cd3f0..3b981529e3d 100644 --- a/torchtitan/experiments/path/comma1m_dataset.py +++ b/torchtitan/experiments/path/comma1m_dataset.py @@ -33,6 +33,7 @@ from openpilot.system.loggerd.config import CAMERA_FPS from scipy.ndimage import gaussian_filter1d from torch.utils.data import get_worker_info + from torchtitan.observability import structured_logger as sl from .dataset import COMMA1M_REPO_ID From 58f4d095397a84dd4aa2bbc9fa1c8eb4ae50373e Mon Sep 17 00:00:00 2001 From: YassineYousfi Date: Tue, 25 Aug 2026 15:35:04 -0700 Subject: [PATCH 9/9] lint! --- torchtitan/distributed/utils.py | 2 +- torchtitan/tools/utils.py | 21 +++++++++------------ 2 files changed, 10 insertions(+), 13 deletions(-) diff --git a/torchtitan/distributed/utils.py b/torchtitan/distributed/utils.py index 398d71402b8..71c3136785a 100644 --- a/torchtitan/distributed/utils.py +++ b/torchtitan/distributed/utils.py @@ -321,7 +321,7 @@ def set_batch_invariance(enable: bool) -> None: # Set NCCL env vars for deterministic inter-GPU collectives. # Must be set BEFORE dist.init_process_group. - # Reference: https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/batch_invariant.py + # Reference: https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/determinism/batch_invariant.py os.environ["NCCL_LAUNCH_MODE"] = "GROUP" # Fixed kernel launch ordering os.environ["NCCL_COLLNET_ENABLE"] = "0" # Disable SHARP (non-deterministic IB HW reduce) os.environ["NCCL_NVLS_ENABLE"] = "0" # Disable NVLink SHARP (non-deterministic NVSwitch HW reduce) diff --git a/torchtitan/tools/utils.py b/torchtitan/tools/utils.py index d0723a8a346..81b732e1bb3 100644 --- a/torchtitan/tools/utils.py +++ b/torchtitan/tools/utils.py @@ -93,11 +93,7 @@ def get_peak_flops(device_name: str) -> float: # Run the lspci command and capture the output result = subprocess.run(["lspci"], stdout=subprocess.PIPE, text=True) # Filter the output for lines containing both "NVIDIA" and "H100" - filtered_lines = [ - line - for line in result.stdout.splitlines() - if "NVIDIA" in line and "H100" in line - ] + filtered_lines = [line for line in result.stdout.splitlines() if "NVIDIA" in line and "H100" in line] # Join all filtered lines into a single string device_name = " ".join(filtered_lines) or device_name except FileNotFoundError as e: @@ -138,7 +134,7 @@ def get_peak_flops(device_name: str) -> float: # GB300 data from https://www.nvidia.com/en-us/data-center/dgx-gb300 return 2.5e15 elif "B300" in device_name or "B200" in device_name: - # data from https://nvdam.widen.net/s/wwnsxrhm2w/blackwell-datasheet-3384703 + # Data from https://www.nvidia.com/en-us/data-center/hgx/ # Checked after GB300 to avoid false match on "GB300" return 2.25e15 elif "MI355X" in device_name: @@ -178,18 +174,19 @@ def get_peak_flops(device_name: str) -> float: # https://awsdocs-neuron.readthedocs-hosted.com/en/latest/about-neuron/arch/neuron-hardware/neuron-core-v4.html return 79e12 * 2 else: - logger.warning( - f"Unknown neuron device: {neuron_device_name}, fallback to trn2/trn3" - ) + logger.warning(f"Unknown neuron device: {neuron_device_name}, fallback to trn2/trn3") return 79e12 * 2 elif "5090" in device_name: - # FP16/FP16 data from https://images.nvidia.com/aem-dam/Solutions/geforce/blackwell/nvidia-rtx-blackwell-gpu-architecture.pdf + # FP16/FP16 data: + # https://images.nvidia.com/aem-dam/Solutions/geforce/blackwell/nvidia-rtx-blackwell-gpu-architecture.pdf return 419.0e12 elif "4090" in device_name: - # FP16/FP16 data from https://images.nvidia.com/aem-dam/Solutions/geforce/blackwell/nvidia-rtx-blackwell-gpu-architecture.pdf + # FP16/FP16 data: + # https://images.nvidia.com/aem-dam/Solutions/geforce/blackwell/nvidia-rtx-blackwell-gpu-architecture.pdf return 330.3e12 elif "3090" in device_name: - # FP16/FP16 data from https://images.nvidia.com/aem-dam/Solutions/geforce/blackwell/nvidia-rtx-blackwell-gpu-architecture.pdf + # FP16/FP16 data: + # https://images.nvidia.com/aem-dam/Solutions/geforce/blackwell/nvidia-rtx-blackwell-gpu-architecture.pdf return 142.4e12 else: # for other GPU types, assume A100 logger.warning(f"Peak flops undefined for: {device_name}, fallback to A100")