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/experiments/path/comma1m_dataset.py b/torchtitan/experiments/path/comma1m_dataset.py new file mode 100644 index 00000000000..3b981529e3d --- /dev/null +++ b/torchtitan/experiments/path/comma1m_dataset.py @@ -0,0 +1,374 @@ +# 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 + +from torchtitan.observability import structured_logger as sl + +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 +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 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] + sl.log_trace_scalar( + { + "comma1m_segments": len(segments), + "local_comma1m_segments": len(self.segments), + } + ) + + 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 _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) + 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.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] + + 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.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(), + } + + 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() + } + 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 new file mode 100644 index 00000000000..2804b5915f2 --- /dev/null +++ b/torchtitan/experiments/path/comma1m_trainer.py @@ -0,0 +1,112 @@ +# 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 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 + +from .dataset import COMMA1M_IMGS_TARGET +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 + + @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() + 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, + input_dict: dict[str, torch.Tensor], + labels: torch.Tensor, + ) -> 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 inputs, targets, {} + + 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 7e5c3a74d54..88c60e6349c 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -8,10 +8,9 @@ import os 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 Literal, TYPE_CHECKING +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 @@ -20,8 +19,7 @@ from torchtitan.distributed.activation_checkpoint import FullAC from torchtitan.protocols.model_spec import ModelSpec -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 ( @@ -32,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: @@ -57,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 = { @@ -150,20 +156,80 @@ def _path(flavor: str) -> PathTrainer.Config: ) +def _comma1m_path(flavor: str) -> Comma1MPathTrainer.Config: + from .comma1m_trainer import Comma1MPathTrainer + from .loss import PathMSELoss + + steps = 16 + 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 = 64 + dataloader.min_mixing = 0 + return Comma1MPathTrainer.Config( + loss=PathMSELoss.Config(), + model_spec=model_registry(flavor), + tokenizer=NoOpTokenizer.Config(), + dataloader=dataloader, + optimizer=_optimizer_config(lr=1e-6), + lr_scheduler=LRSchedulersContainer.Config( + warmup_steps=0, + total_steps=steps, + decay_ratio=0, + decay_type="linear", + min_lr_factor=0, + ), + training=TrainingConfig( + local_batch_size=8, + 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=steps, + enable_first_step_checkpoint=True, + ), + activation_checkpoint=FullAC.Config(), + compile=CompileConfig(enable=True, components=["model", "loss"]), + metrics=MetricsProcessor.Config(log_freq=1, enable_wandb=True), + debug=DebugConfig(seed=0), + ) + + 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, @@ -175,7 +241,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] @@ -215,8 +287,12 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx ) -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", @@ -236,6 +312,8 @@ 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 b74a527cbd0..3b1b8d77c96 100644 --- a/torchtitan/experiments/path/dataset.py +++ b/torchtitan/experiments/path/dataset.py @@ -19,6 +19,10 @@ from .model_constants import FRAME_TYPE, N_FRAMES, SUPERCOMBO_FPS, VisionFrameType +COMMA1M_REPO_ID = "commaai/comma1M" +COMMA1M_IMGS_TARGET = "_comma1m_imgs_target" + + class PathDataLoader(BaseDataLoader): @dataclass(kw_only=True, slots=True) class Config(BaseDataLoader.Config): @@ -44,6 +48,11 @@ def _build_dataset( *, val: bool, ) -> Any: + if config.dataset == COMMA1M_REPO_ID: + from .comma1m_dataset import Comma1MDataset + + return Comma1MDataset(config, val, self.local_rank, self.dp_rank, self.dp_world_size) + if config.pipeline_dir is None: raise ValueError("pipeline_dir is required for internal PATH datasets") @@ -108,11 +117,14 @@ 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: + inputs[COMMA1M_IMGS_TARGET] = targets["imgs"] + 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..37c7ed103ca 100644 --- a/torchtitan/experiments/path/loss.py +++ b/torchtitan/experiments/path/loss.py @@ -6,8 +6,8 @@ from __future__ import annotations +from collections.abc import Callable from dataclasses import dataclass -from xx.training.lib.driving import DrivingLoss, DrivingMetric import torch @@ -17,12 +17,95 @@ from torchtitan.tools.logging import logger +def path_mse_components( + pred: dict[str, torch.Tensor], + targets: dict[str, 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( + 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 + 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): + @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_mse + 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): @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() 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")