diff --git a/src/camera-based-e2e/models/base_model.py b/src/camera-based-e2e/models/base_model.py index 17782e6..c6a760e 100644 --- a/src/camera-based-e2e/models/base_model.py +++ b/src/camera-based-e2e/models/base_model.py @@ -3,6 +3,7 @@ import torch.nn.functional as F import pytorch_lightning as pl import torchvision +from contextlib import nullcontext from dataclasses import asdict, is_dataclass from .losses.depth_loss import DepthLoss @@ -26,6 +27,9 @@ def __init__(self, model: nn.Module, lr: float, lr_vision: float | None = None, super(LitModel, self).__init__() self.model = model + # NVJPEG fall back if we are running on Negishi AMD GPU + self.has_nvjpeg = True + # If we are using ScorerModel, which has a cfg, then save the attributes of the cfg as hparams, so they go into wandb cfg = getattr(model, "cfg", None) if cfg is None: @@ -81,12 +85,16 @@ def decode_batch_jpeg( self, images_jpeg: list[list[torch.Tensor]], device: torch.device | None = None, - ) -> list[torch.Tensor]: + ) -> list[torch.Tensor | None]: + cam_idxs_used = tuple(getattr(self.model.cfg, "cam_idxs_used", range(len(images_jpeg)))) decode_device = self.device if device is None else device + + selected = [(cam_idx, images_jpeg[cam_idx]) for cam_idx in cam_idxs_used] + # Flatten cameras flat_encoded, cam_sizes = [], [] - for cam in images_jpeg: - cam_sizes.append(len(cam)) + for cam_idx, cam in selected: + cam_sizes.append((cam_idx, len(cam))) for jpg in cam: t = jpg if isinstance(jpg, torch.Tensor) else torch.frombuffer(memoryview(jpg), dtype=torch.uint8) # decode_jpeg requires the raw jpeg bytes to be on cpu @@ -94,18 +102,30 @@ def decode_batch_jpeg( t = t.cpu() flat_encoded.append(t) - flat_decoded = torchvision.io.decode_jpeg( - flat_encoded, - mode=torchvision.io.ImageReadMode.UNCHANGED, - device=decode_device, - ) # list of (C, H, W) gpu tensors + try: + flat_decoded = torchvision.io.decode_jpeg( + flat_encoded, + mode=torchvision.io.ImageReadMode.UNCHANGED, + device=decode_device, + ) # list of (C, H, W) gpu tensors + except Exception as e: + if "nvJPEG" not in str(e): + raise e + self.has_nvjpeg = False + flat_decoded = torchvision.io.decode_jpeg( + flat_encoded, + mode=torchvision.io.ImageReadMode.UNCHANGED, + device="cpu", + ) + if torch.device(decode_device).type == "cuda": + flat_decoded = [img.to(decode_device, non_blocking=True) for img in flat_decoded] - out = [] + out = [None] * len(images_jpeg) idx = 0 - for n in cam_sizes: - cam_list = flat_decoded[idx: idx+n] + for cam_idx, n in cam_sizes: + out[cam_idx] = torch.stack(flat_decoded[idx:idx+n], dim=0) idx += n - out.append(torch.stack(cam_list, dim=0)) # (B, C, H, W) + return out def on_fit_start(self) -> None: @@ -211,10 +231,10 @@ def configure_optimizers(self): return optimizer # ---- forward / step ---- - def forward(self, x: torch.Tensor) -> torch.Tensor: + def forward(self, x: dict[str, torch.Tensor]) -> torch.Tensor: return self.model(x) - def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: + def _shared_step(self, batch: dict[str, torch.Tensor | list[torch.Tensor]], stage: str) -> torch.Tensor: past, future, intent = batch['PAST'], batch['FUTURE'], batch['INTENT'] if "IMAGES" in batch: @@ -234,8 +254,15 @@ def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: pred_future = self.forward(model_inputs) # (B, T*2) pred_depth = None pred_scores: torch.Tensor = None + pred_traj_flat: torch.Tensor = None + query_for_score: torch.Tensor = None if isinstance(pred_future, dict): - pred_future, pred_depth, pred_scores = pred_future["trajectory"], pred_future.get("depth", None), pred_future.get("scores", None) + outputs = pred_future + pred_future = outputs["trajectory"] + pred_depth = outputs.get("depth", None) + pred_scores = outputs.get("scores", None) + pred_traj_flat = outputs.get("trajectory_flat", None) + query_for_score = outputs.get("query_for_score", None) pred = pred_future t_steps = future.shape[1] @@ -310,6 +337,67 @@ def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: else: loss_score = torch.tensor(0.0, device=self.device) + # Train-only adversarial scorer supervision in trajectory space. + loss_adv = torch.tensor(0.0, device=self.device) + adv_enabled = bool(getattr(self.hparams, "model_cfg_adv_enabled", False)) + if (stage == "train" and adv_enabled and k_modes > 1 and pred_scores is not None and pred_traj_flat is not None and query_for_score is not None and hasattr(self.model, "score_trajectories")): + epsilon = float(getattr(self.hparams, "model_cfg_adv_epsilon", 0.10)) # max perturbation + adv_steps = max(1, int(getattr(self.hparams, "model_cfg_adv_steps", 3))) # running Projected Gradient Descent + alpha = epsilon / adv_steps + + original_traj = pred_traj_flat.detach() + adv_traj = original_traj + for _ in range(adv_steps): + # run forward pass of scorer on adversarial trajectory + attack_traj = adv_traj.detach().requires_grad_(True) + score_adv = self.model.score_trajectories(attack_traj, query_for_score) + + # compute gradient, gradient ASCENT to MAXIMIZE loss. + grad = torch.autograd.grad(score_adv.sum(), attack_traj, only_inputs=True)[0] + adv_traj = attack_traj + alpha * grad.sign() + delta = torch.clamp(adv_traj - original_traj, min=-epsilon, max=epsilon) + adv_traj = (original_traj + delta).detach() + + bsz, _, t2_adv = adv_traj.shape + t_adv = t2_adv // 2 + adv_ade = torch.norm(adv_traj.view(bsz, k_modes, t_adv, 2) - future[:, None], dim=-1,).mean(dim=-1).detach() + adv_scores = self.model.score_trajectories(adv_traj, query_for_score) + loss_adv = F.mse_loss(adv_scores, adv_ade) + + # Sobolev training. d(Score)/d(traj) ~= d(ADE)/d(traj) + loss_sobolev = torch.tensor(0.0, device=self.device) + sobolev_grad_cosine = torch.tensor(0.0, device=self.device) + sobolev_enabled = bool(getattr(self.hparams, "model_cfg_sobolev_enabled", False)) + if (stage == "train" and sobolev_enabled and k_modes > 1 and pred_scores is not None and pred_traj_flat is not None and query_for_score is not None and hasattr(self.model, "score_trajectories")): + autocast_device = pred_traj_flat.device.type + autocast_ctx = ( + torch.autocast(device_type=autocast_device, enabled=False) + if autocast_device in ("cpu", "cuda") + else nullcontext() + ) + with autocast_ctx: + sobolev_traj = pred_traj_flat.detach().float().requires_grad_(True) + sobolev_query = query_for_score.detach().float() + sobolev_scores = self.model.score_trajectories(sobolev_traj, sobolev_query) + grad_score = torch.autograd.grad( + sobolev_scores.sum(), + sobolev_traj, + create_graph=True, + only_inputs=True, + )[0] + + sobolev_traj_xy = sobolev_traj.view(pred.size(0), k_modes, t_steps, 2) + delta = sobolev_traj_xy - future[:, None].float() + dist = torch.norm(delta, dim=-1, keepdim=True).clamp_min(1e-3) + grad_ade = (delta / (t_steps * dist)).reshape_as(sobolev_traj) + + loss_sobolev = F.smooth_l1_loss(grad_score, grad_ade.detach()) + sobolev_grad_cosine = F.cosine_similarity( + grad_score.reshape(pred.size(0), k_modes, -1), + grad_ade.reshape(pred.size(0), k_modes, -1), + dim=-1, + ).mean() + # Scorer Metrics scorer_metrics = {} if k_modes > 1 and pred_scores is not None: @@ -331,7 +419,8 @@ def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: scorer_metrics[f"{stage}_scorer_spearman"] = rho.mean() # Depth Loss - if pred_depth is not None: + use_depth_loss = bool(getattr(self.hparams, "model_cfg_use_depth_loss", False)) + if pred_depth is not None and use_depth_loss: front_img = images[1] # front camera depth_in = F.interpolate(front_img, size=(128, 128), mode='nearest') loss_depth = self.depth_loss(depth_in, pred_depth, loss_fn=F.l1_loss) @@ -341,13 +430,29 @@ def _shared_step(self, batch: torch.Tensor, stage: str) -> torch.Tensor: loss_depth *= 0.1 # slightly enabled loss_ade *= 1.0 # TODO: tune loss terms loss_score *= 1.0 - total_loss = loss_ade + loss_depth + loss_score + loss_rfs + adv_lambda = float(getattr(self.hparams, "model_cfg_adv_lambda", 0.1)) + sobolev_lambda = float(getattr(self.hparams, "model_cfg_sobolev_lambda", 0.1)) + local_step_lambda = float(getattr(self.hparams, "model_cfg_local_step_lambda", 0.1)) + total_loss = ( + loss_ade + + loss_depth + + loss_score + + loss_rfs + + (adv_lambda * loss_adv) + + (sobolev_lambda * loss_sobolev) + + (local_step_lambda * loss_local_step) + ) # TODO: improve logging both to disk and to console log_payload = { f"{stage}_loss_ade": loss_ade, f"{stage}_loss_score": loss_score, f"{stage}_loss_depth": loss_depth, f"{stage}_loss_rfs": loss_rfs, + f"{stage}_loss_adv": loss_adv, + f"{stage}_loss_sobolev": loss_sobolev, + f"{stage}_sobolev_grad_cosine": sobolev_grad_cosine, + f"{stage}_loss_local_step": loss_local_step, + f"{stage}_local_step_rel_improve": local_step_rel_improve, f"{stage}_rfs_unweighted": rfs_unweighted, f"{stage}_loss": total_loss, } diff --git a/src/camera-based-e2e/models/feature_extractors.py b/src/camera-based-e2e/models/feature_extractors.py index 0864b85..1ed6a9b 100644 --- a/src/camera-based-e2e/models/feature_extractors.py +++ b/src/camera-based-e2e/models/feature_extractors.py @@ -2,6 +2,8 @@ import torch import torch.nn as nn import timm +import torchvision +from torchvision.transforms import v2 class DINOFeatures(nn.Module): def __init__(self, model_name: str = "vit_small_plus_patch16_dinov3.lvd1689m", frozen: bool = True): @@ -14,7 +16,7 @@ def __init__(self, model_name: str = "vit_small_plus_patch16_dinov3.lvd1689m", f for param in self.dino_model.parameters(): param.requires_grad = False - self.dims = [384, 384, 384] # feature dims for each layer + self.dims = [384] # feature dims for last layer self.patch_size = 16 # patch size def forward(self, x: torch.Tensor) -> List[torch.Tensor]: @@ -22,7 +24,7 @@ def forward(self, x: torch.Tensor) -> List[torch.Tensor]: # transforms: resize 256x256, center crop, normalize x_t = self.transforms(x.float().div(255.0)) # preprocess features = self.dino_model(x_t) - return features # 3 x [B, 384, 16, 16] + return features[-1:] # return last layer features as list of 1 tensor (B, C, H', W') class SAMFeatures(nn.Module): def __init__( @@ -52,4 +54,48 @@ def forward(self, x: torch.Tensor) -> List[torch.Tensor]: # x: (B, 3, H, W) x_t = self.transforms(x.float().div(255.0)) # preprocess feats = self.sam_model(x_t) # list of feature maps - return [feats[self.feature_stage]] \ No newline at end of file + return [feats[self.feature_stage]] # (B, C, H', W') + +class EUPEFeatures(nn.Module): + EUPE_DIR = "/depot/mlp/data/robotvision/eupe" # Works on Gilbreth, Gautschi, Negishi, change if on Anvil + CHECKPOINT_PATH = f"{EUPE_DIR}/EUPE-ViT-S.pt" + SIZE = (768, 768) # pe_spatial_... is 512x512 I believe + + def make_transform(self, resize_size: int = 256): + to_tensor = v2.ToImage() + resize = v2.Resize((resize_size, resize_size), antialias=True) + to_float = v2.ToDtype(torch.float32, scale=True) + normalize = v2.Normalize( + mean=(0.485, 0.456, 0.406), + std=(0.229, 0.224, 0.225), + ) + return v2.Compose([to_tensor, resize, to_float, normalize]) + + def __init__(self, model_name: str = "EUPE-ViT-S", frozen: bool = True, feature_stage: int = -1): + super(EUPEFeatures, self).__init__() + if feature_stage != -1: + raise NotImplementedError + self.transform = self.make_transform(resize_size=self.SIZE[0]) # pics are like 900 x 1000 + self.model = torch.hub.load(self.EUPE_DIR, "eupe_vits16", source="local", weights=self.CHECKPOINT_PATH) + if frozen: + for param in self.model.parameters(): + param.requires_grad = False + self.model.eval() + self.dims = [384] # feature dim for the last layer + self.patch_size = 16 + self.n_tokens = (self.SIZE[0] // self.patch_size) ** 2 # 1024 + self.data_config = { + "input_size": (3, self.SIZE[0], self.SIZE[1]) + } + + def forward(self, x: torch.Tensor) -> List[torch.Tensor]: + # x: (B, 3, H, W) + x_t = self.transform(x) # preprocess + with torch.inference_mode(): + with torch.autocast(device_type='cuda', dtype=torch.bfloat16): + features = self.model.forward_features(x_t) # list of feature maps + clstoken, patchtokens = features["x_norm_clstoken"], features["x_norm_patchtokens"] + B, N, C = patchtokens.shape + H = W = int(N**0.5) + patchtokens = patchtokens.transpose(1, 2).reshape(B, C, H, W) # (B, C, H', W') + return [patchtokens] \ No newline at end of file diff --git a/src/camera-based-e2e/models/monocular.py b/src/camera-based-e2e/models/monocular.py index c310027..b10880f 100644 --- a/src/camera-based-e2e/models/monocular.py +++ b/src/camera-based-e2e/models/monocular.py @@ -1,102 +1,49 @@ import torch import torch.nn as nn import torch.nn.functional as F -from math import sqrt +from dataclasses import dataclass from .blocks import TransformerBlock -class MonocularModel(nn.Module): - def __init__( - self, - in_dim: int, - out_dim: int, - feature_extractor: nn.Module - ): - # out_dim: (B, 40) which gets reshaped to (B, 20, 2) later - super(MonocularModel, self).__init__() - self.features = feature_extractor - - # attention - self.feature_dim = sum(self.features.dims) # works for both DINO and SAM - self.key_projection = nn.Linear(in_features=self.feature_dim, out_features=self.feature_dim) # project into "key" space - self.value_projection = nn.Linear(in_features=self.feature_dim, out_features=self.feature_dim) - - # condition the query on intent (B,) and past (B, 16, 6) - query_input_dim = 3 + 16 * 6 # one hot -- concat -- flattened - self.query = nn.Sequential( - nn.Linear(query_input_dim, self.feature_dim), - nn.LeakyReLU(), - nn.Linear(self.feature_dim, self.feature_dim), - ) - - # learnable positional encoding - self.n_tokens = self.features.data_config["input_size"][1] // self.features.patch_size * (self.features.data_config["input_size"][2] // self.features.patch_size) - self.positional_encoding = nn.Parameter(nn.init.trunc_normal_(torch.zeros((1, self.n_tokens, self.feature_dim)), std=0.02)) # (1, N, C) - - # MLP at end rather than directly using softmax as final output - self.decoder = nn.Sequential( - nn.Linear(self.feature_dim, self.feature_dim), - nn.LeakyReLU(), - nn.Linear(self.feature_dim, out_dim), - ) - - # LayerNorms - self.token_norm = nn.LayerNorm(self.feature_dim) - self.query_norm = nn.LayerNorm(self.feature_dim) - self.attn_norm = nn.LayerNorm(self.feature_dim) - - - def forward(self, x: dict) -> torch.Tensor: - # past: (B, 16, 6), intent: int - past, images, intent = x['PAST'], x['IMAGES'], x['INTENT'] - - # Ref: https://github.com/waymo-research/waymo-open-dataset/blob/5f8a1cd42491210e7de629b6f8fc09b65e0cbe99/src/waymo_open_dataset/dataset.proto#L50%20%20order%20=%20[2,%201,%203] - front_cam = images[1] - with torch.no_grad(): - feats = self.features(front_cam) # list or tensor - - # tokens: handle list of features or single tensor - if isinstance(feats, (list, tuple)): - tokens = torch.cat([f.flatten(2) for f in feats], dim=1) # (B, C_total, N) - else: - tokens = feats.flatten(2) # (B, C, N) - tokens = torch.permute(tokens, (0, 2, 1)) + self.positional_encoding # (B, N, C_total) - tokens = self.token_norm(tokens) - - # attention - key = self.key_projection(tokens) # (B, 256, 1152) - value = self.value_projection(tokens) # (B, 256, 40) - - intent_onehot = F.one_hot((intent - 1).long(), num_classes=3).float() # (B, 3). minus 1 --> 0, 1, 2 - past_flat = past.view(past.size(0), -1) # (B, 96) - query = self.query(torch.cat([intent_onehot, past_flat], dim=1)).unsqueeze(1) # (B, 1, 256) - query = self.query_norm(query) - - scores = query @ key.permute((0, 2, 1)) # (B, T, N) - attention = F.softmax(scores / sqrt(key.shape[2]), dim=2) @ value # (B, 1, 40) - attention = self.attn_norm(attention) - return self.decoder(attention.squeeze(1)) # (B, 40) +@dataclass +class DeepMonocularConfig: + # arch + n_blocks: int = 4 + n_proposals: int = 50 + cam_idxs_used: tuple = (1,) # front only + # kinematics + dt: float = 0.25 + max_accel: float = 8.0 + max_omega: float = 1.0 + # training + use_depth_loss: bool = False + # --- PR #16 (Scorer as Optimization Target) --- + # https://github.com/mgagvani/robotvision/pull/16 + # adversarial training + adv_enabled: bool = False + adv_lambda: float = 0.1 + adv_epsilon: float = 0.10 + adv_steps: int = 3 + # sobolev training + sobolev_enabled: bool = True + sobolev_lambda: float = 50 # scale up loss further class DeepMonocularModel(nn.Module): def __init__( self, feature_extractor, out_dim, - n_blocks=1, - n_proposals=50, - dt: float = 0.25, - max_accel: float = 8.0, - max_omega: float = 1.0, ): super().__init__() + self.cfg = DeepMonocularConfig() self.features = feature_extractor self.feature_dim = sum(self.features.dims) if out_dim % 2 != 0: raise ValueError(f"out_dim must be even for (x,y) rollout, got {out_dim}") self.horizon = out_dim // 2 - self.dt = dt - self.max_accel = max_accel - self.max_omega = max_omega + self.dt = self.cfg.dt + self.max_accel = self.cfg.max_accel + self.max_omega = self.cfg.max_omega # Initial Query Projection (Intent + Past -> C) query_input_dim = 3 + 16 * 6 @@ -110,13 +57,16 @@ def __init__( ) # learnable positional encoding - self.n_tokens = self.features.data_config["input_size"][1] // self.features.patch_size * (self.features.data_config["input_size"][2] // self.features.patch_size) + if hasattr(self.features, "n_tokens"): + self.n_tokens = self.features.n_tokens + else: + self.n_tokens = self.features.data_config["input_size"][1] // self.features.patch_size * (self.features.data_config["input_size"][2] // self.features.patch_size) self.positional_encoding = nn.Parameter(nn.init.trunc_normal_(torch.zeros((1, self.n_tokens, self.feature_dim)), std=0.02)) # (1, N, C) # Deep network rather than single attention in MonocularModel self.blocks = nn.ModuleList([ TransformerBlock(self.feature_dim, num_heads=8, mlp_dim=self.feature_dim*4) - for _ in range(n_blocks) + for _ in range(self.cfg.n_blocks) ]) # For Supervised Depth Loss -> (B, 128, 128) @@ -130,7 +80,7 @@ def __init__( nn.Conv2d(32, 1, 1) ) - self.n_proposals = n_proposals + self.n_proposals = self.cfg.n_proposals self.traj_decoder = nn.Sequential( nn.Linear(self.feature_dim, self.feature_dim), nn.GELU(), @@ -152,19 +102,29 @@ def __init__( nn.Linear(self.feature_dim, 1), ) # no softmax, since we use cross entropy later - def bicycle_model(self, control_pred: torch.Tensor, past: torch.Tensor) -> torch.Tensor: - accel = torch.tanh(control_pred[..., 0]) * self.max_accel # (B, K, T) - omega = torch.tanh(control_pred[..., 1]) * self.max_omega # (B, K, T) - - x_state = past[:, -1, 0].unsqueeze(1).expand(-1, self.n_proposals).clone() - y_state = past[:, -1, 1].unsqueeze(1).expand(-1, self.n_proposals).clone() + def rollout_controls( + self, + controls: torch.Tensor, + past: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + if controls.ndim != 4 or controls.size(-1) != 2: + raise ValueError( + f"controls must have shape (B, K, T, 2), got {tuple(controls.shape)}" + ) + + batch_size, n_proposals, horizon, _ = controls.shape + accel = controls[..., 0] + omega = controls[..., 1] + + x_state = past[:, -1, 0].unsqueeze(1).expand(-1, n_proposals).clone() + y_state = past[:, -1, 1].unsqueeze(1).expand(-1, n_proposals).clone() vx0 = past[:, -1, 2] vy0 = past[:, -1, 3] - speed_state = torch.sqrt(vx0 * vx0 + vy0 * vy0 + 1e-6).unsqueeze(1).expand(-1, self.n_proposals).clone() - heading_state = torch.atan2(vy0, vx0).unsqueeze(1).expand(-1, self.n_proposals).clone() + speed_state = torch.sqrt(vx0 * vx0 + vy0 * vy0 + 1e-6).unsqueeze(1).expand(-1, n_proposals).clone() + heading_state = torch.atan2(vy0, vx0).unsqueeze(1).expand(-1, n_proposals).clone() xy_steps = [] - for t in range(self.horizon): + for t in range(horizon): x_state = x_state + speed_state * torch.cos(heading_state) * self.dt y_state = y_state + speed_state * torch.sin(heading_state) * self.dt xy_steps.append(torch.stack([x_state, y_state], dim=-1)) @@ -173,7 +133,27 @@ def bicycle_model(self, control_pred: torch.Tensor, past: torch.Tensor) -> torch speed_state = torch.clamp_min(speed_state + accel[:, :, t] * self.dt, 0.0) traj_xy = torch.stack(xy_steps, dim=2) # (B, K, T, 2) - return traj_xy, traj_xy.reshape(traj_xy.size(0), -1), accel, omega # (B, K*T*2) + return traj_xy, traj_xy.reshape(batch_size, n_proposals, -1) + + def bicycle_model( + self, + control_pred: torch.Tensor, + past: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + accel = torch.tanh(control_pred[..., 0]) * self.max_accel # (B, K, T) + omega = torch.tanh(control_pred[..., 1]) * self.max_omega # (B, K, T) + controls = torch.stack([accel, omega], dim=-1) + traj_xy, traj_flat = self.rollout_controls(controls, past) + return traj_xy, traj_flat, accel, omega + + def score_trajectories( + self, + traj_flat: torch.Tensor, + query_for_score: torch.Tensor, + ) -> torch.Tensor: + traj_feat = self.traj_features(traj_flat) + score_in = torch.cat([query_for_score.to(dtype=traj_feat.dtype), traj_feat], dim=-1) # (B, K, C*2) + return self.score_decoder(score_in).squeeze(-1) # (B, K) def forward(self, x): # Copied from MonocularModel @@ -215,17 +195,18 @@ def forward(self, x): control_pred = self.traj_decoder(query.squeeze(1)).view( query.size(0), self.n_proposals, self.horizon, 2 ) # (B, K, T, 2) - traj_xy, traj_pred, accel, omega = self.bicycle_model(control_pred, past) # (B, K, T*2) + traj_xy, traj_pred_flat, accel, omega = self.bicycle_model(control_pred, past) # (B, K, T*2) - traj_pred_flat = traj_xy.reshape(traj_xy.size(0), self.n_proposals, -1) # (B, K, T*2) - traj_feat: torch.Tensor = self.traj_features(traj_pred_flat.detach()) # (B, K, C) query_for_score = query.squeeze(1).detach()[:, torch.newaxis, :].expand(-1, self.n_proposals, -1) # (B, K, C) - score_in = torch.cat([query_for_score, traj_feat], dim=-1) # (B, K, 2C) - score_pred = self.score_decoder(score_in).squeeze(-1) # (B, K) + score_pred = self.score_trajectories(traj_pred_flat.detach(), query_for_score) # (B, K) + controls_structured = torch.stack([accel, omega], dim=-1) return { - "trajectory": traj_pred, + "trajectory": traj_pred_flat.reshape(traj_pred_flat.size(0), -1), + "trajectory_flat": traj_pred_flat, + "query_for_score": query_for_score, "scores": score_pred, "depth": output_depth, - "controls": torch.stack([accel, omega], dim=-1).reshape(query.size(0), -1), + "controls": controls_structured.reshape(query.size(0), -1), + "controls_structured": controls_structured, } diff --git a/src/camera-based-e2e/train.py b/src/camera-based-e2e/train.py index 3c35363..0dcb515 100644 --- a/src/camera-based-e2e/train.py +++ b/src/camera-based-e2e/train.py @@ -21,7 +21,7 @@ # Replace with your model defined in models/ from models.base_model import LitModel, collate_with_images from models.monocular import DeepMonocularModel -from models.feature_extractors import SAMFeatures +from models.feature_extractors import SAMFeatures, EUPEFeatures class HomogeneousConcatBatchSampler(BatchSampler): @@ -299,11 +299,13 @@ def n_batches(length: int): out_dim = 20 * 2 # Future: (B, 20, 2) model = DeepMonocularModel( - feature_extractor=SAMFeatures( - model_name="timm/vit_pe_spatial_small_patch16_512.fb", frozen=True + # feature_extractor=EUPEFeatures( + # model_name="EUPE-ViT-S", frozen=True + # ), + feature_extractor=EUPEFeatures( + model_name="EUPE-ViT-S", frozen=True ), out_dim=out_dim, - n_blocks=4, ) name = str(model.__class__.__name__.replace("Model", "")).lower() if args.compile: