From de7ff6cbfd6a886845c98e0fd1905d775ab2c51d Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Mon, 17 Aug 2026 19:05:17 -0700 Subject: [PATCH 01/23] path: remove avg-pool, attend over spatial tokens --- .../experiments/path/config_registry.py | 18 ++- torchtitan/experiments/path/model.py | 150 +++++++++++++++--- torchtitan/experiments/path/model_config.py | 90 +++++++++-- 3 files changed, 225 insertions(+), 33 deletions(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 307bfb0e1db..a93c3219e0b 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -23,7 +23,7 @@ from .dataset import PathDataLoader from .loss import PathLoss from .model import parallelize_path -from .model_config import model_config as _model_config +from .model_config import model_config as _model_config, VISION_FEATURES, _spatial_size from .model_constants import ( frame_constants_from_fps, FRAME_TYPE, @@ -37,11 +37,11 @@ from .validate import PathValidator -def model_registry(flavor: str) -> ModelSpec: +def model_registry(flavor: str, *, unvision: bool = False) -> ModelSpec: return ModelSpec( name="path", flavor=flavor, - model=_model_config(flavor), + model=_model_config(flavor, unvision=unvision), parallelize_fn=parallelize_path, pipelining_fn=None, post_optimizer_build_fn=None, @@ -56,7 +56,7 @@ def _dp_degrees() -> tuple[int, int]: return num_nodes, local_world_size -def _path(flavor: str) -> PathTrainer.Config: +def _path(flavor: str, *, unvision: bool = False) -> PathTrainer.Config: steps = 1024 * 55 validation_freq = 1024 reports = { @@ -80,7 +80,7 @@ def _path(flavor: str) -> PathTrainer.Config: plan_only = False return PathTrainer.Config( loss=PathLoss.Config(), - model_spec=model_registry(flavor), + model_spec=model_registry(flavor, unvision=unvision), tokenizer=NoOpTokenizer.Config(), dataloader=_dataloader_config( dataset=DEFAULT_TRAIN_LIST, @@ -92,6 +92,7 @@ def _path(flavor: str) -> PathTrainer.Config: pipeline_dir=BASE_DIR_GT, skip=1, val_skip=1, + unvision=unvision, ), optimizer=_optimizer_config(), lr_scheduler=LRSchedulersContainer.Config( @@ -143,6 +144,7 @@ def _path(flavor: str) -> PathTrainer.Config: pipeline_dir=BASE_DIR_GT, skip=1, val_skip=6, + unvision=unvision, ), mixed_precision_param=mixed_precision_param, reports=reports, @@ -162,6 +164,7 @@ def _dataloader_config( pipeline_dir: str, skip: int, val_skip: int, + unvision: bool | None = None, ) -> PathDataLoader.Config: return PathDataLoader.Config( dataset=dataset, @@ -173,12 +176,14 @@ def _dataloader_config( limit=limit, skip=skip, val_skip=val_skip, + unvision=unvision, ) def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnxCheckpointManager.Config: frame_constants = frame_constants_from_fps(n_frames=N_FRAMES, frame_type=FRAME_TYPE) temporal_len = frame_constants["temporal_len"] + spatial_size = _spatial_size(frame_constants) vision_input_names = [ModelInputs.IMG, ModelInputs.BIG_IMG] temporal_policy_input_names = [ ModelInputs.FEATURES, @@ -193,7 +198,7 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx input_shapes = [ [1, *frame_constants["frame_shapes"][ModelInputs.IMG]], [1, *frame_constants["frame_shapes"][ModelInputs.BIG_IMG]], - [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.FEATURES][0]], + [1, temporal_len, spatial_size, VISION_FEATURES], [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.DESIRE][0]], [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.TRAFFIC][0]], [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.ACTION_T][0]], @@ -245,3 +250,4 @@ def _optimizer_config() -> OptimizersContainer.Config: convnext_thirdxxl = partial(_path, "convnext_thirdxxl") convnext_base = partial(_path, "convnext_base") convnext_xxlarge = partial(_path, "convnext_xxlarge") +convnext_xxlarge_unvision = partial(_path, "convnext_xxlarge", unvision=True) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index f667ccfe778..81cea7db2e9 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -7,9 +7,11 @@ from __future__ import annotations from dataclasses import dataclass +from xx.training.lib.positional_embeddings import get_2d_sincos_pos_embed import torch import torch.nn as nn +import torch.nn.functional as F from einops import rearrange from torch.distributed.device_mesh import DeviceMesh from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy @@ -93,12 +95,14 @@ class Config(Module.Config): head_dim: int dropout: float is_causal: bool = True + causal_block_size: int = 1 def __init__(self, config: Config): super().__init__() self.n_head = config.n_head self.head_dim = config.head_dim self.is_causal = config.is_causal + self.causal_block_size = config.causal_block_size self.norm = config.norm.build() self.q_norm = config.q_norm.build() if config.q_norm is not None else nn.Identity() self.k_norm = config.k_norm.build() if config.k_norm is not None else nn.Identity() @@ -112,7 +116,15 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) q, k, v = qkv.unbind(2) q, k = self.q_norm(q), self.k_norm(k) - x = self.inner_attention(q, k, v, is_causal=self.is_causal) + if self.is_causal and self.causal_block_size > 1: + # Tokens are time-major. Let every spatial token attend within its frame + # while retaining causal attention between frames. + frame_idx = torch.arange(t, device=x.device) // self.causal_block_size + attention_mask = frame_idx[:, None] >= frame_idx[None, :] + q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) + x = F.scaled_dot_product_attention(q, k, v, attn_mask=attention_mask).transpose(1, 2) + else: + x = self.inner_attention(q, k, v, is_causal=self.is_causal) return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) @@ -155,20 +167,92 @@ def apply_fsdp(self, shard, reshard_after_forward: bool) -> None: shard(layer, reshard_after_forward) +class SpatialUnvision(Module): + """Training-only decoder from a spatial ConvNeXt token grid to two RGB views.""" + + OUTPUT_SIZE = (128, 256) + OUTPUT_CHANNELS = 6 + N_EMBD = 256 + N_HEAD = 8 + N_LAYER = 4 + + @dataclass(kw_only=True, slots=True) + class Config(Module.Config): + in_features: int + grid_size: tuple[int, int] + transformer: PathTransformer.Config + + def __init__(self, config: Config): + super().__init__() + self.config = config + grid_h, grid_w = config.grid_size + output_h, output_w = self.OUTPUT_SIZE + self.patch_size = (output_h // grid_h, output_w // grid_w) + dim = self.N_EMBD + self.input_projection = Linear.Config(in_features=config.in_features, out_features=dim, bias=True).build() + self.input_norm = LayerNorm.Config(normalized_shape=dim).build() + self.transformer = config.transformer.build() + self.output_norm = LayerNorm.Config(normalized_shape=dim).build() + self.output_projection = Linear.Config( + in_features=dim, + out_features=self.OUTPUT_CHANNELS * (self.patch_size[0] * self.patch_size[1]), + bias=True, + ).build() + self.register_buffer("pos_embedding", torch.empty(1, grid_h * grid_w, dim), persistent=True) + + def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> None: + device = buffer_device if buffer_device is not None else self.pos_embedding.device + embedding = get_2d_sincos_pos_embed(self.N_EMBD, self.config.grid_size) + self.pos_embedding = torch.from_numpy(embedding).to(device=device, dtype=torch.float32).unsqueeze(0) + + def forward(self, features: torch.Tensor) -> dict[str, torch.Tensor]: + grid_h, grid_w = self.config.grid_size + tokens = self.input_norm(self.input_projection(features)) + tokens = tokens + self.pos_embedding.to(tokens.dtype) + tokens = self.output_projection(self.output_norm(self.transformer(tokens))) + patch_h, patch_w = self.patch_size + images = rearrange( + tokens, + "b (grid_h grid_w) (c patch_h patch_w) -> b c (grid_h patch_h) (grid_w patch_w)", + grid_h=grid_h, + grid_w=grid_w, + patch_h=patch_h, + patch_w=patch_w, + ) + return {"imgs": ((images + 1.0) / 2.0) * 255.0} + + class PointSummarizer(Module): + """Collapse each spatial feature grid with a learned CLS token.""" + @dataclass(kw_only=True, slots=True) class Config(Module.Config): mlp1: PathMLP.Config - mlp2: PathMLP.Config + transformer: PathTransformer.Config + pos_embedding: Embedding.Config + spatial_size: int def __init__(self, config: Config): super().__init__() + self.spatial_size = config.spatial_size + self.n_features = config.pos_embedding.embedding_dim self.mlp1 = config.mlp1.build() - self.mlp2 = config.mlp2.build() + self.transformer = config.transformer.build() + self.pos_embedding = config.pos_embedding.build() + self.cls_token = nn.Parameter(torch.empty(1, 1, self.n_features)) + + def reset_parameters(self) -> None: + nn.init.normal_(self.cls_token, std=0.02) def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.mlp1(x) + x - return self.mlp2(x) + x + pos = self.pos_embedding(torch.arange(self.spatial_size, device=x.device)) + x = x + pos + leading_shape = x.shape[:-2] + cls = self.cls_token.expand(*leading_shape, 1, -1) + x = torch.cat((cls, x), dim=-2).reshape(-1, self.spatial_size + 1, self.n_features) + x = self.transformer(x) + return x.reshape(*leading_shape, self.spatial_size + 1, self.n_features)[..., 0, :] class LinearEncoder(Module): @@ -196,16 +280,19 @@ class Config(Module.Config): traffic_encoder: LinearEncoder.Config action_t_encoder: LinearEncoder.Config transformer: PathTransformer.Config - pos_embedding: Embedding.Config - block_size: int + temporal_pos_embedding: Embedding.Config + spatial_pos_embedding: Embedding.Config + temporal_size: int + spatial_size: int dense_training_outputs: bool def __init__(self, config: Config): super().__init__() - self.block_size = config.block_size + self.temporal_size = config.temporal_size + self.spatial_size = config.spatial_size self.dense_training_outputs = config.dense_training_outputs - if len(config.desire_window_starts) != self.block_size: - raise ValueError(f"Expected {self.block_size} desire window starts, got {len(config.desire_window_starts)}") + if len(config.desire_window_starts) != self.temporal_size: + raise ValueError(f"Expected {self.temporal_size} desire window starts, got {len(config.desire_window_starts)}") self.desire_window_len = config.desire_window_len self.desire_window_starts = config.desire_window_starts self.register_buffer("desire_window_idxs", self._make_desire_window_idxs(), persistent=False) @@ -215,7 +302,8 @@ def __init__(self, config: Config): self.traffic_encoder = config.traffic_encoder.build() self.action_t_encoder = config.action_t_encoder.build() self.transformer = config.transformer.build() - self.pos_embedding = config.pos_embedding.build() + self.temporal_pos_embedding = config.temporal_pos_embedding.build() + self.spatial_pos_embedding = config.spatial_pos_embedding.build() def _make_desire_window_idxs(self, device: torch.device | None = None) -> torch.Tensor: starts = torch.tensor(self.desire_window_starts, dtype=torch.long, device=device) @@ -228,7 +316,7 @@ def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> No def _window_desire(self, desire: torch.Tensor) -> torch.Tensor: desire = desire.index_select(1, self.desire_window_idxs) - return desire.reshape(desire.shape[0], self.block_size, -1) + return desire.reshape(desire.shape[0], self.temporal_size, -1) def forward( self, @@ -239,13 +327,22 @@ def forward( ) -> torch.Tensor: feats = self.mlp1(feats) + feats feats = self.mlp2(feats) + feats + b, t, s, c = feats.shape + # Temporal attention sees every spatial token. Pool only after attention so + # each dense output retains one entry per frame for the existing heads/losses. + feats = feats.reshape(b, t * s, c) desire = self.desire_encoder(self._window_desire(desire)) + desire = desire.repeat_interleave(s, dim=1) traffic_convention = rearrange(self.traffic_encoder(traffic_convention), "b c -> b () c") action_t = rearrange(self.action_t_encoder(action_t), "b c -> b () c") - pos = self.pos_embedding(torch.arange(self.block_size, device=feats.device)) - x = feats + rearrange(pos, "t c -> () t c") + desire + traffic_convention + action_t + temporal_pos = self.temporal_pos_embedding(torch.arange(t, device=feats.device)) + spatial_pos = self.spatial_pos_embedding(torch.arange(s, device=feats.device)) + pos = (temporal_pos[:, None, :] + spatial_pos[None, :, :]).reshape(t * s, c) + x = feats + rearrange(pos, "ts c -> () ts c") + desire + traffic_convention + action_t x = self.transformer(x) - return x if self.dense_training_outputs else x[:, self.block_size - 1] + if self.dense_training_outputs: + return x.reshape(b, t, s, c).mean(dim=2) + return x[:, -s:].mean(dim=1) class Hydra(Module): @@ -334,6 +431,7 @@ class Config(Module.Config): input_frame_names: tuple[str, ...] in_channels: int vision_features: int + grid_size: tuple[int, int] pretrained: bool drop_path_rate: float mean: float @@ -346,9 +444,13 @@ def __init__(self, config: Config): config.flavor, pretrained=False, in_chans=config.in_channels, - num_classes=config.vision_features, + num_classes=0, + global_pool="", drop_path_rate=config.drop_path_rate, ) + # The ConvNeXt head now normalizes without pooling; project its channel-rich + # 2-D map to the policy width at every spatial location. + self.proj = nn.Conv2d(self.encoder.num_features, config.vision_features, kernel_size=1) self.register_buffer("_mean", torch.empty(1, config.in_channels, 1, 1), persistent=True) self.register_buffer("_std", torch.empty(1, config.in_channels, 1, 1), persistent=True) @@ -408,7 +510,9 @@ def forward(self, inputs: dict[str, torch.Tensor]) -> torch.Tensor: x = torch.cat([inputs[name] for name in self.config.input_frame_names], dim=1) dtype = next(self.encoder.parameters()).dtype x = x.to(dtype) - return self.encoder((x - self._mean.to(dtype)) / self._std.to(dtype)) + x = self.encoder((x - self._mean.to(dtype)) / self._std.to(dtype)) + x = self.proj(x) + return x.flatten(2).transpose(1, 2) class PathModel(BaseModel): @@ -420,6 +524,8 @@ class Config(BaseModel.Config): vision: Vision.Config point_policy: Policy.Config temporal_policy: TemporalPolicy.Config + unvision: bool + unvision_decoder: SpatialUnvision.Config def update_from_config(self, *, config, **kwargs) -> None: parallelism = config.parallelism @@ -452,6 +558,7 @@ def __init__(self, config: Config): self.vision = config.vision.build() self.point_policy = config.point_policy.build() self.temporal_policy = config.temporal_policy.build() + self.unvision = config.unvision_decoder.build() if config.unvision else None @staticmethod def input_shapes( @@ -539,13 +646,16 @@ def forward( for name in self.config.vision.input_frame_names } features = self.vision(vision_inputs) - features = rearrange(features, "(b t) c -> b t c", b=b, t=t) - return self.point_policy(features) | self.temporal_policy( + features = rearrange(features, "(b t) s c -> b t s c", b=b, t=t) + outputs = self.point_policy(features) | self.temporal_policy( features, inputs[ModelInputs.DESIRE], inputs[ModelInputs.TRAFFIC], inputs[ModelInputs.ACTION_T], ) + if self.unvision is not None: + outputs |= self.unvision(features[:, -1]) + return outputs def parallelize_path( @@ -618,6 +728,8 @@ def wrap(module: nn.Module, fqn: str) -> nn.Module: wrap, "temporal_policy.temporal_summarizer.transformer", ) + if model.unvision is not None: + model.unvision.transformer.apply_activation_checkpointing(wrap, "unvision.transformer") logger.info(f"Applied {mode} activation checkpointing to the path model") @@ -628,6 +740,8 @@ def _apply_compile(model: PathModel, compile_config: CompileConfig) -> None: model.vision.encoder.compile(backend=compile_config.backend) model.point_policy.compile(backend=compile_config.backend) model.temporal_policy.compile(backend=compile_config.backend) + if model.unvision is not None: + model.unvision.compile(backend=compile_config.backend) logger.info("Compiling path model components with torch.compile") diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index 72eded0a7e9..b58eb474f9a 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -24,6 +24,7 @@ PointSummarizer, Policy, ScaleLayer, + SpatialUnvision, TemporalPolicy, TemporalSummarizer, Vision, @@ -38,15 +39,29 @@ ) +VISION_FEATURES = 512 +VISION_OUTPUT_STRIDE = 32 + 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 +def _vision_grid_size(frame_constants: dict) -> tuple[int, int]: + height, width = frame_constants["frame_shapes"][INPUT_FRAMES_NAMES[0]][-2:] + return height // VISION_OUTPUT_STRIDE, width // VISION_OUTPUT_STRIDE + + +def _spatial_size(frame_constants: dict) -> int: + return math.prod(_vision_grid_size(frame_constants)) + + +def model_config(flavor: str = "convnext_xxlarge", *, unvision: bool = False) -> PathModel.Config: + vision_features = VISION_FEATURES 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) + grid_size = _vision_grid_size(frame_constants) + spatial_size = math.prod(grid_size) return PathModel.Config( n_frames_input=N_FRAMES, @@ -57,6 +72,7 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: input_frame_names=input_frame_names, in_channels=in_channels, vision_features=vision_features, + grid_size=grid_size, pretrained=True, drop_path_rate=0.2, mean=255 / 2, @@ -65,11 +81,26 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: 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), + transformer=PathTransformer.Config( + layers=[ + PathTransformerBlock.Config( + attention=_attention(dim=vision_features, n_head=8, dropout=0.0, is_causal=False), + mlp=_mlp(vision_features, mlp_mult=2, bias=True, dropout=0.0), + ) + for _ in range(2) + ] + ), + pos_embedding=Embedding.Config( + num_embeddings=spatial_size, + embedding_dim=vision_features, + ), + spatial_size=spatial_size, ), hydra=_hydra(POINT_HEADS, in_features=vision_features, mlp_mult=2), ), temporal_policy=temporal_policy_config(), + unvision=unvision, + unvision_decoder=_spatial_unvision_config(in_features=vision_features, grid_size=grid_size), ) @@ -79,11 +110,13 @@ def temporal_policy_config( dropout: float = 0.1, dense_training_outputs: bool = True, ) -> TemporalPolicy.Config: - vision_features = 512 + vision_features = VISION_FEATURES 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) + block_size = len(history_idxs) + spatial_size = _spatial_size(frame_constants) return TemporalPolicy.Config( temporal_summarizer=TemporalSummarizer.Config( mlp1=_mlp(vision_features, mlp_mult=2, bias=False, dropout=0.0), @@ -96,17 +129,27 @@ def temporal_policy_config( transformer=PathTransformer.Config( layers=[ PathTransformerBlock.Config( - attention=_attention(dim=vision_features, n_head=8, dropout=dropout), + attention=_attention( + dim=vision_features, + n_head=8, + dropout=dropout, + causal_block_size=spatial_size, + ), mlp=_mlp(vision_features, mlp_mult=2, bias=True, dropout=dropout), ) for _ in range(4) ] ), - pos_embedding=Embedding.Config( - num_embeddings=len(history_idxs), + temporal_pos_embedding=Embedding.Config( + num_embeddings=block_size, + embedding_dim=vision_features, + ), + spatial_pos_embedding=Embedding.Config( + num_embeddings=spatial_size, embedding_dim=vision_features, ), - block_size=len(history_idxs), + temporal_size=block_size, + spatial_size=spatial_size, dense_training_outputs=dense_training_outputs, ), temporal_hydra=_hydra(heads, in_features=vision_features, mlp_mult=2), @@ -114,6 +157,26 @@ def temporal_policy_config( ) +def _spatial_unvision_config( + *, + in_features: int, + grid_size: tuple[int, int], +) -> SpatialUnvision.Config: + dim = SpatialUnvision.N_EMBD + layers = [ + PathTransformerBlock.Config( + attention=_attention(dim=dim, n_head=SpatialUnvision.N_HEAD, dropout=0.0, is_causal=False), + mlp=_mlp(dim=dim, mlp_mult=8 / 3, bias=False, dropout=0.0), + ) + for _ in range(SpatialUnvision.N_LAYER) + ] + return SpatialUnvision.Config( + in_features=in_features, + grid_size=grid_size, + transformer=PathTransformer.Config(layers=layers), + ) + + 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( @@ -132,7 +195,14 @@ def _encoder(in_features: int, dim: int) -> LinearEncoder.Config: ) -def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Config: +def _attention( + *, + dim: int, + n_head: int, + dropout: float, + is_causal: bool = True, + causal_block_size: int = 1, +) -> PathSelfAttention.Config: head_dim = dim // n_head return PathSelfAttention.Config( norm=LayerNorm.Config(normalized_shape=dim), @@ -144,6 +214,8 @@ def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Co n_head=n_head, head_dim=head_dim, dropout=dropout, + is_causal=is_causal, + causal_block_size=causal_block_size, ) From 6d9686a2f97dff19d1552a23e127e6861edb2f2a Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Tue, 18 Aug 2026 21:31:16 -0700 Subject: [PATCH 02/23] rldriving: train with spatial path models Port the rldriving experiment onto the spatial temporal policy: a shared _temporal_policy_config in path/config_registry guarantees actor parity with the path model warm-start, features are (B,T,S,C), and the dataloader and validator use the current xx APIs. --- .../experiments/rldriving/config_registry.py | 1 + torchtitan/experiments/rldriving/model.py | 44 ++++++++----------- torchtitan/experiments/rldriving/trainer.py | 6 ++- 3 files changed, 23 insertions(+), 28 deletions(-) diff --git a/torchtitan/experiments/rldriving/config_registry.py b/torchtitan/experiments/rldriving/config_registry.py index 61042dff9bc..2a3a49343e6 100644 --- a/torchtitan/experiments/rldriving/config_registry.py +++ b/torchtitan/experiments/rldriving/config_registry.py @@ -183,6 +183,7 @@ def _checkpoint_config( enable=True, checkpoint_base_folder=base_folder, export_onnx=True, + enable_first_step_checkpoint=True, folder=folder, interval=interval, input_names=list(input_shapes), diff --git a/torchtitan/experiments/rldriving/model.py b/torchtitan/experiments/rldriving/model.py index e8fc395ae95..9aa5a451cff 100644 --- a/torchtitan/experiments/rldriving/model.py +++ b/torchtitan/experiments/rldriving/model.py @@ -7,7 +7,6 @@ from __future__ import annotations import copy -import math from dataclasses import dataclass from typing import Any, cast @@ -23,6 +22,12 @@ 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_config import ( + TEMPORAL_HEADS, + _hydra, + _mlp, + temporal_policy_config, +) from torchtitan.experiments.path.model import ( Hydra, LinearEncoder, @@ -31,14 +36,13 @@ TemporalPolicy, TemporalSummarizer, ) -from torchtitan.experiments.path.model_config import TEMPORAL_HEADS, temporal_policy_config -from torchtitan.models.common import LayerNorm, Linear +from torchtitan.models.common import 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. +# B: batch, T: temporal steps, S: spatial tokens, D: model width, A: action components. ACTION_HEAD_NAME = "action" Q_HEAD_NAME = "q" @@ -76,10 +80,10 @@ def __init__(self, config: Config): 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 + features_BTSD = inputs[ModelInputs.FEATURES] + dtype = features_BTSD.dtype critic_features_BD = self.temporal_summarizer( - features_BTD[:, self.history_idxs], + features_BTSD[:, self.history_idxs], inputs[ModelInputs.DESIRE].to(dtype), inputs[ModelInputs.TRAFFIC][:, -1].to(dtype), inputs[ModelInputs.ACTION_T][:, -1].to(dtype), @@ -97,15 +101,7 @@ def actor_config() -> TemporalPolicy.Config: 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, - ) + dim = actor.temporal_summarizer.temporal_pos_embedding.embedding_dim return Critic.Config( temporal_summarizer=copy.deepcopy(actor.temporal_summarizer), history_idxs=actor.history_idxs, @@ -113,14 +109,9 @@ def critic_config(actor: TemporalPolicy.Config) -> Critic.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={}, - ), + post_action_mlp1=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), + post_action_mlp2=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), + q_hydra=_hydra((PathHead(name=Q_HEAD_NAME, output_size=1, mlp=False, scale=False),), in_features=dim, mlp_mult=2), ) @@ -199,7 +190,8 @@ def input_shapes( ModelInputs.FEATURES: ( batch_size, temporal_len, - summarizer.pos_embedding.embedding_dim, + summarizer.spatial_size, + summarizer.temporal_pos_embedding.embedding_dim, ), ModelInputs.DESIRE: (batch_size, temporal_len, desire_dim), ModelInputs.TRAFFIC: ( @@ -262,7 +254,7 @@ def parallelize_rldriving( training: TrainingConfig, parallelism: ParallelismConfig, compile_config: CompileConfig, - ac_config: ActivationCheckpointingConfig, + ac_config: ActivationCheckpointingConfig | None, dump_folder: str, ) -> RLDrivingModel: if compile_config.enable and "model" in compile_config.components: diff --git a/torchtitan/experiments/rldriving/trainer.py b/torchtitan/experiments/rldriving/trainer.py index fe7ab16a7d4..ccffae16911 100644 --- a/torchtitan/experiments/rldriving/trainer.py +++ b/torchtitan/experiments/rldriving/trainer.py @@ -198,7 +198,8 @@ def batch_generator(self, data_iterable: Iterable[Batch]) -> Iterator[Batch]: def prepare_batch(self, batch: Batch) -> PreparedBatch: inputs, targets, metadata = batch - inputs = {name: value.to(self.device) for name, value in inputs.items()} + # 'info' stays on the host; it is a serialized json column, not a tensor + inputs = {name: value.to(self.device) for name, value in inputs.items() if isinstance(value, torch.Tensor)} 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} @@ -215,7 +216,8 @@ def train_step(self, data_iterator: Iterator[Batch]) -> None: 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()) + info = info.cpu().numpy() if isinstance(info, torch.Tensor) else info + self.unique_segment_counter.update(parse_info(value)["name"] for value in info) current_inputs, next_inputs, targets, metadata = self.prepare_batch(batch) batch_size = next(iter(current_inputs.values())).shape[0] self.ntokens_seen += batch_size From af82fb37516f56469e765cc2f51a7785999f1ab4 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Wed, 19 Aug 2026 20:09:48 -0700 Subject: [PATCH 03/23] rldriving supercombo: block-causal mask + spatial hidden_state The temporal summarizer tokens are time-major (t * spatial_size), so the naive supercombo attention needs a frame-granular causal mask. The spatial vision output is (S, C) per frame - flatten it into the 1-D hidden_state so openpilot can feed it back as features_buffer. --- torchtitan/experiments/rldriving/supercombo.py | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/torchtitan/experiments/rldriving/supercombo.py b/torchtitan/experiments/rldriving/supercombo.py index b3372f13ca5..f64957b5271 100644 --- a/torchtitan/experiments/rldriving/supercombo.py +++ b/torchtitan/experiments/rldriving/supercombo.py @@ -54,17 +54,19 @@ def __init__(self) -> None: ) 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: + summarizer = policy.temporal_summarizer + # tokens are time-major (t * spatial); attend within the frame and causally across frames + n_tokens = summarizer.temporal_size * summarizer.spatial_size + for layer in 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) + frame_idx = torch.arange(n_tokens) // attention.causal_block_size + mask = (frame_idx[:, None] >= frame_idx[None, :]).view(1, 1, n_tokens, n_tokens) + attention.register_buffer("_supercombo_mask", mask, 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) + current = self.vision(inputs) # (1, S, C) + features = torch.cat((inputs["features_buffer"], current[:, None]), dim=1) # (1, T, S, C) 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( @@ -74,5 +76,5 @@ def forward(self, inputs: dict[str, torch.Tensor]) -> torch.Tensor: inputs[ModelInputs.ACTION_T][:, None], ) outputs.update({name: policy_outputs[name] for name in names}) - outputs["hidden_state"] = current.detach() + outputs["hidden_state"] = current.detach().reshape(1, -1) return torch.cat([outputs[name] for name in OUTPUT_ORDER] + [self.pad], dim=1) From b855acc3d5b6e4b575ec81eaef3750c23fd2b835 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Thu, 20 Aug 2026 21:15:47 -0700 Subject: [PATCH 04/23] Revert "rldriving supercombo: block-causal mask + spatial hidden_state" This reverts commit 28992598ea8bff6cd87a55fc6a3f834701e138c5. --- torchtitan/experiments/rldriving/supercombo.py | 18 ++++++++---------- 1 file changed, 8 insertions(+), 10 deletions(-) diff --git a/torchtitan/experiments/rldriving/supercombo.py b/torchtitan/experiments/rldriving/supercombo.py index f64957b5271..b3372f13ca5 100644 --- a/torchtitan/experiments/rldriving/supercombo.py +++ b/torchtitan/experiments/rldriving/supercombo.py @@ -54,19 +54,17 @@ def __init__(self) -> None: ) self.register_buffer("pad", torch.zeros(1, -output_size % 4), persistent=False) for policy in (self.off_policy, self.on_policy): - summarizer = policy.temporal_summarizer - # tokens are time-major (t * spatial); attend within the frame and causally across frames - n_tokens = summarizer.temporal_size * summarizer.spatial_size - for layer in summarizer.transformer.layers: + for layer in policy.temporal_summarizer.transformer.layers: attention = layer.attention - frame_idx = torch.arange(n_tokens) // attention.causal_block_size - mask = (frame_idx[:, None] >= frame_idx[None, :]).view(1, 1, n_tokens, n_tokens) - attention.register_buffer("_supercombo_mask", mask, persistent=False) + 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) # (1, S, C) - features = torch.cat((inputs["features_buffer"], current[:, None]), dim=1) # (1, T, S, C) + 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( @@ -76,5 +74,5 @@ def forward(self, inputs: dict[str, torch.Tensor]) -> torch.Tensor: inputs[ModelInputs.ACTION_T][:, None], ) outputs.update({name: policy_outputs[name] for name in names}) - outputs["hidden_state"] = current.detach().reshape(1, -1) + outputs["hidden_state"] = current.detach() return torch.cat([outputs[name] for name in OUTPUT_ORDER] + [self.pad], dim=1) From d746eaecd39e16f740527aed5d67b9c080935313 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Thu, 20 Aug 2026 21:15:47 -0700 Subject: [PATCH 05/23] Revert "rldriving: train with spatial path models" This reverts commit 930e80f16e02f26a7b59b85e83ec02079222e3dc. --- .../experiments/rldriving/config_registry.py | 1 - torchtitan/experiments/rldriving/model.py | 44 +++++++++++-------- torchtitan/experiments/rldriving/trainer.py | 6 +-- 3 files changed, 28 insertions(+), 23 deletions(-) diff --git a/torchtitan/experiments/rldriving/config_registry.py b/torchtitan/experiments/rldriving/config_registry.py index 2a3a49343e6..61042dff9bc 100644 --- a/torchtitan/experiments/rldriving/config_registry.py +++ b/torchtitan/experiments/rldriving/config_registry.py @@ -183,7 +183,6 @@ def _checkpoint_config( enable=True, checkpoint_base_folder=base_folder, export_onnx=True, - enable_first_step_checkpoint=True, folder=folder, interval=interval, input_names=list(input_shapes), diff --git a/torchtitan/experiments/rldriving/model.py b/torchtitan/experiments/rldriving/model.py index 9aa5a451cff..e8fc395ae95 100644 --- a/torchtitan/experiments/rldriving/model.py +++ b/torchtitan/experiments/rldriving/model.py @@ -7,6 +7,7 @@ from __future__ import annotations import copy +import math from dataclasses import dataclass from typing import Any, cast @@ -22,12 +23,6 @@ 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_config import ( - TEMPORAL_HEADS, - _hydra, - _mlp, - temporal_policy_config, -) from torchtitan.experiments.path.model import ( Hydra, LinearEncoder, @@ -36,13 +31,14 @@ TemporalPolicy, TemporalSummarizer, ) -from torchtitan.models.common import Linear +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, S: spatial tokens, D: model width, A: action components. +# B: batch, T: temporal steps, D: model width, A: action components. ACTION_HEAD_NAME = "action" Q_HEAD_NAME = "q" @@ -80,10 +76,10 @@ def __init__(self, config: Config): self.q_hydra = config.q_hydra.build() def forward(self, inputs: TemporalInputs, action: torch.Tensor) -> torch.Tensor: - features_BTSD = inputs[ModelInputs.FEATURES] - dtype = features_BTSD.dtype + features_BTD = inputs[ModelInputs.FEATURES] + dtype = features_BTD.dtype critic_features_BD = self.temporal_summarizer( - features_BTSD[:, self.history_idxs], + features_BTD[:, self.history_idxs], inputs[ModelInputs.DESIRE].to(dtype), inputs[ModelInputs.TRAFFIC][:, -1].to(dtype), inputs[ModelInputs.ACTION_T][:, -1].to(dtype), @@ -101,7 +97,15 @@ def actor_config() -> TemporalPolicy.Config: def critic_config(actor: TemporalPolicy.Config) -> Critic.Config: - dim = actor.temporal_summarizer.temporal_pos_embedding.embedding_dim + 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, @@ -109,9 +113,14 @@ def critic_config(actor: TemporalPolicy.Config) -> Critic.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=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - post_action_mlp2=_mlp(dim, mlp_mult=2, bias=False, dropout=0.0), - q_hydra=_hydra((PathHead(name=Q_HEAD_NAME, output_size=1, mlp=False, scale=False),), in_features=dim, mlp_mult=2), + 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={}, + ), ) @@ -190,8 +199,7 @@ def input_shapes( ModelInputs.FEATURES: ( batch_size, temporal_len, - summarizer.spatial_size, - summarizer.temporal_pos_embedding.embedding_dim, + summarizer.pos_embedding.embedding_dim, ), ModelInputs.DESIRE: (batch_size, temporal_len, desire_dim), ModelInputs.TRAFFIC: ( @@ -254,7 +262,7 @@ def parallelize_rldriving( training: TrainingConfig, parallelism: ParallelismConfig, compile_config: CompileConfig, - ac_config: ActivationCheckpointingConfig | None, + ac_config: ActivationCheckpointingConfig, dump_folder: str, ) -> RLDrivingModel: if compile_config.enable and "model" in compile_config.components: diff --git a/torchtitan/experiments/rldriving/trainer.py b/torchtitan/experiments/rldriving/trainer.py index ccffae16911..fe7ab16a7d4 100644 --- a/torchtitan/experiments/rldriving/trainer.py +++ b/torchtitan/experiments/rldriving/trainer.py @@ -198,8 +198,7 @@ def batch_generator(self, data_iterable: Iterable[Batch]) -> Iterator[Batch]: def prepare_batch(self, batch: Batch) -> PreparedBatch: inputs, targets, metadata = batch - # 'info' stays on the host; it is a serialized json column, not a tensor - inputs = {name: value.to(self.device) for name, value in inputs.items() if isinstance(value, torch.Tensor)} + 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} @@ -216,8 +215,7 @@ def train_step(self, data_iterator: Iterator[Batch]) -> None: batch = next(data_iterator) info = batch[0].get("info") if info is not None: - info = info.cpu().numpy() if isinstance(info, torch.Tensor) else info - self.unique_segment_counter.update(parse_info(value)["name"] for value in info) + 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 From 2944f94782ba087797dfe1426652e2cb543673f2 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Thu, 20 Aug 2026 21:47:06 -0700 Subject: [PATCH 06/23] SpatialUnvision: use learned embedding instead of sincos Drop the fixed 2D sincos positional encoding in favor of a learned nn.Embedding, matching the rest of the path model. Removes the dependency on xx.training.lib.positional_embeddings. --- torchtitan/experiments/path/model.py | 13 +++++-------- 1 file changed, 5 insertions(+), 8 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 81cea7db2e9..0b852e05bd6 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.training.lib.positional_embeddings import get_2d_sincos_pos_embed import torch import torch.nn as nn @@ -198,17 +197,15 @@ def __init__(self, config: Config): out_features=self.OUTPUT_CHANNELS * (self.patch_size[0] * self.patch_size[1]), bias=True, ).build() - self.register_buffer("pos_embedding", torch.empty(1, grid_h * grid_w, dim), persistent=True) - - def _init_self_buffers(self, *, buffer_device: torch.device | None = None) -> None: - device = buffer_device if buffer_device is not None else self.pos_embedding.device - embedding = get_2d_sincos_pos_embed(self.N_EMBD, self.config.grid_size) - self.pos_embedding = torch.from_numpy(embedding).to(device=device, dtype=torch.float32).unsqueeze(0) + self.pos_embedding = Embedding.Config( + num_embeddings=grid_h * grid_w, + embedding_dim=dim, + ).build() def forward(self, features: torch.Tensor) -> dict[str, torch.Tensor]: grid_h, grid_w = self.config.grid_size tokens = self.input_norm(self.input_projection(features)) - tokens = tokens + self.pos_embedding.to(tokens.dtype) + tokens = tokens + self.pos_embedding(torch.arange(grid_h * grid_w, device=features.device)) tokens = self.output_projection(self.output_norm(self.transformer(tokens))) patch_h, patch_w = self.patch_size images = rearrange( From ff18345c3c6bff5789937a2f9d6e407bce2256d9 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 14:47:59 -0700 Subject: [PATCH 07/23] path: drop block-causal attention, use plain causal MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Spatial tokens in the temporal summarizer now attend causally over the time-major (t*s) sequence instead of block-causal. Slightly weird but simpler — removes causal_block_size and the hand-rolled mask. --- torchtitan/experiments/path/model.py | 14 ++------------ torchtitan/experiments/path/model_config.py | 3 --- 2 files changed, 2 insertions(+), 15 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 0b852e05bd6..187bd08c4af 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -10,7 +10,6 @@ import torch import torch.nn as nn -import torch.nn.functional as F from einops import rearrange from torch.distributed.device_mesh import DeviceMesh from torch.distributed.fsdp import CPUOffloadPolicy, fully_shard, MixedPrecisionPolicy @@ -94,14 +93,12 @@ class Config(Module.Config): head_dim: int dropout: float is_causal: bool = True - causal_block_size: int = 1 def __init__(self, config: Config): super().__init__() self.n_head = config.n_head self.head_dim = config.head_dim self.is_causal = config.is_causal - self.causal_block_size = config.causal_block_size self.norm = config.norm.build() self.q_norm = config.q_norm.build() if config.q_norm is not None else nn.Identity() self.k_norm = config.k_norm.build() if config.k_norm is not None else nn.Identity() @@ -115,15 +112,8 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) q, k, v = qkv.unbind(2) q, k = self.q_norm(q), self.k_norm(k) - if self.is_causal and self.causal_block_size > 1: - # Tokens are time-major. Let every spatial token attend within its frame - # while retaining causal attention between frames. - frame_idx = torch.arange(t, device=x.device) // self.causal_block_size - attention_mask = frame_idx[:, None] >= frame_idx[None, :] - q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2) - x = F.scaled_dot_product_attention(q, k, v, attn_mask=attention_mask).transpose(1, 2) - else: - x = self.inner_attention(q, k, v, is_causal=self.is_causal) + # slightly weird to do causal instead of block_causal for spatial tokens + x = self.inner_attention(q, k, v, is_causal=self.is_causal) return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index b58eb474f9a..dac297aca5e 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -133,7 +133,6 @@ def temporal_policy_config( dim=vision_features, n_head=8, dropout=dropout, - causal_block_size=spatial_size, ), mlp=_mlp(vision_features, mlp_mult=2, bias=True, dropout=dropout), ) @@ -201,7 +200,6 @@ def _attention( n_head: int, dropout: float, is_causal: bool = True, - causal_block_size: int = 1, ) -> PathSelfAttention.Config: head_dim = dim // n_head return PathSelfAttention.Config( @@ -215,7 +213,6 @@ def _attention( head_dim=head_dim, dropout=dropout, is_causal=is_causal, - causal_block_size=causal_block_size, ) From 1717f2ae2e704e17a1ec9e3fb1fbceb58801fda2 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:02:17 -0700 Subject: [PATCH 08/23] path: clean up unnecessary reformatting --- torchtitan/experiments/path/model_config.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index dac297aca5e..6f99abe4052 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -129,11 +129,7 @@ def temporal_policy_config( transformer=PathTransformer.Config( layers=[ PathTransformerBlock.Config( - attention=_attention( - dim=vision_features, - n_head=8, - dropout=dropout, - ), + 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) From 66834b6a00ef6c5a8ab5ae480b324ffbc59866d7 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:13:37 -0700 Subject: [PATCH 09/23] SpatialUnvision: replace transformer with conv decoder Drop the 4-layer transformer decoder in favor of a simple conv upsampler that takes the spatial feature grid (b, s, c) -> (b, c, grid_h, grid_w) and upsamples to two RGB views. Matches the xx TinyUnvision approach but keeps the spatial tokens instead of pooling. --- torchtitan/experiments/path/model.py | 60 ++++++++------------- torchtitan/experiments/path/model_config.py | 9 ---- 2 files changed, 23 insertions(+), 46 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 187bd08c4af..d247b860efb 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -157,56 +157,44 @@ def apply_fsdp(self, shard, reshard_after_forward: bool) -> None: class SpatialUnvision(Module): - """Training-only decoder from a spatial ConvNeXt token grid to two RGB views.""" + """Training-only conv decoder from a spatial ConvNeXt token grid to two RGB views.""" OUTPUT_SIZE = (128, 256) OUTPUT_CHANNELS = 6 - N_EMBD = 256 - N_HEAD = 8 - N_LAYER = 4 @dataclass(kw_only=True, slots=True) class Config(Module.Config): in_features: int grid_size: tuple[int, int] - transformer: PathTransformer.Config def __init__(self, config: Config): super().__init__() self.config = config grid_h, grid_w = config.grid_size - output_h, output_w = self.OUTPUT_SIZE - self.patch_size = (output_h // grid_h, output_w // grid_w) - dim = self.N_EMBD - self.input_projection = Linear.Config(in_features=config.in_features, out_features=dim, bias=True).build() - self.input_norm = LayerNorm.Config(normalized_shape=dim).build() - self.transformer = config.transformer.build() - self.output_norm = LayerNorm.Config(normalized_shape=dim).build() - self.output_projection = Linear.Config( - in_features=dim, - out_features=self.OUTPUT_CHANNELS * (self.patch_size[0] * self.patch_size[1]), - bias=True, - ).build() - self.pos_embedding = Embedding.Config( - num_embeddings=grid_h * grid_w, - embedding_dim=dim, - ).build() + dim = config.in_features + self.proj = nn.ConvTranspose2d(dim, dim, kernel_size=4, stride=2, padding=1, bias=False) + self.norm = nn.BatchNorm2d(dim, eps=0.001, momentum=0.01) + self.act = nn.ReLU(inplace=True) + self.upsample = nn.Sequential( + nn.ConvTranspose2d(dim, dim // 2, 4, stride=2, padding=1, bias=False), + nn.BatchNorm2d(dim // 2, eps=0.001, momentum=0.01), + nn.ReLU(inplace=True), + nn.ConvTranspose2d(dim // 2, dim // 4, 4, stride=2, padding=1, bias=False), + nn.BatchNorm2d(dim // 4, eps=0.001, momentum=0.01), + nn.ReLU(inplace=True), + nn.ConvTranspose2d(dim // 4, dim // 8, 4, stride=2, padding=1, bias=False), + nn.BatchNorm2d(dim // 8, eps=0.001, momentum=0.01), + nn.ReLU(inplace=True), + nn.Conv2d(dim // 8, self.OUTPUT_CHANNELS, 7, stride=1, padding=3), + ) def forward(self, features: torch.Tensor) -> dict[str, torch.Tensor]: grid_h, grid_w = self.config.grid_size - tokens = self.input_norm(self.input_projection(features)) - tokens = tokens + self.pos_embedding(torch.arange(grid_h * grid_w, device=features.device)) - tokens = self.output_projection(self.output_norm(self.transformer(tokens))) - patch_h, patch_w = self.patch_size - images = rearrange( - tokens, - "b (grid_h grid_w) (c patch_h patch_w) -> b c (grid_h patch_h) (grid_w patch_w)", - grid_h=grid_h, - grid_w=grid_w, - patch_h=patch_h, - patch_w=patch_w, - ) - return {"imgs": ((images + 1.0) / 2.0) * 255.0} + # (b, s, c) -> (b, c, grid_h, grid_w) + x = features.transpose(1, 2).reshape(features.shape[0], -1, grid_h, grid_w) + x = self.act(self.norm(self.proj(x))) + x = self.upsample(x) + return {"imgs": ((x + 1.0) / 2.0) * 255.0} class PointSummarizer(Module): @@ -715,9 +703,6 @@ def wrap(module: nn.Module, fqn: str) -> nn.Module: wrap, "temporal_policy.temporal_summarizer.transformer", ) - if model.unvision is not None: - model.unvision.transformer.apply_activation_checkpointing(wrap, "unvision.transformer") - logger.info(f"Applied {mode} activation checkpointing to the path model") @@ -729,6 +714,7 @@ def _apply_compile(model: PathModel, compile_config: CompileConfig) -> None: model.temporal_policy.compile(backend=compile_config.backend) if model.unvision is not None: model.unvision.compile(backend=compile_config.backend) + logger.info("Compiling path model components with torch.compile") diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index 6f99abe4052..82debc93136 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -157,18 +157,9 @@ def _spatial_unvision_config( in_features: int, grid_size: tuple[int, int], ) -> SpatialUnvision.Config: - dim = SpatialUnvision.N_EMBD - layers = [ - PathTransformerBlock.Config( - attention=_attention(dim=dim, n_head=SpatialUnvision.N_HEAD, dropout=0.0, is_causal=False), - mlp=_mlp(dim=dim, mlp_mult=8 / 3, bias=False, dropout=0.0), - ) - for _ in range(SpatialUnvision.N_LAYER) - ] return SpatialUnvision.Config( in_features=in_features, grid_size=grid_size, - transformer=PathTransformer.Config(layers=layers), ) From c2821899f283162ca5ee04fb30f6a73b5fae4dd3 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:18:35 -0700 Subject: [PATCH 10/23] path: always enable unvision, remove the option MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Unvision is no longer optional — the decoder is always built and always runs. Removes the unvision flag from model config, config registry, and dataset config. --- torchtitan/experiments/path/config_registry.py | 13 ++++--------- torchtitan/experiments/path/dataset.py | 1 - torchtitan/experiments/path/model.py | 9 +++------ torchtitan/experiments/path/model_config.py | 3 +-- 4 files changed, 8 insertions(+), 18 deletions(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index a93c3219e0b..956df0cdc10 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -37,11 +37,11 @@ from .validate import PathValidator -def model_registry(flavor: str, *, unvision: bool = False) -> ModelSpec: +def model_registry(flavor: str) -> ModelSpec: return ModelSpec( name="path", flavor=flavor, - model=_model_config(flavor, unvision=unvision), + model=_model_config(flavor), parallelize_fn=parallelize_path, pipelining_fn=None, post_optimizer_build_fn=None, @@ -56,7 +56,7 @@ def _dp_degrees() -> tuple[int, int]: return num_nodes, local_world_size -def _path(flavor: str, *, unvision: bool = False) -> PathTrainer.Config: +def _path(flavor: str) -> PathTrainer.Config: steps = 1024 * 55 validation_freq = 1024 reports = { @@ -80,7 +80,7 @@ def _path(flavor: str, *, unvision: bool = False) -> PathTrainer.Config: plan_only = False return PathTrainer.Config( loss=PathLoss.Config(), - model_spec=model_registry(flavor, unvision=unvision), + model_spec=model_registry(flavor), tokenizer=NoOpTokenizer.Config(), dataloader=_dataloader_config( dataset=DEFAULT_TRAIN_LIST, @@ -92,7 +92,6 @@ def _path(flavor: str, *, unvision: bool = False) -> PathTrainer.Config: pipeline_dir=BASE_DIR_GT, skip=1, val_skip=1, - unvision=unvision, ), optimizer=_optimizer_config(), lr_scheduler=LRSchedulersContainer.Config( @@ -144,7 +143,6 @@ def _path(flavor: str, *, unvision: bool = False) -> PathTrainer.Config: pipeline_dir=BASE_DIR_GT, skip=1, val_skip=6, - unvision=unvision, ), mixed_precision_param=mixed_precision_param, reports=reports, @@ -164,7 +162,6 @@ def _dataloader_config( pipeline_dir: str, skip: int, val_skip: int, - unvision: bool | None = None, ) -> PathDataLoader.Config: return PathDataLoader.Config( dataset=dataset, @@ -176,7 +173,6 @@ def _dataloader_config( limit=limit, skip=skip, val_skip=val_skip, - unvision=unvision, ) @@ -250,4 +246,3 @@ def _optimizer_config() -> OptimizersContainer.Config: convnext_thirdxxl = partial(_path, "convnext_thirdxxl") convnext_base = partial(_path, "convnext_base") convnext_xxlarge = partial(_path, "convnext_xxlarge") -convnext_xxlarge_unvision = partial(_path, "convnext_xxlarge", unvision=True) diff --git a/torchtitan/experiments/path/dataset.py b/torchtitan/experiments/path/dataset.py index f904b2f449f..b74a527cbd0 100644 --- a/torchtitan/experiments/path/dataset.py +++ b/torchtitan/experiments/path/dataset.py @@ -35,7 +35,6 @@ class Config(BaseDataLoader.Config): 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 diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index d247b860efb..6f817118a75 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -499,7 +499,6 @@ class Config(BaseModel.Config): vision: Vision.Config point_policy: Policy.Config temporal_policy: TemporalPolicy.Config - unvision: bool unvision_decoder: SpatialUnvision.Config def update_from_config(self, *, config, **kwargs) -> None: @@ -533,7 +532,7 @@ def __init__(self, config: Config): self.vision = config.vision.build() self.point_policy = config.point_policy.build() self.temporal_policy = config.temporal_policy.build() - self.unvision = config.unvision_decoder.build() if config.unvision else None + self.unvision = config.unvision_decoder.build() @staticmethod def input_shapes( @@ -628,8 +627,7 @@ def forward( inputs[ModelInputs.TRAFFIC], inputs[ModelInputs.ACTION_T], ) - if self.unvision is not None: - outputs |= self.unvision(features[:, -1]) + outputs |= self.unvision(features[:, -1]) return outputs @@ -712,8 +710,7 @@ def _apply_compile(model: PathModel, compile_config: CompileConfig) -> None: model.vision.encoder.compile(backend=compile_config.backend) model.point_policy.compile(backend=compile_config.backend) model.temporal_policy.compile(backend=compile_config.backend) - if model.unvision is not None: - model.unvision.compile(backend=compile_config.backend) + model.unvision.compile(backend=compile_config.backend) logger.info("Compiling path model components with torch.compile") diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index 82debc93136..ad8610c9731 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -55,7 +55,7 @@ def _spatial_size(frame_constants: dict) -> int: return math.prod(_vision_grid_size(frame_constants)) -def model_config(flavor: str = "convnext_xxlarge", *, unvision: bool = False) -> PathModel.Config: +def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: vision_features = VISION_FEATURES frame_constants = frame_constants_from_fps(n_frames=N_FRAMES, frame_type=FRAME_TYPE) input_frame_names = tuple(INPUT_FRAMES_NAMES) @@ -99,7 +99,6 @@ def model_config(flavor: str = "convnext_xxlarge", *, unvision: bool = False) -> hydra=_hydra(POINT_HEADS, in_features=vision_features, mlp_mult=2), ), temporal_policy=temporal_policy_config(), - unvision=unvision, unvision_decoder=_spatial_unvision_config(in_features=vision_features, grid_size=grid_size), ) From 71c82ce9657adf7014717cb0c557f15cf171957f Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:23:18 -0700 Subject: [PATCH 11/23] path: move spatial shape into TEMPORAL_INPUTS TEMPORAL_INPUTS[FEATURES] is now (spatial_size, vision_features) instead of (512,). The checkpoint config uses TEMPORAL_INPUTS directly instead of hardcoding spatial_size and VISION_FEATURES. Constants moved to model_constants.py. --- .../experiments/path/config_registry.py | 5 ++--- torchtitan/experiments/path/model_config.py | 21 +++++++------------ .../experiments/path/model_constants.py | 5 ++++- 3 files changed, 13 insertions(+), 18 deletions(-) diff --git a/torchtitan/experiments/path/config_registry.py b/torchtitan/experiments/path/config_registry.py index 956df0cdc10..fe1362cd0ef 100644 --- a/torchtitan/experiments/path/config_registry.py +++ b/torchtitan/experiments/path/config_registry.py @@ -23,7 +23,7 @@ from .dataset import PathDataLoader from .loss import PathLoss from .model import parallelize_path -from .model_config import model_config as _model_config, VISION_FEATURES, _spatial_size +from .model_config import model_config as _model_config from .model_constants import ( frame_constants_from_fps, FRAME_TYPE, @@ -179,7 +179,6 @@ def _dataloader_config( def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnxCheckpointManager.Config: frame_constants = frame_constants_from_fps(n_frames=N_FRAMES, frame_type=FRAME_TYPE) temporal_len = frame_constants["temporal_len"] - spatial_size = _spatial_size(frame_constants) vision_input_names = [ModelInputs.IMG, ModelInputs.BIG_IMG] temporal_policy_input_names = [ ModelInputs.FEATURES, @@ -194,7 +193,7 @@ def _checkpoint_config(folder: str, base_folder: str, interval: int) -> PathOnnx input_shapes = [ [1, *frame_constants["frame_shapes"][ModelInputs.IMG]], [1, *frame_constants["frame_shapes"][ModelInputs.BIG_IMG]], - [1, temporal_len, spatial_size, VISION_FEATURES], + [1, temporal_len, *TEMPORAL_INPUTS[ModelInputs.FEATURES]], [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.DESIRE][0]], [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.TRAFFIC][0]], [1, temporal_len, TEMPORAL_INPUTS[ModelInputs.ACTION_T][0]], diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index ad8610c9731..8b21d4a6cde 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -35,32 +35,25 @@ INPUT_FRAMES_NAMES, ModelInputs, N_FRAMES, + SPATIAL_SIZE, TEMPORAL_INPUTS, + VISION_FEATURES, ) -VISION_FEATURES = 512 -VISION_OUTPUT_STRIDE = 32 - POINT_HEADS = tuple(META_HEADS + POSE_HEADS) TEMPORAL_HEADS = tuple(DRIVING_HEADS + TEMPORAL_META_HEADS) -def _vision_grid_size(frame_constants: dict) -> tuple[int, int]: - height, width = frame_constants["frame_shapes"][INPUT_FRAMES_NAMES[0]][-2:] - return height // VISION_OUTPUT_STRIDE, width // VISION_OUTPUT_STRIDE - - -def _spatial_size(frame_constants: dict) -> int: - return math.prod(_vision_grid_size(frame_constants)) - - def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: vision_features = VISION_FEATURES 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) - grid_size = _vision_grid_size(frame_constants) + grid_size = ( + frame_constants["frame_shapes"][INPUT_FRAMES_NAMES[0]][-2] // 32, + frame_constants["frame_shapes"][INPUT_FRAMES_NAMES[0]][-1] // 32, + ) spatial_size = math.prod(grid_size) return PathModel.Config( @@ -115,7 +108,7 @@ def temporal_policy_config( desire_window_len = frame_constants["desire_window_len"] desire_window_starts = tuple(index - history_idxs[0] for index in history_idxs) block_size = len(history_idxs) - spatial_size = _spatial_size(frame_constants) + spatial_size = SPATIAL_SIZE return TemporalPolicy.Config( temporal_summarizer=TemporalSummarizer.Config( mlp1=_mlp(vision_features, mlp_mult=2, bias=False, dropout=0.0), diff --git a/torchtitan/experiments/path/model_constants.py b/torchtitan/experiments/path/model_constants.py index 4dbec10c3ad..78d39810a8b 100644 --- a/torchtitan/experiments/path/model_constants.py +++ b/torchtitan/experiments/path/model_constants.py @@ -44,6 +44,9 @@ def index_function(idx: int, max_val: float = 192) -> float: META_LEN = 55 DESIRE_LEN = 8 ACTION_LEN = 2 +VISION_FEATURES = 512 +VISION_OUTPUT_STRIDE = 32 +SPATIAL_SIZE = (H // VISION_OUTPUT_STRIDE) * (W // VISION_OUTPUT_STRIDE) LEAD_PRED_DIM = 4 LEAD_TRAJECTORY_DIM = 6 * 4 @@ -108,7 +111,7 @@ class ModelInputs: } TEMPORAL_INPUTS = { - ModelInputs.FEATURES: (512,), + ModelInputs.FEATURES: (SPATIAL_SIZE, VISION_FEATURES), ModelInputs.DESIRE: (DESIRE_LEN,), ModelInputs.TRAFFIC: (2,), ModelInputs.ACTION_T: (ACTION_LEN,), From f23aa8a13fd376e09564cf360ef8506b28380c5c Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:27:36 -0700 Subject: [PATCH 12/23] path: restore _attention to original signature MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Don't add is_causal parameter to _attention — the one non-causal caller (PointSummarizer) builds the config inline with is_causal=False instead. --- torchtitan/experiments/path/model_config.py | 22 ++++++++++++--------- 1 file changed, 13 insertions(+), 9 deletions(-) diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index 8b21d4a6cde..79177645dbb 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -77,7 +77,18 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: transformer=PathTransformer.Config( layers=[ PathTransformerBlock.Config( - attention=_attention(dim=vision_features, n_head=8, dropout=0.0, is_causal=False), + attention=PathSelfAttention.Config( + norm=LayerNorm.Config(normalized_shape=vision_features), + q_norm=LayerNorm.Config(normalized_shape=vision_features // 8), + k_norm=LayerNorm.Config(normalized_shape=vision_features // 8), + c_attn=Linear.Config(in_features=vision_features, out_features=3 * vision_features, bias=True), + c_proj=Linear.Config(in_features=vision_features, out_features=vision_features, bias=True), + inner_attention=ScaledDotProductAttention.Config(), + n_head=8, + head_dim=vision_features // 8, + dropout=0.0, + is_causal=False, + ), mlp=_mlp(vision_features, mlp_mult=2, bias=True, dropout=0.0), ) for _ in range(2) @@ -173,13 +184,7 @@ def _encoder(in_features: int, dim: int) -> LinearEncoder.Config: ) -def _attention( - *, - dim: int, - n_head: int, - dropout: float, - is_causal: bool = True, -) -> PathSelfAttention.Config: +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), @@ -191,7 +196,6 @@ def _attention( n_head=n_head, head_dim=head_dim, dropout=dropout, - is_causal=is_causal, ) From c71b14eeb35367baba7a72ebffed4dbbc16c21eb Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:30:05 -0700 Subject: [PATCH 13/23] path: use VISION_FEATURES directly instead of local alias --- torchtitan/experiments/path/model_config.py | 46 ++++++++++----------- 1 file changed, 22 insertions(+), 24 deletions(-) diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index 79177645dbb..0c08b7e4be4 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -46,7 +46,6 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: - vision_features = VISION_FEATURES 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) @@ -64,7 +63,7 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: flavor=flavor, input_frame_names=input_frame_names, in_channels=in_channels, - vision_features=vision_features, + vision_features=VISION_FEATURES, grid_size=grid_size, pretrained=True, drop_path_rate=0.2, @@ -73,37 +72,37 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: ), point_policy=Policy.Config( summarizer=PointSummarizer.Config( - mlp1=_mlp(vision_features, mlp_mult=2, bias=False, dropout=0.0), + mlp1=_mlp(VISION_FEATURES, mlp_mult=2, bias=False, dropout=0.0), transformer=PathTransformer.Config( layers=[ PathTransformerBlock.Config( attention=PathSelfAttention.Config( - norm=LayerNorm.Config(normalized_shape=vision_features), - q_norm=LayerNorm.Config(normalized_shape=vision_features // 8), - k_norm=LayerNorm.Config(normalized_shape=vision_features // 8), - c_attn=Linear.Config(in_features=vision_features, out_features=3 * vision_features, bias=True), - c_proj=Linear.Config(in_features=vision_features, out_features=vision_features, bias=True), + norm=LayerNorm.Config(normalized_shape=VISION_FEATURES), + q_norm=LayerNorm.Config(normalized_shape=VISION_FEATURES // 8), + k_norm=LayerNorm.Config(normalized_shape=VISION_FEATURES // 8), + c_attn=Linear.Config(in_features=VISION_FEATURES, out_features=3 * VISION_FEATURES, bias=True), + c_proj=Linear.Config(in_features=VISION_FEATURES, out_features=VISION_FEATURES, bias=True), inner_attention=ScaledDotProductAttention.Config(), n_head=8, - head_dim=vision_features // 8, + head_dim=VISION_FEATURES // 8, dropout=0.0, is_causal=False, ), - mlp=_mlp(vision_features, mlp_mult=2, bias=True, dropout=0.0), + mlp=_mlp(VISION_FEATURES, mlp_mult=2, bias=True, dropout=0.0), ) for _ in range(2) ] ), pos_embedding=Embedding.Config( num_embeddings=spatial_size, - embedding_dim=vision_features, + embedding_dim=VISION_FEATURES, ), spatial_size=spatial_size, ), - hydra=_hydra(POINT_HEADS, in_features=vision_features, mlp_mult=2), + hydra=_hydra(POINT_HEADS, in_features=VISION_FEATURES, mlp_mult=2), ), temporal_policy=temporal_policy_config(), - unvision_decoder=_spatial_unvision_config(in_features=vision_features, grid_size=grid_size), + unvision_decoder=_spatial_unvision_config(in_features=VISION_FEATURES, grid_size=grid_size), ) @@ -113,7 +112,6 @@ def temporal_policy_config( dropout: float = 0.1, dense_training_outputs: bool = True, ) -> TemporalPolicy.Config: - vision_features = VISION_FEATURES 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"] @@ -122,35 +120,35 @@ def temporal_policy_config( spatial_size = SPATIAL_SIZE 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), + 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), + 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), + 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) ] ), temporal_pos_embedding=Embedding.Config( num_embeddings=block_size, - embedding_dim=vision_features, + embedding_dim=VISION_FEATURES, ), spatial_pos_embedding=Embedding.Config( num_embeddings=spatial_size, - embedding_dim=vision_features, + embedding_dim=VISION_FEATURES, ), temporal_size=block_size, spatial_size=spatial_size, dense_training_outputs=dense_training_outputs, ), - temporal_hydra=_hydra(heads, in_features=vision_features, mlp_mult=2), + temporal_hydra=_hydra(heads, in_features=VISION_FEATURES, mlp_mult=2), history_idxs=history_idxs, ) From b69f23c87b16c847dff18e6d6ca1c711ada15401 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:37:17 -0700 Subject: [PATCH 14/23] path: restore is_causal on _attention for non-causal point policy --- torchtitan/experiments/path/model_config.py | 16 +++------------- 1 file changed, 3 insertions(+), 13 deletions(-) diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index 0c08b7e4be4..1fbc3abca97 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -76,18 +76,7 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: transformer=PathTransformer.Config( layers=[ PathTransformerBlock.Config( - attention=PathSelfAttention.Config( - norm=LayerNorm.Config(normalized_shape=VISION_FEATURES), - q_norm=LayerNorm.Config(normalized_shape=VISION_FEATURES // 8), - k_norm=LayerNorm.Config(normalized_shape=VISION_FEATURES // 8), - c_attn=Linear.Config(in_features=VISION_FEATURES, out_features=3 * VISION_FEATURES, bias=True), - c_proj=Linear.Config(in_features=VISION_FEATURES, out_features=VISION_FEATURES, bias=True), - inner_attention=ScaledDotProductAttention.Config(), - n_head=8, - head_dim=VISION_FEATURES // 8, - dropout=0.0, - is_causal=False, - ), + attention=_attention(dim=VISION_FEATURES, n_head=8, dropout=0.0, is_causal=False), mlp=_mlp(VISION_FEATURES, mlp_mult=2, bias=True, dropout=0.0), ) for _ in range(2) @@ -182,7 +171,7 @@ def _encoder(in_features: int, dim: int) -> LinearEncoder.Config: ) -def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Config: +def _attention(*, dim: int, n_head: int, dropout: float, is_causal: bool = True) -> PathSelfAttention.Config: head_dim = dim // n_head return PathSelfAttention.Config( norm=LayerNorm.Config(normalized_shape=dim), @@ -194,6 +183,7 @@ def _attention(*, dim: int, n_head: int, dropout: float) -> PathSelfAttention.Co n_head=n_head, head_dim=head_dim, dropout=dropout, + is_causal=is_causal, ) From a7c96866fd08248150d137c30172740d7a583276 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:38:27 -0700 Subject: [PATCH 15/23] path: remove added docstrings --- torchtitan/experiments/path/model.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 6f817118a75..02600500919 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -157,7 +157,6 @@ def apply_fsdp(self, shard, reshard_after_forward: bool) -> None: class SpatialUnvision(Module): - """Training-only conv decoder from a spatial ConvNeXt token grid to two RGB views.""" OUTPUT_SIZE = (128, 256) OUTPUT_CHANNELS = 6 @@ -198,7 +197,6 @@ def forward(self, features: torch.Tensor) -> dict[str, torch.Tensor]: class PointSummarizer(Module): - """Collapse each spatial feature grid with a learned CLS token.""" @dataclass(kw_only=True, slots=True) class Config(Module.Config): From c777d4ed95ab93abfb65b58f224e6d92d70d14b0 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 15:41:10 -0700 Subject: [PATCH 16/23] path: restore transformer unvision, keep always-on Revert the conv decoder back to the 4-layer transformer SpatialUnvision with learned positional embedding. Unvision stays always-on (no option). --- torchtitan/experiments/path/model.py | 61 +++++++++++++-------- torchtitan/experiments/path/model_config.py | 9 +++ 2 files changed, 47 insertions(+), 23 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 02600500919..a2582fa1bc5 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -160,40 +160,52 @@ class SpatialUnvision(Module): OUTPUT_SIZE = (128, 256) OUTPUT_CHANNELS = 6 + N_EMBD = 256 + N_HEAD = 8 + N_LAYER = 4 @dataclass(kw_only=True, slots=True) class Config(Module.Config): in_features: int grid_size: tuple[int, int] + transformer: PathTransformer.Config def __init__(self, config: Config): super().__init__() self.config = config grid_h, grid_w = config.grid_size - dim = config.in_features - self.proj = nn.ConvTranspose2d(dim, dim, kernel_size=4, stride=2, padding=1, bias=False) - self.norm = nn.BatchNorm2d(dim, eps=0.001, momentum=0.01) - self.act = nn.ReLU(inplace=True) - self.upsample = nn.Sequential( - nn.ConvTranspose2d(dim, dim // 2, 4, stride=2, padding=1, bias=False), - nn.BatchNorm2d(dim // 2, eps=0.001, momentum=0.01), - nn.ReLU(inplace=True), - nn.ConvTranspose2d(dim // 2, dim // 4, 4, stride=2, padding=1, bias=False), - nn.BatchNorm2d(dim // 4, eps=0.001, momentum=0.01), - nn.ReLU(inplace=True), - nn.ConvTranspose2d(dim // 4, dim // 8, 4, stride=2, padding=1, bias=False), - nn.BatchNorm2d(dim // 8, eps=0.001, momentum=0.01), - nn.ReLU(inplace=True), - nn.Conv2d(dim // 8, self.OUTPUT_CHANNELS, 7, stride=1, padding=3), - ) + output_h, output_w = self.OUTPUT_SIZE + self.patch_size = (output_h // grid_h, output_w // grid_w) + dim = self.N_EMBD + self.input_projection = Linear.Config(in_features=config.in_features, out_features=dim, bias=True).build() + self.input_norm = LayerNorm.Config(normalized_shape=dim).build() + self.transformer = config.transformer.build() + self.output_norm = LayerNorm.Config(normalized_shape=dim).build() + self.output_projection = Linear.Config( + in_features=dim, + out_features=self.OUTPUT_CHANNELS * (self.patch_size[0] * self.patch_size[1]), + bias=True, + ).build() + self.pos_embedding = Embedding.Config( + num_embeddings=grid_h * grid_w, + embedding_dim=dim, + ).build() def forward(self, features: torch.Tensor) -> dict[str, torch.Tensor]: grid_h, grid_w = self.config.grid_size - # (b, s, c) -> (b, c, grid_h, grid_w) - x = features.transpose(1, 2).reshape(features.shape[0], -1, grid_h, grid_w) - x = self.act(self.norm(self.proj(x))) - x = self.upsample(x) - return {"imgs": ((x + 1.0) / 2.0) * 255.0} + tokens = self.input_norm(self.input_projection(features)) + tokens = tokens + self.pos_embedding(torch.arange(grid_h * grid_w, device=features.device)) + tokens = self.output_projection(self.output_norm(self.transformer(tokens))) + patch_h, patch_w = self.patch_size + images = rearrange( + tokens, + "b (grid_h grid_w) (c patch_h patch_w) -> b c (grid_h patch_h) (grid_w patch_w)", + grid_h=grid_h, + grid_w=grid_w, + patch_h=patch_h, + patch_w=patch_w, + ) + return {"imgs": ((images + 1.0) / 2.0) * 255.0} class PointSummarizer(Module): @@ -699,6 +711,9 @@ def wrap(module: nn.Module, fqn: str) -> nn.Module: wrap, "temporal_policy.temporal_summarizer.transformer", ) + if model.unvision is not None: + model.unvision.transformer.apply_activation_checkpointing(wrap, "unvision.transformer") + logger.info(f"Applied {mode} activation checkpointing to the path model") @@ -708,8 +723,8 @@ def _apply_compile(model: PathModel, compile_config: CompileConfig) -> None: model.vision.encoder.compile(backend=compile_config.backend) model.point_policy.compile(backend=compile_config.backend) model.temporal_policy.compile(backend=compile_config.backend) - model.unvision.compile(backend=compile_config.backend) - + if model.unvision is not None: + model.unvision.compile(backend=compile_config.backend) logger.info("Compiling path model components with torch.compile") diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index 1fbc3abca97..db2426a3bd0 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -147,9 +147,18 @@ def _spatial_unvision_config( in_features: int, grid_size: tuple[int, int], ) -> SpatialUnvision.Config: + dim = SpatialUnvision.N_EMBD + layers = [ + PathTransformerBlock.Config( + attention=_attention(dim=dim, n_head=SpatialUnvision.N_HEAD, dropout=0.0, is_causal=False), + mlp=_mlp(dim=dim, mlp_mult=8 / 3, bias=False, dropout=0.0), + ) + for _ in range(SpatialUnvision.N_LAYER) + ] return SpatialUnvision.Config( in_features=in_features, grid_size=grid_size, + transformer=PathTransformer.Config(layers=layers), ) From c00b7b4a37655c352c027ce2ea5fcf1d3c92c095 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 16:44:18 -0700 Subject: [PATCH 17/23] path: complete spatial runtime compatibility --- torchtitan/experiments/path/model.py | 28 +++++-------------- torchtitan/experiments/path/model_config.py | 22 ++------------- .../experiments/path/model_constants.py | 7 ++++- .../experiments/rldriving/config_registry.py | 3 +- torchtitan/experiments/rldriving/model.py | 13 +++++---- .../experiments/rldriving/supercombo.py | 20 +++++++++---- torchtitan/experiments/rldriving/trainer.py | 6 ++-- 7 files changed, 43 insertions(+), 56 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index a2582fa1bc5..3a16759fba7 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -209,35 +209,19 @@ def forward(self, features: torch.Tensor) -> dict[str, torch.Tensor]: class PointSummarizer(Module): - @dataclass(kw_only=True, slots=True) class Config(Module.Config): mlp1: PathMLP.Config - transformer: PathTransformer.Config - pos_embedding: Embedding.Config - spatial_size: int + mlp2: PathMLP.Config def __init__(self, config: Config): super().__init__() - self.spatial_size = config.spatial_size - self.n_features = config.pos_embedding.embedding_dim self.mlp1 = config.mlp1.build() - self.transformer = config.transformer.build() - self.pos_embedding = config.pos_embedding.build() - self.cls_token = nn.Parameter(torch.empty(1, 1, self.n_features)) - - def reset_parameters(self) -> None: - nn.init.normal_(self.cls_token, std=0.02) + self.mlp2 = config.mlp2.build() def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.mlp1(x) + x - pos = self.pos_embedding(torch.arange(self.spatial_size, device=x.device)) - x = x + pos - leading_shape = x.shape[:-2] - cls = self.cls_token.expand(*leading_shape, 1, -1) - x = torch.cat((cls, x), dim=-2).reshape(-1, self.spatial_size + 1, self.n_features) - x = self.transformer(x) - return x.reshape(*leading_shape, self.spatial_size + 1, self.n_features)[..., 0, :] + return self.mlp2(x) + x class LinearEncoder(Module): @@ -277,7 +261,9 @@ def __init__(self, config: Config): self.spatial_size = config.spatial_size self.dense_training_outputs = config.dense_training_outputs if len(config.desire_window_starts) != self.temporal_size: - raise ValueError(f"Expected {self.temporal_size} desire window starts, got {len(config.desire_window_starts)}") + raise ValueError( + f"Expected {self.temporal_size} desire window starts, got {len(config.desire_window_starts)}" + ) self.desire_window_len = config.desire_window_len self.desire_window_starts = config.desire_window_starts self.register_buffer("desire_window_idxs", self._make_desire_window_idxs(), persistent=False) @@ -631,7 +617,7 @@ def forward( } features = self.vision(vision_inputs) features = rearrange(features, "(b t) s c -> b t s c", b=b, t=t) - outputs = self.point_policy(features) | self.temporal_policy( + outputs = self.point_policy(features.mean(dim=2)) | self.temporal_policy( features, inputs[ModelInputs.DESIRE], inputs[ModelInputs.TRAFFIC], diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index db2426a3bd0..e3bd9bc5f4a 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -38,6 +38,7 @@ SPATIAL_SIZE, TEMPORAL_INPUTS, VISION_FEATURES, + VISION_GRID_SIZE, ) @@ -49,11 +50,7 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: 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) - grid_size = ( - frame_constants["frame_shapes"][INPUT_FRAMES_NAMES[0]][-2] // 32, - frame_constants["frame_shapes"][INPUT_FRAMES_NAMES[0]][-1] // 32, - ) - spatial_size = math.prod(grid_size) + grid_size = VISION_GRID_SIZE return PathModel.Config( n_frames_input=N_FRAMES, @@ -73,20 +70,7 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: point_policy=Policy.Config( summarizer=PointSummarizer.Config( mlp1=_mlp(VISION_FEATURES, mlp_mult=2, bias=False, dropout=0.0), - transformer=PathTransformer.Config( - layers=[ - PathTransformerBlock.Config( - attention=_attention(dim=VISION_FEATURES, n_head=8, dropout=0.0, is_causal=False), - mlp=_mlp(VISION_FEATURES, mlp_mult=2, bias=True, dropout=0.0), - ) - for _ in range(2) - ] - ), - pos_embedding=Embedding.Config( - num_embeddings=spatial_size, - embedding_dim=VISION_FEATURES, - ), - spatial_size=spatial_size, + mlp2=_mlp(VISION_FEATURES, mlp_mult=2, bias=False, dropout=0.0), ), hydra=_hydra(POINT_HEADS, in_features=VISION_FEATURES, mlp_mult=2), ), diff --git a/torchtitan/experiments/path/model_constants.py b/torchtitan/experiments/path/model_constants.py index 78d39810a8b..6611bcae4ae 100644 --- a/torchtitan/experiments/path/model_constants.py +++ b/torchtitan/experiments/path/model_constants.py @@ -46,7 +46,6 @@ def index_function(idx: int, max_val: float = 192) -> float: ACTION_LEN = 2 VISION_FEATURES = 512 VISION_OUTPUT_STRIDE = 32 -SPATIAL_SIZE = (H // VISION_OUTPUT_STRIDE) * (W // VISION_OUTPUT_STRIDE) LEAD_PRED_DIM = 4 LEAD_TRAJECTORY_DIM = 6 * 4 @@ -110,6 +109,12 @@ class ModelInputs: VisionFrameType.YUV: VISION_INPUTS_YUV, } +VISION_GRID_SIZE = ( + VISION_INPUTS[FRAME_TYPE][ModelInputs.IMG][-2] // VISION_OUTPUT_STRIDE, + VISION_INPUTS[FRAME_TYPE][ModelInputs.IMG][-1] // VISION_OUTPUT_STRIDE, +) +SPATIAL_SIZE = VISION_GRID_SIZE[0] * VISION_GRID_SIZE[1] + TEMPORAL_INPUTS = { ModelInputs.FEATURES: (SPATIAL_SIZE, VISION_FEATURES), ModelInputs.DESIRE: (DESIRE_LEN,), diff --git a/torchtitan/experiments/rldriving/config_registry.py b/torchtitan/experiments/rldriving/config_registry.py index 61042dff9bc..15ce8675c64 100644 --- a/torchtitan/experiments/rldriving/config_registry.py +++ b/torchtitan/experiments/rldriving/config_registry.py @@ -83,7 +83,7 @@ def rldriving() -> RLDrivingTrainer.Config: ), warm_start_checkpoint=os.getenv( "RLDRIVING_WARM_START_CHECKPOINT", - "44b83fa5-2a33-7ee7-40f1-e86e3c24ad36/56320", + "849a624a-8a7d-8946-bf04-86148e5e0ef8/44032", ), tokenizer=NoOpTokenizer.Config(), dataloader=RLDrivingDataLoader.Config( @@ -183,6 +183,7 @@ def _checkpoint_config( enable=True, checkpoint_base_folder=base_folder, export_onnx=True, + enable_first_step_checkpoint=True, folder=folder, interval=interval, input_names=list(input_shapes), diff --git a/torchtitan/experiments/rldriving/model.py b/torchtitan/experiments/rldriving/model.py index e8fc395ae95..7904b527006 100644 --- a/torchtitan/experiments/rldriving/model.py +++ b/torchtitan/experiments/rldriving/model.py @@ -38,7 +38,7 @@ from torchtitan.tools.logging import logger -# B: batch, T: temporal steps, D: model width, A: action components. +# B: batch, T: temporal steps, S: spatial tokens, D: model width, A: action components. ACTION_HEAD_NAME = "action" Q_HEAD_NAME = "q" @@ -76,10 +76,10 @@ def __init__(self, config: Config): 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 + features_BTSD = inputs[ModelInputs.FEATURES] + dtype = features_BTSD.dtype critic_features_BD = self.temporal_summarizer( - features_BTD[:, self.history_idxs], + features_BTSD[:, self.history_idxs], inputs[ModelInputs.DESIRE].to(dtype), inputs[ModelInputs.TRAFFIC][:, -1].to(dtype), inputs[ModelInputs.ACTION_T][:, -1].to(dtype), @@ -97,7 +97,7 @@ def actor_config() -> TemporalPolicy.Config: def critic_config(actor: TemporalPolicy.Config) -> Critic.Config: - dim = actor.temporal_summarizer.pos_embedding.embedding_dim + dim = actor.temporal_summarizer.temporal_pos_embedding.embedding_dim hidden = 256 * math.ceil(2 * dim / 256) post_action_mlp = PathMLP.Config( norm=LayerNorm.Config(normalized_shape=dim), @@ -199,7 +199,8 @@ def input_shapes( ModelInputs.FEATURES: ( batch_size, temporal_len, - summarizer.pos_embedding.embedding_dim, + summarizer.spatial_size, + summarizer.temporal_pos_embedding.embedding_dim, ), ModelInputs.DESIRE: (batch_size, temporal_len, desire_dim), ModelInputs.TRAFFIC: ( diff --git a/torchtitan/experiments/rldriving/supercombo.py b/torchtitan/experiments/rldriving/supercombo.py index b3372f13ca5..7d41f003a2c 100644 --- a/torchtitan/experiments/rldriving/supercombo.py +++ b/torchtitan/experiments/rldriving/supercombo.py @@ -43,7 +43,15 @@ def __init__(self) -> None: 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( + spatial_sizes = { + self.vision.config.grid_size[0] * self.vision.config.grid_size[1], + self.off_policy.temporal_summarizer.spatial_size, + self.on_policy.temporal_summarizer.spatial_size, + } + if len(spatial_sizes) != 1: + raise ValueError(f"Supercombo spatial sizes do not match: {sorted(spatial_sizes)}") + spatial_size = spatial_sizes.pop() + output_size = spatial_size * self.vision.config.vision_features + sum( hydra.final_layer[name].out_features for hydra, names in ( (self.point_policy.hydra, VISION_OUTPUT_ORDER), @@ -54,18 +62,18 @@ def __init__(self) -> None: ) self.register_buffer("pad", torch.zeros(1, -output_size % 4), persistent=False) for policy in (self.off_policy, self.on_policy): + summarizer = policy.temporal_summarizer + n_tokens = summarizer.temporal_size * summarizer.spatial_size 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 - ) + mask = torch.ones(1, 1, n_tokens, n_tokens, 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) + outputs = self.point_policy(current.mean(dim=1)) for policy, names in ((self.off_policy, OFF_POLICY_OUTPUT_ORDER), (self.on_policy, ON_POLICY_OUTPUT_ORDER)): policy_outputs = policy( features, @@ -74,5 +82,5 @@ def forward(self, inputs: dict[str, torch.Tensor]) -> torch.Tensor: inputs[ModelInputs.ACTION_T][:, None], ) outputs.update({name: policy_outputs[name] for name in names}) - outputs["hidden_state"] = current.detach() + outputs["hidden_state"] = current.detach().flatten(1) 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 index fe7ab16a7d4..b628153e1cf 100644 --- a/torchtitan/experiments/rldriving/trainer.py +++ b/torchtitan/experiments/rldriving/trainer.py @@ -198,7 +198,8 @@ def batch_generator(self, data_iterable: Iterable[Batch]) -> Iterator[Batch]: def prepare_batch(self, batch: Batch) -> PreparedBatch: inputs, targets, metadata = batch - inputs = {name: value.to(self.device) for name, value in inputs.items()} + # info is a serialized host-side column, not a model tensor. + inputs = {name: value.to(self.device) for name, value in inputs.items() if isinstance(value, torch.Tensor)} 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} @@ -215,7 +216,8 @@ def train_step(self, data_iterator: Iterator[Batch]) -> None: 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()) + info = info.cpu().numpy() if isinstance(info, torch.Tensor) else info + self.unique_segment_counter.update(parse_info(value)["name"] for value in info) current_inputs, next_inputs, targets, metadata = self.prepare_batch(batch) batch_size = next(iter(current_inputs.values())).shape[0] self.ntokens_seen += batch_size From ab62dd52c4e438ef7d2eddf391d94c2996f31c88 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 17:09:03 -0700 Subject: [PATCH 18/23] path: use last causal token as temporal readout --- torchtitan/experiments/path/model.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 3a16759fba7..9c953faba97 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -299,8 +299,8 @@ def forward( feats = self.mlp1(feats) + feats feats = self.mlp2(feats) + feats b, t, s, c = feats.shape - # Temporal attention sees every spatial token. Pool only after attention so - # each dense output retains one entry per frame for the existing heads/losses. + # Tokens are time-major. Under plain causal attention, the last spatial token + # of each frame is its readout because it has seen the complete frame context. feats = feats.reshape(b, t * s, c) desire = self.desire_encoder(self._window_desire(desire)) desire = desire.repeat_interleave(s, dim=1) @@ -312,8 +312,8 @@ def forward( x = feats + rearrange(pos, "ts c -> () ts c") + desire + traffic_convention + action_t x = self.transformer(x) if self.dense_training_outputs: - return x.reshape(b, t, s, c).mean(dim=2) - return x[:, -s:].mean(dim=1) + return x.reshape(b, t, s, c)[:, :, -1] + return x[:, -1] class Hydra(Module): From 782d58d0dafa669816868685b466d0e9fc5c38b8 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 18:11:04 -0700 Subject: [PATCH 19/23] path: simplify spatial vision projection --- torchtitan/experiments/path/model.py | 10 ++-------- 1 file changed, 2 insertions(+), 8 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 9c953faba97..728fd49191a 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -299,8 +299,6 @@ def forward( feats = self.mlp1(feats) + feats feats = self.mlp2(feats) + feats b, t, s, c = feats.shape - # Tokens are time-major. Under plain causal attention, the last spatial token - # of each frame is its readout because it has seen the complete frame context. feats = feats.reshape(b, t * s, c) desire = self.desire_encoder(self._window_desire(desire)) desire = desire.repeat_interleave(s, dim=1) @@ -415,13 +413,10 @@ def __init__(self, config: Config): config.flavor, pretrained=False, in_chans=config.in_channels, - num_classes=0, + num_classes=config.vision_features, global_pool="", drop_path_rate=config.drop_path_rate, ) - # The ConvNeXt head now normalizes without pooling; project its channel-rich - # 2-D map to the policy width at every spatial location. - self.proj = nn.Conv2d(self.encoder.num_features, config.vision_features, kernel_size=1) self.register_buffer("_mean", torch.empty(1, config.in_channels, 1, 1), persistent=True) self.register_buffer("_std", torch.empty(1, config.in_channels, 1, 1), persistent=True) @@ -482,8 +477,7 @@ def forward(self, inputs: dict[str, torch.Tensor]) -> torch.Tensor: dtype = next(self.encoder.parameters()).dtype x = x.to(dtype) x = self.encoder((x - self._mean.to(dtype)) / self._std.to(dtype)) - x = self.proj(x) - return x.flatten(2).transpose(1, 2) + return rearrange(x, "b c h w -> b (h w) c") class PathModel(BaseModel): From c3263d39b79c5ae0871a4ddaae5be880ca52d3cb Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 18:22:05 -0700 Subject: [PATCH 20/23] rldriving: simplify runtime contracts --- torchtitan/experiments/rldriving/supercombo.py | 12 ++---------- torchtitan/experiments/rldriving/trainer.py | 5 ++--- 2 files changed, 4 insertions(+), 13 deletions(-) diff --git a/torchtitan/experiments/rldriving/supercombo.py b/torchtitan/experiments/rldriving/supercombo.py index 7d41f003a2c..41c875b7c87 100644 --- a/torchtitan/experiments/rldriving/supercombo.py +++ b/torchtitan/experiments/rldriving/supercombo.py @@ -5,7 +5,7 @@ # LICENSE file in the root directory of this source tree. from types import MethodType -from xx.ml_tools.constants.model import ModelInputs +from xx.ml_tools.constants.model import ModelInputs, SPATIAL_SIZE import torch @@ -43,15 +43,7 @@ def __init__(self) -> None: self.point_policy = config.point_policy.build() self.off_policy = config.temporal_policy.build() self.on_policy = actor_config().build() - spatial_sizes = { - self.vision.config.grid_size[0] * self.vision.config.grid_size[1], - self.off_policy.temporal_summarizer.spatial_size, - self.on_policy.temporal_summarizer.spatial_size, - } - if len(spatial_sizes) != 1: - raise ValueError(f"Supercombo spatial sizes do not match: {sorted(spatial_sizes)}") - spatial_size = spatial_sizes.pop() - output_size = spatial_size * self.vision.config.vision_features + sum( + output_size = SPATIAL_SIZE * self.vision.config.vision_features + sum( hydra.final_layer[name].out_features for hydra, names in ( (self.point_policy.hydra, VISION_OUTPUT_ORDER), diff --git a/torchtitan/experiments/rldriving/trainer.py b/torchtitan/experiments/rldriving/trainer.py index b628153e1cf..21f2e8ad910 100644 --- a/torchtitan/experiments/rldriving/trainer.py +++ b/torchtitan/experiments/rldriving/trainer.py @@ -199,7 +199,7 @@ def batch_generator(self, data_iterable: Iterable[Batch]) -> Iterator[Batch]: def prepare_batch(self, batch: Batch) -> PreparedBatch: inputs, targets, metadata = batch # info is a serialized host-side column, not a model tensor. - inputs = {name: value.to(self.device) for name, value in inputs.items() if isinstance(value, torch.Tensor)} + inputs = {name: value.to(self.device) for name, value in inputs.items() if name != "info"} 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} @@ -216,8 +216,7 @@ def train_step(self, data_iterator: Iterator[Batch]) -> None: batch = next(data_iterator) info = batch[0].get("info") if info is not None: - info = info.cpu().numpy() if isinstance(info, torch.Tensor) else info - self.unique_segment_counter.update(parse_info(value)["name"] for value in info) + 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 From 2bd9149bc874467c0aba2a12fb52ef2a42468899 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 18:53:14 -0700 Subject: [PATCH 21/23] path: minimize spatial runtime diff --- torchtitan/experiments/path/model.py | 9 ++--- torchtitan/experiments/path/model_config.py | 33 ++++++++++--------- .../experiments/rldriving/config_registry.py | 1 - torchtitan/experiments/rldriving/trainer.py | 3 +- 4 files changed, 20 insertions(+), 26 deletions(-) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index 728fd49191a..a0b0010c1aa 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -112,7 +112,6 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: qkv = self.c_attn(self.norm(x)).view(b, t, 3, self.n_head, self.head_dim) q, k, v = qkv.unbind(2) q, k = self.q_norm(q), self.k_norm(k) - # slightly weird to do causal instead of block_causal for spatial tokens x = self.inner_attention(q, k, v, is_causal=self.is_causal) return self.dropout(self.c_proj(x.reshape(b, t, self.n_head * self.head_dim))) @@ -157,7 +156,6 @@ def apply_fsdp(self, shard, reshard_after_forward: bool) -> None: class SpatialUnvision(Module): - OUTPUT_SIZE = (128, 256) OUTPUT_CHANNELS = 6 N_EMBD = 256 @@ -400,7 +398,6 @@ class Config(Module.Config): input_frame_names: tuple[str, ...] in_channels: int vision_features: int - grid_size: tuple[int, int] pretrained: bool drop_path_rate: float mean: float @@ -691,8 +688,7 @@ def wrap(module: nn.Module, fqn: str) -> nn.Module: wrap, "temporal_policy.temporal_summarizer.transformer", ) - if model.unvision is not None: - model.unvision.transformer.apply_activation_checkpointing(wrap, "unvision.transformer") + model.unvision.transformer.apply_activation_checkpointing(wrap, "unvision.transformer") logger.info(f"Applied {mode} activation checkpointing to the path model") @@ -703,8 +699,7 @@ def _apply_compile(model: PathModel, compile_config: CompileConfig) -> None: model.vision.encoder.compile(backend=compile_config.backend) model.point_policy.compile(backend=compile_config.backend) model.temporal_policy.compile(backend=compile_config.backend) - if model.unvision is not None: - model.unvision.compile(backend=compile_config.backend) + model.unvision.compile(backend=compile_config.backend) logger.info("Compiling path model components with torch.compile") diff --git a/torchtitan/experiments/path/model_config.py b/torchtitan/experiments/path/model_config.py index e3bd9bc5f4a..8da3b05b175 100644 --- a/torchtitan/experiments/path/model_config.py +++ b/torchtitan/experiments/path/model_config.py @@ -47,6 +47,7 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: + vision_features = VISION_FEATURES 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) @@ -60,8 +61,7 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: flavor=flavor, input_frame_names=input_frame_names, in_channels=in_channels, - vision_features=VISION_FEATURES, - grid_size=grid_size, + vision_features=vision_features, pretrained=True, drop_path_rate=0.2, mean=255 / 2, @@ -69,13 +69,13 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config: ), 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), + 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), + hydra=_hydra(POINT_HEADS, in_features=vision_features, mlp_mult=2), ), temporal_policy=temporal_policy_config(), - unvision_decoder=_spatial_unvision_config(in_features=VISION_FEATURES, grid_size=grid_size), + unvision_decoder=_spatial_unvision_config(in_features=vision_features, grid_size=grid_size), ) @@ -85,6 +85,7 @@ def temporal_policy_config( dropout: float = 0.1, dense_training_outputs: bool = True, ) -> TemporalPolicy.Config: + vision_features = VISION_FEATURES 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"] @@ -93,35 +94,35 @@ def temporal_policy_config( spatial_size = SPATIAL_SIZE 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), + 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), + 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), + 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) ] ), temporal_pos_embedding=Embedding.Config( num_embeddings=block_size, - embedding_dim=VISION_FEATURES, + embedding_dim=vision_features, ), spatial_pos_embedding=Embedding.Config( num_embeddings=spatial_size, - embedding_dim=VISION_FEATURES, + embedding_dim=vision_features, ), temporal_size=block_size, spatial_size=spatial_size, dense_training_outputs=dense_training_outputs, ), - temporal_hydra=_hydra(heads, in_features=VISION_FEATURES, mlp_mult=2), + temporal_hydra=_hydra(heads, in_features=vision_features, mlp_mult=2), history_idxs=history_idxs, ) diff --git a/torchtitan/experiments/rldriving/config_registry.py b/torchtitan/experiments/rldriving/config_registry.py index 15ce8675c64..aa5a207b80c 100644 --- a/torchtitan/experiments/rldriving/config_registry.py +++ b/torchtitan/experiments/rldriving/config_registry.py @@ -183,7 +183,6 @@ def _checkpoint_config( enable=True, checkpoint_base_folder=base_folder, export_onnx=True, - enable_first_step_checkpoint=True, folder=folder, interval=interval, input_names=list(input_shapes), diff --git a/torchtitan/experiments/rldriving/trainer.py b/torchtitan/experiments/rldriving/trainer.py index 21f2e8ad910..fe7ab16a7d4 100644 --- a/torchtitan/experiments/rldriving/trainer.py +++ b/torchtitan/experiments/rldriving/trainer.py @@ -198,8 +198,7 @@ def batch_generator(self, data_iterable: Iterable[Batch]) -> Iterator[Batch]: def prepare_batch(self, batch: Batch) -> PreparedBatch: inputs, targets, metadata = batch - # info is a serialized host-side column, not a model tensor. - inputs = {name: value.to(self.device) for name, value in inputs.items() if name != "info"} + 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} From 193a2319b2d66d00abcd6cf65fb19973906ef332 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Fri, 21 Aug 2026 19:07:36 -0700 Subject: [PATCH 22/23] path: shard unvision with FSDP --- torchtitan/experiments/path/model.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/torchtitan/experiments/path/model.py b/torchtitan/experiments/path/model.py index a0b0010c1aa..2dd6d41226f 100644 --- a/torchtitan/experiments/path/model.py +++ b/torchtitan/experiments/path/model.py @@ -738,9 +738,11 @@ def shard(module: nn.Module, reshard: bool) -> None: shard, reshard_after_forward, ) + model.unvision.transformer.apply_fsdp(shard, reshard_after_forward) shard(model.vision.encoder, reshard_after_forward) shard(model.point_policy, reshard_after_forward) shard(model.temporal_policy, reshard_after_forward) + shard(model.unvision, reshard_after_forward) fully_shard(model, **fsdp_config) if enable_symm_mem: From 1ecf95262e68b0e418702196e18ac50320489ed1 Mon Sep 17 00:00:00 2001 From: Bruce Wayne Date: Sat, 22 Aug 2026 22:15:47 -0700 Subject: [PATCH 23/23] fix onnx --- torchtitan/experiments/path/onnx_checkpoint.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/torchtitan/experiments/path/onnx_checkpoint.py b/torchtitan/experiments/path/onnx_checkpoint.py index e5ee56e1743..64dd84e0130 100644 --- a/torchtitan/experiments/path/onnx_checkpoint.py +++ b/torchtitan/experiments/path/onnx_checkpoint.py @@ -29,7 +29,8 @@ def __init__(self, model: PathModel) -> None: def forward(self, inputs: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: features = self.vision(inputs) - outputs = self.point_policy(features) | {"vision_features": features} + point_features = features.float().mean(dim=1).to(features.dtype) + outputs = self.point_policy(point_features) | {"vision_features": features} return {name: value.float() for name, value in outputs.items()}