Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
de7ff6c
path: remove avg-pool, attend over spatial tokens
haraschax Aug 18, 2026
6d9686a
rldriving: train with spatial path models
haraschax Aug 19, 2026
af82fb3
rldriving supercombo: block-causal mask + spatial hidden_state
haraschax Aug 20, 2026
b855acc
Revert "rldriving supercombo: block-causal mask + spatial hidden_state"
haraschax Aug 21, 2026
d746eae
Revert "rldriving: train with spatial path models"
haraschax Aug 21, 2026
2944f94
SpatialUnvision: use learned embedding instead of sincos
haraschax Aug 21, 2026
ff18345
path: drop block-causal attention, use plain causal
haraschax Aug 21, 2026
1717f2a
path: clean up unnecessary reformatting
haraschax Aug 21, 2026
66834b6
SpatialUnvision: replace transformer with conv decoder
haraschax Aug 21, 2026
c282189
path: always enable unvision, remove the option
haraschax Aug 21, 2026
71c82ce
path: move spatial shape into TEMPORAL_INPUTS
haraschax Aug 21, 2026
f23aa8a
path: restore _attention to original signature
haraschax Aug 21, 2026
c71b14e
path: use VISION_FEATURES directly instead of local alias
haraschax Aug 21, 2026
b69f23c
path: restore is_causal on _attention for non-causal point policy
haraschax Aug 21, 2026
a7c9686
path: remove added docstrings
haraschax Aug 21, 2026
c777d4e
path: restore transformer unvision, keep always-on
haraschax Aug 21, 2026
c00b7b4
path: complete spatial runtime compatibility
haraschax Aug 21, 2026
ab62dd5
path: use last causal token as temporal readout
haraschax Aug 22, 2026
782d58d
path: simplify spatial vision projection
haraschax Aug 22, 2026
c3263d3
rldriving: simplify runtime contracts
haraschax Aug 22, 2026
2bd9149
path: minimize spatial runtime diff
haraschax Aug 22, 2026
193a231
path: shard unvision with FSDP
haraschax Aug 22, 2026
1ecf952
fix onnx
haraschax Aug 23, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion torchtitan/experiments/path/config_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -193,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, TEMPORAL_INPUTS[ModelInputs.FEATURES][0]],
[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]],
Expand Down
1 change: 0 additions & 1 deletion torchtitan/experiments/path/dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
100 changes: 87 additions & 13 deletions torchtitan/experiments/path/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -155,6 +155,57 @@ def apply_fsdp(self, shard, reshard_after_forward: bool) -> None:
shard(layer, reshard_after_forward)


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
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
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):
@dataclass(kw_only=True, slots=True)
class Config(Module.Config):
Expand Down Expand Up @@ -196,16 +247,21 @@ 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)
Expand All @@ -215,7 +271,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)
Expand All @@ -228,7 +285,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,
Expand All @@ -239,13 +296,20 @@ def forward(
) -> torch.Tensor:
feats = self.mlp1(feats) + feats
feats = self.mlp2(feats) + feats
b, t, s, c = feats.shape
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)[:, :, -1]
return x[:, -1]


class Hydra(Module):
Expand Down Expand Up @@ -347,6 +411,7 @@ def __init__(self, config: Config):
pretrained=False,
in_chans=config.in_channels,
num_classes=config.vision_features,
global_pool="",
drop_path_rate=config.drop_path_rate,
)
self.register_buffer("_mean", torch.empty(1, config.in_channels, 1, 1), persistent=True)
Expand Down Expand Up @@ -408,7 +473,8 @@ 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))
return rearrange(x, "b c h w -> b (h w) c")


class PathModel(BaseModel):
Expand All @@ -420,6 +486,7 @@ class Config(BaseModel.Config):
vision: Vision.Config
point_policy: Policy.Config
temporal_policy: TemporalPolicy.Config
unvision_decoder: SpatialUnvision.Config

def update_from_config(self, *, config, **kwargs) -> None:
parallelism = config.parallelism
Expand Down Expand Up @@ -452,6 +519,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()

@staticmethod
def input_shapes(
Expand Down Expand Up @@ -539,13 +607,15 @@ 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.mean(dim=2)) | self.temporal_policy(
features,
inputs[ModelInputs.DESIRE],
inputs[ModelInputs.TRAFFIC],
inputs[ModelInputs.ACTION_T],
)
outputs |= self.unvision(features[:, -1])
return outputs


def parallelize_path(
Expand Down Expand Up @@ -618,6 +688,7 @@ def wrap(module: nn.Module, fqn: str) -> nn.Module:
wrap,
"temporal_policy.temporal_summarizer.transformer",
)
model.unvision.transformer.apply_activation_checkpointing(wrap, "unvision.transformer")

logger.info(f"Applied {mode} activation checkpointing to the path model")

Expand All @@ -628,6 +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)
model.unvision.compile(backend=compile_config.backend)
logger.info("Compiling path model components with torch.compile")


Expand Down Expand Up @@ -666,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:
Expand Down
46 changes: 40 additions & 6 deletions torchtitan/experiments/path/model_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@
PointSummarizer,
Policy,
ScaleLayer,
SpatialUnvision,
TemporalPolicy,
TemporalSummarizer,
Vision,
Expand All @@ -34,7 +35,10 @@
INPUT_FRAMES_NAMES,
ModelInputs,
N_FRAMES,
SPATIAL_SIZE,
TEMPORAL_INPUTS,
VISION_FEATURES,
VISION_GRID_SIZE,
)


Expand All @@ -43,10 +47,11 @@


def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config:
vision_features = 512
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

return PathModel.Config(
n_frames_input=N_FRAMES,
Expand All @@ -70,6 +75,7 @@ def model_config(flavor: str = "convnext_xxlarge") -> PathModel.Config:
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),
)


Expand All @@ -79,11 +85,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
return TemporalPolicy.Config(
temporal_summarizer=TemporalSummarizer.Config(
mlp1=_mlp(vision_features, mlp_mult=2, bias=False, dropout=0.0),
Expand All @@ -102,18 +110,43 @@ def temporal_policy_config(
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,
),
block_size=len(history_idxs),
spatial_pos_embedding=Embedding.Config(
num_embeddings=spatial_size,
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),
history_idxs=history_idxs,
)


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(
Expand All @@ -132,7 +165,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),
Expand All @@ -144,6 +177,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,
)


Expand Down
10 changes: 9 additions & 1 deletion torchtitan/experiments/path/model_constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,8 @@ 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

LEAD_PRED_DIM = 4
LEAD_TRAJECTORY_DIM = 6 * 4
Expand Down Expand Up @@ -107,8 +109,14 @@ 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: (512,),
ModelInputs.FEATURES: (SPATIAL_SIZE, VISION_FEATURES),
ModelInputs.DESIRE: (DESIRE_LEN,),
ModelInputs.TRAFFIC: (2,),
ModelInputs.ACTION_T: (ACTION_LEN,),
Expand Down
3 changes: 2 additions & 1 deletion torchtitan/experiments/path/onnx_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()}


Expand Down
2 changes: 1 addition & 1 deletion torchtitan/experiments/rldriving/config_registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
Loading