diff --git a/src/Ablation/ours_radar_no_doppler/inference.py b/src/Ablation/ours_radar_no_doppler/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..76e3b35d6487baea14f290d6e3ac8a5636e33a0f --- /dev/null +++ b/src/Ablation/ours_radar_no_doppler/inference.py @@ -0,0 +1,267 @@ +#!/usr/bin/env python3 +"""Sequence-by-sequence RadarDepth no-Doppler inference on Smoke-Eval. + +Launch with ``accelerate launch inference.py --config ``. +Each output is ``_pred.npy`` with float32 shape ``[N, 1, H, W]`` +and normalized depth clipped to ``[0, 1]``. +""" + +import argparse +import os +import pickle +from typing import Dict, List, Tuple + +import numpy as np +import torch +import yaml +from accelerate import Accelerator +from accelerate.utils import set_seed +from safetensors.torch import load_file +from torch.utils.data import DataLoader +from tqdm import tqdm + +from radar_depth import RadarDepth +from rice_dataset import RiceDataset + + +def _resolve_path(config_path: str, value: str) -> str: + if os.path.isabs(value): + return value + return os.path.normpath( + os.path.join(os.path.dirname(os.path.abspath(config_path)), value) + ) + + +def _validate_prediction_array(predictions: np.ndarray, sequence: str) -> None: + if predictions.ndim != 4 or predictions.shape[1] != 1: + raise RuntimeError( + f"{sequence}: expected prediction shape [N, 1, H, W], " + f"got {predictions.shape}" + ) + if not np.isfinite(predictions).all(): + raise RuntimeError(f"{sequence}: predictions contain NaN or Inf") + if predictions.min() < 0.0 or predictions.max() > 1.0: + raise RuntimeError( + f"{sequence}: normalized predictions are outside [0, 1]: " + f"[{predictions.min()}, {predictions.max()}]" + ) + + +def _merge_rank_results( + gather_dir: str, + sequence: str, + num_processes: int, +) -> Dict[int, np.ndarray]: + safe_sequence = sequence.replace("/", "_").replace("\\", "_").lower() + merged: Dict[int, np.ndarray] = {} + for rank in range(num_processes): + rank_path = os.path.join(gather_dir, f"rank_{rank}_{safe_sequence}.pkl") + with open(rank_path, "rb") as handle: + rank_results = pickle.load(handle) + for frame_idx, prediction in rank_results: + merged.setdefault(int(frame_idx), prediction) + os.remove(rank_path) + return merged + + +def _save_sequence( + output_dir: str, + sequence: str, + predictions: Dict[int, np.ndarray], + expected_frames: List[int], + debug: bool, +) -> np.ndarray: + if not predictions: + raise RuntimeError(f"{sequence}: inference produced no predictions") + + if not debug: + missing = [frame for frame in expected_frames if frame not in predictions] + if missing: + raise RuntimeError( + f"{sequence}: missing {len(missing)} predictions " + f"(first few frame indices: {missing[:5]})" + ) + ordered_frames = expected_frames + else: + ordered_frames = sorted(predictions) + + prediction_array = np.stack( + [predictions[frame] for frame in ordered_frames], axis=0 + ).astype(np.float32, copy=False) + _validate_prediction_array(prediction_array, sequence) + + safe_sequence = sequence.replace("/", "_").replace("\\", "_").lower() + np.save( + os.path.join(output_dir, f"{safe_sequence}_pred.npy"), + prediction_array, + ) + return prediction_array + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Run RadarDepth no-Doppler inference on Smoke-Eval." + ) + parser.add_argument( + "--config", + default="config_stage1_iq1m.yaml", + help="YAML config path", + ) + parser.add_argument("--checkpoint", default=None, help="Override checkpoint path") + parser.add_argument("--output_dir", default=None, help="Override output directory") + parser.add_argument("--debug", action="store_true", help="Process one batch per sequence") + return parser.parse_args() + + +def main() -> None: + cli = parse_args() + with open(cli.config, "r") as handle: + config = yaml.safe_load(handle) or {} + + training_config = config.get("training", {}) + data_config = config.get("data", {}) + inference_config = config.get("inference", {}) + + test_root_value = data_config.get("test_root") + if not test_root_value: + raise ValueError("config['data']['test_root'] is required") + test_root = _resolve_path(cli.config, str(test_root_value)) + + checkpoint_value = cli.checkpoint or inference_config.get("checkpoint_path") + if not checkpoint_value: + raise ValueError( + "Set config['inference']['checkpoint_path'] or pass --checkpoint" + ) + checkpoint_path = _resolve_path(cli.config, str(checkpoint_value)) + if not os.path.isfile(checkpoint_path): + raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") + + output_value = cli.output_dir or inference_config.get( + "output_dir", "inference_results" + ) + output_dir = _resolve_path(cli.config, str(output_value)) + batch_size = int( + inference_config.get("batch_size", training_config.get("batch_size", 1)) + ) + num_workers = int( + inference_config.get("num_workers", data_config.get("num_workers", 0)) + ) + frame_skip = int(inference_config.get("frame_skip", 1)) + mixed_precision = "fp16" + scale_factor = float(data_config.get("scale_factor", 0.001)) + max_depth_m = float(data_config.get("max_depth_m", 11.2)) + depth_resolution = tuple(data_config.get("depth_resolution", [128, 256])) + + accelerator = Accelerator(mixed_precision=mixed_precision) + set_seed(int(training_config.get("seed", 42))) + + discovery_dataset = RiceDataset( + root_dir=test_root, + sequences=None, + frame_skip=frame_skip, + scale_factor=scale_factor, + max_depth_m=max_depth_m, + depth_resolution=depth_resolution, + use_rgb=False, + ) + sequences = discovery_dataset.sequences + if not sequences: + raise ValueError(f"No valid Smoke-Eval sequences found under {test_root}") + del discovery_dataset + + if accelerator.is_main_process: + os.makedirs(output_dir, exist_ok=True) + print(f"Smoke-Eval: {test_root} ({len(sequences)} sequences)") + print(f"Checkpoint: {checkpoint_path}") + print(f"Output: {output_dir}") + print( + f"Mixed precision: {mixed_precision} | " + f"processes: {accelerator.num_processes}" + ) + accelerator.wait_for_everyone() + + gather_dir = os.path.join(output_dir, "_gather") + os.makedirs(gather_dir, exist_ok=True) + + model = RadarDepth( + output_height=int(depth_resolution[0]), + output_width=int(depth_resolution[1]), + ) + model.load_state_dict(load_file(checkpoint_path, device="cpu"), strict=True) + model.eval() + model = accelerator.prepare(model) + + for sequence_index, sequence in enumerate(sequences): + dataset = RiceDataset( + root_dir=test_root, + sequences=[sequence], + frame_skip=frame_skip, + scale_factor=scale_factor, + max_depth_m=max_depth_m, + depth_resolution=depth_resolution, + use_rgb=False, + ) + expected_frames = [int(frame_idx) for _, frame_idx in dataset.index_map] + loader = DataLoader( + dataset, + batch_size=batch_size, + shuffle=False, + num_workers=num_workers, + pin_memory=(accelerator.device.type == "cuda"), + drop_last=False, + ) + loader = accelerator.prepare(loader) + + local_results: List[Tuple[int, np.ndarray]] = [] + with torch.no_grad(): + progress = tqdm( + loader, + desc=f"[{sequence_index + 1}/{len(sequences)}] {sequence}", + disable=not accelerator.is_local_main_process, + dynamic_ncols=True, + leave=False, + ) + for batch in progress: + with accelerator.autocast(): + prediction = model(batch["radar"]).clamp_(0.0, 1.0) + prediction_np = prediction.detach().float().cpu().numpy() + frame_indices = batch["frame_idx"].detach().cpu().tolist() + local_results.extend( + (int(frame_idx), prediction_np[index]) + for index, frame_idx in enumerate(frame_indices) + ) + if cli.debug: + break + + accelerator.wait_for_everyone() + safe_sequence = sequence.replace("/", "_").replace("\\", "_").lower() + rank_path = os.path.join( + gather_dir, + f"rank_{accelerator.process_index}_{safe_sequence}.pkl", + ) + with open(rank_path, "wb") as handle: + pickle.dump(local_results, handle, protocol=pickle.HIGHEST_PROTOCOL) + accelerator.wait_for_everyone() + + if accelerator.is_main_process: + merged = _merge_rank_results( + gather_dir, sequence, accelerator.num_processes + ) + prediction_array = _save_sequence( + output_dir, + sequence, + merged, + expected_frames, + cli.debug, + ) + print(f"{sequence}: saved {prediction_array.shape}") + accelerator.wait_for_everyone() + + if accelerator.is_main_process: + if os.path.isdir(gather_dir) and not os.listdir(gather_dir): + os.rmdir(gather_dir) + print(f"Saved {len(sequences)} sequence predictions to: {output_dir}") + + +if __name__ == "__main__": + main() diff --git a/src/Ablation/ours_radar_no_doppler/iq1m_dataset.py b/src/Ablation/ours_radar_no_doppler/iq1m_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..7e16809f41b8c67c52a71be2533214f0d7e54cc9 --- /dev/null +++ b/src/Ablation/ours_radar_no_doppler/iq1m_dataset.py @@ -0,0 +1,305 @@ +import json +from pathlib import Path +from typing import Dict, List, Optional, Tuple, Any, Union +import cv2 +import numpy as np +import torch +from torch.utils.data import Dataset +from collate_fn_helpers import radar_collator, depth_collator, fisheye_rgb_collator + + +class IQ1MMultiModalDataset(Dataset): + """ + Dataset for loading aligned lidar, radar, and video frames. + + No-doppler ablation: radar is read from + root_dir/radar_no_doppler//amplitude.npy and phase.npy + (single doppler bin), and the doppler axis is repeated 64x so + downstream code sees the standard cube. + + Args: + root_dir: Root directory containing 'lidar', 'radar_no_doppler', + 'video' folders + sequences: Optional list of sequence names to load. If None, loads all. + transform: Optional transform to apply to video frames + """ + + DOPPLER_BINS = 64 + + def __init__( + self, + root_dir: str, + sequences: Optional[List[str]] = None, + frame_skip: int = 1, + split_type: Optional[str] = None, # 'train', 'val', 'test', or None for all + # Processing parameters + scale_factor: float = 0.001, + max_depth_m: float = 11.2, + depth_resolution: Tuple[int, int] = (128, 256), + use_rgb: bool = True, + rgb_resolution: Tuple[int, int] = (128, 256), + ): + self.root_dir = Path(root_dir) + self.depth_dir = self.root_dir / "metric_depth" + self.radar_dir = self.root_dir / "radar_no_doppler" + self.video_dir = self.root_dir / "video" + self.frame_skip = max(1, frame_skip) + self.split_type = split_type + + # Processing parameters + self.proc_params = { + "scale_factor": scale_factor, + "max_depth_m": max_depth_m, + "depth_res": depth_resolution, + "use_rgb": use_rgb, + "rgb_res": rgb_resolution, + } + + # Load split configuration if split_type is specified + if split_type is not None: + split_config = self._load_split_config() + sequences = self._get_sequences_for_split(split_config, sequences) + + # Discover sequences + self.sequences = self._discover_sequences(sequences) + + # Build index mapping (global_idx -> (sequence_name, frame_idx)) + self.index_map: List[Tuple[str, int]] = [] + self.sequence_info: Dict[str, dict] = {} + + # Memory-mapped numpy arrays for efficient loading + self._depth_mmap: Dict[str, np.memmap] = {} + self._radar_amplitude_mmap: Dict[str, np.memmap] = {} + self._radar_phase_mmap: Dict[str, np.memmap] = {} + self._video_captures: Dict[str, cv2.VideoCapture] = {} + + self._build_index() + + def _load_split_config(self) -> Dict: + """Load split configuration from iq1m_split.json""" + split_file = Path(__file__).parent / "iq1m_split.json" + if not split_file.exists(): + raise FileNotFoundError(f"Split configuration not found: {split_file}") + + with open(split_file, "r") as f: + split_config = json.load(f) + + return split_config + + def _get_sequences_for_split( + self, split_config: Dict, requested_sequences: Optional[List[str]] = None + ) -> Optional[List[str]]: + """Get sequences for the specified split type""" + if self.split_type == "test": + sequences = split_config.get("test", []) + elif self.split_type in ["train", "val"]: + # Get all available sequences + all_sequences = self._get_all_available_sequences() + test_sequences = set(split_config.get("test", [])) + # Exclude test sequences + sequences = [s for s in all_sequences if s not in test_sequences] + else: + raise ValueError( + f"Invalid split_type: {self.split_type}. Must be 'train', 'val', 'test', or None" + ) + + # Filter by requested sequences if provided + if requested_sequences is not None: + sequences = [s for s in sequences if s in requested_sequences] + + return sequences + + def _get_all_available_sequences(self) -> List[str]: + """Get all available sequences from the dataset""" + depth_seqs = set( + d.name + for d in self.depth_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + radar_seqs = set( + d.name + for d in self.radar_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + + # Only consider video if use_rgb is True + if self.proc_params["use_rgb"]: + video_seqs = set( + d.name + for d in self.video_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + # Find common sequences across all modalities + common_seqs = depth_seqs & radar_seqs & video_seqs + else: + # Only need depth and radar + common_seqs = depth_seqs & radar_seqs + + return sorted(list(common_seqs)) + + def _discover_sequences(self, sequences: Optional[List[str]] = None) -> List[str]: + """Discover available sequences with required modalities.""" + # Get sequences from each modality folder + depth_seqs = set( + d.name + for d in self.depth_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + radar_seqs = set( + d.name + for d in self.radar_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + + # Only consider video if use_rgb is True + if self.proc_params["use_rgb"]: + video_seqs = set( + d.name + for d in self.video_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + # Find common sequences across all modalities + common_seqs = depth_seqs & radar_seqs & video_seqs + else: + # Only need depth and radar + common_seqs = depth_seqs & radar_seqs + + if sequences is not None: + # Filter to requested sequences + common_seqs = common_seqs & set(sequences) + + return sorted(list(common_seqs)) + + def _build_index(self): + """Build global index mapping and load metadata.""" + for seq_name in self.sequences: + # Load metadata from radar_no_doppler (or radar/lidar if available) + metadata_path = self.radar_dir / seq_name / "metadata.json" + if not metadata_path.exists(): + metadata_path = self.root_dir / "radar" / seq_name / "metadata.json" + if not metadata_path.exists(): + metadata_path = self.root_dir / "lidar" / seq_name / "metadata.json" + with open(metadata_path, "r") as f: + metadata = json.load(f) + + n_frames = metadata["n_frames"] + self.sequence_info[seq_name] = { + "n_frames": n_frames, + "metadata": metadata, + "start_idx": len(self.index_map), + } + + # Add frames to index with skipping + # Range: 0, frame_skip, 2*frame_skip, ... + for frame_idx in range(0, n_frames, self.frame_skip): + self.index_map.append((seq_name, frame_idx)) + + self.sequence_info[seq_name]["end_idx"] = len(self.index_map) + + def _get_depth_mmap(self, seq_name: str) -> np.memmap: + """Get or create memory-mapped metric depth array.""" + if seq_name not in self._depth_mmap: + path = self.depth_dir / seq_name / "metric_depth.npy" + self._depth_mmap[seq_name] = np.load(path, mmap_mode="r") + return self._depth_mmap[seq_name] + + def _get_radar_mmap(self, seq_name: str) -> Tuple[np.memmap, np.memmap]: + """Get or create memory-mapped radar arrays.""" + if seq_name not in self._radar_amplitude_mmap: + amp_path = self.radar_dir / seq_name / "amplitude.npy" + phase_path = self.radar_dir / seq_name / "phase.npy" + self._radar_amplitude_mmap[seq_name] = np.load(amp_path, mmap_mode="r") + self._radar_phase_mmap[seq_name] = np.load(phase_path, mmap_mode="r") + return self._radar_amplitude_mmap[seq_name], self._radar_phase_mmap[seq_name] + + def _get_video_capture(self, seq_name: str) -> cv2.VideoCapture: + """Get or create video capture object.""" + if seq_name not in self._video_captures: + video_path = self.video_dir / seq_name / "video.avi" + cap = cv2.VideoCapture(str(video_path)) + if not cap.isOpened(): + raise RuntimeError(f"Failed to open video: {video_path}") + self._video_captures[seq_name] = cap + return self._video_captures[seq_name] + + def _load_rgb_frame(self, seq_name: str, frame_idx: int) -> np.ndarray: + """Load a specific frame from video.""" + cap = self._get_video_capture(seq_name) + + # Seek to frame + cap.set(cv2.CAP_PROP_POS_FRAMES, frame_idx) + ret, frame = cap.read() + + if not ret: + raise RuntimeError(f"Failed to read frame {frame_idx} from {seq_name}") + + # Convert BGR to RGB + frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) + return frame + + def __len__(self) -> int: + return len(self.index_map) + + def __getitem__(self, idx: int) -> Dict[str, Any]: + seq_name, frame_idx = self.index_map[idx] + + # === Radar (always needed) === + amp_mmap, phase_mmap = self._get_radar_mmap(seq_name) + radar_amp = torch.from_numpy(amp_mmap[frame_idx].copy()).float() + radar_phase = torch.from_numpy(phase_mmap[frame_idx].copy()).float() + # Single doppler bin -> repeat to the standard 64-bin cube so + # downstream code is unchanged + radar_amp = torch.repeat_interleave(radar_amp, self.DOPPLER_BINS, dim=0) + radar_phase = torch.repeat_interleave(radar_phase, self.DOPPLER_BINS, dim=0) + + processed_radar = radar_collator( + radar_amp.unsqueeze(0), + radar_phase.unsqueeze(0), + scale_factor=self.proc_params["scale_factor"], + ).squeeze(0) + + # === Depth (always needed) === + depth_mmap = self._get_depth_mmap(seq_name) + depth = torch.from_numpy(depth_mmap[frame_idx].copy()).float().unsqueeze(0) + + processed_depth = depth_collator( + depth.unsqueeze(0), + max_depth_m=self.proc_params["max_depth_m"], + target_size=self.proc_params["depth_res"], + ).squeeze(0) + + out = { + "radar": processed_radar, + "depth": processed_depth, + "sequence": seq_name, + "frame_idx": frame_idx, + } + + # === RGB (only if use_rgb is True) === + if self.proc_params["use_rgb"]: + rgb = torch.from_numpy( + self._load_rgb_frame(seq_name, frame_idx) + ).float().permute(2, 0, 1) / 255.0 + + out["rgb"] = fisheye_rgb_collator( + rgb.unsqueeze(0), + target_size=self.proc_params["rgb_res"], + ).squeeze(0) + + return out + + def get_sequence_frames(self, seq_name: str) -> List[int]: + """Get global indices for all frames in a sequence.""" + info = self.sequence_info[seq_name] + return list(range(info["start_idx"], info["end_idx"])) + + def close(self): + """Release video capture resources.""" + for cap in self._video_captures.values(): + cap.release() + self._video_captures.clear() + + def __del__(self): + self.close() + + diff --git a/src/Ablation/ours_radar_no_doppler/radar_depth.py b/src/Ablation/ours_radar_no_doppler/radar_depth.py new file mode 100644 index 0000000000000000000000000000000000000000..41442060dcf90792f535795c828879e86a59678f --- /dev/null +++ b/src/Ablation/ours_radar_no_doppler/radar_depth.py @@ -0,0 +1,406 @@ +import torch +import torch.nn as nn +from typing import Tuple + + +class RadarPatchEmbed(nn.Module): + """ + Radar Spectrum Patch Embedding Layer. + + Takes 5D radar spectrum data and converts it into patch embeddings: + 1. Input: [B, 2, 256, 64, 8, 2] where channels are (magnitude, phase) + 2. Patchifies along range and doppler dimensions + 3. Outputs: [B, num_patches, embed_dim] where num_patches = 2048 + + Patch extraction: + - Range dimension (256): patch_size=4, stride=4 -> 64 patches + - Doppler dimension (64): patch_size=2, stride=2 -> 32 patches + - Total patches: 64 × 32 = 2048 + - Each patch: [4 range × 2 doppler × 8 elevation × 2 azimuth] × 2 channels = 256 features + """ + + def __init__( + self, + input_shape: Tuple[int, int, int, int] = ( + 256, + 64, + 8, + 2, + ), # (Range, Doppler, Elevation, Azimuth) + patch_size: Tuple[int, int, int, int] = ( + 4, + 2, + 8, + 2, + ), # (Range, Doppler, Elevation, Azimuth) + stride: Tuple[int, int] = (4, 2), # (Range, Doppler) + embed_dim: int = 256, + in_channels: int = 2, # magnitude + phase + ): + super().__init__() + + self.input_shape = input_shape + self.patch_size = patch_size + self.stride = stride + self.embed_dim = embed_dim + self.in_channels = in_channels + + # Calculate number of patches + range_dim, doppler_dim, elev_dim, azim_dim = input_shape + patch_range, patch_doppler, patch_elev, patch_azim = patch_size + stride_range, stride_doppler = stride + + self.num_patches_range = (range_dim - patch_range) // stride_range + 1 # 64 + self.num_patches_doppler = ( + doppler_dim - patch_doppler + ) // stride_doppler + 1 # 32 + self.num_patches = self.num_patches_range * self.num_patches_doppler # 2048 + + # Each patch has: patch_range × patch_doppler × patch_elev × patch_azim features per channel + patch_volume = ( + patch_range * patch_doppler * patch_elev * patch_azim + ) # 4×2×8×2 = 128 + self.patch_features = patch_volume * in_channels # 128 × 2 = 256 + + # Linear projection from patch features to embedding dimension + self.proj = nn.Linear(self.patch_features, embed_dim) + + print(f"Radar Patch Embedding Configuration:") + print( + f" Input shape: [B, {in_channels}, {range_dim}, {doppler_dim}, {elev_dim}, {azim_dim}]" + ) + print(f" Patch size: {patch_size}") + print(f" Stride: {stride}") + print( + f" Number of patches (range × doppler): {self.num_patches_range} × {self.num_patches_doppler} = {self.num_patches}" + ) + print(f" Patch features per channel: {patch_volume}") + print(f" Total patch features (mag+phase): {self.patch_features}") + print(f" Embedding dimension: {embed_dim}") + + def extract_patches(self, x: torch.Tensor) -> torch.Tensor: + """ + Extract patches from radar spectrum data. + + Args: + x: [B, 2, 256, 64, 8, 2] (magnitude + phase channels) + + Returns: + patches: [B, num_patches, patch_features] + """ + batch_size = x.shape[0] + x_mag = x[:, 0] # [B, 256, 64, 8, 2] + x_phase = x[:, 1] # [B, 256, 64, 8, 2] + + all_patches = [] + + # Extract patches with stride along range and doppler dimensions + for i in range(self.num_patches_range): + for j in range(self.num_patches_doppler): + start_range = i * self.stride[0] + end_range = start_range + self.patch_size[0] + start_doppler = j * self.stride[1] + end_doppler = start_doppler + self.patch_size[1] + + # Extract patch from both channels + patch_mag = x_mag[ + :, start_range:end_range, start_doppler:end_doppler, :, : + ] + patch_phase = x_phase[ + :, start_range:end_range, start_doppler:end_doppler, :, : + ] + + # Flatten patches + patch_mag_flat = patch_mag.flatten(1) # [B, 128] + patch_phase_flat = patch_phase.flatten(1) # [B, 128] + + # Interleave magnitude and phase features + patch_interleaved = torch.stack( + [patch_mag_flat, patch_phase_flat], dim=-1 + ) + patch_interleaved = patch_interleaved.flatten(1, -1) # [B, 256] + + all_patches.append(patch_interleaved) + + # Stack all patches: [B, num_patches, patch_features] + all_patches = torch.stack(all_patches, dim=1) + return all_patches + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """ + Forward pass. + + Args: + x: [B, 2, 256, 64, 8, 2] + + Returns: + embeddings: [B, num_patches, embed_dim] + """ + # Extract patches: [B, 2048, 256] + patches = self.extract_patches(x) + + # Project to embedding dimension: [B, 2048, embed_dim] + embeddings = self.proj(patches) + + return embeddings + + +class RadarEncoder(nn.Module): + """ + Radar Vision Transformer (ViT) Encoder. + """ + + def __init__( + self, + input_shape: Tuple[int, int, int, int] = (256, 64, 8, 2), + patch_size: Tuple[int, int, int, int] = (4, 2, 8, 2), + stride: Tuple[int, int] = (4, 2), + embed_dim: int = 256, + num_heads: int = 8, + num_layers: int = 4, + mlp_ratio: float = 4.0, + dropout: float = 0.1, + ): + super().__init__() + + self.embed_dim = embed_dim + + # Patch embedding layer + self.patch_embed = RadarPatchEmbed( + input_shape=input_shape, + patch_size=patch_size, + stride=stride, + embed_dim=embed_dim, + in_channels=2, + ) + + self.num_patches = self.patch_embed.num_patches + + # Learnable positional embeddings + self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, embed_dim)) + + # Transformer encoder + encoder_layer = nn.TransformerEncoderLayer( + d_model=embed_dim, + nhead=num_heads, + dim_feedforward=int(embed_dim * mlp_ratio), + dropout=dropout, + activation="gelu", + batch_first=True, + norm_first=True, + ) + self.transformer = nn.TransformerEncoder( + encoder_layer=encoder_layer, + num_layers=num_layers, + norm=nn.LayerNorm(embed_dim), + ) + + self._init_weights() + + def _init_weights(self): + """Initialize weights.""" + # Initialize positional embeddings + nn.init.trunc_normal_(self.pos_embed, std=0.02) + + # Initialize patch embedding projection + if hasattr(self.patch_embed.proj, "weight"): + nn.init.xavier_uniform_(self.patch_embed.proj.weight) + if self.patch_embed.proj.bias is not None: + nn.init.zeros_(self.patch_embed.proj.bias) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.patch_embed(x) + x = x + self.pos_embed + x = self.transformer(x) + return x + + +class TransformerDecoderBlock(nn.Module): + """Transformer decoder block with self-attention and feedforward""" + + def __init__(self, embed_dim=384, num_heads=6, mlp_ratio=4.0, dropout=0.0): + super().__init__() + self.norm1 = nn.LayerNorm(embed_dim) + self.attn = nn.MultiheadAttention( + embed_dim, num_heads, dropout=dropout, batch_first=True + ) + self.norm2 = nn.LayerNorm(embed_dim) + self.mlp = nn.Sequential( + nn.Linear(embed_dim, int(embed_dim * mlp_ratio)), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(int(embed_dim * mlp_ratio), embed_dim), + nn.Dropout(dropout), + ) + + def forward(self, x): + # Self-attention with residual + x = x + self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0] + # MLP with residual + x = x + self.mlp(self.norm2(x)) + return x + + +class DepthDecoder(nn.Module): + """ + Hybrid Transformer+CNN decoder for depth image generation. + + Input: [batch_size, num_patches=2048, embed_dim=512] + Output: [batch_size, 1, height=128, width=256] + + Architecture: + 1. Transformer decoder blocks (4 layers) + 2. Reshape to 2D feature map (64x32) + 3. CNN upsampling stages (64x32 -> 128x256) + """ + + def __init__( + self, + embed_dim=256, + num_patches=2048, + patch_grid_size=(64, 32), # Spatial structure from radar encoder + num_decoder_blocks=4, + num_heads=8, + mlp_ratio=4.0, + dropout=0.0, + output_height=128, + output_width=256, + output_channels=1, + ): + super().__init__() + self.embed_dim = embed_dim + self.num_patches = num_patches + self.patch_grid_size = patch_grid_size # (64, 32) spatial grid + self.output_height = output_height + self.output_width = output_width + self.output_channels = output_channels + + # Transformer decoder blocks + self.decoder_blocks = nn.ModuleList( + [ + TransformerDecoderBlock(embed_dim, num_heads, mlp_ratio, dropout) + for _ in range(num_decoder_blocks) + ] + ) + + self.norm = nn.LayerNorm(embed_dim) + + # Projection to intermediate feature map + # From 64x32x256 to 64x32x128 (reduce dimension for upsampling) + self.feature_proj = nn.Conv2d(embed_dim, 128, kernel_size=1) + + # Upsampling network: 64x32 -> 128x256 + # Start from 64x32 (range x doppler), upsample to 128x256 + self.upsample = nn.Sequential( + # Upsample doppler dimension: 64x32 -> 64x64 + nn.Upsample(scale_factor=(1, 2), mode="bilinear", align_corners=False), + nn.Conv2d(128, 64, kernel_size=3, padding=1), + nn.BatchNorm2d(64), + nn.ReLU(inplace=True), + # Upsample both dimensions: 64x64 -> 128x128 + nn.Upsample(scale_factor=(2, 2), mode="bilinear", align_corners=False), + nn.Conv2d(64, 32, kernel_size=3, padding=1), + nn.BatchNorm2d(32), + nn.ReLU(inplace=True), + # Upsample width dimension: 128x128 -> 128x256 + nn.Upsample(scale_factor=(1, 2), mode="bilinear", align_corners=False), + nn.Conv2d(32, output_channels, kernel_size=3, padding=1), + nn.Sigmoid(), # Output in [0, 1] range + ) + + def forward(self, x): + """ + Args: + x: [batch_size, num_patches, embed_dim] + + Returns: + depth: [batch_size, output_channels, output_height, output_width] + """ + batch_size = x.shape[0] + + # Apply transformer decoder blocks + for block in self.decoder_blocks: + x = block(x) + + x = self.norm(x) + + # Reshape to spatial dimensions: [B, 2048, 512] -> [B, 64, 32, 512] + x = x.reshape( + batch_size, + self.patch_grid_size[0], # 64 (range) + self.patch_grid_size[1], # 32 (doppler) + self.embed_dim, + ) + + # Permute to channel-first: [B, H, W, C] -> [B, C, H, W] + x = x.permute(0, 3, 1, 2) + # Shape: [batch, 512, 64, 32] + + # Project features + x = self.feature_proj(x) + # Shape: [batch, 128, 64, 32] + + # Upsample to target resolution + depth = self.upsample(x) + # Shape: [batch, 1, 128, 256] + + return depth + + +class RadarDepth(nn.Module): + """ + End-to-end Radar to Depth model (Doppler-as-Channels). + """ + + def __init__( + self, + # Encoder args + input_shape: Tuple[int, int, int, int] = (256, 64, 8, 2), + patch_size: Tuple[int, int, int, int] = (4, 2, 8, 2), + stride: Tuple[int, int] = (4, 2), + embed_dim: int = 256, + encoder_num_heads: int = 8, + encoder_num_layers: int = 4, + encoder_mlp_ratio: float = 4.0, + encoder_dropout: float = 0.1, + # Decoder args + decoder_num_blocks: int = 4, + decoder_num_heads: int = 8, + output_height: int = 128, + output_width: int = 256, + ): + super().__init__() + + self.encoder = RadarEncoder( + input_shape=input_shape, + patch_size=patch_size, + stride=stride, + embed_dim=embed_dim, + num_heads=encoder_num_heads, + num_layers=encoder_num_layers, + mlp_ratio=encoder_mlp_ratio, + dropout=encoder_dropout, + ) + + # Get patch info from encoder + num_patches = self.encoder.num_patches # 64 + + self.decoder = DepthDecoder( + embed_dim=embed_dim, + num_patches=num_patches, + num_decoder_blocks=decoder_num_blocks, + num_heads=decoder_num_heads, + output_height=output_height, + output_width=output_width, + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.encoder(x) + x = self.decoder(x) + return x + + +def create_radar_encoder(*args, **kwargs): + return RadarEncoder(*args, **kwargs) + + diff --git a/src/Ablation/ours_radar_no_doppler/rice_dataset.py b/src/Ablation/ours_radar_no_doppler/rice_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..b16614cb61ee1cf89aa030b3281a8f56d66fdc44 --- /dev/null +++ b/src/Ablation/ours_radar_no_doppler/rice_dataset.py @@ -0,0 +1,159 @@ +import numpy as np +import torch +from pathlib import Path +from typing import Dict, List, Optional, Tuple +from collate_fn_helpers import dji_rgb_collator, radar_collator, depth_collator +from torch.utils.data import Dataset + + +class RiceDataset(Dataset): + """Dataset for radar, DJI RGB, and ZED depth. + + No-doppler ablation: loads radar_no_doppler.npy (N, 1, elevation, azimuth, + range) and repeats the single doppler bin 64x so downstream code sees the + standard (64, elevation, azimuth, range) cube. + """ + + REQUIRED_FILES = ("radar_no_doppler.npy", "dji_rgb.npy", "zed_depth.npy") + DOPPLER_BINS = 64 + + def __init__( + self, + root_dir: str, + sequences: Optional[List[str]] = None, + frame_skip: int = 1, + depth_in_meters: bool = True, + rgb_normalize: bool = True, + # Processing parameters + scale_factor: float = 0.001, + max_depth_m: float = 11.2, + depth_resolution: Tuple[int, int] = (128, 256), + use_rgb: bool = True, + rgb_resolution: Tuple[int, int] = (128, 256), + ): + self.root_dir = Path(root_dir) + self.frame_skip = max(1, frame_skip) + self.depth_in_meters = depth_in_meters + self.rgb_normalize = rgb_normalize + + # Processing parameters + self.proc_params = { + "scale_factor": scale_factor, + "max_depth_m": max_depth_m, + "depth_res": depth_resolution, + "use_rgb": use_rgb, + "rgb_res": rgb_resolution, + } + + self.sequences = self._discover_sequences(sequences) + self.index_map: List[Tuple[str, int]] = [] + self._seq_arrays: Dict[str, Dict] = {} + + self._build_index() + + def _discover_sequences(self, sequences: Optional[List[str]] = None) -> List[str]: + if not self.root_dir.is_dir(): + raise FileNotFoundError(f"Root directory not found: {self.root_dir}") + + all_seqs = sorted( + d.name + for d in self.root_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + + # Required files always needed + required = ["radar_no_doppler.npy", "zed_depth.npy"] + # Add RGB if use_rgb is True + if self.proc_params["use_rgb"]: + required.append("dji_rgb.npy") + + valid = [ + name + for name in all_seqs + if all((self.root_dir / name / f).exists() for f in required) + ] + if sequences is not None: + valid = [s for s in valid if s in sequences] + return valid + + def _build_index(self) -> None: + self.index_map.clear() + for seq_name in self.sequences: + radar = np.load( + self.root_dir / seq_name / "radar_no_doppler.npy", mmap_mode="r" + ) + for i in range(0, radar.shape[0], self.frame_skip): + self.index_map.append((seq_name, i)) + + def _load_sequence_arrays(self, seq_name: str) -> Dict: + if seq_name not in self._seq_arrays: + seq_dir = self.root_dir / seq_name + arrays = { + "radar": np.load(seq_dir / "radar_no_doppler.npy", mmap_mode="r"), + "depth": np.load(seq_dir / "zed_depth.npy", mmap_mode="r"), + } + # Only load RGB if needed + if self.proc_params["use_rgb"]: + arrays["rgb"] = np.load(seq_dir / "dji_rgb.npy", mmap_mode="r") + self._seq_arrays[seq_name] = arrays + return self._seq_arrays[seq_name] + + def __len__(self) -> int: + return len(self.index_map) + + def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: + seq_name, frame_idx = self.index_map[idx] + arrs = self._load_sequence_arrays(seq_name) + + # === Radar (always needed) === + # (1, elevation, azimuth, range) -> repeat single doppler bin to + # (64, elevation, azimuth, range) so downstream code is unchanged + radar = np.asarray(arrs["radar"][frame_idx].copy()) + radar = np.repeat(radar, self.DOPPLER_BINS, axis=0) + radar_amp = torch.from_numpy(np.abs(radar).astype(np.float32)) + radar_phase = torch.from_numpy((np.angle(radar) / np.pi).astype(np.float32)) + + processed_radar = radar_collator( + radar_amp.unsqueeze(0), + radar_phase.unsqueeze(0), + scale_factor=self.proc_params["scale_factor"], + ).squeeze(0) + + # === Depth (always needed) === + depth = np.asarray(arrs["depth"][frame_idx]).astype(np.float32) + if self.depth_in_meters: + depth = depth / 1000.0 + invalid = ~(np.isfinite(depth) & (depth > 0)) + depth[invalid] = 0.0 + depth = depth[np.newaxis, ...] + depth_tensor = torch.from_numpy(depth).float() + + processed_depth = depth_collator( + depth_tensor.unsqueeze(0), + max_depth_m=self.proc_params["max_depth_m"], + target_size=self.proc_params["depth_res"], + ).squeeze(0) + + # === RGB (only if use_rgb is True) === + out = { + "radar": processed_radar, + "depth": processed_depth, + "sequence": seq_name, + "frame_idx": frame_idx, + } + + if self.proc_params["use_rgb"]: + rgb = np.asarray(arrs["rgb"][frame_idx]) + rgb = np.transpose(rgb, (2, 0, 1)) + if self.rgb_normalize: + rgb = rgb.astype(np.float32) / 255.0 + rgb_tensor = torch.from_numpy(rgb) + + out["rgb"] = dji_rgb_collator( + rgb_tensor.unsqueeze(0), + target_size=self.proc_params["rgb_res"], + ).squeeze(0) + + return out + + diff --git a/src/Ablation/ours_radar_no_doppler/split.json b/src/Ablation/ours_radar_no_doppler/split.json new file mode 100644 index 0000000000000000000000000000000000000000..e20f136687ae69b2f0e9faff1c6a3943afa89795 --- /dev/null +++ b/src/Ablation/ours_radar_no_doppler/split.json @@ -0,0 +1,34 @@ +{ + "test-rice": [ + "Dell-1", + "Dell-2", + "Smoke-Dell-1", + "Smoke-Dell-2", + "Keck-1", + "Keck-2", + "Keck-3", + "Smoke-keck-1", + "Smoke-keck-2", + "Smoke-keck-3" + ], + "test-iq1m": [ + "cfa.cfa.1.fwd", + "cfa.cfa.1.lat", + "cfa.cfa.3.fwd", + "cfa.cfa.3.lat", + "cfa.cfa.a.fwd", + "cfa.cfa.a.lat", + "morrison.morrison.1.fwd", + "morrison.morrison.1.lat", + "morrison.morrison.2.fwd", + "morrison.morrison.2.lat", + "posner.posner.1.fwd", + "posner.posner.1.lat", + "posner.posner.2.fwd", + "posner.posner.2.lat", + "posner.posner.3.fwd", + "posner.posner.3.lat", + "posner.posner.a.fwd", + "posner.posner.a.lat" + ] +} \ No newline at end of file diff --git a/src/Ablation/ours_radar_no_grad/inference.py b/src/Ablation/ours_radar_no_grad/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..6d353d511e4888cef2e9212b259d86ff119bb854 --- /dev/null +++ b/src/Ablation/ours_radar_no_grad/inference.py @@ -0,0 +1,16 @@ +"""Direct Accelerate backend for the no-gradient-loss Stage-1 ablation.""" + +from __future__ import annotations + +import sys +from pathlib import Path + + +STAGE2_DIR = Path(__file__).resolve().parents[2] / "GRADE" / "stage2_diffusion_refinement" +sys.path.insert(0, str(STAGE2_DIR)) + +from inference import main as stage2_main # noqa: E402 + + +if __name__ == "__main__": + stage2_main() diff --git a/src/Baselines/cafnet/collate_fn_helpers.py b/src/Baselines/cafnet/collate_fn_helpers.py new file mode 100644 index 0000000000000000000000000000000000000000..8fa5d76dc84ee17fc5c56642b8de089d452f3e3e --- /dev/null +++ b/src/Baselines/cafnet/collate_fn_helpers.py @@ -0,0 +1,404 @@ +import cv2 +import numpy as np +import torch +from functools import lru_cache +from typing import Callable, Dict, Sequence, Tuple, Union +from torchvision import transforms as T + + +IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32) +IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32) + +# ZED intrinsics at 1280x720 reference resolution. +_K_ZED_REF = np.array( + [ + [521.581604, 0.0, 636.33398438], + [0.0, 521.581604, 373.10964966], + [0.0, 0.0, 1.0], + ], + dtype=np.float64, +) +_ZED_REF_W = 1280 +_ZED_REF_H = 720 + +# DJI calibration constants. +_CALIB_K_DJI = np.array( + [ + [718.48555551, 0.0, 963.36465011], + [0.0, 720.25844189, 537.87569913], + [0.0, 0.0, 1.0], + ], + dtype=np.float64, +) +_CALIB_D_DJI = np.array( + [0.19022699, 0.03466753, 0.05858962, -0.07070669], dtype=np.float64 +) +_CALIB_DEFISH_SHAPE = (1920, 1080) +_CALIB_DEFISH_BALANCE = 0.2 +_CALIB_H_FULL = np.array( + [ + [0.8274446551892256, -0.0742944198979625, 80.23797348979947], + [-0.014725864916652691, 0.8471179917075127, 28.27366063997317], + [-5.083573451500717e-05, -6.846079418201229e-05, 1.0], + ], + dtype=np.float64, +) +_CALIB_OUT_SIZE = (1918, 1105) +_CALIB_CROP = (115, 255, 1400, 760) # top, left, right, bottom + + +@lru_cache(maxsize=1) +def _get_dji_defish_maps() -> Tuple[np.ndarray, np.ndarray]: + r_defish = np.eye(3) + k_new_defish = cv2.fisheye.estimateNewCameraMatrixForUndistortRectify( + _CALIB_K_DJI, + _CALIB_D_DJI, + _CALIB_DEFISH_SHAPE, + r_defish, + balance=_CALIB_DEFISH_BALANCE, + fov_scale=1.0, + ) + map1, map2 = cv2.fisheye.initUndistortRectifyMap( + _CALIB_K_DJI, + _CALIB_D_DJI, + r_defish, + k_new_defish, + _CALIB_DEFISH_SHAPE, + cv2.CV_16SC2, + ) + return map1, map2 + + +def resize_depth_mm(depth_mm: np.ndarray, target_size: Tuple[int, int]) -> np.ndarray: + target_h, target_w = target_size + if depth_mm.shape[:2] == (target_h, target_w): + return depth_mm + return cv2.resize(depth_mm, (target_w, target_h), interpolation=cv2.INTER_NEAREST) + + +def depth_collator( + depth: Union[torch.Tensor, np.ndarray], + max_depth_m: float = 11.2, + target_size: Tuple[int, int] = (128, 256), +) -> Union[torch.Tensor, np.ndarray]: + """Clamp, normalize to [0, 1], and resize depth.""" + is_numpy = isinstance(depth, np.ndarray) + if is_numpy: + depth = torch.from_numpy(depth) + + depth = depth.float() + original_shape = depth.shape + + if depth.dim() == 2: + depth = depth.unsqueeze(0) + elif depth.dim() == 3: + depth = depth.unsqueeze(1) + + invalid_mask = ~(torch.isfinite(depth) & (depth >= 0)) + depth[invalid_mask] = 0.0 + + depth = torch.clamp(depth, min=0.0, max=max_depth_m) + depth = depth / max_depth_m + + invalid_mask = ~torch.isfinite(depth) + depth[invalid_mask] = 0.0 + + resized = T.Resize( + target_size, interpolation=T.InterpolationMode.BILINEAR, antialias=True + )(depth) + + if len(original_shape) == 2: + resized = resized.squeeze(0) + + return resized.numpy() if is_numpy else resized + + +def dji_rgb_collator( + image: torch.Tensor, + target_size: Tuple[int, int] = (128, 256), +) -> torch.Tensor: + """Rectify and resize DJI RGB image batch. + + Args: + image: Tensor with shape (B, C, H, W). + target_size: Target resolution as (height, width). + + Returns: + Tensor in CHW format (B, C, H, W), float32 in [0, 1]. + """ + if not isinstance(image, torch.Tensor): + raise ValueError(f"Expected torch.Tensor, got {type(image)}") + + if image.dim() != 4: + raise ValueError( + f"Expected 4D tensor (B, C, H, W), got {image.dim()}D tensor with shape {image.shape}" + ) + + map1_defish, map2_defish = _get_dji_defish_maps() + target_h, target_w = target_size + + if image.max() <= 1.0: + img_batch = (image.permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8) + else: + img_batch = image.permute(0, 2, 3, 1).cpu().numpy().astype(np.uint8) + + calibrated_images = [] + for img in img_batch: + if img.shape[1] != 1920 or img.shape[0] != 1080: + img = cv2.resize(img, (1920, 1080), interpolation=cv2.INTER_LINEAR) + + img = cv2.remap(img, map1_defish, map2_defish, interpolation=cv2.INTER_LINEAR) + img = cv2.warpPerspective( + img, _CALIB_H_FULL, _CALIB_OUT_SIZE, flags=cv2.INTER_LINEAR + ) + + top, left, right, bottom = _CALIB_CROP + img = img[top:bottom, left:right] + img = cv2.resize(img, (target_w, target_h), interpolation=cv2.INTER_LINEAR) + calibrated_images.append(img) + + out_batch = np.stack(calibrated_images, axis=0) + out_tensor = torch.from_numpy(out_batch).permute(0, 3, 1, 2).float() / 255.0 + return out_tensor + + +def point_cloud_to_sparse_depth( + points_xyz: np.ndarray, + target_shape: Tuple[int, int], + max_depth_m: float, +) -> np.ndarray: + """Project xyz radar points (meters) to a sparse depth image.""" + target_h, target_w = target_shape + sparse_depth = np.zeros((target_h, target_w), dtype=np.float32) + + if points_xyz.size == 0: + return sparse_depth + + pts = np.asarray(points_xyz, dtype=np.float32) + if pts.ndim != 2 or pts.shape[1] != 3: + return sparse_depth + + valid = np.isfinite(pts).all(axis=1) + valid &= pts[:, 2] > 0.0 + valid &= pts[:, 2] <= float(max_depth_m) + pts = pts[valid] + if pts.shape[0] == 0: + return sparse_depth + + sx = target_w / float(_ZED_REF_W) + sy = target_h / float(_ZED_REF_H) + fx = _K_ZED_REF[0, 0] * sx + fy = _K_ZED_REF[1, 1] * sy + cx = _K_ZED_REF[0, 2] * sx + cy = _K_ZED_REF[1, 2] * sy + + z = pts[:, 2] + u = np.rint(pts[:, 0] * fx / z + cx).astype(np.int32) + v = np.rint(pts[:, 1] * fy / z + cy).astype(np.int32) + + in_bounds = (u >= 0) & (u < target_w) & (v >= 0) & (v < target_h) + if not np.any(in_bounds): + return sparse_depth + + u = u[in_bounds] + v = v[in_bounds] + z = z[in_bounds].astype(np.float32) + + min_depth = np.full((target_h, target_w), np.inf, dtype=np.float32) + np.minimum.at(min_depth, (v, u), z) + min_depth[~np.isfinite(min_depth)] = 0.0 + return min_depth + + +def build_radar_gt_map( + depth_m: np.ndarray, + sparse_depth: np.ndarray, + patch_size: Tuple[int, int], + max_dist_correspondence: float, +) -> np.ndarray: + """Build confidence GT using local depth consistency around each radar pixel.""" + h, w = depth_m.shape + radar_gt = np.zeros((h, w), dtype=np.float32) + + ys, xs = np.where(sparse_depth > 0) + if len(ys) == 0: + return radar_gt + + ext_h, ext_w = int(patch_size[0]), int(patch_size[1]) + for y, x in zip(ys, xs): + radar_depth = sparse_depth[y, x] + + delta_x1 = min(x, ext_w) + delta_y1 = min(y, ext_h) + delta_x2 = min(w - x, ext_w) + delta_y2 = min(h - y, ext_h) + + x1 = x - delta_x1 + y1 = y - delta_y1 + x2 = x + delta_x2 + y2 = y + delta_y2 + + distance = np.abs(depth_m[y1:y2, x1:x2] - radar_depth) + gt_label = (distance < float(max_dist_correspondence)).astype(np.float32) + radar_gt[y1:y2, x1:x2] = gt_label + + return radar_gt + + +def make_rice_collate_fn( + input_height: int, + input_width: int, + radar_max_depth_m: float, + max_dist_correspondence: float, + patch_size: Tuple[int, int], +) -> Callable[[Sequence[Dict[str, object]]], Tuple[torch.Tensor, ...]]: + """Create collate_fn for RiceDataset samples. + + Each dataset sample should contain: + - sample_idx: int + - dji_rgb: (H, W, 3) uint8 + - zed_depth_mm: (H, W) uint16 + - radar_pcd_xyz: (N, 3) float32 in meters + """ + + mean = torch.tensor(IMAGENET_MEAN, dtype=torch.float32).view(1, 3, 1, 1) + std = torch.tensor(IMAGENET_STD, dtype=torch.float32).view(1, 3, 1, 1) + + def _collate(batch: Sequence[Dict[str, object]]) -> Tuple[torch.Tensor, ...]: + if len(batch) == 0: + raise ValueError("Received empty batch in collate function") + + sample_indices = [] + rgb_batch = [] + depth_batch = [] + radar_batch = [] + radar_gt_batch = [] + + for sample in batch: + sample_indices.append(int(sample["sample_idx"])) + + rgb = np.asarray(sample["dji_rgb"]).copy() + if rgb.ndim != 3 or rgb.shape[2] != 3: + raise ValueError(f"Expected RGB shape (H, W, 3), got {rgb.shape}") + rgb_batch.append(torch.from_numpy(np.transpose(rgb, (2, 0, 1)))) + + depth_mm = np.asarray(sample["zed_depth_mm"]).copy() + depth_mm = resize_depth_mm(depth_mm, (input_height, input_width)) + depth_m = depth_mm.astype(np.float32) / 1000.0 + invalid = ~(np.isfinite(depth_m) & (depth_m > 0.0)) + depth_m[invalid] = 0.0 + depth_batch.append(depth_m) + + radar_points = np.asarray(sample["radar_pcd_xyz"], dtype=np.float32) + if radar_points.ndim != 2 or radar_points.shape[1] != 3: + radar_points = np.zeros((0, 3), dtype=np.float32) + + if radar_points.shape[0] == 0: + center_v = float(depth_m[input_height // 2, input_width // 2]) + if not np.isfinite(center_v): + center_v = 0.0 + radar_points = np.array([[0.0, 0.0, center_v]], dtype=np.float32) + + sparse_depth = point_cloud_to_sparse_depth( + radar_points, + target_shape=(input_height, input_width), + max_depth_m=radar_max_depth_m, + ) + radar_gt = build_radar_gt_map( + depth_m, + sparse_depth, + patch_size=patch_size, + max_dist_correspondence=max_dist_correspondence, + ) + radar_batch.append(sparse_depth) + radar_gt_batch.append(radar_gt) + + rgb_tensor = torch.stack(rgb_batch, dim=0).float() + rgb_tensor = dji_rgb_collator(rgb_tensor, target_size=(input_height, input_width)) + rgb_tensor = (rgb_tensor - mean) / std + + depth_tensor = torch.from_numpy(np.stack(depth_batch, axis=0)).float().unsqueeze(1) + radar_tensor = torch.from_numpy(np.stack(radar_batch, axis=0)).float().unsqueeze(1) + radar_gt_tensor = ( + torch.from_numpy(np.stack(radar_gt_batch, axis=0)).float().unsqueeze(1) + ) + idx_tensor = torch.tensor(sample_indices, dtype=torch.long) + + return idx_tensor, rgb_tensor, depth_tensor, radar_tensor, radar_gt_tensor + + return _collate + + +# Fisheye RGB Handler Functions ## +def fisheye_rgb_collator( + image: torch.Tensor, + target_size: Tuple[int, int] = (128, 256), +) -> torch.Tensor: + """Calibrate and resize Fisheye RGB image batch. + + Args: + image: Batch of Fisheye RGB images as torch tensor (B, C, H, W) in CHW format + target_size: Target resolution as (height, width) + + Returns: + Batch of calibrated and resized torch tensors in CHW format (B, C, H, W) + """ + IMAGE_WIDTH = 1920 + IMAGE_HEIGHT = 1080 + FOCAL_LENGTH_X = 0.613260 + FOCAL_LENGTH_Y = 0.613260 + CENTER_X = 0.5 + CENTER_Y = 0.5 + K1 = -0.120000 + K2 = -0.015000 + + w, h = IMAGE_WIDTH, IMAGE_HEIGHT + x_out, y_out = np.meshgrid(np.arange(w), np.arange(h)) + x_norm = (x_out - w * CENTER_X) / (w * FOCAL_LENGTH_X) + y_norm = (y_out - h * CENTER_Y) / (h * FOCAL_LENGTH_Y) + r = np.sqrt(x_norm**2 + y_norm**2) + r_distorted = r + K1 * r**2 + K2 * r**3 + r_safe = np.where(r > 0, r, 1.0) + scale = np.where(r > 0, r_distorted / r_safe, 1.0) + x_norm_distorted = x_norm * scale + y_norm_distorted = y_norm * scale + map_x = (x_norm_distorted * (w * FOCAL_LENGTH_X) + w * CENTER_X).astype(np.float32) + map_y = (y_norm_distorted * (h * FOCAL_LENGTH_Y) + h * CENTER_Y).astype(np.float32) + + if not isinstance(image, torch.Tensor): + raise ValueError(f"Expected torch.Tensor, got {type(image)}") + + if image.dim() != 4: + raise ValueError( + f"Expected 4D tensor (B, C, H, W), got {image.dim()}D tensor with shape {image.shape}" + ) + + if image.max() <= 1.0: + img_batch = (image.permute(0, 2, 3, 1).cpu().numpy() * 255).astype(np.uint8) + else: + img_batch = image.permute(0, 2, 3, 1).cpu().numpy().astype(np.uint8) + + calibrated_images = [] + target_h, target_w = target_size + + for img in img_batch: + if img.shape[1] != IMAGE_WIDTH or img.shape[0] != IMAGE_HEIGHT: + img = cv2.resize( + img, (IMAGE_WIDTH, IMAGE_HEIGHT), interpolation=cv2.INTER_LINEAR + ) + + img = cv2.remap( + img, + map_x, + map_y, + interpolation=cv2.INTER_LINEAR, + borderMode=cv2.BORDER_CONSTANT, + borderValue=(0, 0, 0), + ) + + img = cv2.resize(img, (target_w, target_h), interpolation=cv2.INTER_LINEAR) + calibrated_images.append(img) + + out_batch = np.stack(calibrated_images, axis=0) + out_tensor = torch.from_numpy(out_batch).permute(0, 3, 1, 2).float() / 255.0 + return out_tensor diff --git a/src/Baselines/cafnet/dataloader.py b/src/Baselines/cafnet/dataloader.py new file mode 100644 index 0000000000000000000000000000000000000000..3fa0622cf57cba1620e5fa22c39b10176532a84e --- /dev/null +++ b/src/Baselines/cafnet/dataloader.py @@ -0,0 +1,100 @@ +from typing import Optional + +from torch.utils.data import DataLoader + +from collate_fn_helpers import make_rice_collate_fn +from rice_dataset import RiceDataset + + +def _build_dataset( + args, + split: str, + base_dir: Optional[str] = None, + split_json_path: Optional[str] = None, +) -> RiceDataset: + return RiceDataset( + base_dir=base_dir or args.base_dir, + split_json_path=args.split_json if split_json_path is None else split_json_path, + split=split, + input_height=args.input_height, + input_width=args.input_width, + patch_size=args.patch_size, + ) + + +def _build_loader( + args, + split: str, + batch_size: int, + shuffle: bool, + drop_last: bool, + pin_memory: bool, + base_dir: Optional[str] = None, + split_json_path: Optional[str] = None, +): + dataset = _build_dataset( + args, + split=split, + base_dir=base_dir, + split_json_path=split_json_path, + ) + collate_fn = make_rice_collate_fn( + input_height=args.input_height, + input_width=args.input_width, + radar_max_depth_m=args.radar_max_depth_m, + max_dist_correspondence=args.max_dist_correspondence, + patch_size=dataset.patch_size, + ) + return DataLoader( + dataset, + batch_size=batch_size, + shuffle=shuffle, + num_workers=args.num_workers, + pin_memory=pin_memory, + drop_last=drop_last, + collate_fn=collate_fn, + ) + + +def create_train_test_loaders(args, pin_memory: bool = False): + train_loader = _build_loader( + args, + split="train", + batch_size=args.batch_size, + shuffle=True, + drop_last=True, + pin_memory=pin_memory, + ) + test_loader = _build_loader( + args, + split="test", + batch_size=args.batch_size, + shuffle=False, + drop_last=False, + pin_memory=pin_memory, + ) + return train_loader, test_loader + + +def create_inference_loader(args, pin_memory: bool = False): + """Create the single packaged Smoke-Eval loader used for inference.""" + + test_base_dir = getattr(args, "test_base_dir", "") + if not test_base_dir: + raise ValueError("Config must define 'test_base_dir' for inference.") + + test_split = getattr(args, "test_split", "train") + test_split_json = getattr(args, "test_split_json", None) + if not test_split_json: + test_split_json = None + + return _build_loader( + args, + split=test_split, + batch_size=args.batch_size, + shuffle=False, + drop_last=False, + pin_memory=pin_memory, + base_dir=test_base_dir, + split_json_path=test_split_json, + ) diff --git a/src/Baselines/cafnet/extract_pcd_from_depth.py b/src/Baselines/cafnet/extract_pcd_from_depth.py new file mode 100644 index 0000000000000000000000000000000000000000..bab2a51c9c496d10d946836cc112c0410b1e1c9d --- /dev/null +++ b/src/Baselines/cafnet/extract_pcd_from_depth.py @@ -0,0 +1,96 @@ +import cv2 +import numpy as np + +# ZED intrinsics at reference resolution 1280x720 (same values as PointCloudConverter) +_K_ZED_REF = np.array( + [ + [521.581604, 0.0, 636.33398438], + [0.0, 521.581604, 373.10964966], + [0.0, 0.0, 1.0], + ], + dtype=np.float64, +) +_ZED_REF_W = 1280 +_ZED_REF_H = 720 + + +def sample_depth_as_radar( + depth_mm: np.ndarray, + n_samples: int = 100, + target_shape: tuple = (300, 1280), + max_depth_m: float = 11.2, + seed: int | None = None, +) -> tuple: + """ + Randomly sample points from a ground truth ZED depth map and treat them as + radar points, mimicking the sparse depth input the model expects. + + The input depth is resized from its native resolution (e.g. 896x504) to + target_shape using nearest-neighbor interpolation so raw mm values are + preserved. Camera intrinsics are scaled from the 1280x720 ZED reference to + match the target resolution. + + Args: + depth_mm: Ground truth depth map, shape (H, W), dtype uint16, in mm. + n_samples: Number of points to randomly sample (default: 100). + target_shape: (target_H, target_W) to resize to before sampling. + Default (300, 1280) matches the model's required input. + max_depth_m: Maximum valid depth in meters — pixels beyond this are + treated as invalid (default: 11.2 m). + seed: Optional random seed for reproducibility. + + Returns: + points (np.ndarray): (N, 3) float32 array of [X, Y, Z] in meters, + in camera coordinate frame. N <= n_samples. + sparse_depth (np.ndarray): (target_H, target_W) float32 sparse depth map + with only the N sampled pixels filled (meters), + zeros elsewhere. + """ + target_h, target_w = target_shape + + # --- 1. Resize depth map (nearest-neighbor preserves raw mm values) --- + in_h, in_w = depth_mm.shape + if (in_h, in_w) != (target_h, target_w): + depth_resized = cv2.resize( + depth_mm, (target_w, target_h), interpolation=cv2.INTER_NEAREST + ) + else: + depth_resized = depth_mm.copy() + + # --- 2. Scale intrinsics from 1280x720 reference to target resolution --- + sx = target_w / float(_ZED_REF_W) + sy = target_h / float(_ZED_REF_H) + fx = _K_ZED_REF[0, 0] * sx + fy = _K_ZED_REF[1, 1] * sy + cx = _K_ZED_REF[0, 2] * sx + cy = _K_ZED_REF[1, 2] * sy + + # --- 3. Convert to float meters and find valid pixels --- + depth_m = depth_resized.astype(np.float32) / 1000.0 + valid_mask = (depth_m > 0) & (depth_m <= max_depth_m) + valid_v, valid_u = np.where(valid_mask) # row (V), col (U) + + if len(valid_v) == 0: + return ( + np.zeros((0, 3), dtype=np.float32), + np.zeros((target_h, target_w), dtype=np.float32), + ) + + # --- 4. Randomly sample up to n_samples valid pixels --- + rng = np.random.default_rng(seed) + n = min(n_samples, len(valid_v)) + indices = rng.choice(len(valid_v), size=n, replace=False) + sampled_v = valid_v[indices] + sampled_u = valid_u[indices] + sampled_z = depth_m[sampled_v, sampled_u] + + # --- 5. Back-project to 3D camera coordinates (pinhole model) --- + X = (sampled_u - cx) * sampled_z / fx + Y = (sampled_v - cy) * sampled_z / fy + points = np.stack([X, Y, sampled_z], axis=1).astype(np.float32) # (N, 3) + + # --- 6. Build sparse depth map --- + sparse_depth = np.zeros((target_h, target_w), dtype=np.float32) + sparse_depth[sampled_v, sampled_u] = sampled_z + + return points, sparse_depth diff --git a/src/Baselines/cafnet/inference.py b/src/Baselines/cafnet/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..8b073c38797d5d076bdaa82d4ef51979c6fa78e1 --- /dev/null +++ b/src/Baselines/cafnet/inference.py @@ -0,0 +1,224 @@ +import argparse +import os +from typing import Dict, List + +import numpy as np +import torch +import torch.distributed as dist +import yaml +from accelerate import Accelerator +from accelerate.utils import DistributedDataParallelKwargs, set_seed +from safetensors.torch import load_file +from tqdm.auto import tqdm + +from dataloader import create_inference_loader +from models.model import CaFNet + + +DEFAULT_CONFIG = { + # Packaged evaluation dataset. + "base_dir": "", + "split_json": None, + "test_base_dir": None, + "test_split": "train", + "test_split_json": None, + # Input and radar processing + "input_height": 288, + "input_width": 512, + "radar_max_depth_m": 11.2, + "max_dist_correspondence": 0.5, + "patch_size": None, + # Model + "encoder": "resnet34_bts", + "encoder_radar": "resnet18", + "radar_input_channels": 1, + "bts_size": 512, + "max_depth": 11.2, + # Runtime + "batch_size": 8, + # Windows uses spawn-based multiprocessing; keep the public evaluation + # entry point portable and deterministic by default. + "num_workers": 0, + "seed": 42, + "cpu": False, + "mixed_precision": "fp16", + "checkpoint_path": "checkpoints/cafnet.safetensors", + "prediction_dir": "prediction", +} + + +def parse_args(): + parser = argparse.ArgumentParser(description="Run CaFNet inference on Smoke-Eval.") + parser.add_argument("--config", type=str, required=True, help="Path to YAML config") + return parser.parse_args() + + +def load_config(path): + with open(path, "r") as f: + cfg = yaml.safe_load(f) or {} + if not isinstance(cfg, dict): + raise ValueError("Config must be a YAML mapping (key-value pairs).") + + merged = dict(DEFAULT_CONFIG) + merged.update(cfg) + + if not merged["test_base_dir"]: + raise ValueError("Config must define 'test_base_dir'.") + if not merged["checkpoint_path"]: + raise ValueError("Config must define 'checkpoint_path'.") + if not os.path.isfile(merged["checkpoint_path"]): + raise FileNotFoundError(f"Checkpoint not found: {merged['checkpoint_path']}") + if merged.get("radar_input_channels", 1) != 1: + raise ValueError("radar_input_channels must be 1 for this setup.") + + return argparse.Namespace(**merged) + + +def build_model_args(args): + return argparse.Namespace( + encoder=args.encoder, + encoder_radar=args.encoder_radar, + radar_input_channels=args.radar_input_channels, + input_height=args.input_height, + input_width=args.input_width, + max_depth=args.max_depth, + bts_size=args.bts_size, + ) + + +def _extract_model_state(checkpoint): + if isinstance(checkpoint, dict) and isinstance(checkpoint.get("model"), dict): + return checkpoint["model"] + if isinstance(checkpoint, dict): + return checkpoint + raise ValueError("Unsupported checkpoint format.") + + +def _gather_objects(accelerator, obj): + if accelerator.num_processes == 1: + return [obj] + if not dist.is_available() or not dist.is_initialized(): + return [obj] + + gathered = [None for _ in range(accelerator.num_processes)] + dist.all_gather_object(gathered, obj) + return gathered + + +def _merge_predictions(all_rank_predictions): + merged: Dict[str, Dict[int, np.ndarray]] = {} + for rank_dict in all_rank_predictions: + if not rank_dict: + continue + for seq_name, frame_map in rank_dict.items(): + seq_slot = merged.setdefault(seq_name, {}) + for frame_idx, pred in frame_map.items(): + frame_idx = int(frame_idx) + if frame_idx not in seq_slot: + seq_slot[frame_idx] = pred + return merged + + +def _save_sequence_predictions(predictions, out_dir): + os.makedirs(out_dir, exist_ok=True) + for seq_name in sorted(predictions.keys()): + frame_map = predictions[seq_name] + ordered_frames = sorted(frame_map.keys()) + if not ordered_frames: + pred_stack = np.zeros((0,), dtype=np.float32) + else: + pred_stack = np.stack([frame_map[k] for k in ordered_frames], axis=0).astype( + np.float32, + copy=False, + ) + np.save(os.path.join(out_dir, f"{seq_name.lower()}_pred.npy"), pred_stack) + + +def _run_loader_inference(accelerator, model, loader, samples, save_dir, desc): + model.eval() + local_preds: Dict[str, Dict[int, np.ndarray]] = {} + + with torch.no_grad(): + pbar = tqdm( + loader, + desc=desc, + disable=not accelerator.is_local_main_process, + dynamic_ncols=True, + leave=False, + ) + for batch in pbar: + sample_idx, image, depth_gt, radar, radar_gt = batch + + image = image.to(accelerator.device, non_blocking=True) + radar = radar.to(accelerator.device, non_blocking=True) + # Kept for parity with validation loop structure. + _ = depth_gt.to(accelerator.device, non_blocking=True) + _ = radar_gt.to(accelerator.device, non_blocking=True) + + focal = torch.ones((image.size(0),), device=image.device) + _, _, _, _, depth_est, _, _ = model(image, radar, focal) + + pred_np = depth_est.detach().float().cpu().numpy() + if pred_np.ndim == 4 and pred_np.shape[1] == 1: + pred_np = pred_np[:, 0] + + if torch.is_tensor(sample_idx): + sample_idx_list = sample_idx.detach().cpu().tolist() + else: + sample_idx_list = list(sample_idx) + + for local_i, sample_i in enumerate(sample_idx_list): + seq_name, frame_idx = samples[int(sample_i)] + seq_slot = local_preds.setdefault(seq_name, {}) + frame_idx = int(frame_idx) + if frame_idx not in seq_slot: + seq_slot[frame_idx] = pred_np[local_i].astype(np.float32, copy=False) + + gathered = _gather_objects(accelerator, local_preds) + if accelerator.is_main_process: + merged = _merge_predictions(gathered) + _save_sequence_predictions(merged, save_dir) + + accelerator.wait_for_everyone() + + +def main(): + cli = parse_args() + args = load_config(cli.config) + + set_seed(args.seed) + ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + accelerator = Accelerator( + mixed_precision=None if args.mixed_precision in ("no", "none") else args.mixed_precision, + cpu=args.cpu, + kwargs_handlers=[ddp_kwargs], + ) + + test_loader = create_inference_loader( + args, + pin_memory=(accelerator.device.type == "cuda"), + ) + test_samples: List = test_loader.dataset.samples + + model = CaFNet(build_model_args(args)) + + model, test_loader = accelerator.prepare(model, test_loader) + + state_dict = load_file(args.checkpoint_path, device="cpu") + accelerator.unwrap_model(model).load_state_dict(state_dict, strict=True) + + _run_loader_inference( + accelerator=accelerator, + model=model, + loader=test_loader, + samples=test_samples, + save_dir=args.prediction_dir, + desc="Inference", + ) + + if accelerator.is_main_process: + print(f"Saved predictions to: {args.prediction_dir}") + + +if __name__ == "__main__": + main() diff --git a/src/Baselines/cafnet/inference_config.yaml b/src/Baselines/cafnet/inference_config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..12b1e16844e23edfbee3fac192975e7d889f7af5 --- /dev/null +++ b/src/Baselines/cafnet/inference_config.yaml @@ -0,0 +1,29 @@ +# CaFNet inference config for the packaged Smoke-Eval data. +test_base_dir: "../../../evaluation_dataset/Smoke-Eval" +test_split: "train" +test_split_json: null + +# Input and radar preprocessing +input_height: 288 +input_width: 512 +radar_max_depth_m: 11.2 +max_dist_correspondence: 0.5 +patch_size: [64, 128] + +# Model architecture +encoder: resnet34_bts +encoder_radar: resnet18 +radar_input_channels: 1 +bts_size: 512 +max_depth: 11.2 + +# Runtime +batch_size: 32 +num_workers: 0 +seed: 42 +cpu: false +mixed_precision: "fp16" + +# Checkpoint and output root +checkpoint_path: "../../../checkpoints/baselines/cafnet/cafnet.safetensors" +prediction_dir: "prediction" diff --git a/src/Baselines/cafnet/models/bts.py b/src/Baselines/cafnet/models/bts.py new file mode 100644 index 0000000000000000000000000000000000000000..5b3ce68d0a3213bb6ce0ce96e357ed2ea538f3b2 --- /dev/null +++ b/src/Baselines/cafnet/models/bts.py @@ -0,0 +1,367 @@ +# Copyright (C) 2019 Jin Han Lee +# +# This file is a part of BTS. +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see + +import torch +import torch.nn as nn +import torch.nn.functional as torch_nn_func +import math + + +def bn_init_as_tf(m): + if isinstance(m, nn.BatchNorm2d): + m.track_running_stats = True # These two lines enable using stats (moving mean and var) loaded from pretrained model + m.eval() # or zero mean and variance of one if the batch norm layer has no pretrained values + m.affine = True + m.requires_grad = True + + +def weights_init_xavier(m): + if isinstance(m, nn.Conv2d): + torch.nn.init.xavier_uniform_(m.weight) + if m.bias is not None: + torch.nn.init.zeros_(m.bias) + + +class atrous_conv(nn.Sequential): + def __init__(self, in_channels, out_channels, dilation, apply_bn_first=True): + super(atrous_conv, self).__init__() + self.atrous_conv = torch.nn.Sequential() + if apply_bn_first: + self.atrous_conv.add_module('first_bn', nn.BatchNorm2d(in_channels, momentum=0.01, affine=True, track_running_stats=True, eps=1.1e-5)) + + self.atrous_conv.add_module('aconv_sequence', nn.Sequential(nn.ReLU(), + nn.Conv2d(in_channels=in_channels, out_channels=out_channels*2, bias=False, kernel_size=1, stride=1, padding=0), + nn.BatchNorm2d(out_channels*2, momentum=0.01, affine=True, track_running_stats=True), + nn.ReLU(), + nn.Conv2d(in_channels=out_channels * 2, out_channels=out_channels, bias=False, kernel_size=3, stride=1, + padding=(dilation, dilation), dilation=dilation))) + + def forward(self, x): + return self.atrous_conv.forward(x) + +class upconv(nn.Module): + def __init__(self, in_channels, out_channels, ratio=2): + super(upconv, self).__init__() + self.elu = nn.ELU() + self.conv = nn.Conv2d(in_channels=in_channels, out_channels=out_channels, bias=False, kernel_size=3, stride=1, padding=1) + self.ratio = ratio + + def forward(self, x): + up_x = torch_nn_func.interpolate(x, scale_factor=self.ratio, mode='nearest') + out = self.conv(up_x) + out = self.elu(out) + return out + +class reduction_1x1(nn.Sequential): + def __init__(self, num_in_filters, num_out_filters, max_depth, is_final=False): + super(reduction_1x1, self).__init__() + self.max_depth = max_depth + self.is_final = is_final + self.sigmoid = nn.Sigmoid() + self.reduc = torch.nn.Sequential() + + while num_out_filters >= 4: + if num_out_filters < 8: + if self.is_final: + self.reduc.add_module('final', torch.nn.Sequential(nn.Conv2d(num_in_filters, out_channels=1, bias=False, + kernel_size=1, stride=1, padding=0), + nn.Sigmoid())) + else: + self.reduc.add_module('plane_params', torch.nn.Conv2d(num_in_filters, out_channels=3, bias=False, + kernel_size=1, stride=1, padding=0)) + break + else: + self.reduc.add_module('inter_{}_{}'.format(num_in_filters, num_out_filters), + torch.nn.Sequential(nn.Conv2d(in_channels=num_in_filters, out_channels=num_out_filters, + bias=False, kernel_size=1, stride=1, padding=0), + nn.ELU())) + + num_in_filters = num_out_filters + num_out_filters = num_out_filters // 2 + + def forward(self, net): + net = self.reduc.forward(net) + if not self.is_final: + theta = self.sigmoid(net[:, 0, :, :]) * math.pi / 3 + phi = self.sigmoid(net[:, 1, :, :]) * math.pi * 2 + dist = self.sigmoid(net[:, 2, :, :]) * self.max_depth + n1 = torch.mul(torch.sin(theta), torch.cos(phi)).unsqueeze(1) + n2 = torch.mul(torch.sin(theta), torch.sin(phi)).unsqueeze(1) + n3 = torch.cos(theta).unsqueeze(1) + n4 = dist.unsqueeze(1) + net = torch.cat([n1, n2, n3, n4], dim=1) + + return net + +class local_planar_guidance(nn.Module): + def __init__(self, upratio): + super(local_planar_guidance, self).__init__() + self.upratio = upratio + self.u = torch.arange(self.upratio).reshape([1, 1, self.upratio]).float() + self.v = torch.arange(int(self.upratio)).reshape([1, self.upratio, 1]).float() + self.upratio = float(upratio) + + def forward(self, plane_eq, focal): + plane_eq_expanded = torch.repeat_interleave(plane_eq, int(self.upratio), 2) + plane_eq_expanded = torch.repeat_interleave(plane_eq_expanded, int(self.upratio), 3) + n1 = plane_eq_expanded[:, 0, :, :] + n2 = plane_eq_expanded[:, 1, :, :] + n3 = plane_eq_expanded[:, 2, :, :] + n4 = plane_eq_expanded[:, 3, :, :] + + u = self.u.repeat(plane_eq.size(0), plane_eq.size(2) * int(self.upratio), plane_eq.size(3)).cuda() + u = (u - (self.upratio - 1) * 0.5) / self.upratio + + v = self.v.repeat(plane_eq.size(0), plane_eq.size(2), plane_eq.size(3) * int(self.upratio)).cuda() + v = (v - (self.upratio - 1) * 0.5) / self.upratio + + return n4 / (n1 * u + n2 * v + n3) + +class bts_gated_fuse(nn.Module): + def __init__(self, params, feat_out_channels, feat_out_channels_rad, num_features=512): + super(bts_gated_fuse, self).__init__() + self.params = params + self.weight5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[4], feat_out_channels[4], 1, 1, bias=False), + nn.Sigmoid()) + self.project5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[4], feat_out_channels[4], 1, 1, bias=False), + nn.ReLU()) + self.upconv5 = upconv(feat_out_channels[4], num_features) + self.bn5 = nn.BatchNorm2d(num_features, momentum=0.01, affine=True, eps=1.1e-5) + + self.conv5 = torch.nn.Sequential(nn.Conv2d(num_features + feat_out_channels[3], num_features, 3, 1, 1, bias=False), + nn.ELU()) + + self.weight4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[3], num_features, 1, 1, bias=False), + nn.Sigmoid()) + self.project4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[3], num_features, 1, 1, bias=False), + nn.ReLU()) + self.upconv4 = upconv(num_features, num_features // 2) + self.bn4 = nn.BatchNorm2d(num_features // 2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv4 = torch.nn.Sequential(nn.Conv2d(num_features // 2 + feat_out_channels[2], num_features // 2, 3, 1, 1, bias=False), + nn.ELU()) + self.bn4_2 = nn.BatchNorm2d(num_features // 2, momentum=0.01, affine=True, eps=1.1e-5) + + self.daspp_3 = atrous_conv(num_features // 2, num_features // 4, 3, apply_bn_first=False) + self.daspp_6 = atrous_conv(num_features // 2 + num_features // 4 + feat_out_channels[2], num_features // 4, 6) + self.daspp_12 = atrous_conv(num_features + feat_out_channels[2], num_features // 4, 12) + self.daspp_18 = atrous_conv(num_features + num_features // 4 + feat_out_channels[2], num_features // 4, 18) + self.daspp_24 = atrous_conv(num_features + num_features // 2 + feat_out_channels[2], num_features // 4, 24) + self.daspp_conv = torch.nn.Sequential(nn.Conv2d(num_features + num_features // 2 + num_features // 4, num_features // 4, 3, 1, 1, bias=False), + nn.ELU()) + self.reduc8x8 = reduction_1x1(num_features // 4, num_features // 4, self.params.max_depth) + self.lpg8x8 = local_planar_guidance(8) + + self.weight3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[2], num_features // 4, 1, 1, bias=False), + nn.Sigmoid()) + self.project3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[2], num_features // 4, 1, 1, bias=False), + nn.ReLU()) + self.upconv3 = upconv(num_features // 4, num_features // 4) + self.bn3 = nn.BatchNorm2d(num_features // 4, momentum=0.01, affine=True, eps=1.1e-5) + self.conv3 = torch.nn.Sequential(nn.Conv2d(num_features // 4 + feat_out_channels[1] + 1, num_features // 4, 3, 1, 1, bias=False), + nn.ELU()) + self.reduc4x4 = reduction_1x1(num_features // 4, num_features // 8, self.params.max_depth) + self.lpg4x4 = local_planar_guidance(4) + + self.weight2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[1], num_features // 4, 1, 1, bias=False), + nn.Sigmoid()) + self.project2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[1], num_features // 4, 1, 1, bias=False), + nn.ReLU()) + self.upconv2 = upconv(num_features // 4, num_features // 8) + self.bn2 = nn.BatchNorm2d(num_features // 8, momentum=0.01, affine=True, eps=1.1e-5) + self.conv2 = torch.nn.Sequential(nn.Conv2d(num_features // 8 + feat_out_channels[0] + 1, num_features // 8, 3, 1, 1, bias=False), + nn.ELU()) + + self.reduc2x2 = reduction_1x1(num_features // 8, num_features // 16, self.params.max_depth) + self.lpg2x2 = local_planar_guidance(2) + + self.weight1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[0], num_features // 8, 1, 1, bias=False), + nn.Sigmoid()) + self.project1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[0], num_features // 8, 1, 1, bias=False), + nn.ReLU()) + self.upconv1 = upconv(num_features // 8, num_features // 16) + self.reduc1x1 = reduction_1x1(num_features // 16, num_features // 32, self.params.max_depth, is_final=True) + self.conv1 = torch.nn.Sequential(nn.Conv2d(num_features // 16 + 4, num_features // 16, 3, 1, 1, bias=False), + nn.ELU()) + self.get_depth = torch.nn.Sequential(nn.Conv2d(num_features // 16, 1, 3, 1, 1, bias=False), + nn.Sigmoid()) + + self.pool5 = torch.nn.AvgPool2d(32, 32) + self.pool4 = torch.nn.AvgPool2d(16, 16) + self.pool3 = torch.nn.AvgPool2d(8, 8) + self.pool2 = torch.nn.AvgPool2d(4, 4) + self.pool1 = torch.nn.AvgPool2d(2, 2) + + def forward(self, img_features, rad_features, focal, radar_confidence): + skip0, skip1, skip2, skip3 = img_features[0], img_features[1], img_features[2], img_features[3] + rad_skip0, rad_skip1, rad_skip2, rad_skip3 = rad_features[0], rad_features[1], rad_features[2], rad_features[3] + + # prepare radar confidence + radar_confidence5 = self.pool5(radar_confidence) + radar_confidence4 = self.pool4(radar_confidence) + radar_confidence3 = self.pool3(radar_confidence) + radar_confidence2 = self.pool2(radar_confidence) + radar_confidence1 = self.pool1(radar_confidence) + + + rad_weight5 = self.weight5(rad_features[4]) + rad_project5 = self.project5(rad_features[4]) + + dense_features = torch.nn.ReLU()(img_features[4]) + dense_features = dense_features + rad_weight5*rad_project5*radar_confidence5 + upconv5 = self.upconv5(dense_features) # H/16 + upconv5 = self.bn5(upconv5) + concat5 = torch.cat([upconv5, skip3], dim=1) + iconv5 = self.conv5(concat5) + + rad_weight4 = self.weight4(rad_skip3) + rad_project4 = self.project4(rad_skip3) + + iconv5 = iconv5 + rad_weight4*rad_project4*radar_confidence4 + upconv4 = self.upconv4(iconv5) # H/8 + upconv4 = self.bn4(upconv4) + concat4 = torch.cat([upconv4, skip2], dim=1) + iconv4 = self.conv4(concat4) + iconv4 = self.bn4_2(iconv4) + + daspp_3 = self.daspp_3(iconv4) + concat4_2 = torch.cat([concat4, daspp_3], dim=1) + daspp_6 = self.daspp_6(concat4_2) + concat4_3 = torch.cat([concat4_2, daspp_6], dim=1) + daspp_12 = self.daspp_12(concat4_3) + concat4_4 = torch.cat([concat4_3, daspp_12], dim=1) + daspp_18 = self.daspp_18(concat4_4) + concat4_5 = torch.cat([concat4_4, daspp_18], dim=1) + daspp_24 = self.daspp_24(concat4_5) + concat4_daspp = torch.cat([iconv4, daspp_3, daspp_6, daspp_12, daspp_18, daspp_24], dim=1) + daspp_feat = self.daspp_conv(concat4_daspp) + rad_weight3 = self.weight3(rad_skip2) + rad_project3 = self.project3(rad_skip2) + daspp_feat = daspp_feat + rad_weight3*rad_project3*radar_confidence3 + + reduc8x8 = self.reduc8x8(daspp_feat) + plane_normal_8x8 = reduc8x8[:, :3, :, :] + plane_normal_8x8 = torch_nn_func.normalize(plane_normal_8x8, 2, 1) + plane_dist_8x8 = reduc8x8[:, 3, :, :] + plane_eq_8x8 = torch.cat([plane_normal_8x8, plane_dist_8x8.unsqueeze(1)], 1) + depth_8x8 = self.lpg8x8(plane_eq_8x8, focal) + depth_8x8_scaled = depth_8x8.unsqueeze(1) / self.params.max_depth + depth_8x8_scaled_ds = torch_nn_func.interpolate(depth_8x8_scaled, scale_factor=0.25, mode='nearest') + + upconv3 = self.upconv3(daspp_feat) # H/4 + upconv3 = self.bn3(upconv3) + concat3 = torch.cat([upconv3, skip1, depth_8x8_scaled_ds], dim=1) + iconv3 = self.conv3(concat3) + rad_weight2 = self.weight2(rad_skip1) + rad_project2 = self.project2(rad_skip1) + iconv3 = iconv3 + rad_weight2*rad_project2*radar_confidence2 + + reduc4x4 = self.reduc4x4(iconv3) + plane_normal_4x4 = reduc4x4[:, :3, :, :] + plane_normal_4x4 = torch_nn_func.normalize(plane_normal_4x4, 2, 1) + plane_dist_4x4 = reduc4x4[:, 3, :, :] + plane_eq_4x4 = torch.cat([plane_normal_4x4, plane_dist_4x4.unsqueeze(1)], 1) + depth_4x4 = self.lpg4x4(plane_eq_4x4, focal) + depth_4x4_scaled = depth_4x4.unsqueeze(1) / self.params.max_depth + depth_4x4_scaled_ds = torch_nn_func.interpolate(depth_4x4_scaled, scale_factor=0.5, mode='nearest') + + upconv2 = self.upconv2(iconv3) # H/2 + upconv2 = self.bn2(upconv2) + concat2 = torch.cat([upconv2, skip0, depth_4x4_scaled_ds], dim=1) + iconv2 = self.conv2(concat2) + rad_weight1 = self.weight1(rad_skip0) + rad_project1 = self.project1(rad_skip0) + iconv2 = iconv2 + rad_weight1*rad_project1*radar_confidence1 + + reduc2x2 = self.reduc2x2(iconv2) + plane_normal_2x2 = reduc2x2[:, :3, :, :] + plane_normal_2x2 = torch_nn_func.normalize(plane_normal_2x2, 2, 1) + plane_dist_2x2 = reduc2x2[:, 3, :, :] + plane_eq_2x2 = torch.cat([plane_normal_2x2, plane_dist_2x2.unsqueeze(1)], 1) + depth_2x2 = self.lpg2x2(plane_eq_2x2, focal) + depth_2x2_scaled = depth_2x2.unsqueeze(1) / self.params.max_depth + + rad_weight1 = self.weight1(rad_skip0) + rad_project1 = self.project1(rad_skip0) + + upconv1 = self.upconv1(iconv2) + reduc1x1 = self.reduc1x1(upconv1) + concat1 = torch.cat([upconv1, reduc1x1, depth_2x2_scaled, depth_4x4_scaled, depth_8x8_scaled], dim=1) + iconv1 = self.conv1(concat1) + final_depth = self.params.max_depth * self.get_depth(iconv1) + + return depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth + +class encoder_image(nn.Module): + def __init__(self, params): + super(encoder_image, self).__init__() + self.params = params + import torchvision.models as models + if params.encoder == 'densenet121_bts': + self.base_model = models.densenet121(pretrained=False).features + self.feat_names = ['relu0', 'pool0', 'transition1', 'transition2', 'norm5'] + self.feat_out_channels = [64, 64, 128, 256, 1024] + elif params.encoder == 'densenet161_bts': + self.base_model = models.densenet161(pretrained=False).features + self.feat_names = ['relu0', 'pool0', 'transition1', 'transition2', 'norm5'] + self.feat_out_channels = [96, 96, 192, 384, 2208] + elif params.encoder == 'resnet50_bts': + self.base_model = models.resnet50(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 256, 512, 1024, 2048] + elif params.encoder == 'resnet34_bts': + self.base_model = models.resnet34(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + elif params.encoder == 'resnet18_bts': + self.base_model = models.resnet18(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + elif params.encoder == 'resnet101_bts': + self.base_model = models.resnet101(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 256, 512, 1024, 2048] + elif params.encoder == 'resnext50_bts': + self.base_model = models.resnext50_32x4d(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 256, 512, 1024, 2048] + elif params.encoder == 'resnext101_bts': + self.base_model = models.resnext101_32x8d(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 256, 512, 1024, 2048] + elif params.encoder == 'mobilenetv2_bts': + self.base_model = models.mobilenet_v2(pretrained=False).features + self.feat_inds = [2, 4, 7, 11, 19] + self.feat_out_channels = [16, 24, 32, 64, 1280] + self.feat_names = [] + else: + print('Not supported encoder: {}'.format(params.encoder)) + + def forward(self, x): + feature = x + skip_feat = [] + i = 1 + for k, v in self.base_model._modules.items(): + if 'fc' in k or 'avgpool' in k: + continue + feature = v(feature) + if self.params.encoder == 'mobilenetv2_bts': + if i == 2 or i == 4 or i == 7 or i == 11 or i == 19: + skip_feat.append(feature) + else: + if any(x in k for x in self.feat_names): + skip_feat.append(feature) + i = i + 1 + return skip_feat diff --git a/src/Baselines/cafnet/models/model.py b/src/Baselines/cafnet/models/model.py new file mode 100644 index 0000000000000000000000000000000000000000..4eb94690d8ca447ff50a7c966793b62fb16dd544 --- /dev/null +++ b/src/Baselines/cafnet/models/model.py @@ -0,0 +1,28 @@ +import torch +import torch.nn as nn +from models.bts import encoder_image, bts_gated_fuse +from models.radar import encoder_radar_sparse_conv, encoder_radar_sub, decoder_radar + +class CaFNet(nn.Module): + def __init__(self, params, threshold=0.4): + super(CaFNet, self).__init__() + self.threshold = threshold + self.encoder = encoder_image(params) + self.encoder_radar1 = encoder_radar_sparse_conv(params) + self.decoder_radar = decoder_radar(params, self.encoder.feat_out_channels, self.encoder_radar1.feat_out_channels) + self.encoder_radar2 = encoder_radar_sub(params) + self.decoder = bts_gated_fuse(params, self.encoder.feat_out_channels, self.encoder_radar2.feat_out_channels, params.bts_size) + + + def forward(self, x, radar, focal): + + skip_feat = self.encoder(x) + skip_feat_radar = self.encoder_radar1(radar) + rad_confidence, rad_depth = self.decoder_radar(skip_feat, skip_feat_radar) + mask = (rad_confidence > self.threshold).float() + radar_new_input = torch.cat([mask*rad_depth, radar], axis=1) + skip_feat_radar_new = self.encoder_radar2(radar_new_input) + + depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth = self.decoder(skip_feat, skip_feat_radar_new, focal, rad_confidence) + + return depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth, rad_confidence, rad_depth diff --git a/src/Baselines/cafnet/models/radar.py b/src/Baselines/cafnet/models/radar.py new file mode 100644 index 0000000000000000000000000000000000000000..95729efa551242c9eb2c00c714b5ab056c7345c7 --- /dev/null +++ b/src/Baselines/cafnet/models/radar.py @@ -0,0 +1,212 @@ +from models.bts import upconv +import torch +import torch.nn as nn +import torchvision.models as models + +class encoder_radar_sparse_conv(nn.Module): + def __init__(self, params): + # radar encoder for the first stage + super(encoder_radar_sparse_conv, self).__init__() + + self.params = params + self.sparse_conv1 = SparseConv(params.radar_input_channels, 16, 7, activation='elu') + self.sparse_conv2 = SparseConv(16, 16, 5, activation='elu') + self.sparse_conv3 = SparseConv(16, 16, 3, activation='elu') + self.sparse_conv4 = SparseConv(16, 3, 3, activation='elu') + + if params.encoder_radar == 'resnet34': + self.base_model_radar = models.resnet34(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + elif params.encoder_radar == 'resnet18': + self.base_model_radar = models.resnet18(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + else: + print('Not supported encoder: {}'.format(params.encoder)) + + def forward(self, x): + mask = (x[:, 0] > 0).float().unsqueeze(1) + feature = x + feature, mask = self.sparse_conv1(feature, mask) + feature, mask = self.sparse_conv2(feature, mask) + feature, mask = self.sparse_conv3(feature, mask) + feature, mask = self.sparse_conv4(feature, mask) + + skip_feat = [] + i = 1 + for k, v in self.base_model_radar._modules.items(): + if 'fc' in k or 'avgpool' in k: + continue + feature = v(feature) + if any(x in k for x in self.feat_names): + skip_feat.append(feature) + i = i + 1 + return skip_feat + +class encoder_radar_sub(nn.Module): + def __init__(self, params): + # radar encoder for the second stage + super(encoder_radar_sub, self).__init__() + + self.params = params + import torchvision.models as models + self.conv = torch.nn.Sequential(nn.Conv2d(params.radar_input_channels+1, 3, 3, 1, 1, bias=False), + nn.ELU()) + + if params.encoder_radar == 'resnet34': + self.base_model_radar = models.resnet34(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + elif params.encoder_radar == 'resnet18': + self.base_model_radar = models.resnet18(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + else: + print('Not supported encoder: {}'.format(params.encoder)) + def forward(self, x): + feature = x + feature = self.conv(feature) + skip_feat = [] + i = 1 + for k, v in self.base_model_radar._modules.items(): + if 'fc' in k or 'avgpool' in k: + continue + feature = v(feature) + if any(x in k for x in self.feat_names): + skip_feat.append(feature) + i = i + 1 + return skip_feat + + +class decoder_radar(nn.Module): + def __init__(self, params, feat_out_channels_img, feat_out_channels_radar): + super(decoder_radar, self).__init__() + self.params = params + self.upconv5 = upconv(feat_out_channels_img[4]+feat_out_channels_radar[4], feat_out_channels_radar[4]//2) + self.bn5 = nn.BatchNorm2d(feat_out_channels_radar[4]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[4]//2, feat_out_channels_radar[4]//2, 3, 1, 1, bias=False), + nn.ELU()) + + self.upconv4 = upconv(feat_out_channels_img[3]+feat_out_channels_radar[3]+feat_out_channels_radar[4]//2, feat_out_channels_radar[3]//2) + self.bn4 = nn.BatchNorm2d(feat_out_channels_radar[3]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[3]//2, feat_out_channels_radar[3]//2, 3, 1, 1, bias=False), + nn.ELU()) + + self.upconv3 = upconv(feat_out_channels_img[2]+feat_out_channels_radar[2]+feat_out_channels_radar[3]//2, feat_out_channels_radar[2]//2) + self.bn3 = nn.BatchNorm2d(feat_out_channels_radar[2]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[2]//2, feat_out_channels_radar[2]//2, 3, 1, 1, bias=False), + nn.ELU()) + + self.upconv2 = upconv(feat_out_channels_img[1]+feat_out_channels_radar[1]+feat_out_channels_radar[2]//2, feat_out_channels_radar[1]//2) + self.bn2 = nn.BatchNorm2d(feat_out_channels_radar[1]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[1]//2, feat_out_channels_radar[1]//2, 3, 1, 1, bias=False), + nn.ELU()) + + self.upconv1 = upconv(feat_out_channels_img[0]+feat_out_channels_radar[0]+feat_out_channels_radar[1]//2, feat_out_channels_radar[0]//2) + self.bn1 = nn.BatchNorm2d(feat_out_channels_radar[0]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, feat_out_channels_radar[0]//2, 3, 1, 1, bias=False), + nn.ELU()) + + # self.get_depth = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, 1, 3, 1, 1, bias=False), + # nn.Sigmoid()) + + self.get_depth = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, 2, 3, 1, 1, bias=False), + nn.Sigmoid()) + + def forward(self, image_features, radar_features): + img_skip0, img_skip1, img_skip2, img_skip3, img_final = image_features[0], image_features[1], image_features[2], image_features[3], image_features[4] + rad_skip0, rad_skip1, rad_skip2, rad_skip3, rad_final = radar_features[0], radar_features[1], radar_features[2], radar_features[3], radar_features[4] + final = torch.cat([img_final, rad_final], axis=1) + upconv5 = self.upconv5(final) + upconv5 = self.bn5(upconv5) + upconv5 = self.conv5(upconv5) + upconv5 = torch.cat([img_skip3, rad_skip3, upconv5], axis=1) + + upconv4 = self.upconv4(upconv5) + upconv4 = self.bn4(upconv4) + upconv4 = self.conv4(upconv4) + upconv4 = torch.cat([img_skip2, rad_skip2, upconv4], axis=1) + + upconv3 = self.upconv3(upconv4) + upconv3 = self.bn3(upconv3) + upconv3 = self.conv3(upconv3) + upconv3 = torch.cat([img_skip1, rad_skip1, upconv3], axis=1) + + upconv2 = self.upconv2(upconv3) + upconv2 = self.bn2(upconv2) + upconv2 = self.conv2(upconv2) + upconv2 = torch.cat([img_skip0, rad_skip0, upconv2], axis=1) + + upconv1 = self.upconv1(upconv2) + upconv1 = self.bn1(upconv1) + upconv1 = self.conv1(upconv1) + + # confidence = self.get_depth(upconv1) + # depth = self.params.max_depth * confidence + depth_conf = self.get_depth(upconv1) + depth = self.params.max_depth * depth_conf[:, 0:1] + confidence = depth_conf[:, 1:2] + + return confidence, depth + + +class SparseConv(nn.Module): + + def __init__(self, + in_channels, + out_channels, + kernel_size, + activation='relu'): + super().__init__() + + padding = kernel_size//2 + + self.conv = nn.Conv2d( + in_channels, + out_channels, + kernel_size=kernel_size, + padding=padding, + bias=False) + + self.bias = nn.Parameter( + torch.zeros(out_channels), + requires_grad=True) + + self.sparsity = nn.Conv2d( + in_channels, + out_channels, + kernel_size=kernel_size, + padding=padding, + bias=False) + + kernel = torch.FloatTensor(torch.ones([kernel_size, kernel_size])).unsqueeze(0).unsqueeze(0) + + self.sparsity.weight = nn.Parameter( + data=kernel, + requires_grad=False) + + if activation == 'relu': + self.act = nn.ReLU(inplace=False) + elif activation == 'sigmoid': + self.act = nn.Sigmoid() + elif activation == 'elu': + self.act = nn.ELU() + + self.max_pool = nn.MaxPool2d( + kernel_size, + stride=1, + padding=padding) + + + + def forward(self, x, mask): + x = x*mask + x = self.conv(x) + normalizer = 1/(self.sparsity(mask)+1e-8) + x = x * normalizer + self.bias.unsqueeze(0).unsqueeze(2).unsqueeze(3) + x = self.act(x) + + mask = self.max_pool(mask) + + return x, mask diff --git a/src/Baselines/cafnet/rice_dataset.py b/src/Baselines/cafnet/rice_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..5d764f29ee853eda7ae653b77dfe7a45027d2f45 --- /dev/null +++ b/src/Baselines/cafnet/rice_dataset.py @@ -0,0 +1,121 @@ +import json +import os +from typing import Dict, List, Optional, Tuple + +import numpy as np +from torch.utils.data import Dataset + + +class RiceDataset(Dataset): + """Raw Rice dataset reader for DJI RGB, ZED depth and radar point clouds. + + This dataset returns raw per-frame arrays and leaves geometric processing to + `collate_fn_helpers.make_rice_collate_fn`. + """ + + def __init__( + self, + base_dir: str, + split_json_path: Optional[str] = None, + split: str = "train", + input_height: int = 288, + input_width: int = 512, + patch_size: Optional[Tuple[int, int]] = None, + ): + self.base_dir = base_dir + self.split = split + self.input_height = int(input_height) + self.input_width = int(input_width) + self.patch_size = self._resolve_patch_size(patch_size) + + test_sequences = self._load_test_split(split_json_path) + + all_sequences = sorted( + d + for d in os.listdir(base_dir) + if os.path.isdir(os.path.join(base_dir, d)) and not d.startswith(".") + ) + + self.sequences: List[str] = [] + for seq in all_sequences: + if split == "train" and seq in test_sequences: + continue + if split == "test" and seq not in test_sequences: + continue + if self._is_valid_sequence(os.path.join(base_dir, seq)): + self.sequences.append(seq) + + self.dji_rgb_mmaps: Dict[str, np.memmap] = {} + self.zed_depth_mmaps: Dict[str, np.memmap] = {} + self.samples: List[Tuple[str, int]] = [] + + for seq in self.sequences: + seq_dir = os.path.join(self.base_dir, seq) + dji_rgb_path = os.path.join(seq_dir, "dji_rgb.npy") + zed_depth_path = os.path.join(seq_dir, "zed_depth.npy") + + self.dji_rgb_mmaps[seq] = np.load(dji_rgb_path, mmap_mode="r") + self.zed_depth_mmaps[seq] = np.load(zed_depth_path, mmap_mode="r") + + n_frames = min( + len(self.dji_rgb_mmaps[seq]), + len(self.zed_depth_mmaps[seq]), + ) + for frame_idx in range(n_frames): + self.samples.append((seq, frame_idx)) + + def _resolve_patch_size( + self, patch_size: Optional[Tuple[int, int]] + ) -> Tuple[int, int]: + if patch_size is not None: + return int(patch_size[0]), int(patch_size[1]) + + # Scale default CaFNet patch size (50, 150) from 352x704. + base_h, base_w = 352, 704 + scale_h = self.input_height / float(base_h) + scale_w = self.input_width / float(base_w) + ext_h = max(1, int(round(50 * scale_h))) + ext_w = max(1, int(round(150 * scale_w))) + return ext_h, ext_w + + def _load_test_split(self, split_json_path: Optional[str]) -> set: + if not split_json_path or not os.path.exists(split_json_path): + return set() + with open(split_json_path, "r") as f: + payload = json.load(f) + return set(payload.get("test", [])) + + def _is_valid_sequence(self, seq_dir: str) -> bool: + dji_rgb_path = os.path.join(seq_dir, "dji_rgb.npy") + zed_depth_path = os.path.join(seq_dir, "zed_depth.npy") + pcd_dir = os.path.join(seq_dir, "pcd") + return ( + os.path.exists(dji_rgb_path) + and os.path.exists(zed_depth_path) + and os.path.isdir(pcd_dir) + ) + + def __len__(self) -> int: + return len(self.samples) + + def __getitem__(self, idx: int) -> Dict[str, object]: + seq, frame_idx = self.samples[idx] + seq_dir = os.path.join(self.base_dir, seq) + + dji_rgb = np.asarray(self.dji_rgb_mmaps[seq][frame_idx]).copy() + zed_depth_mm = np.asarray(self.zed_depth_mmaps[seq][frame_idx]).copy() + + pcd_path = os.path.join(seq_dir, "pcd", f"pcd_{frame_idx}.npy") + if os.path.exists(pcd_path): + radar_pcd_xyz = np.asarray(np.load(pcd_path), dtype=np.float32) + else: + radar_pcd_xyz = np.zeros((0, 3), dtype=np.float32) + + return { + "sample_idx": idx, + "sequence": seq, + "frame_idx": frame_idx, + "dji_rgb": dji_rgb, + "zed_depth_mm": zed_depth_mm, + "radar_pcd_xyz": radar_pcd_xyz, + } diff --git a/src/Baselines/cafnet/split.json b/src/Baselines/cafnet/split.json new file mode 100644 index 0000000000000000000000000000000000000000..f6e7a920c8068dcf74ab6acb272480fdca99e1d1 --- /dev/null +++ b/src/Baselines/cafnet/split.json @@ -0,0 +1,14 @@ +{ + "test": [ + "Dell-1", + "Dell-2", + "Smoke-Dell-1", + "Smoke-Dell-2", + "Keck-1", + "Keck-2", + "Keck-3", + "Smoke-keck-1", + "Smoke-keck-2", + "Smoke-keck-3" + ] +} \ No newline at end of file diff --git a/src/Baselines/cafnet_no_smoke/collate_fn_helpers.py b/src/Baselines/cafnet_no_smoke/collate_fn_helpers.py new file mode 100644 index 0000000000000000000000000000000000000000..8fa5d76dc84ee17fc5c56642b8de089d452f3e3e --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/collate_fn_helpers.py @@ -0,0 +1,404 @@ +import cv2 +import numpy as np +import torch +from functools import lru_cache +from typing import Callable, Dict, Sequence, Tuple, Union +from torchvision import transforms as T + + +IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32) +IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32) + +# ZED intrinsics at 1280x720 reference resolution. +_K_ZED_REF = np.array( + [ + [521.581604, 0.0, 636.33398438], + [0.0, 521.581604, 373.10964966], + [0.0, 0.0, 1.0], + ], + dtype=np.float64, +) +_ZED_REF_W = 1280 +_ZED_REF_H = 720 + +# DJI calibration constants. +_CALIB_K_DJI = np.array( + [ + [718.48555551, 0.0, 963.36465011], + [0.0, 720.25844189, 537.87569913], + [0.0, 0.0, 1.0], + ], + dtype=np.float64, +) +_CALIB_D_DJI = np.array( + [0.19022699, 0.03466753, 0.05858962, -0.07070669], dtype=np.float64 +) +_CALIB_DEFISH_SHAPE = (1920, 1080) +_CALIB_DEFISH_BALANCE = 0.2 +_CALIB_H_FULL = np.array( + [ + [0.8274446551892256, -0.0742944198979625, 80.23797348979947], + [-0.014725864916652691, 0.8471179917075127, 28.27366063997317], + [-5.083573451500717e-05, -6.846079418201229e-05, 1.0], + ], + dtype=np.float64, +) +_CALIB_OUT_SIZE = (1918, 1105) +_CALIB_CROP = (115, 255, 1400, 760) # top, left, right, bottom + + +@lru_cache(maxsize=1) +def _get_dji_defish_maps() -> Tuple[np.ndarray, np.ndarray]: + r_defish = np.eye(3) + k_new_defish = cv2.fisheye.estimateNewCameraMatrixForUndistortRectify( + _CALIB_K_DJI, + _CALIB_D_DJI, + _CALIB_DEFISH_SHAPE, + r_defish, + balance=_CALIB_DEFISH_BALANCE, + fov_scale=1.0, + ) + map1, map2 = cv2.fisheye.initUndistortRectifyMap( + _CALIB_K_DJI, + _CALIB_D_DJI, + r_defish, + k_new_defish, + _CALIB_DEFISH_SHAPE, + cv2.CV_16SC2, + ) + return map1, map2 + + +def resize_depth_mm(depth_mm: np.ndarray, target_size: Tuple[int, int]) -> np.ndarray: + target_h, target_w = target_size + if depth_mm.shape[:2] == (target_h, target_w): + return depth_mm + return cv2.resize(depth_mm, (target_w, target_h), interpolation=cv2.INTER_NEAREST) + + +def depth_collator( + depth: Union[torch.Tensor, np.ndarray], + max_depth_m: float = 11.2, + target_size: Tuple[int, int] = (128, 256), +) -> Union[torch.Tensor, np.ndarray]: + """Clamp, normalize to [0, 1], and resize depth.""" + is_numpy = isinstance(depth, np.ndarray) + if is_numpy: + depth = torch.from_numpy(depth) + + depth = depth.float() + original_shape = depth.shape + + if depth.dim() == 2: + depth = depth.unsqueeze(0) + elif depth.dim() == 3: + depth = depth.unsqueeze(1) + + invalid_mask = ~(torch.isfinite(depth) & (depth >= 0)) + depth[invalid_mask] = 0.0 + + depth = torch.clamp(depth, min=0.0, max=max_depth_m) + depth = depth / max_depth_m + + invalid_mask = ~torch.isfinite(depth) + depth[invalid_mask] = 0.0 + + resized = T.Resize( + target_size, interpolation=T.InterpolationMode.BILINEAR, antialias=True + )(depth) + + if len(original_shape) == 2: + resized = resized.squeeze(0) + + return resized.numpy() if is_numpy else resized + + +def dji_rgb_collator( + image: torch.Tensor, + target_size: Tuple[int, int] = (128, 256), +) -> torch.Tensor: + """Rectify and resize DJI RGB image batch. + + Args: + image: Tensor with shape (B, C, H, W). + target_size: Target resolution as (height, width). + + Returns: + Tensor in CHW format (B, C, H, W), float32 in [0, 1]. + """ + if not isinstance(image, torch.Tensor): + raise ValueError(f"Expected torch.Tensor, got {type(image)}") + + if image.dim() != 4: + raise ValueError( + f"Expected 4D tensor (B, C, H, W), got {image.dim()}D tensor with shape {image.shape}" + ) + + map1_defish, map2_defish = _get_dji_defish_maps() + target_h, target_w = target_size + + if image.max() <= 1.0: + img_batch = (image.permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8) + else: + img_batch = image.permute(0, 2, 3, 1).cpu().numpy().astype(np.uint8) + + calibrated_images = [] + for img in img_batch: + if img.shape[1] != 1920 or img.shape[0] != 1080: + img = cv2.resize(img, (1920, 1080), interpolation=cv2.INTER_LINEAR) + + img = cv2.remap(img, map1_defish, map2_defish, interpolation=cv2.INTER_LINEAR) + img = cv2.warpPerspective( + img, _CALIB_H_FULL, _CALIB_OUT_SIZE, flags=cv2.INTER_LINEAR + ) + + top, left, right, bottom = _CALIB_CROP + img = img[top:bottom, left:right] + img = cv2.resize(img, (target_w, target_h), interpolation=cv2.INTER_LINEAR) + calibrated_images.append(img) + + out_batch = np.stack(calibrated_images, axis=0) + out_tensor = torch.from_numpy(out_batch).permute(0, 3, 1, 2).float() / 255.0 + return out_tensor + + +def point_cloud_to_sparse_depth( + points_xyz: np.ndarray, + target_shape: Tuple[int, int], + max_depth_m: float, +) -> np.ndarray: + """Project xyz radar points (meters) to a sparse depth image.""" + target_h, target_w = target_shape + sparse_depth = np.zeros((target_h, target_w), dtype=np.float32) + + if points_xyz.size == 0: + return sparse_depth + + pts = np.asarray(points_xyz, dtype=np.float32) + if pts.ndim != 2 or pts.shape[1] != 3: + return sparse_depth + + valid = np.isfinite(pts).all(axis=1) + valid &= pts[:, 2] > 0.0 + valid &= pts[:, 2] <= float(max_depth_m) + pts = pts[valid] + if pts.shape[0] == 0: + return sparse_depth + + sx = target_w / float(_ZED_REF_W) + sy = target_h / float(_ZED_REF_H) + fx = _K_ZED_REF[0, 0] * sx + fy = _K_ZED_REF[1, 1] * sy + cx = _K_ZED_REF[0, 2] * sx + cy = _K_ZED_REF[1, 2] * sy + + z = pts[:, 2] + u = np.rint(pts[:, 0] * fx / z + cx).astype(np.int32) + v = np.rint(pts[:, 1] * fy / z + cy).astype(np.int32) + + in_bounds = (u >= 0) & (u < target_w) & (v >= 0) & (v < target_h) + if not np.any(in_bounds): + return sparse_depth + + u = u[in_bounds] + v = v[in_bounds] + z = z[in_bounds].astype(np.float32) + + min_depth = np.full((target_h, target_w), np.inf, dtype=np.float32) + np.minimum.at(min_depth, (v, u), z) + min_depth[~np.isfinite(min_depth)] = 0.0 + return min_depth + + +def build_radar_gt_map( + depth_m: np.ndarray, + sparse_depth: np.ndarray, + patch_size: Tuple[int, int], + max_dist_correspondence: float, +) -> np.ndarray: + """Build confidence GT using local depth consistency around each radar pixel.""" + h, w = depth_m.shape + radar_gt = np.zeros((h, w), dtype=np.float32) + + ys, xs = np.where(sparse_depth > 0) + if len(ys) == 0: + return radar_gt + + ext_h, ext_w = int(patch_size[0]), int(patch_size[1]) + for y, x in zip(ys, xs): + radar_depth = sparse_depth[y, x] + + delta_x1 = min(x, ext_w) + delta_y1 = min(y, ext_h) + delta_x2 = min(w - x, ext_w) + delta_y2 = min(h - y, ext_h) + + x1 = x - delta_x1 + y1 = y - delta_y1 + x2 = x + delta_x2 + y2 = y + delta_y2 + + distance = np.abs(depth_m[y1:y2, x1:x2] - radar_depth) + gt_label = (distance < float(max_dist_correspondence)).astype(np.float32) + radar_gt[y1:y2, x1:x2] = gt_label + + return radar_gt + + +def make_rice_collate_fn( + input_height: int, + input_width: int, + radar_max_depth_m: float, + max_dist_correspondence: float, + patch_size: Tuple[int, int], +) -> Callable[[Sequence[Dict[str, object]]], Tuple[torch.Tensor, ...]]: + """Create collate_fn for RiceDataset samples. + + Each dataset sample should contain: + - sample_idx: int + - dji_rgb: (H, W, 3) uint8 + - zed_depth_mm: (H, W) uint16 + - radar_pcd_xyz: (N, 3) float32 in meters + """ + + mean = torch.tensor(IMAGENET_MEAN, dtype=torch.float32).view(1, 3, 1, 1) + std = torch.tensor(IMAGENET_STD, dtype=torch.float32).view(1, 3, 1, 1) + + def _collate(batch: Sequence[Dict[str, object]]) -> Tuple[torch.Tensor, ...]: + if len(batch) == 0: + raise ValueError("Received empty batch in collate function") + + sample_indices = [] + rgb_batch = [] + depth_batch = [] + radar_batch = [] + radar_gt_batch = [] + + for sample in batch: + sample_indices.append(int(sample["sample_idx"])) + + rgb = np.asarray(sample["dji_rgb"]).copy() + if rgb.ndim != 3 or rgb.shape[2] != 3: + raise ValueError(f"Expected RGB shape (H, W, 3), got {rgb.shape}") + rgb_batch.append(torch.from_numpy(np.transpose(rgb, (2, 0, 1)))) + + depth_mm = np.asarray(sample["zed_depth_mm"]).copy() + depth_mm = resize_depth_mm(depth_mm, (input_height, input_width)) + depth_m = depth_mm.astype(np.float32) / 1000.0 + invalid = ~(np.isfinite(depth_m) & (depth_m > 0.0)) + depth_m[invalid] = 0.0 + depth_batch.append(depth_m) + + radar_points = np.asarray(sample["radar_pcd_xyz"], dtype=np.float32) + if radar_points.ndim != 2 or radar_points.shape[1] != 3: + radar_points = np.zeros((0, 3), dtype=np.float32) + + if radar_points.shape[0] == 0: + center_v = float(depth_m[input_height // 2, input_width // 2]) + if not np.isfinite(center_v): + center_v = 0.0 + radar_points = np.array([[0.0, 0.0, center_v]], dtype=np.float32) + + sparse_depth = point_cloud_to_sparse_depth( + radar_points, + target_shape=(input_height, input_width), + max_depth_m=radar_max_depth_m, + ) + radar_gt = build_radar_gt_map( + depth_m, + sparse_depth, + patch_size=patch_size, + max_dist_correspondence=max_dist_correspondence, + ) + radar_batch.append(sparse_depth) + radar_gt_batch.append(radar_gt) + + rgb_tensor = torch.stack(rgb_batch, dim=0).float() + rgb_tensor = dji_rgb_collator(rgb_tensor, target_size=(input_height, input_width)) + rgb_tensor = (rgb_tensor - mean) / std + + depth_tensor = torch.from_numpy(np.stack(depth_batch, axis=0)).float().unsqueeze(1) + radar_tensor = torch.from_numpy(np.stack(radar_batch, axis=0)).float().unsqueeze(1) + radar_gt_tensor = ( + torch.from_numpy(np.stack(radar_gt_batch, axis=0)).float().unsqueeze(1) + ) + idx_tensor = torch.tensor(sample_indices, dtype=torch.long) + + return idx_tensor, rgb_tensor, depth_tensor, radar_tensor, radar_gt_tensor + + return _collate + + +# Fisheye RGB Handler Functions ## +def fisheye_rgb_collator( + image: torch.Tensor, + target_size: Tuple[int, int] = (128, 256), +) -> torch.Tensor: + """Calibrate and resize Fisheye RGB image batch. + + Args: + image: Batch of Fisheye RGB images as torch tensor (B, C, H, W) in CHW format + target_size: Target resolution as (height, width) + + Returns: + Batch of calibrated and resized torch tensors in CHW format (B, C, H, W) + """ + IMAGE_WIDTH = 1920 + IMAGE_HEIGHT = 1080 + FOCAL_LENGTH_X = 0.613260 + FOCAL_LENGTH_Y = 0.613260 + CENTER_X = 0.5 + CENTER_Y = 0.5 + K1 = -0.120000 + K2 = -0.015000 + + w, h = IMAGE_WIDTH, IMAGE_HEIGHT + x_out, y_out = np.meshgrid(np.arange(w), np.arange(h)) + x_norm = (x_out - w * CENTER_X) / (w * FOCAL_LENGTH_X) + y_norm = (y_out - h * CENTER_Y) / (h * FOCAL_LENGTH_Y) + r = np.sqrt(x_norm**2 + y_norm**2) + r_distorted = r + K1 * r**2 + K2 * r**3 + r_safe = np.where(r > 0, r, 1.0) + scale = np.where(r > 0, r_distorted / r_safe, 1.0) + x_norm_distorted = x_norm * scale + y_norm_distorted = y_norm * scale + map_x = (x_norm_distorted * (w * FOCAL_LENGTH_X) + w * CENTER_X).astype(np.float32) + map_y = (y_norm_distorted * (h * FOCAL_LENGTH_Y) + h * CENTER_Y).astype(np.float32) + + if not isinstance(image, torch.Tensor): + raise ValueError(f"Expected torch.Tensor, got {type(image)}") + + if image.dim() != 4: + raise ValueError( + f"Expected 4D tensor (B, C, H, W), got {image.dim()}D tensor with shape {image.shape}" + ) + + if image.max() <= 1.0: + img_batch = (image.permute(0, 2, 3, 1).cpu().numpy() * 255).astype(np.uint8) + else: + img_batch = image.permute(0, 2, 3, 1).cpu().numpy().astype(np.uint8) + + calibrated_images = [] + target_h, target_w = target_size + + for img in img_batch: + if img.shape[1] != IMAGE_WIDTH or img.shape[0] != IMAGE_HEIGHT: + img = cv2.resize( + img, (IMAGE_WIDTH, IMAGE_HEIGHT), interpolation=cv2.INTER_LINEAR + ) + + img = cv2.remap( + img, + map_x, + map_y, + interpolation=cv2.INTER_LINEAR, + borderMode=cv2.BORDER_CONSTANT, + borderValue=(0, 0, 0), + ) + + img = cv2.resize(img, (target_w, target_h), interpolation=cv2.INTER_LINEAR) + calibrated_images.append(img) + + out_batch = np.stack(calibrated_images, axis=0) + out_tensor = torch.from_numpy(out_batch).permute(0, 3, 1, 2).float() / 255.0 + return out_tensor diff --git a/src/Baselines/cafnet_no_smoke/dataloader.py b/src/Baselines/cafnet_no_smoke/dataloader.py new file mode 100644 index 0000000000000000000000000000000000000000..3fa0622cf57cba1620e5fa22c39b10176532a84e --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/dataloader.py @@ -0,0 +1,100 @@ +from typing import Optional + +from torch.utils.data import DataLoader + +from collate_fn_helpers import make_rice_collate_fn +from rice_dataset import RiceDataset + + +def _build_dataset( + args, + split: str, + base_dir: Optional[str] = None, + split_json_path: Optional[str] = None, +) -> RiceDataset: + return RiceDataset( + base_dir=base_dir or args.base_dir, + split_json_path=args.split_json if split_json_path is None else split_json_path, + split=split, + input_height=args.input_height, + input_width=args.input_width, + patch_size=args.patch_size, + ) + + +def _build_loader( + args, + split: str, + batch_size: int, + shuffle: bool, + drop_last: bool, + pin_memory: bool, + base_dir: Optional[str] = None, + split_json_path: Optional[str] = None, +): + dataset = _build_dataset( + args, + split=split, + base_dir=base_dir, + split_json_path=split_json_path, + ) + collate_fn = make_rice_collate_fn( + input_height=args.input_height, + input_width=args.input_width, + radar_max_depth_m=args.radar_max_depth_m, + max_dist_correspondence=args.max_dist_correspondence, + patch_size=dataset.patch_size, + ) + return DataLoader( + dataset, + batch_size=batch_size, + shuffle=shuffle, + num_workers=args.num_workers, + pin_memory=pin_memory, + drop_last=drop_last, + collate_fn=collate_fn, + ) + + +def create_train_test_loaders(args, pin_memory: bool = False): + train_loader = _build_loader( + args, + split="train", + batch_size=args.batch_size, + shuffle=True, + drop_last=True, + pin_memory=pin_memory, + ) + test_loader = _build_loader( + args, + split="test", + batch_size=args.batch_size, + shuffle=False, + drop_last=False, + pin_memory=pin_memory, + ) + return train_loader, test_loader + + +def create_inference_loader(args, pin_memory: bool = False): + """Create the single packaged Smoke-Eval loader used for inference.""" + + test_base_dir = getattr(args, "test_base_dir", "") + if not test_base_dir: + raise ValueError("Config must define 'test_base_dir' for inference.") + + test_split = getattr(args, "test_split", "train") + test_split_json = getattr(args, "test_split_json", None) + if not test_split_json: + test_split_json = None + + return _build_loader( + args, + split=test_split, + batch_size=args.batch_size, + shuffle=False, + drop_last=False, + pin_memory=pin_memory, + base_dir=test_base_dir, + split_json_path=test_split_json, + ) diff --git a/src/Baselines/cafnet_no_smoke/extract_pcd_from_depth.py b/src/Baselines/cafnet_no_smoke/extract_pcd_from_depth.py new file mode 100644 index 0000000000000000000000000000000000000000..bab2a51c9c496d10d946836cc112c0410b1e1c9d --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/extract_pcd_from_depth.py @@ -0,0 +1,96 @@ +import cv2 +import numpy as np + +# ZED intrinsics at reference resolution 1280x720 (same values as PointCloudConverter) +_K_ZED_REF = np.array( + [ + [521.581604, 0.0, 636.33398438], + [0.0, 521.581604, 373.10964966], + [0.0, 0.0, 1.0], + ], + dtype=np.float64, +) +_ZED_REF_W = 1280 +_ZED_REF_H = 720 + + +def sample_depth_as_radar( + depth_mm: np.ndarray, + n_samples: int = 100, + target_shape: tuple = (300, 1280), + max_depth_m: float = 11.2, + seed: int | None = None, +) -> tuple: + """ + Randomly sample points from a ground truth ZED depth map and treat them as + radar points, mimicking the sparse depth input the model expects. + + The input depth is resized from its native resolution (e.g. 896x504) to + target_shape using nearest-neighbor interpolation so raw mm values are + preserved. Camera intrinsics are scaled from the 1280x720 ZED reference to + match the target resolution. + + Args: + depth_mm: Ground truth depth map, shape (H, W), dtype uint16, in mm. + n_samples: Number of points to randomly sample (default: 100). + target_shape: (target_H, target_W) to resize to before sampling. + Default (300, 1280) matches the model's required input. + max_depth_m: Maximum valid depth in meters — pixels beyond this are + treated as invalid (default: 11.2 m). + seed: Optional random seed for reproducibility. + + Returns: + points (np.ndarray): (N, 3) float32 array of [X, Y, Z] in meters, + in camera coordinate frame. N <= n_samples. + sparse_depth (np.ndarray): (target_H, target_W) float32 sparse depth map + with only the N sampled pixels filled (meters), + zeros elsewhere. + """ + target_h, target_w = target_shape + + # --- 1. Resize depth map (nearest-neighbor preserves raw mm values) --- + in_h, in_w = depth_mm.shape + if (in_h, in_w) != (target_h, target_w): + depth_resized = cv2.resize( + depth_mm, (target_w, target_h), interpolation=cv2.INTER_NEAREST + ) + else: + depth_resized = depth_mm.copy() + + # --- 2. Scale intrinsics from 1280x720 reference to target resolution --- + sx = target_w / float(_ZED_REF_W) + sy = target_h / float(_ZED_REF_H) + fx = _K_ZED_REF[0, 0] * sx + fy = _K_ZED_REF[1, 1] * sy + cx = _K_ZED_REF[0, 2] * sx + cy = _K_ZED_REF[1, 2] * sy + + # --- 3. Convert to float meters and find valid pixels --- + depth_m = depth_resized.astype(np.float32) / 1000.0 + valid_mask = (depth_m > 0) & (depth_m <= max_depth_m) + valid_v, valid_u = np.where(valid_mask) # row (V), col (U) + + if len(valid_v) == 0: + return ( + np.zeros((0, 3), dtype=np.float32), + np.zeros((target_h, target_w), dtype=np.float32), + ) + + # --- 4. Randomly sample up to n_samples valid pixels --- + rng = np.random.default_rng(seed) + n = min(n_samples, len(valid_v)) + indices = rng.choice(len(valid_v), size=n, replace=False) + sampled_v = valid_v[indices] + sampled_u = valid_u[indices] + sampled_z = depth_m[sampled_v, sampled_u] + + # --- 5. Back-project to 3D camera coordinates (pinhole model) --- + X = (sampled_u - cx) * sampled_z / fx + Y = (sampled_v - cy) * sampled_z / fy + points = np.stack([X, Y, sampled_z], axis=1).astype(np.float32) # (N, 3) + + # --- 6. Build sparse depth map --- + sparse_depth = np.zeros((target_h, target_w), dtype=np.float32) + sparse_depth[sampled_v, sampled_u] = sampled_z + + return points, sparse_depth diff --git a/src/Baselines/cafnet_no_smoke/inference.py b/src/Baselines/cafnet_no_smoke/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..15506ee6bac4e881a40006ccae5cd555c32cb987 --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/inference.py @@ -0,0 +1,224 @@ +import argparse +import os +from typing import Dict, List + +import numpy as np +import torch +import torch.distributed as dist +import yaml +from accelerate import Accelerator +from accelerate.utils import DistributedDataParallelKwargs, set_seed +from safetensors.torch import load_file +from tqdm.auto import tqdm + +from dataloader import create_inference_loader +from models.model import CaFNet + + +DEFAULT_CONFIG = { + # Packaged evaluation dataset. + "base_dir": "", + "split_json": None, + "test_base_dir": None, + "test_split": "train", + "test_split_json": None, + # Input and radar processing + "input_height": 288, + "input_width": 512, + "radar_max_depth_m": 11.2, + "max_dist_correspondence": 0.5, + "patch_size": None, + # Model + "encoder": "resnet34_bts", + "encoder_radar": "resnet18", + "radar_input_channels": 1, + "bts_size": 512, + "max_depth": 11.2, + # Runtime + "batch_size": 8, + # Windows uses spawn-based multiprocessing; keep the public evaluation + # entry point portable and deterministic by default. + "num_workers": 0, + "seed": 42, + "cpu": False, + "mixed_precision": "fp16", + "checkpoint_path": "checkpoints/cafnet_no_smoke.safetensors", + "prediction_dir": "prediction", +} + + +def parse_args(): + parser = argparse.ArgumentParser(description="Run CaFNet inference on Smoke-Eval.") + parser.add_argument("--config", type=str, required=True, help="Path to YAML config") + return parser.parse_args() + + +def load_config(path): + with open(path, "r") as f: + cfg = yaml.safe_load(f) or {} + if not isinstance(cfg, dict): + raise ValueError("Config must be a YAML mapping (key-value pairs).") + + merged = dict(DEFAULT_CONFIG) + merged.update(cfg) + + if not merged["test_base_dir"]: + raise ValueError("Config must define 'test_base_dir'.") + if not merged["checkpoint_path"]: + raise ValueError("Config must define 'checkpoint_path'.") + if not os.path.isfile(merged["checkpoint_path"]): + raise FileNotFoundError(f"Checkpoint not found: {merged['checkpoint_path']}") + if merged.get("radar_input_channels", 1) != 1: + raise ValueError("radar_input_channels must be 1 for this setup.") + + return argparse.Namespace(**merged) + + +def build_model_args(args): + return argparse.Namespace( + encoder=args.encoder, + encoder_radar=args.encoder_radar, + radar_input_channels=args.radar_input_channels, + input_height=args.input_height, + input_width=args.input_width, + max_depth=args.max_depth, + bts_size=args.bts_size, + ) + + +def _extract_model_state(checkpoint): + if isinstance(checkpoint, dict) and isinstance(checkpoint.get("model"), dict): + return checkpoint["model"] + if isinstance(checkpoint, dict): + return checkpoint + raise ValueError("Unsupported checkpoint format.") + + +def _gather_objects(accelerator, obj): + if accelerator.num_processes == 1: + return [obj] + if not dist.is_available() or not dist.is_initialized(): + return [obj] + + gathered = [None for _ in range(accelerator.num_processes)] + dist.all_gather_object(gathered, obj) + return gathered + + +def _merge_predictions(all_rank_predictions): + merged: Dict[str, Dict[int, np.ndarray]] = {} + for rank_dict in all_rank_predictions: + if not rank_dict: + continue + for seq_name, frame_map in rank_dict.items(): + seq_slot = merged.setdefault(seq_name, {}) + for frame_idx, pred in frame_map.items(): + frame_idx = int(frame_idx) + if frame_idx not in seq_slot: + seq_slot[frame_idx] = pred + return merged + + +def _save_sequence_predictions(predictions, out_dir): + os.makedirs(out_dir, exist_ok=True) + for seq_name in sorted(predictions.keys()): + frame_map = predictions[seq_name] + ordered_frames = sorted(frame_map.keys()) + if not ordered_frames: + pred_stack = np.zeros((0,), dtype=np.float32) + else: + pred_stack = np.stack([frame_map[k] for k in ordered_frames], axis=0).astype( + np.float32, + copy=False, + ) + np.save(os.path.join(out_dir, f"{seq_name.lower()}_pred.npy"), pred_stack) + + +def _run_loader_inference(accelerator, model, loader, samples, save_dir, desc): + model.eval() + local_preds: Dict[str, Dict[int, np.ndarray]] = {} + + with torch.no_grad(): + pbar = tqdm( + loader, + desc=desc, + disable=not accelerator.is_local_main_process, + dynamic_ncols=True, + leave=False, + ) + for batch in pbar: + sample_idx, image, depth_gt, radar, radar_gt = batch + + image = image.to(accelerator.device, non_blocking=True) + radar = radar.to(accelerator.device, non_blocking=True) + # Kept for parity with validation loop structure. + _ = depth_gt.to(accelerator.device, non_blocking=True) + _ = radar_gt.to(accelerator.device, non_blocking=True) + + focal = torch.ones((image.size(0),), device=image.device) + _, _, _, _, depth_est, _, _ = model(image, radar, focal) + + pred_np = depth_est.detach().float().cpu().numpy() + if pred_np.ndim == 4 and pred_np.shape[1] == 1: + pred_np = pred_np[:, 0] + + if torch.is_tensor(sample_idx): + sample_idx_list = sample_idx.detach().cpu().tolist() + else: + sample_idx_list = list(sample_idx) + + for local_i, sample_i in enumerate(sample_idx_list): + seq_name, frame_idx = samples[int(sample_i)] + seq_slot = local_preds.setdefault(seq_name, {}) + frame_idx = int(frame_idx) + if frame_idx not in seq_slot: + seq_slot[frame_idx] = pred_np[local_i].astype(np.float32, copy=False) + + gathered = _gather_objects(accelerator, local_preds) + if accelerator.is_main_process: + merged = _merge_predictions(gathered) + _save_sequence_predictions(merged, save_dir) + + accelerator.wait_for_everyone() + + +def main(): + cli = parse_args() + args = load_config(cli.config) + + set_seed(args.seed) + ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + accelerator = Accelerator( + mixed_precision=None if args.mixed_precision in ("no", "none") else args.mixed_precision, + cpu=args.cpu, + kwargs_handlers=[ddp_kwargs], + ) + + test_loader = create_inference_loader( + args, + pin_memory=(accelerator.device.type == "cuda"), + ) + test_samples: List = test_loader.dataset.samples + + model = CaFNet(build_model_args(args)) + + model, test_loader = accelerator.prepare(model, test_loader) + + state_dict = load_file(args.checkpoint_path, device="cpu") + accelerator.unwrap_model(model).load_state_dict(state_dict, strict=True) + + _run_loader_inference( + accelerator=accelerator, + model=model, + loader=test_loader, + samples=test_samples, + save_dir=args.prediction_dir, + desc="Inference", + ) + + if accelerator.is_main_process: + print(f"Saved predictions to: {args.prediction_dir}") + + +if __name__ == "__main__": + main() diff --git a/src/Baselines/cafnet_no_smoke/inference_config.yaml b/src/Baselines/cafnet_no_smoke/inference_config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..5bc70aaa292ecf66062250f699a758373ee41304 --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/inference_config.yaml @@ -0,0 +1,29 @@ +# CaFNet-no-smoke inference config for the packaged Smoke-Eval data. +test_base_dir: "../../../evaluation_dataset/Smoke-Eval" +test_split: "train" +test_split_json: null + +# Input and radar preprocessing +input_height: 288 +input_width: 512 +radar_max_depth_m: 11.2 +max_dist_correspondence: 0.5 +patch_size: [64, 128] + +# Model architecture +encoder: resnet34_bts +encoder_radar: resnet18 +radar_input_channels: 1 +bts_size: 512 +max_depth: 11.2 + +# Runtime +batch_size: 32 +num_workers: 0 +seed: 42 +cpu: false +mixed_precision: "fp16" + +# Checkpoint and output root +checkpoint_path: "../../../checkpoints/baselines/cafnet_no_smoke/cafnet_no_smoke.safetensors" +prediction_dir: "prediction" diff --git a/src/Baselines/cafnet_no_smoke/models/bts.py b/src/Baselines/cafnet_no_smoke/models/bts.py new file mode 100644 index 0000000000000000000000000000000000000000..5b3ce68d0a3213bb6ce0ce96e357ed2ea538f3b2 --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/models/bts.py @@ -0,0 +1,367 @@ +# Copyright (C) 2019 Jin Han Lee +# +# This file is a part of BTS. +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see + +import torch +import torch.nn as nn +import torch.nn.functional as torch_nn_func +import math + + +def bn_init_as_tf(m): + if isinstance(m, nn.BatchNorm2d): + m.track_running_stats = True # These two lines enable using stats (moving mean and var) loaded from pretrained model + m.eval() # or zero mean and variance of one if the batch norm layer has no pretrained values + m.affine = True + m.requires_grad = True + + +def weights_init_xavier(m): + if isinstance(m, nn.Conv2d): + torch.nn.init.xavier_uniform_(m.weight) + if m.bias is not None: + torch.nn.init.zeros_(m.bias) + + +class atrous_conv(nn.Sequential): + def __init__(self, in_channels, out_channels, dilation, apply_bn_first=True): + super(atrous_conv, self).__init__() + self.atrous_conv = torch.nn.Sequential() + if apply_bn_first: + self.atrous_conv.add_module('first_bn', nn.BatchNorm2d(in_channels, momentum=0.01, affine=True, track_running_stats=True, eps=1.1e-5)) + + self.atrous_conv.add_module('aconv_sequence', nn.Sequential(nn.ReLU(), + nn.Conv2d(in_channels=in_channels, out_channels=out_channels*2, bias=False, kernel_size=1, stride=1, padding=0), + nn.BatchNorm2d(out_channels*2, momentum=0.01, affine=True, track_running_stats=True), + nn.ReLU(), + nn.Conv2d(in_channels=out_channels * 2, out_channels=out_channels, bias=False, kernel_size=3, stride=1, + padding=(dilation, dilation), dilation=dilation))) + + def forward(self, x): + return self.atrous_conv.forward(x) + +class upconv(nn.Module): + def __init__(self, in_channels, out_channels, ratio=2): + super(upconv, self).__init__() + self.elu = nn.ELU() + self.conv = nn.Conv2d(in_channels=in_channels, out_channels=out_channels, bias=False, kernel_size=3, stride=1, padding=1) + self.ratio = ratio + + def forward(self, x): + up_x = torch_nn_func.interpolate(x, scale_factor=self.ratio, mode='nearest') + out = self.conv(up_x) + out = self.elu(out) + return out + +class reduction_1x1(nn.Sequential): + def __init__(self, num_in_filters, num_out_filters, max_depth, is_final=False): + super(reduction_1x1, self).__init__() + self.max_depth = max_depth + self.is_final = is_final + self.sigmoid = nn.Sigmoid() + self.reduc = torch.nn.Sequential() + + while num_out_filters >= 4: + if num_out_filters < 8: + if self.is_final: + self.reduc.add_module('final', torch.nn.Sequential(nn.Conv2d(num_in_filters, out_channels=1, bias=False, + kernel_size=1, stride=1, padding=0), + nn.Sigmoid())) + else: + self.reduc.add_module('plane_params', torch.nn.Conv2d(num_in_filters, out_channels=3, bias=False, + kernel_size=1, stride=1, padding=0)) + break + else: + self.reduc.add_module('inter_{}_{}'.format(num_in_filters, num_out_filters), + torch.nn.Sequential(nn.Conv2d(in_channels=num_in_filters, out_channels=num_out_filters, + bias=False, kernel_size=1, stride=1, padding=0), + nn.ELU())) + + num_in_filters = num_out_filters + num_out_filters = num_out_filters // 2 + + def forward(self, net): + net = self.reduc.forward(net) + if not self.is_final: + theta = self.sigmoid(net[:, 0, :, :]) * math.pi / 3 + phi = self.sigmoid(net[:, 1, :, :]) * math.pi * 2 + dist = self.sigmoid(net[:, 2, :, :]) * self.max_depth + n1 = torch.mul(torch.sin(theta), torch.cos(phi)).unsqueeze(1) + n2 = torch.mul(torch.sin(theta), torch.sin(phi)).unsqueeze(1) + n3 = torch.cos(theta).unsqueeze(1) + n4 = dist.unsqueeze(1) + net = torch.cat([n1, n2, n3, n4], dim=1) + + return net + +class local_planar_guidance(nn.Module): + def __init__(self, upratio): + super(local_planar_guidance, self).__init__() + self.upratio = upratio + self.u = torch.arange(self.upratio).reshape([1, 1, self.upratio]).float() + self.v = torch.arange(int(self.upratio)).reshape([1, self.upratio, 1]).float() + self.upratio = float(upratio) + + def forward(self, plane_eq, focal): + plane_eq_expanded = torch.repeat_interleave(plane_eq, int(self.upratio), 2) + plane_eq_expanded = torch.repeat_interleave(plane_eq_expanded, int(self.upratio), 3) + n1 = plane_eq_expanded[:, 0, :, :] + n2 = plane_eq_expanded[:, 1, :, :] + n3 = plane_eq_expanded[:, 2, :, :] + n4 = plane_eq_expanded[:, 3, :, :] + + u = self.u.repeat(plane_eq.size(0), plane_eq.size(2) * int(self.upratio), plane_eq.size(3)).cuda() + u = (u - (self.upratio - 1) * 0.5) / self.upratio + + v = self.v.repeat(plane_eq.size(0), plane_eq.size(2), plane_eq.size(3) * int(self.upratio)).cuda() + v = (v - (self.upratio - 1) * 0.5) / self.upratio + + return n4 / (n1 * u + n2 * v + n3) + +class bts_gated_fuse(nn.Module): + def __init__(self, params, feat_out_channels, feat_out_channels_rad, num_features=512): + super(bts_gated_fuse, self).__init__() + self.params = params + self.weight5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[4], feat_out_channels[4], 1, 1, bias=False), + nn.Sigmoid()) + self.project5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[4], feat_out_channels[4], 1, 1, bias=False), + nn.ReLU()) + self.upconv5 = upconv(feat_out_channels[4], num_features) + self.bn5 = nn.BatchNorm2d(num_features, momentum=0.01, affine=True, eps=1.1e-5) + + self.conv5 = torch.nn.Sequential(nn.Conv2d(num_features + feat_out_channels[3], num_features, 3, 1, 1, bias=False), + nn.ELU()) + + self.weight4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[3], num_features, 1, 1, bias=False), + nn.Sigmoid()) + self.project4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[3], num_features, 1, 1, bias=False), + nn.ReLU()) + self.upconv4 = upconv(num_features, num_features // 2) + self.bn4 = nn.BatchNorm2d(num_features // 2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv4 = torch.nn.Sequential(nn.Conv2d(num_features // 2 + feat_out_channels[2], num_features // 2, 3, 1, 1, bias=False), + nn.ELU()) + self.bn4_2 = nn.BatchNorm2d(num_features // 2, momentum=0.01, affine=True, eps=1.1e-5) + + self.daspp_3 = atrous_conv(num_features // 2, num_features // 4, 3, apply_bn_first=False) + self.daspp_6 = atrous_conv(num_features // 2 + num_features // 4 + feat_out_channels[2], num_features // 4, 6) + self.daspp_12 = atrous_conv(num_features + feat_out_channels[2], num_features // 4, 12) + self.daspp_18 = atrous_conv(num_features + num_features // 4 + feat_out_channels[2], num_features // 4, 18) + self.daspp_24 = atrous_conv(num_features + num_features // 2 + feat_out_channels[2], num_features // 4, 24) + self.daspp_conv = torch.nn.Sequential(nn.Conv2d(num_features + num_features // 2 + num_features // 4, num_features // 4, 3, 1, 1, bias=False), + nn.ELU()) + self.reduc8x8 = reduction_1x1(num_features // 4, num_features // 4, self.params.max_depth) + self.lpg8x8 = local_planar_guidance(8) + + self.weight3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[2], num_features // 4, 1, 1, bias=False), + nn.Sigmoid()) + self.project3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[2], num_features // 4, 1, 1, bias=False), + nn.ReLU()) + self.upconv3 = upconv(num_features // 4, num_features // 4) + self.bn3 = nn.BatchNorm2d(num_features // 4, momentum=0.01, affine=True, eps=1.1e-5) + self.conv3 = torch.nn.Sequential(nn.Conv2d(num_features // 4 + feat_out_channels[1] + 1, num_features // 4, 3, 1, 1, bias=False), + nn.ELU()) + self.reduc4x4 = reduction_1x1(num_features // 4, num_features // 8, self.params.max_depth) + self.lpg4x4 = local_planar_guidance(4) + + self.weight2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[1], num_features // 4, 1, 1, bias=False), + nn.Sigmoid()) + self.project2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[1], num_features // 4, 1, 1, bias=False), + nn.ReLU()) + self.upconv2 = upconv(num_features // 4, num_features // 8) + self.bn2 = nn.BatchNorm2d(num_features // 8, momentum=0.01, affine=True, eps=1.1e-5) + self.conv2 = torch.nn.Sequential(nn.Conv2d(num_features // 8 + feat_out_channels[0] + 1, num_features // 8, 3, 1, 1, bias=False), + nn.ELU()) + + self.reduc2x2 = reduction_1x1(num_features // 8, num_features // 16, self.params.max_depth) + self.lpg2x2 = local_planar_guidance(2) + + self.weight1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[0], num_features // 8, 1, 1, bias=False), + nn.Sigmoid()) + self.project1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[0], num_features // 8, 1, 1, bias=False), + nn.ReLU()) + self.upconv1 = upconv(num_features // 8, num_features // 16) + self.reduc1x1 = reduction_1x1(num_features // 16, num_features // 32, self.params.max_depth, is_final=True) + self.conv1 = torch.nn.Sequential(nn.Conv2d(num_features // 16 + 4, num_features // 16, 3, 1, 1, bias=False), + nn.ELU()) + self.get_depth = torch.nn.Sequential(nn.Conv2d(num_features // 16, 1, 3, 1, 1, bias=False), + nn.Sigmoid()) + + self.pool5 = torch.nn.AvgPool2d(32, 32) + self.pool4 = torch.nn.AvgPool2d(16, 16) + self.pool3 = torch.nn.AvgPool2d(8, 8) + self.pool2 = torch.nn.AvgPool2d(4, 4) + self.pool1 = torch.nn.AvgPool2d(2, 2) + + def forward(self, img_features, rad_features, focal, radar_confidence): + skip0, skip1, skip2, skip3 = img_features[0], img_features[1], img_features[2], img_features[3] + rad_skip0, rad_skip1, rad_skip2, rad_skip3 = rad_features[0], rad_features[1], rad_features[2], rad_features[3] + + # prepare radar confidence + radar_confidence5 = self.pool5(radar_confidence) + radar_confidence4 = self.pool4(radar_confidence) + radar_confidence3 = self.pool3(radar_confidence) + radar_confidence2 = self.pool2(radar_confidence) + radar_confidence1 = self.pool1(radar_confidence) + + + rad_weight5 = self.weight5(rad_features[4]) + rad_project5 = self.project5(rad_features[4]) + + dense_features = torch.nn.ReLU()(img_features[4]) + dense_features = dense_features + rad_weight5*rad_project5*radar_confidence5 + upconv5 = self.upconv5(dense_features) # H/16 + upconv5 = self.bn5(upconv5) + concat5 = torch.cat([upconv5, skip3], dim=1) + iconv5 = self.conv5(concat5) + + rad_weight4 = self.weight4(rad_skip3) + rad_project4 = self.project4(rad_skip3) + + iconv5 = iconv5 + rad_weight4*rad_project4*radar_confidence4 + upconv4 = self.upconv4(iconv5) # H/8 + upconv4 = self.bn4(upconv4) + concat4 = torch.cat([upconv4, skip2], dim=1) + iconv4 = self.conv4(concat4) + iconv4 = self.bn4_2(iconv4) + + daspp_3 = self.daspp_3(iconv4) + concat4_2 = torch.cat([concat4, daspp_3], dim=1) + daspp_6 = self.daspp_6(concat4_2) + concat4_3 = torch.cat([concat4_2, daspp_6], dim=1) + daspp_12 = self.daspp_12(concat4_3) + concat4_4 = torch.cat([concat4_3, daspp_12], dim=1) + daspp_18 = self.daspp_18(concat4_4) + concat4_5 = torch.cat([concat4_4, daspp_18], dim=1) + daspp_24 = self.daspp_24(concat4_5) + concat4_daspp = torch.cat([iconv4, daspp_3, daspp_6, daspp_12, daspp_18, daspp_24], dim=1) + daspp_feat = self.daspp_conv(concat4_daspp) + rad_weight3 = self.weight3(rad_skip2) + rad_project3 = self.project3(rad_skip2) + daspp_feat = daspp_feat + rad_weight3*rad_project3*radar_confidence3 + + reduc8x8 = self.reduc8x8(daspp_feat) + plane_normal_8x8 = reduc8x8[:, :3, :, :] + plane_normal_8x8 = torch_nn_func.normalize(plane_normal_8x8, 2, 1) + plane_dist_8x8 = reduc8x8[:, 3, :, :] + plane_eq_8x8 = torch.cat([plane_normal_8x8, plane_dist_8x8.unsqueeze(1)], 1) + depth_8x8 = self.lpg8x8(plane_eq_8x8, focal) + depth_8x8_scaled = depth_8x8.unsqueeze(1) / self.params.max_depth + depth_8x8_scaled_ds = torch_nn_func.interpolate(depth_8x8_scaled, scale_factor=0.25, mode='nearest') + + upconv3 = self.upconv3(daspp_feat) # H/4 + upconv3 = self.bn3(upconv3) + concat3 = torch.cat([upconv3, skip1, depth_8x8_scaled_ds], dim=1) + iconv3 = self.conv3(concat3) + rad_weight2 = self.weight2(rad_skip1) + rad_project2 = self.project2(rad_skip1) + iconv3 = iconv3 + rad_weight2*rad_project2*radar_confidence2 + + reduc4x4 = self.reduc4x4(iconv3) + plane_normal_4x4 = reduc4x4[:, :3, :, :] + plane_normal_4x4 = torch_nn_func.normalize(plane_normal_4x4, 2, 1) + plane_dist_4x4 = reduc4x4[:, 3, :, :] + plane_eq_4x4 = torch.cat([plane_normal_4x4, plane_dist_4x4.unsqueeze(1)], 1) + depth_4x4 = self.lpg4x4(plane_eq_4x4, focal) + depth_4x4_scaled = depth_4x4.unsqueeze(1) / self.params.max_depth + depth_4x4_scaled_ds = torch_nn_func.interpolate(depth_4x4_scaled, scale_factor=0.5, mode='nearest') + + upconv2 = self.upconv2(iconv3) # H/2 + upconv2 = self.bn2(upconv2) + concat2 = torch.cat([upconv2, skip0, depth_4x4_scaled_ds], dim=1) + iconv2 = self.conv2(concat2) + rad_weight1 = self.weight1(rad_skip0) + rad_project1 = self.project1(rad_skip0) + iconv2 = iconv2 + rad_weight1*rad_project1*radar_confidence1 + + reduc2x2 = self.reduc2x2(iconv2) + plane_normal_2x2 = reduc2x2[:, :3, :, :] + plane_normal_2x2 = torch_nn_func.normalize(plane_normal_2x2, 2, 1) + plane_dist_2x2 = reduc2x2[:, 3, :, :] + plane_eq_2x2 = torch.cat([plane_normal_2x2, plane_dist_2x2.unsqueeze(1)], 1) + depth_2x2 = self.lpg2x2(plane_eq_2x2, focal) + depth_2x2_scaled = depth_2x2.unsqueeze(1) / self.params.max_depth + + rad_weight1 = self.weight1(rad_skip0) + rad_project1 = self.project1(rad_skip0) + + upconv1 = self.upconv1(iconv2) + reduc1x1 = self.reduc1x1(upconv1) + concat1 = torch.cat([upconv1, reduc1x1, depth_2x2_scaled, depth_4x4_scaled, depth_8x8_scaled], dim=1) + iconv1 = self.conv1(concat1) + final_depth = self.params.max_depth * self.get_depth(iconv1) + + return depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth + +class encoder_image(nn.Module): + def __init__(self, params): + super(encoder_image, self).__init__() + self.params = params + import torchvision.models as models + if params.encoder == 'densenet121_bts': + self.base_model = models.densenet121(pretrained=False).features + self.feat_names = ['relu0', 'pool0', 'transition1', 'transition2', 'norm5'] + self.feat_out_channels = [64, 64, 128, 256, 1024] + elif params.encoder == 'densenet161_bts': + self.base_model = models.densenet161(pretrained=False).features + self.feat_names = ['relu0', 'pool0', 'transition1', 'transition2', 'norm5'] + self.feat_out_channels = [96, 96, 192, 384, 2208] + elif params.encoder == 'resnet50_bts': + self.base_model = models.resnet50(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 256, 512, 1024, 2048] + elif params.encoder == 'resnet34_bts': + self.base_model = models.resnet34(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + elif params.encoder == 'resnet18_bts': + self.base_model = models.resnet18(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + elif params.encoder == 'resnet101_bts': + self.base_model = models.resnet101(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 256, 512, 1024, 2048] + elif params.encoder == 'resnext50_bts': + self.base_model = models.resnext50_32x4d(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 256, 512, 1024, 2048] + elif params.encoder == 'resnext101_bts': + self.base_model = models.resnext101_32x8d(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 256, 512, 1024, 2048] + elif params.encoder == 'mobilenetv2_bts': + self.base_model = models.mobilenet_v2(pretrained=False).features + self.feat_inds = [2, 4, 7, 11, 19] + self.feat_out_channels = [16, 24, 32, 64, 1280] + self.feat_names = [] + else: + print('Not supported encoder: {}'.format(params.encoder)) + + def forward(self, x): + feature = x + skip_feat = [] + i = 1 + for k, v in self.base_model._modules.items(): + if 'fc' in k or 'avgpool' in k: + continue + feature = v(feature) + if self.params.encoder == 'mobilenetv2_bts': + if i == 2 or i == 4 or i == 7 or i == 11 or i == 19: + skip_feat.append(feature) + else: + if any(x in k for x in self.feat_names): + skip_feat.append(feature) + i = i + 1 + return skip_feat diff --git a/src/Baselines/cafnet_no_smoke/models/model.py b/src/Baselines/cafnet_no_smoke/models/model.py new file mode 100644 index 0000000000000000000000000000000000000000..4eb94690d8ca447ff50a7c966793b62fb16dd544 --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/models/model.py @@ -0,0 +1,28 @@ +import torch +import torch.nn as nn +from models.bts import encoder_image, bts_gated_fuse +from models.radar import encoder_radar_sparse_conv, encoder_radar_sub, decoder_radar + +class CaFNet(nn.Module): + def __init__(self, params, threshold=0.4): + super(CaFNet, self).__init__() + self.threshold = threshold + self.encoder = encoder_image(params) + self.encoder_radar1 = encoder_radar_sparse_conv(params) + self.decoder_radar = decoder_radar(params, self.encoder.feat_out_channels, self.encoder_radar1.feat_out_channels) + self.encoder_radar2 = encoder_radar_sub(params) + self.decoder = bts_gated_fuse(params, self.encoder.feat_out_channels, self.encoder_radar2.feat_out_channels, params.bts_size) + + + def forward(self, x, radar, focal): + + skip_feat = self.encoder(x) + skip_feat_radar = self.encoder_radar1(radar) + rad_confidence, rad_depth = self.decoder_radar(skip_feat, skip_feat_radar) + mask = (rad_confidence > self.threshold).float() + radar_new_input = torch.cat([mask*rad_depth, radar], axis=1) + skip_feat_radar_new = self.encoder_radar2(radar_new_input) + + depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth = self.decoder(skip_feat, skip_feat_radar_new, focal, rad_confidence) + + return depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth, rad_confidence, rad_depth diff --git a/src/Baselines/cafnet_no_smoke/models/radar.py b/src/Baselines/cafnet_no_smoke/models/radar.py new file mode 100644 index 0000000000000000000000000000000000000000..95729efa551242c9eb2c00c714b5ab056c7345c7 --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/models/radar.py @@ -0,0 +1,212 @@ +from models.bts import upconv +import torch +import torch.nn as nn +import torchvision.models as models + +class encoder_radar_sparse_conv(nn.Module): + def __init__(self, params): + # radar encoder for the first stage + super(encoder_radar_sparse_conv, self).__init__() + + self.params = params + self.sparse_conv1 = SparseConv(params.radar_input_channels, 16, 7, activation='elu') + self.sparse_conv2 = SparseConv(16, 16, 5, activation='elu') + self.sparse_conv3 = SparseConv(16, 16, 3, activation='elu') + self.sparse_conv4 = SparseConv(16, 3, 3, activation='elu') + + if params.encoder_radar == 'resnet34': + self.base_model_radar = models.resnet34(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + elif params.encoder_radar == 'resnet18': + self.base_model_radar = models.resnet18(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + else: + print('Not supported encoder: {}'.format(params.encoder)) + + def forward(self, x): + mask = (x[:, 0] > 0).float().unsqueeze(1) + feature = x + feature, mask = self.sparse_conv1(feature, mask) + feature, mask = self.sparse_conv2(feature, mask) + feature, mask = self.sparse_conv3(feature, mask) + feature, mask = self.sparse_conv4(feature, mask) + + skip_feat = [] + i = 1 + for k, v in self.base_model_radar._modules.items(): + if 'fc' in k or 'avgpool' in k: + continue + feature = v(feature) + if any(x in k for x in self.feat_names): + skip_feat.append(feature) + i = i + 1 + return skip_feat + +class encoder_radar_sub(nn.Module): + def __init__(self, params): + # radar encoder for the second stage + super(encoder_radar_sub, self).__init__() + + self.params = params + import torchvision.models as models + self.conv = torch.nn.Sequential(nn.Conv2d(params.radar_input_channels+1, 3, 3, 1, 1, bias=False), + nn.ELU()) + + if params.encoder_radar == 'resnet34': + self.base_model_radar = models.resnet34(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + elif params.encoder_radar == 'resnet18': + self.base_model_radar = models.resnet18(pretrained=False) + self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4'] + self.feat_out_channels = [64, 64, 128, 256, 512] + else: + print('Not supported encoder: {}'.format(params.encoder)) + def forward(self, x): + feature = x + feature = self.conv(feature) + skip_feat = [] + i = 1 + for k, v in self.base_model_radar._modules.items(): + if 'fc' in k or 'avgpool' in k: + continue + feature = v(feature) + if any(x in k for x in self.feat_names): + skip_feat.append(feature) + i = i + 1 + return skip_feat + + +class decoder_radar(nn.Module): + def __init__(self, params, feat_out_channels_img, feat_out_channels_radar): + super(decoder_radar, self).__init__() + self.params = params + self.upconv5 = upconv(feat_out_channels_img[4]+feat_out_channels_radar[4], feat_out_channels_radar[4]//2) + self.bn5 = nn.BatchNorm2d(feat_out_channels_radar[4]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[4]//2, feat_out_channels_radar[4]//2, 3, 1, 1, bias=False), + nn.ELU()) + + self.upconv4 = upconv(feat_out_channels_img[3]+feat_out_channels_radar[3]+feat_out_channels_radar[4]//2, feat_out_channels_radar[3]//2) + self.bn4 = nn.BatchNorm2d(feat_out_channels_radar[3]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[3]//2, feat_out_channels_radar[3]//2, 3, 1, 1, bias=False), + nn.ELU()) + + self.upconv3 = upconv(feat_out_channels_img[2]+feat_out_channels_radar[2]+feat_out_channels_radar[3]//2, feat_out_channels_radar[2]//2) + self.bn3 = nn.BatchNorm2d(feat_out_channels_radar[2]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[2]//2, feat_out_channels_radar[2]//2, 3, 1, 1, bias=False), + nn.ELU()) + + self.upconv2 = upconv(feat_out_channels_img[1]+feat_out_channels_radar[1]+feat_out_channels_radar[2]//2, feat_out_channels_radar[1]//2) + self.bn2 = nn.BatchNorm2d(feat_out_channels_radar[1]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[1]//2, feat_out_channels_radar[1]//2, 3, 1, 1, bias=False), + nn.ELU()) + + self.upconv1 = upconv(feat_out_channels_img[0]+feat_out_channels_radar[0]+feat_out_channels_radar[1]//2, feat_out_channels_radar[0]//2) + self.bn1 = nn.BatchNorm2d(feat_out_channels_radar[0]//2, momentum=0.01, affine=True, eps=1.1e-5) + self.conv1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, feat_out_channels_radar[0]//2, 3, 1, 1, bias=False), + nn.ELU()) + + # self.get_depth = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, 1, 3, 1, 1, bias=False), + # nn.Sigmoid()) + + self.get_depth = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, 2, 3, 1, 1, bias=False), + nn.Sigmoid()) + + def forward(self, image_features, radar_features): + img_skip0, img_skip1, img_skip2, img_skip3, img_final = image_features[0], image_features[1], image_features[2], image_features[3], image_features[4] + rad_skip0, rad_skip1, rad_skip2, rad_skip3, rad_final = radar_features[0], radar_features[1], radar_features[2], radar_features[3], radar_features[4] + final = torch.cat([img_final, rad_final], axis=1) + upconv5 = self.upconv5(final) + upconv5 = self.bn5(upconv5) + upconv5 = self.conv5(upconv5) + upconv5 = torch.cat([img_skip3, rad_skip3, upconv5], axis=1) + + upconv4 = self.upconv4(upconv5) + upconv4 = self.bn4(upconv4) + upconv4 = self.conv4(upconv4) + upconv4 = torch.cat([img_skip2, rad_skip2, upconv4], axis=1) + + upconv3 = self.upconv3(upconv4) + upconv3 = self.bn3(upconv3) + upconv3 = self.conv3(upconv3) + upconv3 = torch.cat([img_skip1, rad_skip1, upconv3], axis=1) + + upconv2 = self.upconv2(upconv3) + upconv2 = self.bn2(upconv2) + upconv2 = self.conv2(upconv2) + upconv2 = torch.cat([img_skip0, rad_skip0, upconv2], axis=1) + + upconv1 = self.upconv1(upconv2) + upconv1 = self.bn1(upconv1) + upconv1 = self.conv1(upconv1) + + # confidence = self.get_depth(upconv1) + # depth = self.params.max_depth * confidence + depth_conf = self.get_depth(upconv1) + depth = self.params.max_depth * depth_conf[:, 0:1] + confidence = depth_conf[:, 1:2] + + return confidence, depth + + +class SparseConv(nn.Module): + + def __init__(self, + in_channels, + out_channels, + kernel_size, + activation='relu'): + super().__init__() + + padding = kernel_size//2 + + self.conv = nn.Conv2d( + in_channels, + out_channels, + kernel_size=kernel_size, + padding=padding, + bias=False) + + self.bias = nn.Parameter( + torch.zeros(out_channels), + requires_grad=True) + + self.sparsity = nn.Conv2d( + in_channels, + out_channels, + kernel_size=kernel_size, + padding=padding, + bias=False) + + kernel = torch.FloatTensor(torch.ones([kernel_size, kernel_size])).unsqueeze(0).unsqueeze(0) + + self.sparsity.weight = nn.Parameter( + data=kernel, + requires_grad=False) + + if activation == 'relu': + self.act = nn.ReLU(inplace=False) + elif activation == 'sigmoid': + self.act = nn.Sigmoid() + elif activation == 'elu': + self.act = nn.ELU() + + self.max_pool = nn.MaxPool2d( + kernel_size, + stride=1, + padding=padding) + + + + def forward(self, x, mask): + x = x*mask + x = self.conv(x) + normalizer = 1/(self.sparsity(mask)+1e-8) + x = x * normalizer + self.bias.unsqueeze(0).unsqueeze(2).unsqueeze(3) + x = self.act(x) + + mask = self.max_pool(mask) + + return x, mask diff --git a/src/Baselines/cafnet_no_smoke/rice_dataset.py b/src/Baselines/cafnet_no_smoke/rice_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..1a45d711b32ae557847b80e0ba10ccdf5f7a1f83 --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/rice_dataset.py @@ -0,0 +1,123 @@ +import json +import os +from typing import Dict, List, Optional, Tuple + +import numpy as np +from torch.utils.data import Dataset + + +class RiceDataset(Dataset): + """Raw Rice dataset reader for DJI RGB, ZED depth and radar point clouds. + + This dataset returns raw per-frame arrays and leaves geometric processing to + `collate_fn_helpers.make_rice_collate_fn`. + """ + + def __init__( + self, + base_dir: str, + split_json_path: Optional[str] = None, + split: str = "train", + input_height: int = 288, + input_width: int = 512, + patch_size: Optional[Tuple[int, int]] = None, + ): + self.base_dir = base_dir + self.split = split + self.input_height = int(input_height) + self.input_width = int(input_width) + self.patch_size = self._resolve_patch_size(patch_size) + + test_sequences = self._load_test_split(split_json_path) + + all_sequences = sorted( + d + for d in os.listdir(base_dir) + if os.path.isdir(os.path.join(base_dir, d)) and not d.startswith(".") + ) + + self.sequences: List[str] = [] + for seq in all_sequences: + if split == "train" and seq in test_sequences: + continue + # if split == "train" and seq.lower().startswith("smoke"): + # continue + if split == "test" and seq not in test_sequences: + continue + if self._is_valid_sequence(os.path.join(base_dir, seq)): + self.sequences.append(seq) + + self.dji_rgb_mmaps: Dict[str, np.memmap] = {} + self.zed_depth_mmaps: Dict[str, np.memmap] = {} + self.samples: List[Tuple[str, int]] = [] + + for seq in self.sequences: + seq_dir = os.path.join(self.base_dir, seq) + dji_rgb_path = os.path.join(seq_dir, "dji_rgb.npy") + zed_depth_path = os.path.join(seq_dir, "zed_depth.npy") + + self.dji_rgb_mmaps[seq] = np.load(dji_rgb_path, mmap_mode="r") + self.zed_depth_mmaps[seq] = np.load(zed_depth_path, mmap_mode="r") + + n_frames = min( + len(self.dji_rgb_mmaps[seq]), + len(self.zed_depth_mmaps[seq]), + ) + for frame_idx in range(n_frames): + self.samples.append((seq, frame_idx)) + + def _resolve_patch_size( + self, patch_size: Optional[Tuple[int, int]] + ) -> Tuple[int, int]: + if patch_size is not None: + return int(patch_size[0]), int(patch_size[1]) + + # Scale default CaFNet patch size (50, 150) from 352x704. + base_h, base_w = 352, 704 + scale_h = self.input_height / float(base_h) + scale_w = self.input_width / float(base_w) + ext_h = max(1, int(round(50 * scale_h))) + ext_w = max(1, int(round(150 * scale_w))) + return ext_h, ext_w + + def _load_test_split(self, split_json_path: Optional[str]) -> set: + if not split_json_path or not os.path.exists(split_json_path): + return set() + with open(split_json_path, "r") as f: + payload = json.load(f) + return set(payload.get("test", [])) + + def _is_valid_sequence(self, seq_dir: str) -> bool: + dji_rgb_path = os.path.join(seq_dir, "dji_rgb.npy") + zed_depth_path = os.path.join(seq_dir, "zed_depth.npy") + pcd_dir = os.path.join(seq_dir, "pcd") + return ( + os.path.exists(dji_rgb_path) + and os.path.exists(zed_depth_path) + and os.path.isdir(pcd_dir) + ) + + def __len__(self) -> int: + return len(self.samples) + + def __getitem__(self, idx: int) -> Dict[str, object]: + seq, frame_idx = self.samples[idx] + seq_dir = os.path.join(self.base_dir, seq) + + dji_rgb = np.asarray(self.dji_rgb_mmaps[seq][frame_idx]).copy() + zed_depth_mm = np.asarray(self.zed_depth_mmaps[seq][frame_idx]).copy() + + pcd_path = os.path.join(seq_dir, "pcd", f"pcd_{frame_idx}.npy") + if os.path.exists(pcd_path): + radar_pcd_xyz = np.asarray(np.load(pcd_path), dtype=np.float32) + else: + radar_pcd_xyz = np.zeros((0, 3), dtype=np.float32) + + return { + "sample_idx": idx, + "sequence": seq, + "frame_idx": frame_idx, + "dji_rgb": dji_rgb, + "zed_depth_mm": zed_depth_mm, + "radar_pcd_xyz": radar_pcd_xyz, + } diff --git a/src/Baselines/cafnet_no_smoke/split.json b/src/Baselines/cafnet_no_smoke/split.json new file mode 100644 index 0000000000000000000000000000000000000000..f6e7a920c8068dcf74ab6acb272480fdca99e1d1 --- /dev/null +++ b/src/Baselines/cafnet_no_smoke/split.json @@ -0,0 +1,14 @@ +{ + "test": [ + "Dell-1", + "Dell-2", + "Smoke-Dell-1", + "Smoke-Dell-2", + "Keck-1", + "Keck-2", + "Keck-3", + "Smoke-keck-1", + "Smoke-keck-2", + "Smoke-keck-3" + ] +} \ No newline at end of file diff --git a/src/Baselines/da3/inference.py b/src/Baselines/da3/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..ece75a9818ff1346d20b0e502e02cff1d2ecef07 --- /dev/null +++ b/src/Baselines/da3/inference.py @@ -0,0 +1,179 @@ +"""Depth Anything 3 metric-depth inference for Smoke-Eval sequences. + +This inference-only adapter follows the official ByteDance-Seed +Depth-Anything-3 Python API. The upstream package supplies the model +architecture; this file supplies the artifact's local weights, camera +calibration, sequence sharding, and output contract. +""" + +import argparse +from pathlib import Path + +import cv2 +import numpy as np +import torch +from accelerate import Accelerator +from safetensors.torch import load_file +from tqdm.auto import tqdm + + +INTRINSICS = np.array( + [[365.13, 0.0, 445.43], [0.0, 365.13, 261.18], [0.0, 0.0, 1.0]], + dtype=np.float32, +) +SCALE_FACTOR = 1.15 * 365.13 / 300.0 +TARGET_SIZE = (896, 504) + + +class Calibrator: + """Defish DJI frames and map them to the ZED-aligned view.""" + + def __init__(self): + k_dji = np.array( + [ + [718.48555551, 0.0, 963.36465011], + [0.0, 720.25844189, 537.87569913], + [0.0, 0.0, 1.0], + ], + dtype=np.float64, + ) + d_dji = np.array( + [0.19022699, 0.03466753, 0.05858962, -0.07070669], + dtype=np.float64, + ) + new_k = cv2.fisheye.estimateNewCameraMatrixForUndistortRectify( + k_dji, + d_dji, + (1920, 1080), + np.eye(3), + balance=0.2, + fov_scale=1.0, + ) + self.map1, self.map2 = cv2.fisheye.initUndistortRectifyMap( + k_dji, + d_dji, + np.eye(3), + new_k, + (1920, 1080), + cv2.CV_16SC2, + ) + self.homography = np.array( + [ + [ + 0.8274446551892256, + -0.0742944198979625, + 80.23797348979947, + ], + [ + -0.014725864916652691, + 0.8471179917075127, + 28.27366063997317, + ], + [ + -5.083573451500717e-05, + -6.846079418201229e-05, + 1.0, + ], + ], + dtype=np.float64, + ) + + def __call__(self, rgb: np.ndarray) -> np.ndarray: + bgr = cv2.cvtColor(rgb, cv2.COLOR_RGB2BGR) + if bgr.shape[:2] != (1080, 1920): + bgr = cv2.resize(bgr, (1920, 1080), interpolation=cv2.INTER_LINEAR) + bgr = cv2.remap(bgr, self.map1, self.map2, cv2.INTER_LINEAR) + bgr = cv2.warpPerspective(bgr, self.homography, (1918, 1105)) + bgr = bgr[115:760, 255:1400] + bgr = cv2.resize(bgr, TARGET_SIZE, interpolation=cv2.INTER_AREA) + return cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--data_root", required=True) + parser.add_argument("--checkpoint", required=True) + parser.add_argument("--output_dir", required=True) + parser.add_argument("--model_name", default="da3metric-large") + parser.add_argument("--batch_size", type=int, default=16) + parser.add_argument("--sequences", nargs="*", default=None) + return parser.parse_args() + + +@torch.no_grad() +def main() -> None: + args = parse_args() + from depth_anything_3.api import DepthAnything3 + + class AccelerateFP16DepthAnything3(DepthAnything3): + """Use the official API while leaving autocast to Accelerate.""" + + @torch.inference_mode() + def forward( + self, + image, + extrinsics=None, + intrinsics=None, + export_feat_layers=None, + infer_gs=False, + use_ray_pose=False, + ref_view_strategy="saddle_balanced", + ): + return self.model( + image, + extrinsics, + intrinsics, + export_feat_layers, + infer_gs, + use_ray_pose, + ref_view_strategy, + ) + + accelerator = Accelerator(mixed_precision="fp16") + data_root = Path(args.data_root) + output_dir = Path(args.output_dir) + sequences = sorted(path for path in data_root.iterdir() if path.is_dir()) + if args.sequences: + requested = set(args.sequences) + sequences = [path for path in sequences if path.name in requested] + local_sequences = sequences[ + accelerator.process_index :: accelerator.num_processes + ] + + model = AccelerateFP16DepthAnything3(model_name=args.model_name) + model.load_state_dict(load_file(args.checkpoint, device="cpu"), strict=True) + model = model.to(accelerator.device).eval() + calibrate = Calibrator() + if accelerator.is_main_process: + output_dir.mkdir(parents=True, exist_ok=True) + accelerator.wait_for_everyone() + + for sequence in local_sequences: + rgb = np.load(sequence / "dji_rgb.npy", mmap_mode="r") + depth_chunks = [] + for start in tqdm( + range(0, len(rgb), args.batch_size), + desc=sequence.name, + disable=not accelerator.is_local_main_process, + ): + end = min(start + args.batch_size, len(rgb)) + images = [ + calibrate(np.asarray(rgb[index])) for index in range(start, end) + ] + intrinsics = np.repeat(INTRINSICS[None], len(images), axis=0) + with accelerator.autocast(): + prediction = model.inference( + images, + intrinsics=intrinsics, + process_res=896, + process_res_method="upper_bound_resize", + ) + depth_chunks.append(prediction.depth * SCALE_FACTOR) + depth = np.concatenate(depth_chunks).astype(np.float32, copy=False) + np.save(output_dir / f"{sequence.name.lower()}_pred.npy", depth) + + accelerator.wait_for_everyone() + + +if __name__ == "__main__": + main() diff --git a/src/Baselines/grt/augmentations.py b/src/Baselines/grt/augmentations.py new file mode 100644 index 0000000000000000000000000000000000000000..9c1c00304dde36b7318ede3bff733f05dee4906d --- /dev/null +++ b/src/Baselines/grt/augmentations.py @@ -0,0 +1,193 @@ +import torch +import torchvision.transforms.functional as TF +from torchvision.transforms import Resize, InterpolationMode +from typing import Union +import numpy as np + +AZIMUTH_RESOLUTION = 128 +ELEVATION_RESOLUTION = 64 + +# Depth output resolution: height=64, width=128 +DEPTH_TARGET_HEIGHT = 64 +DEPTH_TARGET_WIDTH = 128 + +resize_transform = Resize( + size=[ELEVATION_RESOLUTION, AZIMUTH_RESOLUTION], + interpolation=InterpolationMode.BILINEAR, + antialias=True, +) + +depth_resize_transform = Resize( + size=(DEPTH_TARGET_HEIGHT, DEPTH_TARGET_WIDTH), + interpolation=InterpolationMode.BILINEAR, + antialias=True, +) + + +def translate_radar(radar_data): + """ + Applies normalization to radar data after batching from dataloader. + Called before passing data into the model. + + Args: + radar_data: Batched radar tensor from dataloader + Shape: [B, 64, 8, 2, 256, 2] (batch, doppler, azimuth, elevation, range, channels) + - Channel 0: raw amplitude values + - Channel 1: phase normalized to [-1, 1] (divided by π) + + Returns: + Processed radar tensor with same shape [B, 64, 8, 2, 256, 2] + - Channel 0: sqrt(amplitude * 1e-3) for magnitude normalization + - Channel 1: phase * π (converted back to radians [-π, π]) + """ + radar_mag = radar_data[..., 0] # [B, 64, 8, 2, 256] - Extract raw amplitude + radar_phase = radar_data[..., 1] # [B, 64, 8, 2, 256] - Extract normalized phase + + # Normalize amplitude: scale then sqrt + radar_mag_processed = torch.sqrt(radar_mag * 1e-6) + + # Convert phase back to radians: [-1, 1] -> [-π, π] + radar_phase_processed = radar_phase * torch.pi + + # Stack channels back together: [B, 64, 8, 2, 256, 2] + radar_data_translated = torch.stack( + [radar_mag_processed, radar_phase_processed], dim=-1 + ) + return radar_data_translated + + +def resize_depth( + depth_map: Union[torch.Tensor, np.ndarray], +) -> Union[torch.Tensor, np.ndarray]: + """ + Process depth map from dataloader (same pipeline as denoiser/control crop_depth): + mm -> meters, clamp [0, 11.2] m, normalize to [0, 1], resize to (64, 128) (h, w). + + Args: + depth_map: Depth in millimeters. Torch or numpy. + Shapes: (H, W), (B, H, W), or (B, 1, H, W). + + Returns: + Depth in [0, 1], spatial size (64, 128). Shape [B, 64, 128] for batched input. + """ + is_numpy = isinstance(depth_map, np.ndarray) + if is_numpy: + depth_map = torch.from_numpy(depth_map) + + depth_map = depth_map.float() + original_shape = depth_map.shape + + if depth_map.dim() == 2: + depth_map = depth_map.unsqueeze(0) # (H, W) -> (1, H, W) + elif depth_map.dim() == 3: + depth_map = depth_map.unsqueeze(1) # (B, H, W) -> (B, 1, H, W) + elif depth_map.dim() != 4: + raise ValueError(f"Unexpected depth shape: {original_shape}") + + invalid_mask = ~(torch.isfinite(depth_map) & (depth_map >= 0)) + depth_map[invalid_mask] = 0.0 + + depth_map = depth_map / 1000.0 # mm -> meters + max_depth_m = 11.2 + depth_map = torch.clamp(depth_map, min=0.0, max=max_depth_m) + depth_map = depth_map / max_depth_m # [0, 1] + + invalid_mask = ~torch.isfinite(depth_map) + depth_map[invalid_mask] = 0.0 + + depth_map = depth_resize_transform(depth_map) # (..., 64, 128) + depth_values = depth_map.squeeze(1) # [B, 64, 128] or [1, 64, 128] + + if len(original_shape) == 2: + depth_values = depth_values.squeeze(0) # (64, 128) + + if is_numpy: + depth_values = depth_values.numpy() + return depth_values + + +def quantize_depth_to_occupancy(depth_values, num_range_bins=64): + """ + Quantizes 2D depth values into 3D binary occupancy grid. + + Args: + depth_values: Resized depth tensor + Shape: [B, elevation, azimuth] + Values: normalized to [0, 1] range + num_range_bins: Number of range bins for quantization (default: 64) + + Returns: + Binary 3D occupancy grid + Shape: [B, elevation, azimuth, num_range_bins] + Values: binary (0 or 1) indicating occupied bins + """ + B, elevation, azimuth = depth_values.shape + + # Quantize normalized depth [0, 1] directly to range bins [0, num_range_bins-1] + # Each bin represents 1/num_range_bins of the normalized depth range + bin_indices = torch.floor( + depth_values / (1.0 / num_range_bins) + ).long() # [B, elevation, azimuth] + bin_indices = torch.clamp( + bin_indices, 0, num_range_bins - 1 + ) # Handle edge case where depth_values = 1.0 + + # Create binary 3D occupancy grid + occupancy_grid = torch.zeros( + B, + elevation, + azimuth, + num_range_bins, + dtype=torch.float32, + device=depth_values.device, + ) # [B, elevation, azimuth, num_range_bins] + + # Set occupied bins to 1 + # Use advanced indexing to mark the appropriate range bin for each (elevation, azimuth) cell + batch_idx = torch.arange(B, device=depth_values.device)[:, None, None].expand( + B, elevation, azimuth + ) + elevation_idx = torch.arange(elevation, device=depth_values.device)[ + None, :, None + ].expand(B, elevation, azimuth) + azimuth_idx = torch.arange(azimuth, device=depth_values.device)[ + None, None, : + ].expand(B, elevation, azimuth) + + occupancy_grid[batch_idx, elevation_idx, azimuth_idx, bin_indices] = 1.0 + + return occupancy_grid # [B, elevation, azimuth, num_range_bins] + + +def dequantize_depth(occupancy_grid): + """ + Converts 3D binary occupancy grid back to 2D depth map. + This is the inverse operation of quantize_depth_to_occupancy. + + Args: + occupancy_grid: Binary 3D occupancy grid + Shape: [B, 64, 128, 64] (batch, elevation, azimuth, range) + Values: binary (0 or 1) or continuous (predicted probabilities) + + Returns: + Reconstructed depth map + Shape: [B, 1, 64, 128] (batch, channel, elevation, azimuth) + Values: normalized to [0, 1] range + """ + num_range_bins = occupancy_grid.shape[3] + + # Find the range bin with maximum value for each (elevation, azimuth) cell + # For binary: finds the occupied bin + # For continuous: finds the most likely bin + bin_indices = torch.argmax(occupancy_grid, dim=3) # [B, 64, 128] + + # Convert bin indices back to normalized depth values [0, 1] + # Use bin center: (bin_idx + 0.5) / num_bins + depth_values = (bin_indices.float() + 1) / num_range_bins # [B, 64, 128] + + # Add channel dimension: [B, 64, 128] -> [B, 1, 64, 128] + depth_map = depth_values.unsqueeze(1) # [B, 1, 64, 128] + + return depth_map + + diff --git a/src/Baselines/grt/dataloader.py b/src/Baselines/grt/dataloader.py new file mode 100644 index 0000000000000000000000000000000000000000..8c05a3c139796595c955d8aaf9cd756f181aabe5 --- /dev/null +++ b/src/Baselines/grt/dataloader.py @@ -0,0 +1,330 @@ +""" +Dataloader for MobiCom processed dataset (output of processor.py). + +Uses the optimized format produced by processor.py: +- radar.npy: (N, doppler, elevation, azimuth, range) complex64 +- dji_rgb.avi: DJI RGB video (FFV1), (N, H, W, 3) uint8 +- zed_depth.npy: (N, H, W) uint16, depth in millimeters + +This module provides: +- `RiceDataset`: frame-level dataset returning radar amplitude/phase, DJI RGB, + and ZED depth (ground truth). +- `create_rice_dataloader`: generic dataloader for an arbitrary set of sequences. +- `create_split_dataloaders`: reads train/val/test split from split.json + (default: radar_model/split.json) and returns train/val/test dataloaders. +""" + +import json +from pathlib import Path +from typing import Dict, List, Optional, Tuple + +import cv2 +import numpy as np +import torch +from torch.utils.data import Dataset, DataLoader, random_split + + +class RiceDataset(Dataset): + """ + Dataset for processor.py output: radar, DJI RGB, and ZED depth per frame. + + Args: + root_dir: Root directory containing sequence subdirs (e.g. processed/), + each with radar.npy, dji_rgb.avi, zed_depth.npy. + sequences: Optional list of sequence names to load. If None, loads all + subdirs that contain the three required files. + frame_skip: Sample every frame_skip frames (1 = all frames). + return_radar_complex: If True, return radar as complex tensor; if False, + return radar_amplitude and radar_phase as separate float tensors. + depth_in_meters: If True, convert depth from mm to meters. + rgb_normalize: If True, return RGB in [0, 1] float; else uint8 [0, 255]. + """ + + # GRT inference consumes radar and depth only. Smoke-Eval packages RGB + # frames as ``dji_rgb.npy`` rather than the original training video, so + # requiring the unused video would incorrectly discard every sequence. + REQUIRED_FILES = ("radar.npy", "zed_depth.npy") + + def __init__( + self, + root_dir: str, + sequences: Optional[List[str]] = None, + frame_skip: int = 1, + return_radar_complex: bool = False, + depth_in_meters: bool = True, + rgb_normalize: bool = True, + ): + self.root_dir = Path(root_dir) + self.frame_skip = max(1, frame_skip) + self.return_radar_complex = return_radar_complex + self.depth_in_meters = depth_in_meters + self.rgb_normalize = rgb_normalize + + self.sequences = self._discover_sequences(sequences) + self.index_map: List[Tuple[str, int]] = [] # (seq_name, frame_idx) + self._seq_arrays: Dict[str, Dict] = {} # seq -> {radar, depth, dji_rgb} + + self._build_index() + + def _discover_sequences(self, sequences: Optional[List[str]] = None) -> List[str]: + """Return list of sequence names that have all required files.""" + if not self.root_dir.is_dir(): + raise FileNotFoundError(f"Root directory not found: {self.root_dir}") + + all_seqs = sorted( + d.name + for d in self.root_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + valid = [] + for name in all_seqs: + seq_dir = self.root_dir / name + if all((seq_dir / f).exists() for f in self.REQUIRED_FILES): + valid.append(name) + if sequences is not None: + valid = [s for s in valid if s in sequences] + return valid + + def _build_index(self) -> None: + """Build (seq_name, frame_idx) index, using radar.npy for frame count.""" + self.index_map.clear() + for seq_name in self.sequences: + seq_dir = self.root_dir / seq_name + radar_path = seq_dir / "radar.npy" + radar = np.load(radar_path, mmap_mode="r") + n_frames = radar.shape[0] + for i in range(0, n_frames, self.frame_skip): + self.index_map.append((seq_name, i)) + + # def _load_video_rgb(self, path: Path) -> np.ndarray: + # """Load RGB AVI (e.g. FFV1) as (N, H, W, 3) uint8 RGB.""" + # cap = cv2.VideoCapture(str(path)) + # if not cap.isOpened(): + # raise RuntimeError(f"Failed to open video: {path}") + # frames = [] + # while True: + # ret, frame = cap.read() + # if not ret: + # break + # rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) + # frames.append(rgb) + # cap.release() + # if not frames: + # return np.empty((0, 0, 0, 3), dtype=np.uint8) + # return np.stack(frames, axis=0) + + def _load_sequence_arrays(self, seq_name: str) -> Dict: + """Lazy-load or return cached arrays for a sequence.""" + if seq_name not in self._seq_arrays: + seq_dir = self.root_dir / seq_name + # dji_rgb = self._load_video_rgb(seq_dir / "dji_rgb.avi") + self._seq_arrays[seq_name] = { + "radar": np.load(seq_dir / "radar.npy", mmap_mode="r"), + "depth": np.load(seq_dir / "zed_depth.npy", mmap_mode="r"), + # "dji_rgb": dji_rgb, + } + return self._seq_arrays[seq_name] + + def __len__(self) -> int: + return len(self.index_map) + + def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: + seq_name, frame_idx = self.index_map[idx] + arrs = self._load_sequence_arrays(seq_name) + + # (H, W, 3) uint8 + # rgb = np.asarray(arrs["dji_rgb"][frame_idx]) + # (H, W) uint16 mm (processor saves as uint16) + depth = np.asarray(arrs["depth"][frame_idx]).astype(np.float32) + # (doppler, elevation, azimuth, range) complex64 + radar = np.asarray(arrs["radar"][frame_idx]).copy() + + # Depth: uint16 mm -> float; optional mm -> m; handle invalid + if self.depth_in_meters: + depth = depth / 1000.0 + invalid = ~(np.isfinite(depth) & (depth > 0)) + depth[invalid] = 0.0 + depth = depth[np.newaxis, ...] # (1, H, W) + + # RGB: (H, W, 3) -> (3, H, W) + # rgb = np.transpose(rgb, (2, 0, 1)) + # if self.rgb_normalize: + # rgb = rgb.astype(np.float32) / 255.0 + + # Radar: amplitude and phase + radar_amplitude = np.abs(radar).astype(np.float32) + radar_phase = np.angle(radar).astype(np.float32) / np.pi + out = { + "radar_amplitude": torch.from_numpy(radar_amplitude), + "radar_phase": torch.from_numpy(radar_phase), + # "rgb": torch.from_numpy(rgb), + "depth": torch.from_numpy(depth), + "sequence": seq_name, + "frame_idx": frame_idx, + } + if self.return_radar_complex: + out["radar_cube"] = torch.from_numpy(radar.copy()) + # Depth in mm for optional use (1, H, W) float32 + depth_mm = np.asarray(arrs["depth"][frame_idx]).astype(np.float32) + out["depth_mm"] = torch.from_numpy(depth_mm[np.newaxis, ...]) + return out + + +def create_rice_dataloader( + root_dir: str, + batch_size: int = 8, + num_workers: int = 0, + frame_skip: int = 1, + sequences: Optional[List[str]] = None, + return_radar_complex: bool = False, + depth_in_meters: bool = True, + rgb_normalize: bool = True, + shuffle: bool = True, +) -> DataLoader: + """Create a DataLoader for the Rice (processor output) dataset.""" + dataset = RiceDataset( + root_dir=root_dir, + sequences=sequences, + frame_skip=frame_skip, + return_radar_complex=return_radar_complex, + depth_in_meters=depth_in_meters, + rgb_normalize=rgb_normalize, + ) + return DataLoader( + dataset, + batch_size=batch_size, + shuffle=shuffle, + num_workers=num_workers, + pin_memory=True, + ) + + +def create_split_dataloaders( + root_dir: str, + split_json_path: Optional[str] = None, + batch_size: int = 8, + num_workers: int = 0, + frame_skip: int = 1, + return_radar_complex: bool = False, + depth_in_meters: bool = True, + rgb_normalize: bool = True, + val_ratio: float = 0.2, + seed: Optional[int] = 42, +) -> Tuple[DataLoader, DataLoader, DataLoader]: + """ + Create train/val/test dataloaders using split.json. + + Reads the dataset split from split.json. If split_json_path is None, + uses radar_model/split.json (same directory as this module). + + Split JSON format: + { "test": ["seq_x", ...], "train": ["seq_a", ...] } // "train" optional + If "train" is present and non-empty, only those sequences are used for train/val. + Otherwise, all sequences under root_dir with required files that are not in "test" are used for training. + Validation is a random fraction (val_ratio) of the training samples. + + Returns: + train_loader, val_loader, test_loader + """ + if split_json_path is None: + split_path = Path(__file__).resolve().parent / "split.json" + else: + split_path = Path(split_json_path) + if not split_path.exists() and not split_path.is_absolute(): + # Resolve relative path from this module's directory (e.g. radar_model/) + fallback = Path(__file__).resolve().parent / split_path.name + if fallback.exists(): + split_path = fallback + + with split_path.open("r") as f: + split = json.load(f) + + test_sequences = split.get("test", []) + train_sequences_json = split.get("train", None) + + # Discover all valid sequences in root_dir + _discover = RiceDataset( + root_dir=root_dir, + sequences=None, + frame_skip=frame_skip, + return_radar_complex=return_radar_complex, + depth_in_meters=depth_in_meters, + rgb_normalize=rgb_normalize, + ) + test_set = set(test_sequences) + if train_sequences_json is not None and len(train_sequences_json) > 0: + # Use explicit train list (intersect with discovered so only valid seqs are used) + train_sequences = [s for s in train_sequences_json if s in _discover.sequences] + else: + # No "train" key: use all discovered sequences not in test + train_sequences = [s for s in _discover.sequences if s not in test_set] + + full_train_dataset = RiceDataset( + root_dir=root_dir, + sequences=train_sequences, + frame_skip=frame_skip, + return_radar_complex=return_radar_complex, + depth_in_meters=depth_in_meters, + rgb_normalize=rgb_normalize, + ) + + # Random split of training data for validation + n_total = len(full_train_dataset) + n_val = int(n_total * val_ratio) + if n_val == 0 and n_total > 0: + n_val = 1 + n_train = n_total - n_val + + if n_total == 0: + train_dataset = full_train_dataset + val_dataset = RiceDataset( + root_dir=root_dir, + sequences=[], + frame_skip=frame_skip, + return_radar_complex=return_radar_complex, + depth_in_meters=depth_in_meters, + rgb_normalize=rgb_normalize, + ) + elif seed is None: + train_dataset, val_dataset = random_split( + full_train_dataset, [n_train, n_val] + ) + else: + generator = torch.Generator() + generator.manual_seed(seed) + train_dataset, val_dataset = random_split( + full_train_dataset, [n_train, n_val], generator=generator + ) + + test_dataset = RiceDataset( + root_dir=root_dir, + sequences=test_sequences, + frame_skip=frame_skip, + return_radar_complex=return_radar_complex, + depth_in_meters=depth_in_meters, + rgb_normalize=rgb_normalize, + ) + + train_loader = DataLoader( + train_dataset, + batch_size=batch_size, + shuffle=True, + num_workers=num_workers, + pin_memory=True, + ) + val_loader = DataLoader( + val_dataset, + batch_size=batch_size, + shuffle=False, + num_workers=num_workers, + pin_memory=True, + ) + test_loader = DataLoader( + test_dataset, + batch_size=batch_size, + shuffle=False, + num_workers=num_workers, + pin_memory=True, + ) + + return train_loader, val_loader, test_loader diff --git a/src/Baselines/grt/grt_model.py b/src/Baselines/grt/grt_model.py new file mode 100644 index 0000000000000000000000000000000000000000..b5325432377cf68e41946884ccc342031b213825 --- /dev/null +++ b/src/Baselines/grt/grt_model.py @@ -0,0 +1,585 @@ +"""GRT-Small Model - from official codebase. + +This implementation directly copies necessary modules from the official GRT codebase +(grt/deepradar/modules). +""" + +import torch +import torch.nn as nn +from typing import Literal, Optional, Sequence +import numpy as np +from einops import rearrange + +# ============================================================================ +# Official GRT Modules (copied from grt/deepradar/modules/*.py) +# ============================================================================ + + +class PatchMerge(nn.Module): + """Merge patches with normalization and nominally reduced projection. + + From: grt/deepradar/modules/patch.py + """ + + def __init__( + self, d_in: int, d_out: int, scale: Sequence[int] = [], norm: bool = True + ) -> None: + super().__init__() + + self.scale = scale + d_merge = d_in * int(np.prod(scale)) + self.linear = nn.Linear(d_merge, d_out, bias=False) + self.norm = nn.LayerNorm(d_merge) if norm else None + + def _merge(self, x: torch.Tensor) -> torch.Tensor: + """Perform patch merging.""" + n, *t, c = x.shape + dims = sum(([d // s, s] for d, s in zip(t, self.scale)), start=[n]) + order = ( + [0] + + [2 * i + 1 for i in range(len(self.scale))] + + [2 * i + 2 for i in range(len(self.scale))] + + [-1] + ) + t2 = [d // s for d, s in zip(t, self.scale)] + return x.reshape(dims + [c]).permute(order).reshape(n, *t2, -1) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Merge and project.""" + merged = self._merge(x) + if self.norm is not None: + merged = self.norm(merged) + return self.linear(merged) + + +class Sinusoid(nn.Module): + """Centered N-dimensional sinusoidal positional embedding. + + From: grt/deepradar/modules/position.py + """ + + def __init__( + self, + scale: Optional[Sequence[float]] = None, + global_scale: float = 1.0, + coef: float = 10000.0, + ) -> None: + super().__init__() + if scale is None: + self.scale = [global_scale] + else: + self.scale = [s * global_scale for s in scale] + self.coef = coef + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Apply sinusoidal embedding.""" + # w = coef ** (-i / c) + nd = len(x.shape) - 2 + c = x.shape[-1] // 2 // nd + i = torch.arange(c, device=x.device) + w = self.coef ** (-i / c) + + start_dim = 0 + for axis, (d, scale) in enumerate(zip(x.shape[1:-1], self.scale * nd)): + # t = scale * (j - d/2) / (d/2) = scale * (2j / d - 1) + t = scale * (2 * (torch.arange(d, device=x.device) + 0.5) / d - 1) + wt = t[:, None] * w[None, :] + + p_slice = [None] * (len(x.shape) - 1) + [slice(None)] + p_slice[axis + 1] = slice(None) + + # pos[2 * i] = sin(w * t) + x_sin_slice = [slice(None)] * len(x.shape) + x_sin_slice[-1] = slice(start_dim, start_dim + c * 2, 2) + x_sin_slice = tuple(x_sin_slice) + p_slice_tuple = tuple(p_slice) + x[x_sin_slice] = x[x_sin_slice] + torch.sin(wt)[p_slice_tuple] + + # pos[2 * i + 1] = cos(w * t) + x_cos_slice = [slice(None)] * len(x.shape) + x_cos_slice[-1] = slice(start_dim + 1, start_dim + c * 2 + 1, 2) + x_cos_slice = tuple(x_cos_slice) + x[x_cos_slice] = x[x_cos_slice] + torch.cos(wt)[p_slice_tuple] + + start_dim += c * 2 + + return x + + +class Readout(nn.Module): + """Add readout token (concatenating along the spatial axis). + + From: grt/deepradar/modules/position.py + """ + + def __init__(self, d_model: int = 512) -> None: + super().__init__() + self.readout = nn.Parameter(data=torch.normal(0, 0.02, (d_model,))) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Concatenate readout token.""" + readout = torch.tile(self.readout[None, None, :], (x.shape[0], 1, 1)) + return torch.concatenate((x, readout), dim=1) + + +def transformer_mlp( + d_model: int = 512, + d_feedforward: int = 2048, + activation: str = "GELU", + dropout: float = 0.0, + eps: float = 1e-5, +) -> nn.Module: + """Create transformer MLP. + + From: grt/deepradar/modules/transformer.py + """ + return nn.Sequential( + nn.LayerNorm(d_model, eps=eps, bias=True), + nn.Linear(d_model, d_feedforward, bias=True), + getattr(nn, activation)(), + nn.Dropout(dropout), + nn.Linear(d_feedforward, d_model, bias=True), + nn.Dropout(dropout), + ) + + +class TransformerLayer(nn.Module): + """Single transformer (encoder) layer. + + Uses PyTorch's naming convention to match checkpoint: + - self_attn (not attn) + - linear1, linear2 (not feedforward.0, feedforward.4) + - norm1, norm2 (for attention and feedforward) + """ + + def __init__( + self, + d_model: int = 512, + n_head: int = 8, + d_feedforward: int = 2048, + dropout: float = 0.0, + activation: str = "GELU", + ) -> None: + super().__init__() + + # Attention with PyTorch naming + self.self_attn = nn.MultiheadAttention( + d_model, n_head, dropout=dropout, bias=True, batch_first=True + ) + self.dropout1 = nn.Dropout(dropout) + + # Feedforward with PyTorch naming + self.linear1 = nn.Linear(d_model, d_feedforward, bias=True) + self.dropout = nn.Dropout(dropout) + self.linear2 = nn.Linear(d_feedforward, d_model, bias=True) + self.dropout2 = nn.Dropout(dropout) + + # Norms + self.norm1 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + self.norm2 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + + # Activation + self.activation = getattr(nn, activation)() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Apply transformer with pre-norm (norm_first=True style).""" + # Self attention block + x2 = self.norm1(x) + x2 = self.self_attn(x2, x2, x2, need_weights=False)[0] + x = x + self.dropout1(x2) + + # Feedforward block + x2 = self.norm2(x) + x2 = self.linear1(x2) + x2 = self.activation(x2) + x2 = self.dropout(x2) + x2 = self.linear2(x2) + x = x + self.dropout2(x2) + + return x + + +class TransformerDecoder(nn.Module): + """Single transformer (decoder) layer. + + Uses PyTorch's naming convention to match checkpoint: + - self_attn, multihead_attn (not attn, attn2) + - linear1, linear2 (not feedforward.0, feedforward.4) + - norm1, norm2, norm3 (for self-attn, cross-attn, and feedforward) + """ + + def __init__( + self, + d_model: int = 512, + n_head: int = 8, + d_feedforward: int = 2048, + dropout: float = 0.0, + activation: str = "GELU", + ) -> None: + super().__init__() + + # Self attention with PyTorch naming + self.self_attn = nn.MultiheadAttention( + d_model, n_head, dropout=dropout, bias=True, batch_first=True + ) + self.dropout1 = nn.Dropout(dropout) + + # Cross attention with PyTorch naming (multihead_attn, not attn2) + self.multihead_attn = nn.MultiheadAttention( + d_model, n_head, dropout=dropout, bias=True, batch_first=True + ) + self.dropout2 = nn.Dropout(dropout) + + # Feedforward with PyTorch naming + self.linear1 = nn.Linear(d_model, d_feedforward, bias=True) + self.dropout = nn.Dropout(dropout) + self.linear2 = nn.Linear(d_feedforward, d_model, bias=True) + self.dropout3 = nn.Dropout(dropout) + + # Norms (note: norm2 is for cross-attention) + self.norm1 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + self.norm2 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + self.norm3 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + + # Activation + self.activation = getattr(nn, activation)() + + def forward(self, x: torch.Tensor, x_enc: torch.Tensor) -> torch.Tensor: + """Apply transformer decoder with pre-norm.""" + # Self attention block + x2 = self.norm1(x) + x2 = self.self_attn(x2, x2, x2, need_weights=False)[0] + x = x + self.dropout1(x2) + + # Cross attention block + x2 = self.norm2(x) + x2 = self.multihead_attn(x2, x_enc, x_enc, need_weights=False)[0] + x = x + self.dropout2(x2) + + # Feedforward block + x2 = self.norm3(x) + x2 = self.linear1(x2) + x2 = self.activation(x2) + x2 = self.dropout(x2) + x2 = self.linear2(x2) + x = x + self.dropout3(x2) + + return x + + +class BasisChange(nn.Module): + """Create "change-of-basis" query. + + From: grt/deepradar/modules/transformer.py + """ + + def __init__( + self, + shape: Sequence[int] = [], + flatten: bool = True, + scale: Optional[Sequence[float]] = None, + global_scale: float = 1.0, + ) -> None: + super().__init__() + + self.pos = Sinusoid(scale=scale, global_scale=global_scale) + self.shape = shape + self.flatten = flatten + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Apply change of basis.""" + idxs = tuple([slice(None)] + [None] * len(self.shape) + [slice(None)]) + query = self.pos(torch.tile(x[idxs], (1, *self.shape, 1))) + + if self.flatten: + query = query.reshape(x.shape[0], -1, x.shape[-1]) + return query + + +class Unpatch(nn.Module): + """Unpatch data. + + Args: + output_size: output 2D shape. + features: number of input features; should be `>= size * size`. + size: patch size as (width, height, channels). + """ + + def __init__( + self, + output_size: Sequence[int], + features: int = 512, + size: Sequence[int] = (16, 16), + ) -> None: + super().__init__() + + self.linear = nn.Linear(features, output_size[-1] * int(np.prod(size))) + self.size = size + self.output_size = output_size + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Perform 2D unpatching. + + Operates in batch-spatial-feature order; spatial axes are flattened on + the input, and unflattened in the output. + """ + embedding = self.linear(x) + + if len(self.size) == 2: + return rearrange( + embedding, + "n (x1 x2) (s1 s2 c) -> n (x1 s1) (x2 s2) c", + x1=self.output_size[0] // self.size[0], + x2=self.output_size[1] // self.size[1], + s1=self.size[0], + s2=self.size[1], + c=self.output_size[-1], + ) + elif len(self.size) == 3: + return rearrange( + embedding, + "n (x1 x2 x3) (s1 s2 s3 c) -> n (x1 s1) (x2 s2) (x3 s3) c", + x1=self.output_size[0] // self.size[0], + x2=self.output_size[1] // self.size[1], + x3=self.output_size[2] // self.size[2], + s1=self.size[0], + s2=self.size[1], + s3=self.size[2], + c=self.output_size[-1], + ) + else: + raise ValueError("Unpatch is only implemented for 2D and 3D tensors.") + + +# ============================================================================ +# GRT Model Components +# ============================================================================ + + +class GRTEncoder(nn.Module): + """GRT Transformer Encoder matching official implementation.""" + + def __init__( + self, + layers: int = 4, + dim: int = 512, + ff_ratio: float = 4.0, + head_dim: int = 64, + dropout: float = 0.1, + activation: str = "GELU", + patch: list[int] = [2, 8, 2, 4], + pos_scale: list[float] = [1.0, 1.0, 1.0, 1.0], + global_scale: float = 16.0, + input_channels: int = 2, + positions: Literal["flat", "nd"] = "nd", + ): + super().__init__() + + # Patch embedding + self.patch = PatchMerge(d_in=input_channels, d_out=dim, scale=patch, norm=False) + + # Position embedding + self.positions = positions + self.pos = Sinusoid(scale=pos_scale, global_scale=global_scale) + + # Readout token + self.readout = Readout(d_model=dim) + + # Encoder layers + self.layers = nn.ModuleList( + [ + TransformerLayer( + d_feedforward=int(ff_ratio * dim), + d_model=dim, + n_head=dim // head_dim, + dropout=dropout, + activation=activation, + ) + for _ in range(layers) + ] + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Forward pass.""" + # Patch embedding + embedded = self.patch(x) + + # Apply positional encoding + if self.positions == "nd": + embedded = self.pos(embedded) + + # Flatten spatial dimensions + flat = embedded.reshape(embedded.shape[0], -1, embedded.shape[-1]) + + # Apply flat positional encoding if needed + if self.positions == "flat": + flat = self.pos(flat) + + # Add readout token + x = self.readout(flat) + + # Apply encoder layers + for layer in self.layers: + x = layer(x) + + return x + + +class GRTDecoder3D(nn.Module): + """GRT 3D Transformer Decoder matching official implementation.""" + + def __init__( + self, + key: str = "map", + layers: int = 4, + dim: int = 512, + ff_ratio: float = 4.0, + head_dim: int = 64, + dropout: float = 0.1, + activation: str = "GELU", + shape: list[int] = [64, 128, 64], + pos_scale: list[float] = [1.0, 1.0, 1.0], + global_scale: float = 16.0, + patch: list[int] = [8, 8, 8], + out_dim: int = 0, + positions: Literal["flat", "nd"] = "nd", + mode: Literal["last", "pool"] = "last", + ): + super().__init__() + + self.key = key + self.out_dim = out_dim + self.mode = mode + + # Decoder layers + self.layers = nn.ModuleList( + [ + TransformerDecoder( + d_feedforward=int(ff_ratio * dim), + d_model=dim, + n_head=dim // head_dim, + dropout=dropout, + activation=activation, + ) + for _ in range(layers) + ] + ) + + # Query generation with position encoding + query_shape = [s // p for s, p in zip(shape, patch)] + if positions == "flat": + query_shape = [int(np.prod(query_shape))] + + self.query = BasisChange( + shape=query_shape, scale=pos_scale, global_scale=global_scale, flatten=True + ) + + # Unpatch to reconstruct output + self.unpatch = Unpatch( + output_size=(*shape, max(1, self.out_dim)), features=dim, size=patch + ) + + def forward(self, encoded: torch.Tensor) -> dict[str, torch.Tensor]: + """Forward pass.""" + # Extract readout token or pool + if self.mode == "last": + x = encoded[:, -1, :] + else: + x = torch.mean(encoded, dim=1) + + # Generate query with positional encoding + x = self.query(x) + + # Encoded features without readout token + enc = encoded[:, :-1, :] + + # Apply decoder layers + for layer in self.layers: + x = layer(x, enc) + + # Unpatch to 3D output + out = self.unpatch(x) + + # Squeeze channel dimension if binary output + if self.out_dim == 0: + out = out[..., 0] + + return {self.key: out} + + +# ============================================================================ +# Complete GRT-Small Model +# ============================================================================ + + +class GRTSmall(nn.Module): + """GRT-Small model for 3D occupancy mapping. + + Input: (batch, doppler, azimuth, elevation, range, 2) + - doppler: 64 + - azimuth: 8 + - elevation: 2 + - range: 256 + - channels: 2 (I/Q) + + Output: (batch, elevation, azimuth, range) + - elevation: 64 + - azimuth: 128 + - range: 64 + + ~29M parameters for GRT-small variant. + """ + + def __init__(self): + super().__init__() + + dim = 512 + layers = 4 + + # Create encoder - stored as "tokenizer" + "encoder" in checkpoint + # But we organize logically here and handle mapping in load_checkpoint + self.tokenizer = GRTEncoder( + layers=layers, + dim=dim, + ff_ratio=4.0, + head_dim=64, + dropout=0.1, + activation="GELU", + patch=[2, 8, 2, 4], + pos_scale=[1.0, 1.0, 1.0, 1.0], + global_scale=16.0, + input_channels=2, + positions="nd", + ) + + # Create decoder wrapper + self.decoder = nn.Module() + self.decoder.occ3d = GRTDecoder3D( + key="map", + layers=layers, + dim=dim, + ff_ratio=4.0, + head_dim=64, + dropout=0.1, + activation="GELU", + shape=[64, 128, 64], + pos_scale=[1.0, 1.0, 1.0], + global_scale=16.0, + patch=[8, 8, 8], + out_dim=0, + positions="nd", + mode="last", + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Forward pass.""" + # Encode + encoded = self.tokenizer(x) + + # Decode + output = self.decoder.occ3d(encoded) + + # Return just the occupancy map tensor + return output["map"] + + diff --git a/src/Baselines/grt/inference.py b/src/Baselines/grt/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..572ecaca594a8728070a7b79831ab75270cf328c --- /dev/null +++ b/src/Baselines/grt/inference.py @@ -0,0 +1,220 @@ +#!/usr/bin/env python3 +""" +Inference Script for GRT-Small (finetuned weights) + +Runs inference on all valid sequences in the configured Smoke-Eval root by +default, or on an explicit list supplied with ``--sequences``. +using weights trained by grt_finetune/train.py. +For each sequence, saves one .npy file: pred_depth.npy (dequantized predicted depth [T, 64, 128], values in [0, 1]). + +Single GPU: Each frame is seen exactly once; no duplication or incompleteness. +Multi-GPU (DDP): Dataloader is sharded; each rank writes its results to a file, then +main process merges with deduplication by frame_idx (keeps first occurrence) and saves. +""" + +import os +import torch +import numpy as np +import argparse +import yaml +import pickle +from tqdm import tqdm +from accelerate import Accelerator +from accelerate.utils import set_seed +from collections import defaultdict +from safetensors.torch import load_file + +from grt_model import GRTSmall +from dataloader import create_rice_dataloader +from augmentations import ( + translate_radar, + dequantize_depth, +) + +def batch_radar_to_spectrum( + radar_amplitude: torch.Tensor, radar_phase: torch.Tensor +) -> torch.Tensor: + """Restore the GRT spectrum layout from the packaged Smoke-Eval tensors.""" + + amplitude = radar_amplitude.permute(0, 1, 3, 2, 4) + phase = radar_phase.permute(0, 1, 3, 2, 4) + return torch.stack((amplitude, phase), dim=-1) + + +def main(): + parser = argparse.ArgumentParser( + description="Run GRT inference on Smoke-Eval." + ) + parser.add_argument( + "--config", type=str, default="config.yaml", help="Path to config file" + ) + parser.add_argument( + "--checkpoint", + type=str, + required=True, + help="Path to weights-only GRT .safetensors file", + ) + parser.add_argument( + "--output_dir", + type=str, + default="inference_results", + help="Directory to save results", + ) + parser.add_argument( + "--sequences", + type=str, + nargs="+", + default=None, + help="Optional sequence names; default discovers all valid sequences.", + ) + parser.add_argument( + "--debug", action="store_true", help="Run in debug mode (process only 1 batch)" + ) + args = parser.parse_args() + + # Load config + with open(args.config, "r") as f: + config = yaml.safe_load(f) + + # Initialize accelerator + accelerator = Accelerator(mixed_precision="fp16") + set_seed(config["training"].get("seed", 42)) + + # Create output directory (all ranks so DDP gather_dir can be created) + os.makedirs(args.output_dir, exist_ok=True) + + # Create model + accelerator.print("Creating GRT-Small model...") + model = GRTSmall() + + # Safetensors files contain only the model state dictionary. + accelerator.print(f"Loading checkpoint from {args.checkpoint}") + model.load_state_dict(load_file(args.checkpoint, device="cpu"), strict=True) + + # With ``sequences=None`` the public dataset loader discovers every valid + # sequence under the configured Smoke-Eval root. + accelerator.print(f"Inference sequences: {args.sequences}") + inference_loader = create_rice_dataloader( + root_dir=config["paths"]["data_root"], + batch_size=config["training"]["batch_size"], + num_workers=0, + frame_skip=1, + sequences=args.sequences, + shuffle=False, + ) + + # Prepare model and dataloader + model, inference_loader = accelerator.prepare(model, inference_loader) + model.eval() + + # Dictionary to aggregate results by sequence: sequence_id -> list of (frame_idx, pred_depth) + results_by_sequence = defaultdict(list) + + accelerator.print("Starting inference...") + + with torch.no_grad(): + for batch in tqdm( + inference_loader, disable=not accelerator.is_local_main_process + ): + # Extract data + rsp_data = batch_radar_to_spectrum( + batch["radar_amplitude"], batch["radar_phase"] + ) + sequences = batch["sequence"] + frame_indices = batch["frame_idx"] + + # Apply radar augmentation + rsp_data = translate_radar(rsp_data) + + # Forward pass + occupancy_pred_logits = model(rsp_data) # [B, 64, 128, 64] + + # Dequantize predicted occupancy to depth [B, 1, 64, 128], values in [0, 1] + pred_depth = dequantize_depth(occupancy_pred_logits) + pred_depth_np = ( + pred_depth.cpu().numpy().astype(np.float32) + ) # [B, 1, 64, 128] + + # Collect results (frame_idx, pred_depth per sample) + for i in range(len(sequences)): + seq_id = sequences[i] + f_idx = frame_indices[i].item() + # Store [1, 64, 128] per frame; will stack to [T, 64, 128] when saving + results_by_sequence[seq_id].append( + { + "frame_idx": f_idx, + "pred_depth": pred_depth_np[i], + } + ) + + if args.debug: + break + + # Single GPU: save directly (each frame seen once, no duplication) + # Multi-GPU: gather via files, merge with dedupe by frame_idx, then save + if accelerator.num_processes == 1: + if accelerator.is_main_process: + accelerator.print("Saving results (single process)...") + for seq_id, frames in tqdm( + results_by_sequence.items(), desc="Saving sequences" + ): + frames.sort(key=lambda x: x["frame_idx"]) + pred_depth_stack = np.stack([f["pred_depth"] for f in frames], axis=0) + pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 64, 128] + np.save( + os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"), + pred_depth_stack, + ) + accelerator.print( + f" {seq_id}: saved {pred_depth_stack.shape[0]} frames" + ) + accelerator.print(f"Processed {len(results_by_sequence)} sequences.") + accelerator.print(f"Results saved to {args.output_dir}") + else: + # DDP: gather results from all ranks via files, dedupe by frame_idx, save on main + accelerator.wait_for_everyone() + gather_dir = os.path.join(args.output_dir, "_gather") + os.makedirs(gather_dir, exist_ok=True) + rank = accelerator.process_index + rank_file = os.path.join(gather_dir, f"rank_{rank}_results.pkl") + with open(rank_file, "wb") as f: + pickle.dump(dict(results_by_sequence), f, protocol=pickle.HIGHEST_PROTOCOL) + accelerator.wait_for_everyone() + + if accelerator.is_main_process: + accelerator.print("Merging and deduplicating results from all ranks...") + merged_results = defaultdict(dict) # seq_id -> {frame_idx: pred_depth} + for r in range(accelerator.num_processes): + pkl_path = os.path.join(gather_dir, f"rank_{r}_results.pkl") + with open(pkl_path, "rb") as f: + rank_results = pickle.load(f) + for seq_id, frames in rank_results.items(): + for frame_data in frames: + f_idx = frame_data["frame_idx"] + if f_idx not in merged_results[seq_id]: + merged_results[seq_id][f_idx] = frame_data["pred_depth"] + os.remove(pkl_path) + + for seq_id, frame_dict in tqdm( + merged_results.items(), desc="Saving sequences" + ): + sorted_items = sorted(frame_dict.items(), key=lambda x: x[0]) + pred_depth_stack = np.stack([item[1] for item in sorted_items], axis=0) + pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 64, 128] + np.save( + os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"), + pred_depth_stack, + ) + accelerator.print( + f" {seq_id}: saved {pred_depth_stack.shape[0]} frames" + ) + if os.path.isdir(gather_dir) and not os.listdir(gather_dir): + os.rmdir(gather_dir) + accelerator.print(f"Processed {len(merged_results)} sequences.") + accelerator.print(f"Results saved to {args.output_dir}") + + accelerator.wait_for_everyone() + + +if __name__ == "__main__": + main() diff --git a/src/Baselines/grt/split.json b/src/Baselines/grt/split.json new file mode 100644 index 0000000000000000000000000000000000000000..6d4c8d97a758d12498d4d8983e44e8b32c1e77f8 --- /dev/null +++ b/src/Baselines/grt/split.json @@ -0,0 +1,16 @@ +{ + "test": [ + "Dell-1", + "Dell-2", + "Smoke-Dell-1", + "Smoke-Dell-2", + "brk-2", + "brk-3", + "Brk-b", + "brk-basement", + "Brk-stair", + "Smoke-brk-2", + "Smoke-brk-3", + "Smoke-brk-b" + ] +} \ No newline at end of file diff --git a/src/Baselines/grt_image/augmentations.py b/src/Baselines/grt_image/augmentations.py new file mode 100644 index 0000000000000000000000000000000000000000..529c17ae25984dacc39aacee7753198075566d02 --- /dev/null +++ b/src/Baselines/grt_image/augmentations.py @@ -0,0 +1,193 @@ +import torch +import torchvision.transforms.functional as TF +from torchvision.transforms import Resize, InterpolationMode +from typing import Union +import numpy as np + +AZIMUTH_RESOLUTION = 256 +ELEVATION_RESOLUTION = 128 + +# Depth output resolution: height=128, width=256 +DEPTH_TARGET_HEIGHT = 128 +DEPTH_TARGET_WIDTH = 256 + +resize_transform = Resize( + size=[ELEVATION_RESOLUTION, AZIMUTH_RESOLUTION], + interpolation=InterpolationMode.BILINEAR, + antialias=True, +) + +depth_resize_transform = Resize( + size=(DEPTH_TARGET_HEIGHT, DEPTH_TARGET_WIDTH), + interpolation=InterpolationMode.BILINEAR, + antialias=True, +) + + +def translate_radar(radar_data): + """ + Applies normalization to radar data after batching from dataloader. + Called before passing data into the model. + + Args: + radar_data: Batched radar tensor from dataloader + Shape: [B, 64, 8, 2, 256, 2] (batch, doppler, azimuth, elevation, range, channels) + - Channel 0: raw amplitude values + - Channel 1: phase normalized to [-1, 1] (divided by π) + + Returns: + Processed radar tensor with same shape [B, 64, 8, 2, 256, 2] + - Channel 0: sqrt(amplitude * 1e-3) for magnitude normalization + - Channel 1: phase * π (converted back to radians [-π, π]) + """ + radar_mag = radar_data[..., 0] # [B, 64, 8, 2, 256] - Extract raw amplitude + radar_phase = radar_data[..., 1] # [B, 64, 8, 2, 256] - Extract normalized phase + + # Normalize amplitude: scale then sqrt + radar_mag_processed = torch.sqrt(radar_mag * 1e-6) + + # Convert phase back to radians: [-1, 1] -> [-π, π] + radar_phase_processed = radar_phase * torch.pi + + # Stack channels back together: [B, 64, 8, 2, 256, 2] + radar_data_translated = torch.stack( + [radar_mag_processed, radar_phase_processed], dim=-1 + ) + return radar_data_translated + + +def resize_depth( + depth_map: Union[torch.Tensor, np.ndarray], +) -> Union[torch.Tensor, np.ndarray]: + """ + Process depth map from dataloader (same pipeline as denoiser/control crop_depth): + mm -> meters, clamp [0, 11.2] m, normalize to [0, 1], resize to (128, 256) (h, w). + + Args: + depth_map: Depth in millimeters. Torch or numpy. + Shapes: (H, W), (B, H, W), or (B, 1, H, W). + + Returns: + Depth in [0, 1], spatial size (128, 256). Shape [B, 128, 256] for batched input. + """ + is_numpy = isinstance(depth_map, np.ndarray) + if is_numpy: + depth_map = torch.from_numpy(depth_map) + + depth_map = depth_map.float() + original_shape = depth_map.shape + + if depth_map.dim() == 2: + depth_map = depth_map.unsqueeze(0) # (H, W) -> (1, H, W) + elif depth_map.dim() == 3: + depth_map = depth_map.unsqueeze(1) # (B, H, W) -> (B, 1, H, W) + elif depth_map.dim() != 4: + raise ValueError(f"Unexpected depth shape: {original_shape}") + + invalid_mask = ~(torch.isfinite(depth_map) & (depth_map >= 0)) + depth_map[invalid_mask] = 0.0 + + depth_map = depth_map / 1000.0 # mm -> meters + max_depth_m = 11.2 + depth_map = torch.clamp(depth_map, min=0.0, max=max_depth_m) + depth_map = depth_map / max_depth_m # [0, 1] + + invalid_mask = ~torch.isfinite(depth_map) + depth_map[invalid_mask] = 0.0 + + depth_map = depth_resize_transform(depth_map) # (..., 128, 256) + depth_values = depth_map.squeeze(1) # [B, 128, 256] or [1, 128, 256] + + if len(original_shape) == 2: + depth_values = depth_values.squeeze(0) # (128, 256) + + if is_numpy: + depth_values = depth_values.numpy() + return depth_values + + +def quantize_depth_to_occupancy(depth_values, num_range_bins=64): + """ + Quantizes 2D depth values into 3D binary occupancy grid. + + Args: + depth_values: Resized depth tensor + Shape: [B, elevation, azimuth] + Values: normalized to [0, 1] range + num_range_bins: Number of range bins for quantization (default: 64) + + Returns: + Binary 3D occupancy grid + Shape: [B, elevation, azimuth, num_range_bins] + Values: binary (0 or 1) indicating occupied bins + """ + B, elevation, azimuth = depth_values.shape + + # Quantize normalized depth [0, 1] directly to range bins [0, num_range_bins-1] + # Each bin represents 1/num_range_bins of the normalized depth range + bin_indices = torch.floor( + depth_values / (1.0 / num_range_bins) + ).long() # [B, elevation, azimuth] + bin_indices = torch.clamp( + bin_indices, 0, num_range_bins - 1 + ) # Handle edge case where depth_values = 1.0 + + # Create binary 3D occupancy grid + occupancy_grid = torch.zeros( + B, + elevation, + azimuth, + num_range_bins, + dtype=torch.float32, + device=depth_values.device, + ) # [B, elevation, azimuth, num_range_bins] + + # Set occupied bins to 1 + # Use advanced indexing to mark the appropriate range bin for each (elevation, azimuth) cell + batch_idx = torch.arange(B, device=depth_values.device)[:, None, None].expand( + B, elevation, azimuth + ) + elevation_idx = torch.arange(elevation, device=depth_values.device)[ + None, :, None + ].expand(B, elevation, azimuth) + azimuth_idx = torch.arange(azimuth, device=depth_values.device)[ + None, None, : + ].expand(B, elevation, azimuth) + + occupancy_grid[batch_idx, elevation_idx, azimuth_idx, bin_indices] = 1.0 + + return occupancy_grid # [B, elevation, azimuth, num_range_bins] + + +def dequantize_depth(occupancy_grid): + """ + Converts 3D binary occupancy grid back to 2D depth map. + This is the inverse operation of quantize_depth_to_occupancy. + + Args: + occupancy_grid: Binary 3D occupancy grid + Shape: [B, 128, 256, 64] (batch, elevation, azimuth, range) + Values: binary (0 or 1) or continuous (predicted probabilities) + + Returns: + Reconstructed depth map + Shape: [B, 1, 128, 256] (batch, channel, elevation, azimuth) + Values: normalized to [0, 1] range + """ + num_range_bins = occupancy_grid.shape[3] + + # Find the range bin with maximum value for each (elevation, azimuth) cell + # For binary: finds the occupied bin + # For continuous: finds the most likely bin + bin_indices = torch.argmax(occupancy_grid, dim=3) # [B, 128, 256] + + # Convert bin indices back to normalized depth values [0, 1] + # Use bin center: (bin_idx + 0.5) / num_bins + depth_values = (bin_indices.float() + 1) / num_range_bins # [B, 128, 256] + + # Add channel dimension: [B, 128, 256] -> [B, 1, 128, 256] + depth_map = depth_values.unsqueeze(1) # [B, 1, 128, 256] + + return depth_map + + diff --git a/src/Baselines/grt_image/dataloader.py b/src/Baselines/grt_image/dataloader.py new file mode 100644 index 0000000000000000000000000000000000000000..f3c7264ef01d2e0be7e87045fc18aeb0030bfe45 --- /dev/null +++ b/src/Baselines/grt_image/dataloader.py @@ -0,0 +1,344 @@ +""" +Dataloader for MobiCom processed dataset (output of processor.py). + +Uses the optimized format produced by processor.py: +- radar.npy: (N, doppler, elevation, azimuth, range) complex64 +- dji_rgb.npy: (N, H, W, 3) uint8 +- zed_depth.npy: (N, H, W) uint16, depth in millimeters + +This module provides: +- `RiceDataset`: frame-level dataset returning radar amplitude/phase, DJI RGB, + and ZED depth (ground truth). +- `create_rice_dataloader`: generic dataloader for an arbitrary set of sequences. +- `create_train_val_test_loaders`: uses the configured split file for fixed + validation sequences and a separate Smoke-Eval root for testing. +""" + +import json +from pathlib import Path +from typing import Dict, List, Optional, Tuple + +import numpy as np +import torch +import torch.nn.functional as F +from torch.utils.data import Dataset, DataLoader + + +class RiceDataset(Dataset): + """ + Dataset for processor.py output: radar, DJI RGB, and ZED depth per frame. + + Args: + root_dir: Root directory containing sequence subdirs (e.g. processed/), + each with radar.npy, dji_rgb.npy, zed_depth.npy. + sequences: Optional list of sequence names to load. If None, loads all + subdirs that contain the three required files. + frame_skip: Sample every frame_skip frames (1 = all frames). + return_radar_complex: If True, return radar as complex tensor; if False, + return radar_amplitude and radar_phase as separate float tensors. + depth_in_meters: If True, convert depth from mm to meters. + rgb_normalize: If True, return RGB in [0, 1] float; else uint8 [0, 255]. + """ + + REQUIRED_FILES = ("radar.npy", "dji_rgb.npy", "zed_depth.npy") + + def __init__( + self, + root_dir: str, + sequences: Optional[List[str]] = None, + frame_skip: int = 1, + return_radar_complex: bool = False, + depth_in_meters: bool = True, + rgb_normalize: bool = True, + image_height: int = 288, + image_width: int = 512, + ): + self.root_dir = Path(root_dir) + self.frame_skip = max(1, frame_skip) + self.return_radar_complex = return_radar_complex + self.depth_in_meters = depth_in_meters + self.rgb_normalize = rgb_normalize + self.image_height = int(image_height) + self.image_width = int(image_width) + if self.image_height <= 0 or self.image_width <= 0: + raise ValueError("image_height and image_width must be positive") + + self.sequences = self._discover_sequences(sequences) + self.index_map: List[Tuple[str, int]] = [] # (seq_name, frame_idx) + self._seq_arrays: Dict[str, Dict] = {} # seq -> {radar, depth, dji_rgb} + + self._build_index() + + def _discover_sequences(self, sequences: Optional[List[str]] = None) -> List[str]: + """Return list of sequence names that have all required files.""" + if not self.root_dir.is_dir(): + raise FileNotFoundError(f"Root directory not found: {self.root_dir}") + + all_seqs = sorted( + d.name + for d in self.root_dir.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + valid = [] + for name in all_seqs: + seq_dir = self.root_dir / name + if all((seq_dir / f).exists() for f in self.REQUIRED_FILES): + valid.append(name) + if sequences is not None: + valid = [s for s in valid if s in sequences] + return valid + + def _build_index(self) -> None: + """Build (seq_name, frame_idx) index, using radar.npy for frame count.""" + self.index_map.clear() + for seq_name in self.sequences: + seq_dir = self.root_dir / seq_name + radar_path = seq_dir / "radar.npy" + arrays = self._load_sequence_arrays(seq_name) + n_frames = min(array.shape[0] for array in arrays.values()) + for i in range(0, n_frames, self.frame_skip): + self.index_map.append((seq_name, i)) + + def _load_sequence_arrays(self, seq_name: str) -> Dict: + """Lazy-load or return cached arrays for a sequence.""" + if seq_name not in self._seq_arrays: + seq_dir = self.root_dir / seq_name + self._seq_arrays[seq_name] = { + "radar": np.load(seq_dir / "radar.npy", mmap_mode="r"), + "rgb": np.load(seq_dir / "dji_rgb.npy", mmap_mode="r"), + "depth": np.load(seq_dir / "zed_depth.npy", mmap_mode="r"), + } + return self._seq_arrays[seq_name] + + def __len__(self) -> int: + return len(self.index_map) + + def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]: + seq_name, frame_idx = self.index_map[idx] + arrs = self._load_sequence_arrays(seq_name) + + rgb = np.asarray(arrs["rgb"][frame_idx]).copy() + if rgb.ndim != 3 or rgb.shape[-1] != 3: + raise ValueError(f"Expected RGB frame shaped [H, W, 3], got {rgb.shape}") + # (H, W) uint16 mm (processor saves as uint16) + depth = np.asarray(arrs["depth"][frame_idx]).astype(np.float32) + # (doppler, elevation, azimuth, range) complex64 + radar = np.asarray(arrs["radar"][frame_idx]).copy() + + # Depth: uint16 mm -> float; optional mm -> m; handle invalid + if self.depth_in_meters: + depth = depth / 1000.0 + invalid = ~(np.isfinite(depth) & (depth > 0)) + depth[invalid] = 0.0 + depth = depth[np.newaxis, ...] # (1, H, W) + + # RGB: [H, W, 3] uint8 -> resized [3, image_height, image_width] float. + image = torch.from_numpy(np.transpose(rgb, (2, 0, 1)).copy()).float() + if self.rgb_normalize: + image = image / 255.0 + image = F.interpolate( + image.unsqueeze(0), + size=(self.image_height, self.image_width), + mode="bilinear", + align_corners=False, + ).squeeze(0) + + # Radar: amplitude and phase + radar_amplitude = np.abs(radar).astype(np.float32) + radar_phase = np.angle(radar).astype(np.float32) / np.pi + out = { + "radar_amplitude": torch.from_numpy(radar_amplitude), + "radar_phase": torch.from_numpy(radar_phase), + "image": image, + "depth": torch.from_numpy(depth), + "sequence": seq_name, + "frame_idx": frame_idx, + } + if self.return_radar_complex: + out["radar_cube"] = torch.from_numpy(radar.copy()) + # Depth in mm for optional use (1, H, W) float32 + depth_mm = np.asarray(arrs["depth"][frame_idx]).astype(np.float32) + out["depth_mm"] = torch.from_numpy(depth_mm[np.newaxis, ...]) + return out + + +def create_rice_dataloader( + root_dir: str, + batch_size: int = 8, + num_workers: int = 0, + frame_skip: int = 1, + sequences: Optional[List[str]] = None, + return_radar_complex: bool = False, + depth_in_meters: bool = True, + rgb_normalize: bool = True, + image_height: int = 288, + image_width: int = 512, + shuffle: bool = True, +) -> DataLoader: + """Create a DataLoader for the Rice (processor output) dataset.""" + dataset = RiceDataset( + root_dir=root_dir, + sequences=sequences, + frame_skip=frame_skip, + return_radar_complex=return_radar_complex, + depth_in_meters=depth_in_meters, + rgb_normalize=rgb_normalize, + image_height=image_height, + image_width=image_width, + ) + return DataLoader( + dataset, + batch_size=batch_size, + shuffle=shuffle, + num_workers=num_workers, + pin_memory=True, + ) + + +def create_train_val_test_loaders( + train_root: str, + split_json_path: Optional[str], + test_root: str, + batch_size: int = 8, + num_workers: int = 0, + frame_skip: int = 1, + return_radar_complex: bool = False, + depth_in_meters: bool = True, + rgb_normalize: bool = True, + image_height: int = 288, + image_width: int = 512, +) -> Tuple[DataLoader, DataLoader, DataLoader]: + """Create fixed training/validation and Smoke-Eval test loaders. + + The ``test`` list in the configured split file is treated as a fixed + validation sequence list. All other valid training sequences are used + for training. ``test_root`` is a separately structured Smoke-Eval tree; + every valid sequence it contains is evaluated only as the test set. + """ + if split_json_path is None: + split_path = Path(__file__).resolve().parent / "split.json" + else: + split_path = Path(split_json_path) + if not split_path.exists() and not split_path.is_absolute(): + fallback = Path(__file__).resolve().parent / split_path.name + if fallback.exists(): + split_path = fallback + + with split_path.open("r") as f: + split = json.load(f) + validation_sequences = split.get("test", []) + + discovered_train = RiceDataset( + root_dir=train_root, + frame_skip=frame_skip, + return_radar_complex=return_radar_complex, + depth_in_meters=depth_in_meters, + rgb_normalize=rgb_normalize, + image_height=image_height, + image_width=image_width, + ) + validation_set = set(validation_sequences) + train_sequences = [ + sequence + for sequence in discovered_train.sequences + if sequence not in validation_set + ] + resolved_validation_sequences = [ + sequence + for sequence in validation_sequences + if sequence in discovered_train.sequences + ] + + dataset_kwargs = { + "frame_skip": frame_skip, + "return_radar_complex": return_radar_complex, + "depth_in_meters": depth_in_meters, + "rgb_normalize": rgb_normalize, + "image_height": image_height, + "image_width": image_width, + } + train_dataset = RiceDataset( + root_dir=train_root, sequences=train_sequences, **dataset_kwargs + ) + val_dataset = RiceDataset( + root_dir=train_root, + sequences=resolved_validation_sequences, + **dataset_kwargs, + ) + test_dataset = RiceDataset(root_dir=test_root, sequences=None, **dataset_kwargs) + + loader_kwargs = {"batch_size": batch_size, "num_workers": num_workers, "pin_memory": True} + train_loader = DataLoader(train_dataset, shuffle=True, **loader_kwargs) + val_loader = DataLoader(val_dataset, shuffle=False, **loader_kwargs) + test_loader = DataLoader(test_dataset, shuffle=False, **loader_kwargs) + return train_loader, val_loader, test_loader + + +def create_train_val_loaders( + train_root: str, + split_json_path: Optional[str], + batch_size: int = 8, + num_workers: int = 0, + frame_skip: int = 1, + return_radar_complex: bool = False, + depth_in_meters: bool = True, + rgb_normalize: bool = True, + image_height: int = 288, + image_width: int = 512, +) -> Tuple[DataLoader, DataLoader]: + """Create training and fixed validation loaders only.""" + if split_json_path is None: + split_path = Path(__file__).resolve().parent / "split.json" + else: + split_path = Path(split_json_path) + if not split_path.exists() and not split_path.is_absolute(): + fallback = Path(__file__).resolve().parent / split_path.name + if fallback.exists(): + split_path = fallback + + with split_path.open("r") as f: + split = json.load(f) + validation_sequences = split.get("test", []) + + discovered = RiceDataset( + root_dir=train_root, + frame_skip=frame_skip, + return_radar_complex=return_radar_complex, + depth_in_meters=depth_in_meters, + rgb_normalize=rgb_normalize, + image_height=image_height, + image_width=image_width, + ) + validation_set = set(validation_sequences) + train_sequences = [ + sequence for sequence in discovered.sequences if sequence not in validation_set + ] + resolved_validation_sequences = [ + sequence for sequence in validation_sequences if sequence in discovered.sequences + ] + + dataset_kwargs = { + "frame_skip": frame_skip, + "return_radar_complex": return_radar_complex, + "depth_in_meters": depth_in_meters, + "rgb_normalize": rgb_normalize, + "image_height": image_height, + "image_width": image_width, + } + train_dataset = RiceDataset( + root_dir=train_root, sequences=train_sequences, **dataset_kwargs + ) + val_dataset = RiceDataset( + root_dir=train_root, + sequences=resolved_validation_sequences, + **dataset_kwargs, + ) + loader_kwargs = { + "batch_size": batch_size, + "num_workers": num_workers, + "pin_memory": True, + } + return ( + DataLoader(train_dataset, shuffle=True, **loader_kwargs), + DataLoader(val_dataset, shuffle=False, **loader_kwargs), + ) diff --git a/src/Baselines/grt_image/grt_image_resnet_inference.example.yaml b/src/Baselines/grt_image/grt_image_resnet_inference.example.yaml new file mode 100644 index 0000000000000000000000000000000000000000..fb9e03c4ec06268e36caa6eadaab7b56593a118c --- /dev/null +++ b/src/Baselines/grt_image/grt_image_resnet_inference.example.yaml @@ -0,0 +1,16 @@ +# Anonymous, release-relative configuration for the paper's GRT+Image baseline. +# This is the ResNet-18 implementation in Baselines/grt_image. +paths: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + +training: + batch_size: 1 + mixed_precision: fp16 + seed: 42 + +data: + image_height: 288 + image_width: 512 + +model: + resnet18_pretrained: true diff --git a/src/Baselines/grt_image/grt_model.py b/src/Baselines/grt_image/grt_model.py new file mode 100644 index 0000000000000000000000000000000000000000..c5fa1ca8b224cf67f92e0dba7d39b6e608620706 --- /dev/null +++ b/src/Baselines/grt_image/grt_model.py @@ -0,0 +1,799 @@ +"""GRT-Small Model - from official codebase. + +This implementation directly copies necessary modules from the official GRT codebase +(grt/deepradar/modules). +""" + +import torch +import torch.nn as nn +from torchvision.models import ResNet18_Weights, resnet18 +from typing import Literal, Optional, Sequence +import numpy as np +from einops import rearrange +from safetensors.torch import load_file + +# ============================================================================ +# Official GRT Modules (copied from grt/deepradar/modules/*.py) +# ============================================================================ + + +class PatchMerge(nn.Module): + """Merge patches with normalization and nominally reduced projection. + + From: grt/deepradar/modules/patch.py + """ + + def __init__( + self, d_in: int, d_out: int, scale: Sequence[int] = [], norm: bool = True + ) -> None: + super().__init__() + + self.scale = scale + d_merge = d_in * int(np.prod(scale)) + self.linear = nn.Linear(d_merge, d_out, bias=False) + self.norm = nn.LayerNorm(d_merge) if norm else None + + def _merge(self, x: torch.Tensor) -> torch.Tensor: + """Perform patch merging.""" + n, *t, c = x.shape + dims = sum(([d // s, s] for d, s in zip(t, self.scale)), start=[n]) + order = ( + [0] + + [2 * i + 1 for i in range(len(self.scale))] + + [2 * i + 2 for i in range(len(self.scale))] + + [-1] + ) + t2 = [d // s for d, s in zip(t, self.scale)] + return x.reshape(dims + [c]).permute(order).reshape(n, *t2, -1) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Merge and project.""" + merged = self._merge(x) + if self.norm is not None: + merged = self.norm(merged) + return self.linear(merged) + + +class Sinusoid(nn.Module): + """Centered N-dimensional sinusoidal positional embedding. + + From: grt/deepradar/modules/position.py + """ + + def __init__( + self, + scale: Optional[Sequence[float]] = None, + global_scale: float = 1.0, + coef: float = 10000.0, + ) -> None: + super().__init__() + if scale is None: + self.scale = [global_scale] + else: + self.scale = [s * global_scale for s in scale] + self.coef = coef + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Apply sinusoidal embedding.""" + # w = coef ** (-i / c) + nd = len(x.shape) - 2 + c = x.shape[-1] // 2 // nd + i = torch.arange(c, device=x.device) + w = self.coef ** (-i / c) + + start_dim = 0 + for axis, (d, scale) in enumerate(zip(x.shape[1:-1], self.scale * nd)): + # t = scale * (j - d/2) / (d/2) = scale * (2j / d - 1) + t = scale * (2 * (torch.arange(d, device=x.device) + 0.5) / d - 1) + wt = t[:, None] * w[None, :] + + p_slice = [None] * (len(x.shape) - 1) + [slice(None)] + p_slice[axis + 1] = slice(None) + + # pos[2 * i] = sin(w * t) + x_sin_slice = [slice(None)] * len(x.shape) + x_sin_slice[-1] = slice(start_dim, start_dim + c * 2, 2) + x_sin_slice = tuple(x_sin_slice) + p_slice_tuple = tuple(p_slice) + x[x_sin_slice] = x[x_sin_slice] + torch.sin(wt)[p_slice_tuple] + + # pos[2 * i + 1] = cos(w * t) + x_cos_slice = [slice(None)] * len(x.shape) + x_cos_slice[-1] = slice(start_dim + 1, start_dim + c * 2 + 1, 2) + x_cos_slice = tuple(x_cos_slice) + x[x_cos_slice] = x[x_cos_slice] + torch.cos(wt)[p_slice_tuple] + + start_dim += c * 2 + + return x + + +class Readout(nn.Module): + """Add readout token (concatenating along the spatial axis). + + From: grt/deepradar/modules/position.py + """ + + def __init__(self, d_model: int = 512) -> None: + super().__init__() + self.readout = nn.Parameter(data=torch.normal(0, 0.02, (d_model,))) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Concatenate readout token.""" + readout = torch.tile(self.readout[None, None, :], (x.shape[0], 1, 1)) + return torch.concatenate((x, readout), dim=1) + + +def transformer_mlp( + d_model: int = 512, + d_feedforward: int = 2048, + activation: str = "GELU", + dropout: float = 0.0, + eps: float = 1e-5, +) -> nn.Module: + """Create transformer MLP. + + From: grt/deepradar/modules/transformer.py + """ + return nn.Sequential( + nn.LayerNorm(d_model, eps=eps, bias=True), + nn.Linear(d_model, d_feedforward, bias=True), + getattr(nn, activation)(), + nn.Dropout(dropout), + nn.Linear(d_feedforward, d_model, bias=True), + nn.Dropout(dropout), + ) + + +class TransformerLayer(nn.Module): + """Single transformer (encoder) layer. + + Uses PyTorch's naming convention to match checkpoint: + - self_attn (not attn) + - linear1, linear2 (not feedforward.0, feedforward.4) + - norm1, norm2 (for attention and feedforward) + """ + + def __init__( + self, + d_model: int = 512, + n_head: int = 8, + d_feedforward: int = 2048, + dropout: float = 0.0, + activation: str = "GELU", + ) -> None: + super().__init__() + + # Attention with PyTorch naming + self.self_attn = nn.MultiheadAttention( + d_model, n_head, dropout=dropout, bias=True, batch_first=True + ) + self.dropout1 = nn.Dropout(dropout) + + # Feedforward with PyTorch naming + self.linear1 = nn.Linear(d_model, d_feedforward, bias=True) + self.dropout = nn.Dropout(dropout) + self.linear2 = nn.Linear(d_feedforward, d_model, bias=True) + self.dropout2 = nn.Dropout(dropout) + + # Norms + self.norm1 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + self.norm2 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + + # Activation + self.activation = getattr(nn, activation)() + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Apply transformer with pre-norm (norm_first=True style).""" + # Self attention block + x2 = self.norm1(x) + x2 = self.self_attn(x2, x2, x2, need_weights=False)[0] + x = x + self.dropout1(x2) + + # Feedforward block + x2 = self.norm2(x) + x2 = self.linear1(x2) + x2 = self.activation(x2) + x2 = self.dropout(x2) + x2 = self.linear2(x2) + x = x + self.dropout2(x2) + + return x + + +class TransformerDecoder(nn.Module): + """Single transformer (decoder) layer. + + Uses PyTorch's naming convention to match checkpoint: + - self_attn, multihead_attn (not attn, attn2) + - linear1, linear2 (not feedforward.0, feedforward.4) + - norm1, norm2, norm3 (for self-attn, cross-attn, and feedforward) + """ + + def __init__( + self, + d_model: int = 512, + n_head: int = 8, + d_feedforward: int = 2048, + dropout: float = 0.0, + activation: str = "GELU", + ) -> None: + super().__init__() + + # Self attention with PyTorch naming + self.self_attn = nn.MultiheadAttention( + d_model, n_head, dropout=dropout, bias=True, batch_first=True + ) + self.dropout1 = nn.Dropout(dropout) + + # Cross attention with PyTorch naming (multihead_attn, not attn2) + self.multihead_attn = nn.MultiheadAttention( + d_model, n_head, dropout=dropout, bias=True, batch_first=True + ) + self.dropout2 = nn.Dropout(dropout) + + # Feedforward with PyTorch naming + self.linear1 = nn.Linear(d_model, d_feedforward, bias=True) + self.dropout = nn.Dropout(dropout) + self.linear2 = nn.Linear(d_feedforward, d_model, bias=True) + self.dropout3 = nn.Dropout(dropout) + + # Norms (note: norm2 is for cross-attention) + self.norm1 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + self.norm2 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + self.norm3 = nn.LayerNorm(d_model, eps=1e-5, bias=True) + + # Activation + self.activation = getattr(nn, activation)() + + def forward(self, x: torch.Tensor, x_enc: torch.Tensor) -> torch.Tensor: + """Apply transformer decoder with pre-norm.""" + # Self attention block + x2 = self.norm1(x) + x2 = self.self_attn(x2, x2, x2, need_weights=False)[0] + x = x + self.dropout1(x2) + + # Cross attention block + x2 = self.norm2(x) + x2 = self.multihead_attn(x2, x_enc, x_enc, need_weights=False)[0] + x = x + self.dropout2(x2) + + # Feedforward block + x2 = self.norm3(x) + x2 = self.linear1(x2) + x2 = self.activation(x2) + x2 = self.dropout(x2) + x2 = self.linear2(x2) + x = x + self.dropout3(x2) + + return x + + +class BasisChange(nn.Module): + """Create "change-of-basis" query. + + From: grt/deepradar/modules/transformer.py + """ + + def __init__( + self, + shape: Sequence[int] = [], + flatten: bool = True, + scale: Optional[Sequence[float]] = None, + global_scale: float = 1.0, + ) -> None: + super().__init__() + + self.pos = Sinusoid(scale=scale, global_scale=global_scale) + self.shape = shape + self.flatten = flatten + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Apply change of basis.""" + idxs = tuple([slice(None)] + [None] * len(self.shape) + [slice(None)]) + query = self.pos(torch.tile(x[idxs], (1, *self.shape, 1))) + + if self.flatten: + query = query.reshape(x.shape[0], -1, x.shape[-1]) + return query + + +class Unpatch(nn.Module): + """Unpatch data. + + Args: + output_size: output 2D shape. + features: number of input features; should be `>= size * size`. + size: patch size as (width, height, channels). + """ + + def __init__( + self, + output_size: Sequence[int], + features: int = 512, + size: Sequence[int] = (16, 16), + ) -> None: + super().__init__() + + self.linear = nn.Linear(features, output_size[-1] * int(np.prod(size))) + self.size = size + self.output_size = output_size + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Perform 2D unpatching. + + Operates in batch-spatial-feature order; spatial axes are flattened on + the input, and unflattened in the output. + """ + embedding = self.linear(x) + + if len(self.size) == 2: + return rearrange( + embedding, + "n (x1 x2) (s1 s2 c) -> n (x1 s1) (x2 s2) c", + x1=self.output_size[0] // self.size[0], + x2=self.output_size[1] // self.size[1], + s1=self.size[0], + s2=self.size[1], + c=self.output_size[-1], + ) + elif len(self.size) == 3: + return rearrange( + embedding, + "n (x1 x2 x3) (s1 s2 s3 c) -> n (x1 s1) (x2 s2) (x3 s3) c", + x1=self.output_size[0] // self.size[0], + x2=self.output_size[1] // self.size[1], + x3=self.output_size[2] // self.size[2], + s1=self.size[0], + s2=self.size[1], + s3=self.size[2], + c=self.output_size[-1], + ) + else: + raise ValueError("Unpatch is only implemented for 2D and 3D tensors.") + + +# ============================================================================ +# GRT Model Components +# ============================================================================ + + +class GRTEncoder(nn.Module): + """GRT Transformer Encoder matching official implementation.""" + + def __init__( + self, + layers: int = 4, + dim: int = 512, + ff_ratio: float = 4.0, + head_dim: int = 64, + dropout: float = 0.1, + activation: str = "GELU", + patch: list[int] = [2, 8, 2, 4], + pos_scale: list[float] = [1.0, 1.0, 1.0, 1.0], + global_scale: float = 16.0, + input_channels: int = 2, + positions: Literal["flat", "nd"] = "nd", + ): + super().__init__() + + # Patch embedding + self.patch = PatchMerge(d_in=input_channels, d_out=dim, scale=patch, norm=False) + + # Position embedding + self.positions = positions + self.pos = Sinusoid(scale=pos_scale, global_scale=global_scale) + + # Readout token + self.readout = Readout(d_model=dim) + + # Encoder layers + self.layers = nn.ModuleList( + [ + TransformerLayer( + d_feedforward=int(ff_ratio * dim), + d_model=dim, + n_head=dim // head_dim, + dropout=dropout, + activation=activation, + ) + for _ in range(layers) + ] + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Forward pass.""" + # Patch embedding + embedded = self.patch(x) + + # Apply positional encoding + if self.positions == "nd": + embedded = self.pos(embedded) + + # Flatten spatial dimensions + flat = embedded.reshape(embedded.shape[0], -1, embedded.shape[-1]) + + # Apply flat positional encoding if needed + if self.positions == "flat": + flat = self.pos(flat) + + # Add readout token + x = self.readout(flat) + + # Apply encoder layers + for layer in self.layers: + x = layer(x) + + return x + + +class GRTDecoder3D(nn.Module): + """GRT 3D Transformer Decoder matching official implementation.""" + + def __init__( + self, + key: str = "map", + layers: int = 4, + dim: int = 512, + ff_ratio: float = 4.0, + head_dim: int = 64, + dropout: float = 0.1, + activation: str = "GELU", + shape: list[int] = [64, 128, 64], + pos_scale: list[float] = [1.0, 1.0, 1.0], + global_scale: float = 16.0, + patch: list[int] = [8, 8, 8], + out_dim: int = 0, + positions: Literal["flat", "nd"] = "nd", + mode: Literal["last", "pool"] = "last", + ): + super().__init__() + + self.key = key + self.out_dim = out_dim + self.mode = mode + + # Decoder layers + self.layers = nn.ModuleList( + [ + TransformerDecoder( + d_feedforward=int(ff_ratio * dim), + d_model=dim, + n_head=dim // head_dim, + dropout=dropout, + activation=activation, + ) + for _ in range(layers) + ] + ) + + # Query generation with position encoding + query_shape = [s // p for s, p in zip(shape, patch)] + if positions == "flat": + query_shape = [int(np.prod(query_shape))] + + self.query = BasisChange( + shape=query_shape, scale=pos_scale, global_scale=global_scale, flatten=True + ) + + # Unpatch to reconstruct output + self.unpatch = Unpatch( + output_size=(*shape, max(1, self.out_dim)), features=dim, size=patch + ) + + def forward(self, encoded: torch.Tensor) -> dict[str, torch.Tensor]: + """Forward pass.""" + # Extract readout token or pool + if self.mode == "last": + x = encoded[:, -1, :] + else: + x = torch.mean(encoded, dim=1) + + # Generate query with positional encoding + x = self.query(x) + + # Encoded features without readout token + enc = encoded[:, :-1, :] + + # Apply decoder layers + for layer in self.layers: + x = layer(x, enc) + + # Unpatch to 3D output + out = self.unpatch(x) + + # Squeeze channel dimension if binary output + if self.out_dim == 0: + out = out[..., 0] + + return {self.key: out} + + +# ============================================================================ +# Complete GRT-Small Model +# ============================================================================ + + +class GRTSmall(nn.Module): + """GRT-Small model for 3D occupancy mapping. + + Input: (batch, doppler, azimuth, elevation, range, 2) + - doppler: 64 + - azimuth: 8 + - elevation: 2 + - range: 256 + - channels: 2 (I/Q) + + Output: (batch, elevation, azimuth, range) + - elevation: 64 + - azimuth: 128 + - range: 64 + + ~29M parameters for GRT-small variant. + """ + + def __init__(self): + super().__init__() + + dim = 512 + layers = 4 + + # Create encoder - stored as "tokenizer" + "encoder" in checkpoint + # But we organize logically here and handle mapping in load_checkpoint + self.tokenizer = GRTEncoder( + layers=layers, + dim=dim, + ff_ratio=4.0, + head_dim=64, + dropout=0.1, + activation="GELU", + patch=[2, 8, 2, 4], + pos_scale=[1.0, 1.0, 1.0, 1.0], + global_scale=16.0, + input_channels=2, + positions="nd", + ) + + # Create decoder wrapper + self.decoder = nn.Module() + self.decoder.occ3d = GRTDecoder3D( + key="map", + layers=layers, + dim=dim, + ff_ratio=4.0, + head_dim=64, + dropout=0.1, + activation="GELU", + shape=[64, 128, 64], + pos_scale=[1.0, 1.0, 1.0], + global_scale=16.0, + patch=[8, 8, 8], + out_dim=0, + positions="nd", + mode="last", + ) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Forward pass.""" + # Encode + encoded = self.tokenizer(x) + + # Decode + output = self.decoder.occ3d(encoded) + + # Return just the occupancy map tensor + return output["map"] + + +class ResNet18ImageTokenizer(nn.Module): + """Coarse ResNet-18 spatial tokens projected into GRT's 512-D memory.""" + + def __init__( + self, + pretrained: bool, + image_height: int, + image_width: int, + output_dim: int = 512, + ): + super().__init__() + self.pretrained = bool(pretrained) + self.image_height = int(image_height) + self.image_width = int(image_width) + self.output_stride = 32 + if ( + self.image_height % self.output_stride + or self.image_width % self.output_stride + ): + raise ValueError( + "ResNet-18 tokenization requires image dimensions divisible by 32, " + f"got {(self.image_height, self.image_width)}" + ) + + weights = ResNet18_Weights.DEFAULT if self.pretrained else None + resnet = resnet18(weights=weights) + self.backbone = nn.Sequential( + resnet.conv1, + resnet.bn1, + resnet.relu, + resnet.maxpool, + resnet.layer1, + resnet.layer2, + resnet.layer3, + resnet.layer4, + ) + self.register_buffer( + "image_mean", + torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1), + ) + self.register_buffer( + "image_std", + torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1), + ) + self.projection = nn.Sequential( + nn.LayerNorm(512), + nn.Linear(512, output_dim), + ) + self.modality = nn.Parameter(torch.empty(1, 1, output_dim)) + nn.init.normal_(self.modality, mean=0.0, std=0.02) + + if self.pretrained: + for parameter in self.backbone.parameters(): + parameter.requires_grad_(False) + self.backbone.eval() + + def train(self, mode: bool = True): + super().train(mode) + if self.pretrained: + self.backbone.eval() + return self + + def forward(self, image: torch.Tensor) -> torch.Tensor: + """Return layer-4 spatial features as [B, H/32 * W/32, 512].""" + if image.ndim != 4 or image.shape[1] != 3: + raise ValueError( + "ResNet18ImageTokenizer expects RGB images shaped [B, 3, H, W], " + f"got {tuple(image.shape)}" + ) + height, width = image.shape[-2:] + if height != self.image_height or width != self.image_width: + raise ValueError( + "Image size must match the configured ResNet-18 size " + f"{(self.image_height, self.image_width)}, got {(height, width)}" + ) + + image = (image - self.image_mean) / self.image_std + if self.pretrained: + with torch.no_grad(): + features = self.backbone(image) + else: + features = self.backbone(image) + + spatial_tokens = features.flatten(2).transpose(1, 2) + return self.projection(spatial_tokens) + self.modality + + +def fuse_decoder_memory( + radar_encoded: torch.Tensor, image_tokens: torch.Tensor +) -> torch.Tensor: + """Insert image memory before GRT's final readout token. + + GRTDecoder3D uses the final token as its query seed and every preceding + token as cross-attention memory. Keeping the readout last is therefore a + required part of the fusion contract. + """ + if radar_encoded.ndim != 3 or image_tokens.ndim != 3: + raise ValueError("radar_encoded and image_tokens must both be [B, N, C]") + if radar_encoded.shape[1] < 1: + raise ValueError("radar_encoded must contain the GRT readout token") + if ( + radar_encoded.shape[0] != image_tokens.shape[0] + or radar_encoded.shape[2] != image_tokens.shape[2] + ): + raise ValueError( + "radar and image token batches must have matching batch and channel dimensions" + ) + return torch.cat( + [radar_encoded[:, :-1, :], image_tokens, radar_encoded[:, -1:, :]], + dim=1, + ) + + +class GRTImageNaiveSmall(nn.Module): + """Naive GRT+Image model with a fresh joint occupancy decoder.""" + + def __init__( + self, + resnet18_pretrained: bool = False, + image_height: int = 288, + image_width: int = 512, + ): + super().__init__() + + dim = 512 + layers = 4 + self.tokenizer = GRTEncoder( + layers=layers, + dim=dim, + ff_ratio=4.0, + head_dim=64, + dropout=0.1, + activation="GELU", + patch=[2, 8, 2, 4], + pos_scale=[1.0, 1.0, 1.0, 1.0], + global_scale=16.0, + input_channels=2, + positions="nd", + ) + self.image_tokenizer = ResNet18ImageTokenizer( + pretrained=resnet18_pretrained, + image_height=image_height, + image_width=image_width, + output_dim=dim, + ) + + self.decoder = nn.Module() + self.decoder.occ3d = GRTDecoder3D( + key="map", + layers=layers, + dim=dim, + ff_ratio=4.0, + head_dim=64, + dropout=0.1, + activation="GELU", + shape=[128, 256, 64], + pos_scale=[1.0, 1.0, 1.0], + global_scale=16.0, + patch=[8, 8, 8], + out_dim=0, + positions="nd", + mode="last", + ) + self._radar_encoder_frozen = False + + def freeze_radar_encoder(self) -> None: + """Freeze GRT feature extraction and keep its dropout disabled.""" + self._radar_encoder_frozen = True + for parameter in self.tokenizer.parameters(): + parameter.requires_grad_(False) + self.tokenizer.eval() + + def train(self, mode: bool = True): + super().train(mode) + if self._radar_encoder_frozen: + self.tokenizer.eval() + return self + + def forward(self, radar: torch.Tensor, image: torch.Tensor) -> torch.Tensor: + radar_encoded = self.tokenizer(radar) + image_tokens = self.image_tokenizer(image) + fused_encoded = fuse_decoder_memory(radar_encoded, image_tokens) + return self.decoder.occ3d(fused_encoded)["map"] + + +def load_radar_encoder_checkpoint( + model: GRTImageNaiveSmall, checkpoint_path, map_location="cpu" +) -> dict: + """Load only the pretrained GRT tokenizer/encoder and leave fusion fresh.""" + state_dict = load_file(checkpoint_path, device="cpu") + + encoder_state = { + key: value for key, value in state_dict.items() if key.startswith("tokenizer.") + } + if not encoder_state: + raise RuntimeError( + "Radar checkpoint does not contain any tokenizer.* encoder parameters" + ) + + missing_keys, unexpected_keys = model.load_state_dict(encoder_state, strict=False) + missing_encoder_keys = [ + key for key in missing_keys if key.startswith("tokenizer.") + ] + if missing_encoder_keys or unexpected_keys: + raise RuntimeError( + "Radar checkpoint is not compatible with the GRT encoder: " + f"missing encoder keys {missing_encoder_keys}; " + f"unexpected keys {list(unexpected_keys)}" + ) + return checkpoint + + diff --git a/src/Baselines/grt_image/inference.py b/src/Baselines/grt_image/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..fffa475f9596d4a606df579876a0175580058366 --- /dev/null +++ b/src/Baselines/grt_image/inference.py @@ -0,0 +1,224 @@ +#!/usr/bin/env python3 +""" +Inference for the naive GRT+Image baseline. + +Runs inference on specified sequences (default: brk_3rd, brk_3rd_fog, brk_3rd_fog2) +using weights trained by train.py. +For each sequence, saves pred_depth.npy with shape [T, 128, 256] in [0, 1]. + +Single GPU: Each frame is seen exactly once; no duplication or incompleteness. +Multi-GPU (DDP): Dataloader is sharded; each rank writes its results to a file, then +main process merges with deduplication by frame_idx (keeps first occurrence) and saves. +""" + +import os +import torch +import numpy as np +import argparse +import yaml +import pickle +from tqdm import tqdm +from accelerate import Accelerator +from accelerate.utils import set_seed +from collections import defaultdict +from safetensors.torch import load_file + +from grt_model import GRTImageNaiveSmall +from dataloader import create_rice_dataloader +from augmentations import ( + translate_radar, + dequantize_depth, +) + +def batch_radar_to_spectrum( + radar_amplitude: torch.Tensor, radar_phase: torch.Tensor +) -> torch.Tensor: + """Build GRT's [B, D, A, E, R, 2] spectrum from loader tensors.""" + amplitude = radar_amplitude.permute(0, 1, 3, 2, 4) + phase = radar_phase.permute(0, 1, 3, 2, 4) + return torch.stack([amplitude, phase], dim=-1) + + +def main(): + parser = argparse.ArgumentParser( + description="Run naive GRT+Image inference on Smoke-Eval sequences" + ) + parser.add_argument( + "--config", type=str, default="config.yaml", help="Path to config file" + ) + parser.add_argument( + "--checkpoint", + type=str, + required=True, + help="Path to validation-selected GRT+Image .safetensors file", + ) + parser.add_argument( + "--output_dir", + type=str, + default="inference_results", + help="Directory to save results", + ) + parser.add_argument( + "--sequences", + type=str, + nargs="+", + default=None, + help="Optional Smoke-Eval sequence subset (default: every valid sequence)", + ) + parser.add_argument( + "--debug", action="store_true", help="Run in debug mode (process only 1 batch)" + ) + args = parser.parse_args() + + # Load config + with open(args.config, "r") as f: + config = yaml.safe_load(f) + + # Initialize accelerator + accelerator = Accelerator(mixed_precision="fp16") + set_seed(config["training"].get("seed", 42)) + + # Create output directory (all ranks so DDP gather_dir can be created) + os.makedirs(args.output_dir, exist_ok=True) + + # Create model + accelerator.print("Creating naive GRT+Image model...") + model = GRTImageNaiveSmall( + resnet18_pretrained=config["model"].get("resnet18_pretrained", True), + image_height=config["data"].get("image_height", 288), + image_width=config["data"].get("image_width", 512), + ) + + # Safetensors files contain only the model state dictionary. + accelerator.print(f"Loading checkpoint from {args.checkpoint}") + model.load_state_dict(load_file(args.checkpoint, device="cpu"), strict=True) + + sequence_description = args.sequences if args.sequences else "all valid Smoke-Eval sequences" + accelerator.print(f"Inference sequences: {sequence_description}") + inference_loader = create_rice_dataloader( + root_dir=config["paths"]["smoke_eval_root"], + batch_size=config["training"]["batch_size"], + num_workers=0, + frame_skip=1, + sequences=args.sequences, + image_height=config["data"].get("image_height", 288), + image_width=config["data"].get("image_width", 512), + shuffle=False, + ) + + # Prepare model and dataloader + model, inference_loader = accelerator.prepare(model, inference_loader) + model.eval() + + # Dictionary to aggregate results by sequence: sequence_id -> list of (frame_idx, pred_depth) + results_by_sequence = defaultdict(list) + + accelerator.print("Starting inference...") + + with torch.no_grad(): + for batch in tqdm( + inference_loader, disable=not accelerator.is_local_main_process + ): + # Extract data + rsp_data = batch_radar_to_spectrum( + batch["radar_amplitude"], batch["radar_phase"] + ) + image = batch["image"] + sequences = batch["sequence"] + frame_indices = batch["frame_idx"] + + # Apply radar augmentation + rsp_data = translate_radar(rsp_data) + + # Forward pass + occupancy_pred_logits = model(rsp_data, image) # [B, 128, 256, 64] + + # Dequantize to depth [B, 1, 128, 256], values in [0, 1]. + pred_depth = dequantize_depth(occupancy_pred_logits) + pred_depth_np = ( + pred_depth.cpu().numpy().astype(np.float32) + ) # [B, 1, 128, 256] + + # Collect results (frame_idx, pred_depth per sample) + for i in range(len(sequences)): + seq_id = sequences[i] + f_idx = frame_indices[i].item() + # Store [1, 128, 256] per frame. + results_by_sequence[seq_id].append( + { + "frame_idx": f_idx, + "pred_depth": pred_depth_np[i], + } + ) + + if args.debug: + break + + # Single GPU: save directly (each frame seen once, no duplication) + # Multi-GPU: gather via files, merge with dedupe by frame_idx, then save + if accelerator.num_processes == 1: + if accelerator.is_main_process: + accelerator.print("Saving results (single process)...") + for seq_id, frames in tqdm( + results_by_sequence.items(), desc="Saving sequences" + ): + frames.sort(key=lambda x: x["frame_idx"]) + pred_depth_stack = np.stack([f["pred_depth"] for f in frames], axis=0) + pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 128, 256] + np.save( + os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"), + pred_depth_stack, + ) + accelerator.print( + f" {seq_id}: saved {pred_depth_stack.shape[0]} frames" + ) + accelerator.print(f"Processed {len(results_by_sequence)} sequences.") + accelerator.print(f"Results saved to {args.output_dir}") + else: + # DDP: gather results from all ranks via files, dedupe by frame_idx, save on main + accelerator.wait_for_everyone() + gather_dir = os.path.join(args.output_dir, "_gather") + os.makedirs(gather_dir, exist_ok=True) + rank = accelerator.process_index + rank_file = os.path.join(gather_dir, f"rank_{rank}_results.pkl") + with open(rank_file, "wb") as f: + pickle.dump(dict(results_by_sequence), f, protocol=pickle.HIGHEST_PROTOCOL) + accelerator.wait_for_everyone() + + if accelerator.is_main_process: + accelerator.print("Merging and deduplicating results from all ranks...") + merged_results = defaultdict(dict) # seq_id -> {frame_idx: pred_depth} + for r in range(accelerator.num_processes): + pkl_path = os.path.join(gather_dir, f"rank_{r}_results.pkl") + with open(pkl_path, "rb") as f: + rank_results = pickle.load(f) + for seq_id, frames in rank_results.items(): + for frame_data in frames: + f_idx = frame_data["frame_idx"] + if f_idx not in merged_results[seq_id]: + merged_results[seq_id][f_idx] = frame_data["pred_depth"] + os.remove(pkl_path) + + for seq_id, frame_dict in tqdm( + merged_results.items(), desc="Saving sequences" + ): + sorted_items = sorted(frame_dict.items(), key=lambda x: x[0]) + pred_depth_stack = np.stack([item[1] for item in sorted_items], axis=0) + pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 128, 256] + np.save( + os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"), + pred_depth_stack, + ) + accelerator.print( + f" {seq_id}: saved {pred_depth_stack.shape[0]} frames" + ) + if os.path.isdir(gather_dir) and not os.listdir(gather_dir): + os.rmdir(gather_dir) + accelerator.print(f"Processed {len(merged_results)} sequences.") + accelerator.print(f"Results saved to {args.output_dir}") + + accelerator.wait_for_everyone() + + +if __name__ == "__main__": + main() diff --git a/src/Baselines/grt_image/split.json b/src/Baselines/grt_image/split.json new file mode 100644 index 0000000000000000000000000000000000000000..6d4c8d97a758d12498d4d8983e44e8b32c1e77f8 --- /dev/null +++ b/src/Baselines/grt_image/split.json @@ -0,0 +1,16 @@ +{ + "test": [ + "Dell-1", + "Dell-2", + "Smoke-Dell-1", + "Smoke-Dell-2", + "brk-2", + "brk-3", + "Brk-b", + "brk-basement", + "Brk-stair", + "Smoke-brk-2", + "Smoke-brk-3", + "Smoke-brk-b" + ] +} \ No newline at end of file diff --git a/src/Baselines/radarcam-depth/data/SML_dataset.py b/src/Baselines/radarcam-depth/data/SML_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..1866a6c67af432926234596e61f06e5d0e4c66d3 --- /dev/null +++ b/src/Baselines/radarcam-depth/data/SML_dataset.py @@ -0,0 +1,83 @@ +import torch.utils.data +import numpy as np +import modules.midas.utils as utils +from PIL import Image + +def load_input_image(input_image_fp): + return utils.read_image(input_image_fp) + + +def load_sparse_depth(input_sparse_depth_fp): + input_sparse_depth = np.array(Image.open(input_sparse_depth_fp), dtype=np.float32) / 256.0 + input_sparse_depth[input_sparse_depth <= 0] = 0.0 + return input_sparse_depth + + +class SML_dataset(torch.utils.data.Dataset): + def __init__(self, + image_paths, + radar_paths, + gt_paths, + sparse_gt_paths, + rcnet_paths, + mono_pred_paths = None, + mono_ga_paths = None, + ): + + self.n_sample = len(image_paths) + + for paths in [image_paths, radar_paths, gt_paths, sparse_gt_paths, + rcnet_paths, mono_pred_paths, mono_ga_paths]: + if paths is not None: + assert len(paths) == self.n_sample + + self.image_paths = image_paths + self.radar_paths = radar_paths + self.gt_paths = gt_paths + self.sparse_gt_paths = sparse_gt_paths + self.rcnet_paths = rcnet_paths + self.mono_pred_paths = mono_pred_paths + self.mono_ga_paths = mono_ga_paths + + + def __getitem__(self, index): + image = load_input_image(self.image_paths[index]) + radar = load_sparse_depth(self.radar_paths[index]) + gt = load_sparse_depth(self.gt_paths[index]) + sparse_gt = load_sparse_depth(self.sparse_gt_paths[index]) + rcnet = load_sparse_depth(self.rcnet_paths[index]) + + image, radar, gt, sparse_gt, rcnet = [ + T.astype(np.float32) + for T in [image, radar, gt, sparse_gt, rcnet] + ] + + # Crop the image for ZJU dataset + if image.shape[0] == 720: + image = image[720 // 3: 720 // 4 * 3, :, :] + radar = radar[720 // 3: 720 // 4 * 3, :] + gt = gt[720 // 3: 720 // 4 * 3, :] + sparse_gt = sparse_gt[720 // 3: 720 // 4 * 3, :] + + + if self.mono_ga_paths is not None: + mono_pred = load_sparse_depth(self.mono_ga_paths[index]) + mono_pred = mono_pred.astype(np.float32) + if mono_pred.shape[0] == 720: + mono_pred = mono_pred[720 // 3: 720 // 4 * 3, :] + else: + mono_pred = None + + if self.mono_ga_paths is not None: + mono_ga = load_sparse_depth(self.mono_ga_paths[index]) + mono_ga = mono_ga.astype(np.float32) + if mono_ga.shape[0] == 720: + mono_ga = mono_ga[720 // 3: 720 // 4 * 3, :] + else: + mono_ga = None + + return image, mono_pred, radar, gt, sparse_gt, rcnet, mono_ga + + + def __len__(self): + return self.n_sample \ No newline at end of file diff --git a/src/Baselines/radarcam-depth/data/data_utils.py b/src/Baselines/radarcam-depth/data/data_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..8c177d7351693ccb2b29752822e1651787223a92 --- /dev/null +++ b/src/Baselines/radarcam-depth/data/data_utils.py @@ -0,0 +1,326 @@ +import numpy as np +from scipy.interpolate import LinearNDInterpolator +from PIL import Image +import matplotlib.pyplot as plt + + + +def load_data_path(root, file_name_txt, data_type): + with open(file_name_txt, 'r') as f: + data_path = f.readlines() + data_path = [root + x.strip() + data_type for x in data_path] + return data_path + + +def load_data_path_nu(root, name_list, data_type): + data_path = [root + x.strip() + data_type for x in name_list] + return data_path + + +def read_paths(filepath): + ''' + Reads a newline delimited file containing paths + + Arg(s): + filepath : str + path to file to be read + Return: + list[str] : list of paths + ''' + + path_list = [] + with open(filepath) as f: + while True: + path = f.readline().rstrip('\n') + + # If there was nothing to read + if path == '': + break + + path_list.append(path) + + return path_list + + +def write_paths(filepath, paths): + ''' + Stores line delimited paths into file + + Arg(s): + filepath : str + path to file to save paths + paths : list[str] + paths to write into file + ''' + + with open(filepath, 'w') as o: + for idx in range(len(paths)): + o.write(paths[idx] + '\n') + + +def load_image(path, normalize=False, data_format='HWC'): + ''' + Loads an RGB image + + Arg(s): + path : str + path to RGB image + normalize : bool + if set, then normalize image between [0, 1] + data_format : str + 'CHW', or 'HWC' + Returns: + numpy[float32] : H x W x C or C x H x W image + ''' + + # Load image + image = Image.open(path).convert('RGB') + + # Convert to numpy + image = np.asarray(image, np.float32) + + if data_format == 'HWC': + pass + elif data_format == 'CHW': + image = np.transpose(image, (2, 0, 1)) + else: + raise ValueError('Unsupported data format: {}'.format(data_format)) + + # Normalize + image = image / 255.0 if normalize else image #255.0 + + return image + + + +def load_depth(path, multiplier=256.0, data_format='HW'): + ''' + Loads a depth map from a 16-bit PNG file + + Arg(s): + path : str + path to 16-bit PNG file + multiplier : float + multiplier for encoding float as 16/32 bit unsigned integer + data_format : str + HW, CHW, HWC + Returns: + numpy[float32] : depth map + ''' + + # Loads depth map from 16-bit PNG file + z = np.array(Image.open(path), dtype=np.float32) + + # Assert 16-bit (not 8-bit) depth map + z = z / multiplier + z[z <= 0] = 0.0 + + if data_format == 'HW': + pass + elif data_format == 'CHW': + z = np.expand_dims(z, axis=0) + elif data_format == 'HWC': + z = np.expand_dims(z, axis=-1) + else: + raise ValueError('Unsupported data format: {}'.format(data_format)) + + return z + + +def save_depth(z, path, multiplier=256.0): + ''' + Saves a depth map to a 16-bit PNG file + + Arg(s): + z : numpy[float32] + depth map + path : str + path to store depth map + multiplier : float + multiplier for encoding float as 16/32 bit unsigned integer + ''' + + z = np.uint32(z * multiplier) + z = Image.fromarray(z, mode='I') + z.save(path) + + +def save_color_depth(z, path): + ''' + Saves a color depth map to a 16-bit PNG file + + Arg(s): + z : numpy[float32] + depth map + path : str + path to store depth map + multiplier : float + multiplier for encoding float as 16/32 bit unsigned integer + ''' + + # Normalize depth map to the range [0, 1] + z_normalized = (z - np.min(z)) / (np.max(z) - np.min(z)) + + # Convert depth map to color + # colormap = plt.cm.jet # Choose a colormap (e.g., jet) + colormap = plt.cm.viridis + z_color = colormap(z_normalized) + + # Scale color values to the range [0, 255] and convert to uint8 + z_color = np.uint8(z_color * 255) + + # Save color depth map as an image + image = Image.fromarray(z_color) + image.save(path) + + +def load_response(path, multiplier=2**14, data_format='HW'): + ''' + Loads a response map from a 16-bit PNG file + + Arg(s): + path : str + path to 16-bit PNG file + multiplier : float + multiplier for encoding float as 16/32 bit unsigned integer + data_format : str + HW, CHW, HWC + Returns: + numpy[float32] : response map + ''' + + # Loads response map from 16-bit PNG file + response = np.array(Image.open(path), dtype=np.float32) + + # Convert using encodering multiplier + response = response / multiplier + + if data_format == 'HW': + pass + elif data_format == 'CHW': + response = np.expand_dims(response, axis=0) + elif data_format == 'HWC': + response = np.expand_dims(response, axis=-1) + else: + raise ValueError('Unsupported data format: {}'.format(data_format)) + + return response + + +def save_response(response, path, multiplier=2**14): + ''' + Saves a response map to a 16-bit PNG file + + Arg(s): + response : numpy[float32] + depth map + path : str + path to store depth map + multiplier : float + multiplier for encoding float as 16/32 bit unsigned integer + ''' + + response = np.uint32(response * multiplier) + response = Image.fromarray(response, mode='I') + response.save(path) + + +def interpolate_depth(depth_map, validity_map, log_space=False): + ''' + Interpolate sparse depth with barycentric coordinates + + Arg(s): + depth_map : np.float32 + H x W depth map + validity_map : np.float32 + H x W depth map + log_space : bool + if set then produce in log space + Returns: + np.float32 : H x W interpolated depth map + ''' + + assert depth_map.ndim == 2 and validity_map.ndim == 2 + + rows, cols = depth_map.shape + data_row_idx, data_col_idx = np.where(validity_map) + depth_values = depth_map[data_row_idx, data_col_idx] + + # Perform linear interpolation in log space + if log_space: + depth_values = np.log(depth_values) + + interpolator = LinearNDInterpolator( + # points=Delaunay(np.stack([data_row_idx, data_col_idx], axis=1).astype(np.float32)), + points=np.stack([data_row_idx, data_col_idx], axis=1), + values=depth_values, + fill_value=0 if not log_space else np.log(1e-3)) + + query_row_idx, query_col_idx = np.meshgrid( + np.arange(rows), np.arange(cols), indexing='ij') + + query_coord = np.stack( + [query_row_idx.ravel(), query_col_idx.ravel()], axis=1) + + Z = interpolator(query_coord).reshape([rows, cols]) + + if log_space: + Z = np.exp(Z) + Z[Z < 1e-1] = 0.0 + + return Z + + +def interpolate_depth_ZJU(depth_map, validity_map=None, log_space=False, window_size=12): + ''' + Interpolate sparse depth with barycentric coordinates + Args: + depth_map : np.float32 + H x W depth map + validity_map : np.float32 + H x W depth map + log_space : bool + if set then produce in log space + window_size : int + size of the window for checking validity + Returns: + np.float32 : H x W interpolated depth map + ''' + assert depth_map.ndim == 2 + if validity_map is None: + validity_map = depth_map > 0.0 + rows, cols = depth_map.shape + data_row_idx, data_col_idx = np.where(validity_map) + depth_values = depth_map[data_row_idx, data_col_idx] + # Perform linear interpolation in log space + if log_space: + depth_values = np.log(depth_values) + interpolator = LinearNDInterpolator( + points=np.stack([data_row_idx, data_col_idx], axis=1), + values=depth_values, + fill_value=0 if not log_space else np.log(1e-3)) + query_row_idx, query_col_idx = np.meshgrid(np.arange(rows), np.arange(cols), indexing='ij') + Z = np.zeros_like(depth_map) + + # Create window indices for each query point + query_indices = np.stack([query_row_idx.ravel(), query_col_idx.ravel()], axis=1) + window_indices = np.indices((window_size, window_size)).reshape(2, -1) - window_size // 2 + + # Calculate window indices for each query point + window_row_indices = np.clip(query_indices[:, 0, None] + window_indices[0], 0, rows - 1) + window_col_indices = np.clip(query_indices[:, 1, None] + window_indices[1], 0, cols - 1) + + # Get window values and check validity + window_values = depth_map[window_row_indices, window_col_indices] + valid_indices = np.any(window_values > 0, axis=1) + + # Interpolate for valid query points + valid_query_indices = np.where(valid_indices)[0] + valid_query_coords = query_indices[valid_query_indices] + Z.ravel()[valid_query_indices] = interpolator(valid_query_coords) + + if log_space: + Z = np.exp(Z) + Z[Z < 1e-1] = 0.0 + + return Z \ No newline at end of file diff --git a/src/Baselines/radarcam-depth/data/datasets.py b/src/Baselines/radarcam-depth/data/datasets.py new file mode 100644 index 0000000000000000000000000000000000000000..9f946a0a6379c57e6ca01656ff6a61a53f585342 --- /dev/null +++ b/src/Baselines/radarcam-depth/data/datasets.py @@ -0,0 +1,392 @@ +import torch +import torch.utils.data +from torch.utils.data import Dataset +import numpy as np +import data.data_utils as data_utils +import random +import os +from PIL import Image +from data.data_utils import load_depth + + +def random_sample(T): + ''' + Arg(s): + T : numpy[float32] + C x N array + Returns: + numpy[float32] : random sample from T + ''' + + index = np.random.randint(0, T.shape[0]) + return T[index, :] + + +def random_crop(inputs, shape, crop_type=['none']): + ''' + Apply crop to inputs e.g. images, depth + + Arg(s): + inputs : list[numpy[float32]] + list of numpy arrays e.g. images, depth, and validity maps + shape : list[int] + shape (height, width) to crop inputs + crop_type : str + none, horizontal, vertical, anchored, top, bottom, left, right, center + Return: + list[numpy[float32]] : list of cropped inputs + ''' + + n_height, n_width = shape + _, o_height, o_width = inputs[0].shape + + # Get delta of crop and original height and width + + d_height = o_height - n_height + d_width = o_width - n_width + + # By default, perform center crop + y_start = d_height // 2 + x_start = d_width // 2 + + # If left alignment, then set starting height to 0 + if 'left' in crop_type: + x_start = 0 + + # If right alignment, then set starting height to right most position + elif 'right' in crop_type: + x_start = d_width + + elif 'horizontal' in crop_type: + + # Select from one of the pre-defined anchored locations + if 'anchored' in crop_type: + # Create anchor positions + crop_anchors = [ + 0.0, 0.50, 1.0 + ] + + widths = [ + anchor * d_width for anchor in crop_anchors + ] + x_start = int(widths[np.random.randint(low=0, high=len(widths))]) + + # Randomly select a crop location + else: + x_start = np.random.randint(low=0, high=d_width) + + # If top alignment, then set starting height to 0 + if 'top' in crop_type: + y_start = 0 + + # If bottom alignment, then set starting height to lowest position + elif 'bottom' in crop_type: + y_start = d_height + + elif 'vertical' in crop_type and np.random.rand() <= 0.30: + + # Select from one of the pre-defined anchored locations + if 'anchored' in crop_type: + # Create anchor positions + crop_anchors = [ + 0.0, 0.50, 1.0 + ] + + heights = [ + anchor * d_height for anchor in crop_anchors + ] + y_start = int(heights[np.random.randint(low=0, high=len(heights))]) + + # Randomly select a crop location + else: + y_start = np.random.randint(low=0, high=d_height) + + elif 'center' in crop_type: + pass + + # Crop each input into (n_height, n_width) + y_end = y_start + n_height + x_end = x_start + n_width + + outputs = [ + T[:, y_start:y_end, x_start:x_end] for T in inputs + ] + + return outputs + + +class RCNetTrainingDataset(torch.utils.data.Dataset): + ''' + Dataset for fetching: + (1) image + (2) radar point + (3) ground truth + (4) bounding boxes for the points + (5) image crops for summary part of the code + + Arg(s): + image_paths : list[str] + paths to images + radar_paths : list[str] + paths to radar points + ground_truth_paths : list[str] + paths to ground truth depth maps + crop_width : int + width of crop centered at the radar point + total_points_sampled: int + total number of points sampled from the total radar points available. Repeats the same points multiple times if total points in the frame is less than total sampled points + sample_probability_of_lidar: int + randomly sample lidar with this probability and add noise to it instead of using radar points + min_radar_depth_m: float + minimum depth accepted for synthetic radar sampling + max_radar_depth_m: float + maximum depth accepted for synthetic radar sampling + ''' + + def __init__(self, + image_paths, + radar_paths, + ground_truth_paths, + patch_size, + total_points_sampled, + sample_probability_of_lidar, + min_radar_depth_m=0.05, + max_radar_depth_m=11.2): + + self.n_sample = len(image_paths) + + assert self.n_sample == len(ground_truth_paths) + assert self.n_sample == len(radar_paths) + + self.image_paths = image_paths + self.radar_paths = radar_paths + self.ground_truth_paths = ground_truth_paths + + self.patch_size = patch_size + self.pad_size_x = patch_size[1] // 2 + self.padding = ((0, 0), (0, 0), (self.pad_size_x, self.pad_size_x)) + + self.data_format = 'CHW' + self.total_points_sampled = total_points_sampled + self.sample_probability_of_lidar = sample_probability_of_lidar + self.min_radar_depth_m = min_radar_depth_m + self.max_radar_depth_m = max_radar_depth_m + + def __getitem__(self, index): + + # Load image + image = data_utils.load_image( + self.image_paths[index], + normalize=False, + data_format=self.data_format) + + height, width = image.shape[1:] + if height == 720: # ZJU dataset + image = image[:, 720 // 3: 720 // 4 * 3, :] + + image = np.pad( + image, + pad_width=self.padding, + mode='edge') + + # Load radar points N x 3 + radar_points = np.load(self.radar_paths[index]) + + if height == 720: + radar_points = radar_points[radar_points[:, 1] < 720 // 4 * 3] + radar_points[:, 1] = radar_points[:, 1] - 720 // 3 + radar_points = radar_points[radar_points[:, 1] >= 0] + + if radar_points.ndim == 1: + # Only one point (,3), expand to 1 x 3 + radar_points = np.expand_dims(radar_points, axis=0) + + # Store bounding boxes for all radar points + bounding_boxes_list = [] + + # randomly sample radar points to output + if radar_points.shape[0] <= self.total_points_sampled: + radar_points = np.repeat(radar_points, 100, axis=0) + random_idx = np.random.randint(radar_points.shape[0], size=self.total_points_sampled) + radar_points = radar_points[random_idx, :] + + # Load ground truth depth + ground_truth = data_utils.load_depth( + self.ground_truth_paths[index], + data_format=self.data_format) + + if height == 720: + ground_truth = ground_truth[:, 720 // 3: 720 // 4 * 3] + + if random.random() < self.sample_probability_of_lidar: + ground_truth_for_sampling = np.copy(ground_truth) + ground_truth_for_sampling = ground_truth_for_sampling.squeeze() + valid_lidar = np.isfinite(ground_truth_for_sampling) + valid_lidar &= ground_truth_for_sampling >= self.min_radar_depth_m + valid_lidar &= ground_truth_for_sampling <= self.max_radar_depth_m + idx_lidar_samples = np.where(valid_lidar) + n_lidar_samples = len(idx_lidar_samples[0]) + + if n_lidar_samples > 0: + # Keep the fixed point count required by RC-Net. Replacement + # handles frames with fewer valid GT pixels than requested. + if n_lidar_samples >= self.total_points_sampled: + random_indices = random.sample( + range(n_lidar_samples), self.total_points_sampled + ) + else: + random_indices = np.random.choice( + n_lidar_samples, + size=self.total_points_sampled, + replace=True, + ) + + points_x = idx_lidar_samples[1][random_indices] + points_y = idx_lidar_samples[0][random_indices] + points_z = ground_truth_for_sampling[points_y, points_x] + + noise_for_fake_radar_x = np.random.normal(0, 25, radar_points.shape[0]) + noise_for_fake_radar_z = np.random.uniform(low=0.0, high=0.4, size=radar_points.shape[0]) + + fake_radar_points = np.copy(radar_points) + fake_radar_points[:, 0] = points_x + noise_for_fake_radar_x + fake_radar_points[:, 0] = np.clip(fake_radar_points[:, 0], 0, ground_truth_for_sampling.shape[1]) + fake_radar_points[:, 2] = points_z + noise_for_fake_radar_z + # we keep the y as the same it is since it is erroneous + + # convert x and y indices back to int after adding noise + fake_radar_points[:, 0] = fake_radar_points[:, 0].astype(int) + fake_radar_points[:, 1] = fake_radar_points[:, 1].astype(int) + + radar_points = np.copy(fake_radar_points) + + # get the shifted radar points after padding + for radar_point_idx in range(0, radar_points.shape[0]): + # Set radar point to the center of the patch + radar_points[radar_point_idx, 0] = radar_points[radar_point_idx, 0] + self.pad_size_x + + bounding_box = [0, 0, 0, 0] + bounding_box[0] = radar_points[radar_point_idx, 0] - self.pad_size_x + bounding_box[1] = 0 + bounding_box[2] = radar_points[radar_point_idx, 0] + self.pad_size_x + bounding_box[3] = self.patch_size[0] + bounding_boxes_list.append(np.asarray(bounding_box)) + + ground_truth = np.pad( + ground_truth, + pad_width=self.padding, + mode='constant', + constant_values=0) + + ground_truth_crops = [] + + # Crop image and ground truth + for radar_point_idx in range(0, radar_points.shape[0]): + start_x = int(radar_points[radar_point_idx, 0] - self.pad_size_x) + end_x = int(radar_points[radar_point_idx, 0] + self.pad_size_x) + start_y = image.shape[-2] - self.patch_size[0] + + ground_truth_cropped = ground_truth[:, start_y:, start_x:end_x] + ground_truth_crops.append(ground_truth_cropped) + + image = image[:, start_y:, ...] + + ground_truth = np.asarray(ground_truth_crops) + + # Convert to float32 + image, radar_points, ground_truth = [ + T.astype(np.float32) + for T in [image, radar_points, ground_truth] + ] + + bounding_boxes_list = [T.astype(np.float32) for T in bounding_boxes_list] + + bounding_boxes_list = np.stack(bounding_boxes_list, axis=0) + + return image, radar_points, bounding_boxes_list, ground_truth + + def __len__(self): + return self.n_sample + + +class RCNetInferenceDataset(torch.utils.data.Dataset): + ''' + Dataset for fetching: + (1) image + (2) radar points + (3) ground truth (if available) + + Arg(s): + image_paths : list[str] + paths to images + radar_paths : list[str] + paths to radar points + ground_truth_paths : list[str] + paths to ground truth paths + ''' + + def __init__(self, image_paths, radar_paths, ground_truth_paths=None): + + self.n_sample = len(image_paths) + + assert self.n_sample == len(radar_paths) + + self.image_paths = image_paths + self.radar_paths = radar_paths + + if ground_truth_paths is not None and None not in ground_truth_paths: + assert self.n_sample == len(ground_truth_paths) + self.ground_truth_available = True + else: + self.ground_truth_available = False + + self.ground_truth_paths = ground_truth_paths + + self.data_format = 'CHW' + + def __getitem__(self, index): + + # Load image + image = data_utils.load_image( + self.image_paths[index], + normalize=False, + data_format=self.data_format) + + height, width = image.shape[1:] + if height == 720: # ZJU dataset + image = image[:, 720 // 3: 720 // 4 * 3, :] + + # Load radar points N x 3 + radar_points = np.load(self.radar_paths[index]) + + if height == 720: + radar_points = radar_points[radar_points[:, 1] < 720 // 4 * 3] + radar_points[:, 1] = radar_points[:, 1] - 720 // 3 + radar_points = radar_points[radar_points[:, 1] >= 0] + + if radar_points.ndim == 1: + # Expand to 1 x 3 + radar_points = np.expand_dims(radar_points, axis=0) + + inputs = [image, radar_points] + + if self.ground_truth_available: + # Load ground truth depth + ground_truth = data_utils.load_depth( + self.ground_truth_paths[index], + data_format=self.data_format) + if height == 720: + ground_truth = ground_truth[:, 720 // 3: 720 // 4 * 3] + + inputs.append(ground_truth) + + # Convert to float32 + inputs = [ + T.astype(np.float32) + for T in inputs + ] + + return inputs + + def __len__(self): + return self.n_sample diff --git a/src/Baselines/radarcam-depth/linear_attention.py b/src/Baselines/radarcam-depth/linear_attention.py new file mode 100644 index 0000000000000000000000000000000000000000..7f4a5646f99f30fffd73013fb2650bebaa95ec06 --- /dev/null +++ b/src/Baselines/radarcam-depth/linear_attention.py @@ -0,0 +1,184 @@ +import torch +from torch.nn import Module, Dropout +import torch.nn as nn +import copy + + +def elu_feature_map(x): + return torch.nn.functional.elu(x) + 1 + + + +class LinearAttention(Module): + def __init__(self, eps=1e-6): + super().__init__() + self.feature_map = elu_feature_map + self.eps = eps + + def forward(self, queries, keys, values, q_mask=None, kv_mask=None): + """ Multi-Head linear attention proposed in "Transformers are RNNs" + Args: + queries: [N, L, H, D] + keys: [N, S, H, D] + values: [N, S, H, D] + q_mask: [N, L] + kv_mask: [N, S] + Returns: + queried_values: (N, L, H, D) + """ + Q = self.feature_map(queries) + K = self.feature_map(keys) + + # set padded position to zero + if q_mask is not None: + Q = Q * q_mask[:, :, None, None] + if kv_mask is not None: + K = K * kv_mask[:, :, None, None] + values = values * kv_mask[:, :, None, None] + + v_length = values.size(1) + values = values / v_length # prevent fp16 overflow + KV = torch.einsum("nshd,nshv->nhdv", K, values) # (S,D)' @ S,V + Z = 1 / (torch.einsum("nlhd,nhd->nlh", Q, K.sum(dim=1)) + self.eps) + queried_values = torch.einsum("nlhd,nhdv,nlh->nlhv", Q, KV, Z) * v_length + + return queried_values.contiguous() + + + +class FullAttention(Module): + def __init__(self, use_dropout=False, attention_dropout=0.1): + super().__init__() + self.use_dropout = use_dropout + self.dropout = Dropout(attention_dropout) + + def forward(self, queries, keys, values, q_mask=None, kv_mask=None): + """ Multi-head scaled dot-product attention, a.k.a full attention. + Args: + queries: [N, L, H, D] + keys: [N, S, H, D] + values: [N, S, H, D] + q_mask: [N, L] + kv_mask: [N, S] + Returns: + queried_values: (N, L, H, D) + """ + + # Compute the unnormalized attention and apply the masks + QK = torch.einsum("nlhd,nshd->nlsh", queries, keys) + if kv_mask is not None: + QK.masked_fill_(~(q_mask[:, :, None, None] * kv_mask[:, None, :, None]), float('-inf')) + + # Compute the attention and the weighted average + softmax_temp = 1. / queries.size(3)**.5 # sqrt(D) + A = torch.softmax(softmax_temp * QK, dim=2) + if self.use_dropout: + A = self.dropout(A) + + queried_values = torch.einsum("nlsh,nshd->nlhd", A, values) + + return queried_values.contiguous() + + + +class LoFTREncoderLayer(nn.Module): + def __init__(self, + d_model, + nhead, + attention='linear'): + super(LoFTREncoderLayer, self).__init__() + + self.dim = d_model // nhead + self.nhead = nhead + + # multi-head attention + self.q_proj = nn.Linear(d_model, d_model, bias=False) + self.k_proj = nn.Linear(d_model, d_model, bias=False) + self.v_proj = nn.Linear(d_model, d_model, bias=False) + self.attention = LinearAttention() if attention == 'linear' else FullAttention() + self.merge = nn.Linear(d_model, d_model, bias=False) + + # feed-forward network + self.mlp = nn.Sequential( + nn.Linear(d_model*2, d_model*2, bias=False), + nn.ReLU(True), + nn.Linear(d_model*2, d_model, bias=False), + ) + + # norm and dropout + self.norm1 = nn.LayerNorm(d_model) + self.norm2 = nn.LayerNorm(d_model) + + def forward(self, x, source, x_mask=None, source_mask=None): + """ + Args: + x (torch.Tensor): [N, L, C] + source (torch.Tensor): [N, S, C] + x_mask (torch.Tensor): [N, L] (optional) + source_mask (torch.Tensor): [N, S] (optional) + """ + bs = x.size(0) + query, key, value = x, source, source + + # multi-head attention + query = self.q_proj(query).view(bs, -1, self.nhead, self.dim) # [N, L, (H, D)] + key = self.k_proj(key).view(bs, -1, self.nhead, self.dim) # [N, S, (H, D)] + value = self.v_proj(value).view(bs, -1, self.nhead, self.dim) + message = self.attention(query, key, value, q_mask=x_mask, kv_mask=source_mask) # [N, L, (H, D)] + message = self.merge(message.view(bs, -1, self.nhead*self.dim)) # [N, L, C] + message = self.norm1(message) + + # feed-forward network + message = self.mlp(torch.cat([x, message], dim=2)) + message = self.norm2(message) + + return x + message + + + +class LocalFeatureTransformer(nn.Module): + """A Local Feature Transformer (LoFTR) module.""" + + def __init__(self, type, n_layers=1, d_model=256, nhead=8, attention='linear'): + super(LocalFeatureTransformer, self).__init__() + + self.d_model = d_model + self.nhead = nhead + self.layer_names = type * n_layers + self.attention = attention + encoder_layer = LoFTREncoderLayer(self.d_model, self.nhead, self.attention) + + self.layers = nn.ModuleList([copy.deepcopy(encoder_layer) for _ in range(len(self.layer_names))]) + self._reset_parameters() + + def _reset_parameters(self): + for p in self.parameters(): + if p.dim() > 1: + nn.init.xavier_uniform_(p) + + def forward(self, feat0, feat1, mask0=None, mask1=None): + """ + Args: + feat0 (torch.Tensor): [N, L, C] + feat1 (torch.Tensor): [N, S, C] + mask0 (torch.Tensor): [N, L] (optional) + mask1 (torch.Tensor): [N, S] (optional) + """ + + assert self.d_model == feat0.size(2), "the feature number of src and transformer must be equal" + + for layer, name in zip(self.layers, self.layer_names): + # if name == 'self0': + # feat0 = layer(feat0, feat0, mask0, mask0) + # elif name == 'self1': + # feat1 = layer(feat1, feat1, mask1, mask1) + if name == 'self': + feat0 = layer(feat0, feat0, mask0, mask0) + feat1 = layer(feat1, feat1, mask1, mask1) + elif name == 'cross': + feat0 = layer(feat0, feat1, mask0, mask1) + feat1 = layer(feat1, feat0, mask1, mask0) + else: + raise KeyError + + return feat0, feat1 \ No newline at end of file diff --git a/src/Baselines/radarcam-depth/modules/estimator.py b/src/Baselines/radarcam-depth/modules/estimator.py new file mode 100644 index 0000000000000000000000000000000000000000..331c9f4c5b5a10ff6d9e5e4b616370923695bbd6 --- /dev/null +++ b/src/Baselines/radarcam-depth/modules/estimator.py @@ -0,0 +1,188 @@ +import numpy as np +import time +from scipy.optimize import minimize_scalar + +def compute_scale_and_shift_ls(prediction, target, mask): + # tuple specifying with axes to sum + sum_axes = (0, 1) + + # system matrix: A = [[a_00, a_01], [a_10, a_11]] + a_00 = np.sum(mask * prediction * prediction, sum_axes) + a_01 = np.sum(mask * prediction, sum_axes) + a_11 = np.sum(mask, sum_axes) + + # right hand side: b = [b_0, b_1] + b_0 = np.sum(mask * prediction * target, sum_axes) + b_1 = np.sum(mask * target, sum_axes) + + # solution: x = A^-1 . b = [[a_11, -a_01], [-a_10, a_00]] / (a_00 * a_11 - a_01 * a_10) . b + x_0 = np.zeros_like(b_0) + x_1 = np.zeros_like(b_1) + + det = a_00 * a_11 - a_01 * a_01 + # A needs to be a positive definite matrix. + valid = det > 0 + + x_0[valid] = (a_11[valid] * b_0[valid] - a_01[valid] * b_1[valid]) / det[valid] + x_1[valid] = (-a_01[valid] * b_0[valid] + a_00[valid] * b_1[valid]) / det[valid] + + return x_0, x_1 + + + +def compute_scale_and_shift_ransac(prediction, target, mask, + num_iterations, sample_size, + inlier_threshold, inlier_ratio_threshold): + # start = time.time() + best_scale = 0.0 + best_shift = 0.0 + best_inlier_count = 0 + + valid_indices = np.where(mask) + valid_count = len(valid_indices[0]) + # print('valid_count: ', valid_count) + + for _ in range(num_iterations): + if valid_count < sample_size: + break + + # Randomly sample from valid indices + indices = np.random.choice(valid_count, size=sample_size, replace=False) + mask_sample = np.zeros_like(mask) + mask_sample[valid_indices[0][indices], valid_indices[1][indices]] = 1 + + # Calculate x_0 and x_1 for the sampled data + sum_axes = (0, 1) + a_00 = np.sum(mask_sample * prediction * prediction, sum_axes) + a_01 = np.sum(mask_sample * prediction, sum_axes) + a_11 = np.sum(mask_sample, sum_axes) + b_0 = np.sum(mask_sample * prediction * target, sum_axes) + b_1 = np.sum(mask_sample * target, sum_axes) + det = a_00 * a_11 - a_01 * a_01 + valid = det > 0 + x_0 = np.zeros_like(b_0) + x_1 = np.zeros_like(b_1) + x_0[valid] = (a_11[valid] * b_0[valid] - a_01[valid] * b_1[valid]) / det[valid] + x_1[valid] = (-a_01[valid] * b_0[valid] + a_00[valid] * b_1[valid]) / det[valid] + + # Calculate residuals and count inliers + residuals = np.abs(mask * prediction * x_0 + x_1 - mask * target) + residuals = residuals[mask] + + inlier_count = np.sum(residuals < inlier_threshold) + + # Update best model if current model has more inliers + if inlier_count > best_inlier_count: + best_scale = x_0 + best_shift = x_1 + best_inlier_count = inlier_count + inlier_ratio = inlier_count / valid_count + if inlier_ratio > inlier_ratio_threshold: + break + + print('best_inlier_count: ', best_inlier_count) + print('inlier_ratio: ', best_inlier_count / valid_count) + # print('best_scale: ', best_scale) + # print('best_shift: ', best_shift) + # print('time', time.time() - start) + return best_scale, best_shift + + + +class LeastSquaresEstimator(object): + def __init__(self, estimate, target, valid): + self.estimate = estimate + self.target = target + self.valid = valid + + # to be computed + self.scale = 1.0 + self.shift = 0.0 + self.output = None + + def compute_scale_and_shift_ran(self, + num_iterations=60, sample_size=5, + inlier_threshold=0.02, inlier_ratio_threshold=0.8): + self.scale, self.shift = compute_scale_and_shift_ransac(self.estimate, self.target, self.valid, + num_iterations, sample_size, + inlier_threshold, inlier_ratio_threshold) + + def compute_scale_and_shift(self): + self.scale, self.shift = compute_scale_and_shift_ls(self.estimate, self.target, self.valid) + + + def apply_scale_and_shift(self): + self.output = self.estimate * self.scale + self.shift + + def clamp_min_max(self, clamp_min=None, clamp_max=None): + if clamp_min is not None: + if clamp_min > 0: + clamp_min_inv = 1.0/clamp_min + self.output[self.output > clamp_min_inv] = clamp_min_inv + assert np.max(self.output) <= clamp_min_inv + else: # divide by zero, so skip + pass + if clamp_max is not None: + clamp_max_inv = 1.0/clamp_max + self.output[self.output < clamp_max_inv] = clamp_max_inv + + + +def objective_function(x_0, prediction, target, mask): + # Calculate x_0 * prediction + x_0_prediction = x_0 * prediction + # Calculate the error between x_0 * prediction and target, using the mask + error = np.sum(mask * abs(x_0_prediction - target)) + return error + + + +class Optimizer(object): + def __init__(self, estimate, target, valid, depth_type): + self.estimate = estimate + self.target = target + self.valid = valid + self.depth_type = depth_type + # to be computed + self.scale = 1.0 + self.output = None + + def optimize_scale(self): + if self.depth_type == 'inv': + bounds = (0.0003, 0.01) + else: + bounds = (0.5, 1.6) # pos + + # Minimize the objective function using scipy.optimize.minimize_scalar + result = minimize_scalar( + objective_function, args=(self.estimate, self.target, self.valid), + bounds=bounds + ) + + # Extract the optimized x_0 value from the result + optimized_x_0 = result.x + self.scale = optimized_x_0 + + def apply_scale(self): + self.output = self.estimate * self.scale + + def clamp_min_max(self, clamp_min=None, clamp_max=None): + if clamp_min is not None: + if clamp_min > 0: + clamp_min_inv = 1.0/clamp_min + self.output[self.output > clamp_min_inv] = clamp_min_inv + assert np.max(self.output) <= clamp_min_inv + else: # divide by zero, so skip + pass + if clamp_max is not None: + clamp_max_inv = 1.0/clamp_max + self.output[self.output < clamp_max_inv] = clamp_max_inv + + def clamp_min_max_pos(self, clamp_min=None, clamp_max=None): + if clamp_min is not None: + if clamp_min >= 0: + self.output[self.output < clamp_min] = clamp_min + else: + pass + if clamp_max is not None: + self.output[self.output > clamp_max] = clamp_max \ No newline at end of file diff --git a/src/Baselines/radarcam-depth/modules/midas/base_model.py b/src/Baselines/radarcam-depth/modules/midas/base_model.py new file mode 100644 index 0000000000000000000000000000000000000000..774297eef0d2d83ec2e224baedd9daf5f16dc725 --- /dev/null +++ b/src/Baselines/radarcam-depth/modules/midas/base_model.py @@ -0,0 +1,12 @@ +import torch +from safetensors.torch import load_file + + +class BaseModel(torch.nn.Module): + def load(self, path): + """Load model from file. + + Args: + path (str): file path + """ + self.load_state_dict(load_file(path, device="cpu"), strict=True) diff --git a/src/Baselines/radarcam-depth/modules/midas/blocks.py b/src/Baselines/radarcam-depth/modules/midas/blocks.py new file mode 100644 index 0000000000000000000000000000000000000000..0158524e217d3f7b7abe9f4aebd071f6fbb5ea89 --- /dev/null +++ b/src/Baselines/radarcam-depth/modules/midas/blocks.py @@ -0,0 +1,197 @@ +import torch +import torch.nn as nn + +def _make_encoder(backbone, features, use_pretrained, groups=1, expand=False, exportable=True): + if backbone == "efficientnet_lite3": + pretrained = _make_pretrained_efficientnet_lite3(use_pretrained, exportable=exportable) + scratch = _make_scratch([32, 48, 136, 384], features, groups=groups, expand=expand) # efficientnet_lite3 + else: + print(f"Backbone '{backbone}' not implemented") + assert False + + return pretrained, scratch + + +def _make_scratch(in_shape, out_shape, groups=1, expand=False): + scratch = nn.Module() + + out_shape1 = out_shape + out_shape2 = out_shape + out_shape3 = out_shape + out_shape4 = out_shape + if expand==True: + out_shape1 = out_shape + out_shape2 = out_shape*2 + out_shape3 = out_shape*4 + out_shape4 = out_shape*8 + + scratch.layer1_rn = nn.Conv2d( + in_shape[0], out_shape1, kernel_size=3, stride=1, padding=1, bias=False, groups=groups + ) + scratch.layer2_rn = nn.Conv2d( + in_shape[1], out_shape2, kernel_size=3, stride=1, padding=1, bias=False, groups=groups + ) + scratch.layer3_rn = nn.Conv2d( + in_shape[2], out_shape3, kernel_size=3, stride=1, padding=1, bias=False, groups=groups + ) + scratch.layer4_rn = nn.Conv2d( + in_shape[3], out_shape4, kernel_size=3, stride=1, padding=1, bias=False, groups=groups + ) + + return scratch + + +def _make_pretrained_efficientnet_lite3(use_pretrained, exportable=False): + efficientnet = torch.hub.load( + "rwightman/gen-efficientnet-pytorch", + "tf_efficientnet_lite3", + pretrained=use_pretrained, + exportable=exportable, + trust_repo=True, + ) + return _make_efficientnet_backbone(efficientnet) + + +def _make_efficientnet_backbone(effnet): + pretrained = nn.Module() + + pretrained.layer1 = nn.Sequential( + effnet.conv_stem, effnet.bn1, effnet.act1, *effnet.blocks[0:2] + ) + pretrained.layer2 = nn.Sequential(*effnet.blocks[2:3]) + pretrained.layer3 = nn.Sequential(*effnet.blocks[3:5]) + pretrained.layer4 = nn.Sequential(*effnet.blocks[5:9]) + + return pretrained + + +class ResidualConvUnit_custom(nn.Module): + """Residual convolution module. + """ + + def __init__(self, features, activation, bn): + """Init. + + Args: + features (int): number of features + """ + super().__init__() + + self.bn = bn + + self.groups=1 + + self.conv1 = nn.Conv2d( + features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups + ) + + self.conv2 = nn.Conv2d( + features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups + ) + + if self.bn==True: + self.bn1 = nn.BatchNorm2d(features) + self.bn2 = nn.BatchNorm2d(features) + + self.activation = activation + + self.skip_add = nn.quantized.FloatFunctional() + + def forward(self, x): + """Forward pass. + + Args: + x (tensor): input + + Returns: + tensor: output + """ + + out = self.activation(x) + out = self.conv1(out) + if self.bn==True: + out = self.bn1(out) + + out = self.activation(out) + out = self.conv2(out) + if self.bn==True: + out = self.bn2(out) + + if self.groups > 1: + out = self.conv_merge(out) + + return self.skip_add.add(out, x) + + +class FeatureFusionBlock_custom(nn.Module): + """Feature fusion block. + """ + + def __init__(self, features, activation, deconv=False, bn=False, expand=False, align_corners=True): + """Init. + + Args: + features (int): number of features + """ + super(FeatureFusionBlock_custom, self).__init__() + + self.deconv = deconv + self.align_corners = align_corners + + self.groups=1 + + self.expand = expand + out_features = features + if self.expand==True: + out_features = features//2 + + self.out_conv = nn.Conv2d(features, out_features, kernel_size=1, stride=1, padding=0, bias=True, groups=1) + + self.resConfUnit1 = ResidualConvUnit_custom(features, activation, bn) + self.resConfUnit2 = ResidualConvUnit_custom(features, activation, bn) + + self.skip_add = nn.quantized.FloatFunctional() + + def forward(self, *xs): + """Forward pass. + + Returns: + tensor: output + """ + output = xs[0] + + if len(xs) == 2: + res = self.resConfUnit1(xs[1]) + output = self.skip_add.add(output, res) + + output = self.resConfUnit2(output) + + output = nn.functional.interpolate( + output, scale_factor=2, mode="bilinear", align_corners=self.align_corners + ) + + output = self.out_conv(output) + + return output + + +class OutputConv(nn.Module): + """Output conv block. + """ + + def __init__(self, features, groups, activation, non_negative): + + super(OutputConv, self).__init__() + + self.output_conv = nn.Sequential( + nn.Conv2d(features, features//2, kernel_size=3, stride=1, padding=1, groups=groups), + nn.Upsample(scale_factor=2, mode="bilinear"), + nn.Conv2d(features//2, 32, kernel_size=3, stride=1, padding=1), + activation, + nn.Conv2d(32, 1, kernel_size=1, stride=1, padding=0), + nn.ReLU(True) if non_negative else nn.Identity(), + nn.Identity(), + ) + + def forward(self, x): + return self.output_conv(x) diff --git a/src/Baselines/radarcam-depth/modules/midas/midas_net_custom.py b/src/Baselines/radarcam-depth/modules/midas/midas_net_custom.py new file mode 100644 index 0000000000000000000000000000000000000000..18027e22a73cf1bbdd6de9686ea5ddb4834b50a2 --- /dev/null +++ b/src/Baselines/radarcam-depth/modules/midas/midas_net_custom.py @@ -0,0 +1,138 @@ +import torch +import torch.nn as nn + +from torch.nn import functional as F + +from .base_model import BaseModel +from .blocks import FeatureFusionBlock_custom, _make_encoder, OutputConv + +def weights_init(m): + import math + # initialize from normal (Gaussian) distribution + if isinstance(m, nn.Conv2d): + n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels + m.weight.data.normal_(0, math.sqrt(2.0 / n)) + if m.bias is not None: + m.bias.data.zero_() + elif isinstance(m, nn.BatchNorm2d): + m.weight.data.fill_(1) + m.bias.data.zero_() + + +class MidasNet_small_videpth(BaseModel): + """Network for monocular depth estimation. + """ + + def __init__(self, device = 'cpu', path=None, features=64, backbone="efficientnet_lite3", non_negative=False, exportable=True, channels_last=False, align_corners=True, + blocks={'expand': True}, in_channels=2, regress='r', min_pred=None, max_pred=None): + """Init. + + Args: + path (str, optional): Path to saved model. Defaults to None. + features (int, optional): Number of features. Defaults to 64. + backbone (str, optional): Backbone network for encoder. Defaults to efficientnet_lite3. + """ + print("Loading weights: ", path) + + super(MidasNet_small_videpth, self).__init__() + + use_pretrained = False + + self.channels_last = channels_last + self.blocks = blocks + self.backbone = backbone + + self.groups = 1 + + # for model output + self.regress = regress + self.min_pred = min_pred + self.max_pred = max_pred + + features1=features + features2=features + features3=features + features4=features + self.expand = False + if "expand" in self.blocks and self.blocks['expand'] == True: + self.expand = True + features1=features + features2=features*2 + features3=features*4 + features4=features*8 + + self.first = nn.Sequential( + nn.Conv2d(in_channels, 3, kernel_size=3, stride=1, padding=1), + nn.BatchNorm2d(3), + nn.ReLU(inplace=True) + ) + self.first.apply(weights_init) + + self.pretrained, self.scratch = _make_encoder(self.backbone, features, use_pretrained, groups=self.groups, expand=self.expand, exportable=exportable) + + self.scratch.activation = nn.ReLU(False) + + self.scratch.refinenet4 = FeatureFusionBlock_custom(features4, self.scratch.activation, deconv=False, bn=False, expand=self.expand, align_corners=align_corners) + self.scratch.refinenet3 = FeatureFusionBlock_custom(features3, self.scratch.activation, deconv=False, bn=False, expand=self.expand, align_corners=align_corners) + self.scratch.refinenet2 = FeatureFusionBlock_custom(features2, self.scratch.activation, deconv=False, bn=False, expand=self.expand, align_corners=align_corners) + self.scratch.refinenet1 = FeatureFusionBlock_custom(features1, self.scratch.activation, deconv=False, bn=False, align_corners=align_corners) + + self.scratch.output_conv = OutputConv(features, self.groups, self.scratch.activation, non_negative) + + if path: + self.load(path) + + self.to(device) + + + def forward(self, x, d): + """Forward pass. + + Args: + x (tensor): input data (image) + d (tensor): unalterated input depth + + Returns: + tensor: depth + """ + if self.channels_last==True: + print("self.channels_last = ", self.channels_last) + x.contiguous(memory_format=torch.channels_last) + + layer_0 = self.first(x) + + layer_1 = self.pretrained.layer1(layer_0) + layer_2 = self.pretrained.layer2(layer_1) + layer_3 = self.pretrained.layer3(layer_2) + layer_4 = self.pretrained.layer4(layer_3) + + layer_1_rn = self.scratch.layer1_rn(layer_1) + layer_2_rn = self.scratch.layer2_rn(layer_2) + layer_3_rn = self.scratch.layer3_rn(layer_3) + layer_4_rn = self.scratch.layer4_rn(layer_4) + + path_4 = self.scratch.refinenet4(layer_4_rn) + path_3 = self.scratch.refinenet3(path_4, layer_3_rn) + path_2 = self.scratch.refinenet2(path_3, layer_2_rn) + path_1 = self.scratch.refinenet1(path_2, layer_1_rn) + + out = self.scratch.output_conv(path_1) + + scales = F.relu(1.0 + out) + pred = d * scales + + # clamp pred to min and max + if self.min_pred is not None: + min_pred_inv = 1.0/self.min_pred + pred[pred > min_pred_inv] = min_pred_inv + if self.max_pred is not None: + max_pred_inv = 1.0/self.max_pred + pred[pred < max_pred_inv] = max_pred_inv + + # also return scales + return (pred, scales) + + + + + diff --git a/src/Baselines/radarcam-depth/modules/midas/normalization.py b/src/Baselines/radarcam-depth/modules/midas/normalization.py new file mode 100644 index 0000000000000000000000000000000000000000..6810e21c2f3ff97f38c7268b2d6ec96d786954c7 --- /dev/null +++ b/src/Baselines/radarcam-depth/modules/midas/normalization.py @@ -0,0 +1,109 @@ +VOID_INTERMEDIATE = { + + "dpt_beit_large_512" : { + "void_150" : { + "mean" : {"int_depth" : 0.730, "int_scales" : 0.380}, + "std" : {"int_depth" : 0.226, "int_scales" : 0.102}, + }, + "void_500" : { + "mean" : {"int_depth" : 0.736, "int_scales" : 0.366}, + "std" : {"int_depth" : 0.232, "int_scales" : 0.099}, + }, + "void_1500" : { + "mean" : {"int_depth" : 0.730, "int_scales" : 0.355}, + "std" : {"int_depth" : 0.232, "int_scales" : 0.096}, + }, + }, + + "dpt_swin2_large_384" : { + "void_150" : { + "mean" : {"int_depth" : 0.730, "int_scales" : 0.402}, + "std" : {"int_depth" : 0.219, "int_scales" : 0.107}, + }, + "void_500" : { + "mean" : {"int_depth" : 0.736, "int_scales" : 0.389}, + "std" : {"int_depth" : 0.224, "int_scales" : 0.106}, + }, + "void_1500" : { + "mean" : {"int_depth" : 0.730, "int_scales" : 0.377}, + "std" : {"int_depth" : 0.226, "int_scales" : 0.103}, + }, + }, + + "dpt_large" : { + "void_150" : { + "mean" : {"int_depth" : 0.729, "int_scales" : 0.403}, + "std" : {"int_depth" : 0.213, "int_scales" : 0.116}, + }, + "void_500" : { + "mean" : {"int_depth" : 0.735, "int_scales" : 0.390}, + "std" : {"int_depth" : 0.219, "int_scales" : 0.116}, + }, + "void_1500" : { + "mean" : {"int_depth" : 0.730, "int_scales" : 0.380}, + "std" : {"int_depth" : 0.221, "int_scales" : 0.116}, + }, + }, + + "dpt_hybrid": { + "void_150" : { + "mean" : {"int_depth" : 0.729, "int_scales" : 0.404}, + "std" : {"int_depth" : 0.210, "int_scales" : 0.117}, + }, + "void_500" : { + "mean" : {"int_depth" : 0.735, "int_scales" : 0.392}, + "std" : {"int_depth" : 0.215, "int_scales" : 0.118}, + }, + "void_1500" : { + "mean" : {"int_depth" : 0.730, "int_scales" : 0.381}, + "std" : {"int_depth" : 0.218, "int_scales" : 0.117}, + }, + }, + + "dpt_swin2_tiny_256" : { + "void_150" : { + "mean" : {"int_depth" : 0.735, "int_scales" : 0.419}, + "std" : {"int_depth" : 0.207, "int_scales" : 0.122}, + }, + "void_500" : { + "mean" : {"int_depth" : 0.741, "int_scales" : 0.406}, + "std" : {"int_depth" : 0.212, "int_scales" : 0.124}, + }, + "void_1500" : { + "mean" : {"int_depth" : 0.733, "int_scales" : 0.396}, + "std" : {"int_depth" : 0.213, "int_scales" : 0.125}, + }, + }, + + "dpt_levit_224" : { + "void_150" : { + "mean" : {"int_depth" : 0.734, "int_scales" : 0.421}, + "std" : {"int_depth" : 0.198, "int_scales" : 0.129}, + }, + "void_500" : { + "mean" : {"int_depth" : 0.740, "int_scales" : 0.410}, + "std" : {"int_depth" : 0.202, "int_scales" : 0.134}, + }, + "void_1500" : { + "mean" : {"int_depth" : 0.734, "int_scales" : 0.400}, + "std" : {"int_depth" : 0.204, "int_scales" : 0.137}, + }, + }, + + "midas_small" : { + "void_150" : { + "mean" : {"int_depth" : 0.723, "int_scales" : 0.402}, + "std" : {"int_depth" : 0.190, "int_scales" : 0.132}, + }, + "void_500" : { + "mean" : {"int_depth" : 0.731, "int_scales" : 0.393}, + "std" : {"int_depth" : 0.196, "int_scales" : 0.136}, + }, + "void_1500" : { + "mean" : {"int_depth" : 0.728, "int_scales" : 0.385}, + "std" : {"int_depth" : 0.199, "int_scales" : 0.140}, + }, + }, + +} + diff --git a/src/Baselines/radarcam-depth/modules/midas/transforms.py b/src/Baselines/radarcam-depth/modules/midas/transforms.py new file mode 100644 index 0000000000000000000000000000000000000000..cadc8c555543fa226c86124b08a0ff1813ef2978 --- /dev/null +++ b/src/Baselines/radarcam-depth/modules/midas/transforms.py @@ -0,0 +1,263 @@ +import numpy as np +import cv2 +import math +import torch +import torchvision.transforms as transforms + +from modules.midas.utils import normalize_unit_range +import modules.midas.normalization as normalization + +class Resize(object): + """Resize sample to given size (width, height). + """ + + def __init__( + self, + width, + height, + resize_target=True, + keep_aspect_ratio=False, + ensure_multiple_of=1, + resize_method="lower_bound", + image_interpolation_method=cv2.INTER_AREA, + ): + """Init. + + Args: + width (int): desired output width + height (int): desired output height + resize_target (bool, optional): + True: Resize the full sample (image, mask, target). + False: Resize image only. + Defaults to True. + keep_aspect_ratio (bool, optional): + True: Keep the aspect ratio of the input sample. + Output sample might not have the given width and height, and + resize behaviour depends on the parameter 'resize_method'. + Defaults to False. + ensure_multiple_of (int, optional): + Output width and height is constrained to be multiple of this parameter. + Defaults to 1. + resize_method (str, optional): + "lower_bound": Output will be at least as large as the given size. + "upper_bound": Output will be at max as large as the given size. (Output size might be smaller than given size.) + "minimal": Scale as least as possible. (Output size might be smaller than given size.) + Defaults to "lower_bound". + """ + self.__width = width + self.__height = height + + self.__resize_target = resize_target + self.__keep_aspect_ratio = keep_aspect_ratio + self.__multiple_of = ensure_multiple_of + self.__resize_method = resize_method + self.__image_interpolation_method = image_interpolation_method + + def constrain_to_multiple_of(self, x, min_val=0, max_val=None): + y = (np.round(x / self.__multiple_of) * self.__multiple_of).astype(int) + + if max_val is not None and y > max_val: + y = (np.floor(x / self.__multiple_of) * self.__multiple_of).astype(int) + + if y < min_val: + y = (np.ceil(x / self.__multiple_of) * self.__multiple_of).astype(int) + + return y + + def get_size(self, width, height): + # determine new height and width + scale_height = self.__height / height + scale_width = self.__width / width + + if self.__keep_aspect_ratio: + if self.__resize_method == "lower_bound": + # scale such that output size is lower bound + if scale_width > scale_height: + # fit width + scale_height = scale_width + else: + # fit height + scale_width = scale_height + elif self.__resize_method == "upper_bound": + # scale such that output size is upper bound + if scale_width < scale_height: + # fit width + scale_height = scale_width + else: + # fit height + scale_width = scale_height + elif self.__resize_method == "minimal": + # scale as least as possbile + if abs(1 - scale_width) < abs(1 - scale_height): + # fit width + scale_height = scale_width + else: + # fit height + scale_width = scale_height + else: + raise ValueError( + f"resize_method {self.__resize_method} not implemented" + ) + + if self.__resize_method == "lower_bound": + new_height = self.constrain_to_multiple_of( + scale_height * height, min_val=self.__height + ) + new_width = self.constrain_to_multiple_of( + scale_width * width, min_val=self.__width + ) + elif self.__resize_method == "upper_bound": + new_height = self.constrain_to_multiple_of( + scale_height * height, max_val=self.__height + ) + new_width = self.constrain_to_multiple_of( + scale_width * width, max_val=self.__width + ) + elif self.__resize_method == "minimal": + new_height = self.constrain_to_multiple_of(scale_height * height) + new_width = self.constrain_to_multiple_of(scale_width * width) + else: + raise ValueError(f"resize_method {self.__resize_method} not implemented") + + return (new_width, new_height) + + def __call__(self, sample): + width, height = self.get_size( + sample["image"].shape[1], sample["image"].shape[0] + ) + + # resize sample + for item in sample.keys(): + interpolation_method = self.__image_interpolation_method + sample[item] = cv2.resize( + sample[item], + (width, height), + interpolation=interpolation_method, + ) + + if self.__resize_target: + + if "gt" in sample: + sample["gt"] = cv2.resize( + sample["gt"], + (width, height), + interpolation=cv2.INTER_NEAREST + ) + + if "sparse_gt" in sample: + sample["sparse_gt"] = cv2.resize( + sample["sparse_gt"], + (width, height), + interpolation=cv2.INTER_NEAREST + ) + if "gt_sky" in sample: + sample["gt_sky"] = cv2.resize( + sample["gt_sky"], + (width, height), + interpolation=cv2.INTER_NEAREST + ) + + return sample + + + +class NormalizeIntermediate(object): + """Normalize intermediate data by given mean and std. + """ + + def __init__(self, mean, std): + + self.__int_depth_mean = mean["int_depth"] + self.__int_depth_std = std["int_depth"] + + self.__int_scales_mean = mean["int_scales"] + self.__int_scales_std = std["int_scales"] + + def __call__(self, sample): + + if "int_depth" in sample and sample["int_depth"] is not None: + sample["int_depth"] = (sample["int_depth"] - self.__int_depth_mean) / self.__int_depth_std + + if "int_scales" in sample and sample["int_scales"] is not None: + sample["int_scales"] = (sample["int_scales"] - self.__int_scales_mean) / self.__int_scales_std + + return sample + + +class PrepareForNet(object): + """Prepare sample for usage as network input. + """ + + def __init__(self): + pass + + def __call__(self, sample): + + for item in sample.keys(): + + if sample[item] is None: + pass + elif item == "image": + image = np.transpose(sample["image"], (2, 0, 1)) + sample["image"] = np.ascontiguousarray(image).astype(np.float32) + else: + array = sample[item].astype(np.float32) + array = np.expand_dims(array, axis=0) # add channel dim + sample[item] = np.ascontiguousarray(array) + + return sample + + +class Tensorize(object): + """Convert sample to tensor. + """ + + def __init__(self): + pass + + def __call__(self, sample): + + for item in sample.keys(): + + if sample[item] is None: + pass + else: + # before tensorizing, verify that data is clean + assert not np.any(np.isnan(sample[item])) + sample[item] = torch.Tensor(sample[item]) + + return sample + + +def get_transforms(depth_predictor, sparsifier, nsamples): + + resize_method_dict = { + "dpt_beit_large_512" : "minimal", + "dpt_swin2_large_384" : "minimal", + "dpt_large" : "minimal", + "dpt_hybrid" : "minimal", + "dpt_swin2_tiny_256" : "minimal", + "dpt_levit_224" : "minimal", + "midas_small" : "upper_bound", + } + + sml_model_transform_steps = [ + Resize( + width=288, + height=288, + resize_target=False, + keep_aspect_ratio=True, + ensure_multiple_of=32, + resize_method=resize_method_dict["dpt_hybrid"], + image_interpolation_method=cv2.INTER_NEAREST, + ), + NormalizeIntermediate( + mean=normalization.VOID_INTERMEDIATE[depth_predictor][f"{sparsifier}_{nsamples}"]["mean"], + std=normalization.VOID_INTERMEDIATE[depth_predictor][f"{sparsifier}_{nsamples}"]["std"], + ), + PrepareForNet(), + Tensorize(), + ] + + return transforms.Compose(sml_model_transform_steps) + diff --git a/src/Baselines/radarcam-depth/modules/midas/utils.py b/src/Baselines/radarcam-depth/modules/midas/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..9064a625a26b16317b8719d3054fdfb16832b99c --- /dev/null +++ b/src/Baselines/radarcam-depth/modules/midas/utils.py @@ -0,0 +1,237 @@ +"""Utils for monoDepth. +""" +import sys +import re +import numpy as np +import cv2 +import torch + + +def read_pfm(path): + """Read pfm file. + + Args: + path (str): path to file + + Returns: + tuple: (data, scale) + """ + with open(path, "rb") as file: + + color = None + width = None + height = None + scale = None + endian = None + + header = file.readline().rstrip() + if header.decode("ascii") == "PF": + color = True + elif header.decode("ascii") == "Pf": + color = False + else: + raise Exception("Not a PFM file: " + path) + + dim_match = re.match(r"^(\d+)\s(\d+)\s$", file.readline().decode("ascii")) + if dim_match: + width, height = list(map(int, dim_match.groups())) + else: + raise Exception("Malformed PFM header.") + + scale = float(file.readline().decode("ascii").rstrip()) + if scale < 0: + # little-endian + endian = "<" + scale = -scale + else: + # big-endian + endian = ">" + + data = np.fromfile(file, endian + "f") + shape = (height, width, 3) if color else (height, width) + + data = np.reshape(data, shape) + data = np.flipud(data) + + return data, scale + + +def write_pfm(path, image, scale=1): + """Write pfm file. + + Args: + path (str): pathto file + image (array): data + scale (int, optional): Scale. Defaults to 1. + """ + + with open(path, "wb") as file: + color = None + + if image.dtype.name != "float32": + raise Exception("Image dtype must be float32.") + + image = np.flipud(image) + + if len(image.shape) == 3 and image.shape[2] == 3: # color image + color = True + elif ( + len(image.shape) == 2 or len(image.shape) == 3 and image.shape[2] == 1 + ): # greyscale + color = False + else: + raise Exception("Image must have H x W x 3, H x W x 1 or H x W dimensions.") + + file.write("PF\n" if color else "Pf\n".encode()) + file.write("%d %d\n".encode() % (image.shape[1], image.shape[0])) + + endian = image.dtype.byteorder + + if endian == "<" or endian == "=" and sys.byteorder == "little": + scale = -scale + + file.write("%f\n".encode() % scale) + + image.tofile(file) + + +def read_image(path): + """Read image and output RGB image (0-1). + + Args: + path (str): path to file + + Returns: + array: RGB image (0-1) + """ + img = cv2.imread(path) + + if img.ndim == 2: + img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR) + + img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) / 255.0 + + return img + + +def resize_image(img): + """Resize image and make it fit for network. + + Args: + img (array): image + + Returns: + tensor: data ready for network + """ + height_orig = img.shape[0] + width_orig = img.shape[1] + + if width_orig > height_orig: + scale = width_orig / 384 + else: + scale = height_orig / 384 + + height = (np.ceil(height_orig / scale / 32) * 32).astype(int) + width = (np.ceil(width_orig / scale / 32) * 32).astype(int) + + img_resized = cv2.resize(img, (width, height), interpolation=cv2.INTER_AREA) + + img_resized = ( + torch.from_numpy(np.transpose(img_resized, (2, 0, 1))).contiguous().float() + ) + img_resized = img_resized.unsqueeze(0) + + return img_resized + + +def resize_depth(depth, width, height): + """Resize depth map and bring to CPU (numpy). + + Args: + depth (tensor): depth + width (int): image width + height (int): image height + + Returns: + array: processed depth + """ + depth = torch.squeeze(depth[0, :, :, :]).to("cpu") + + depth_resized = cv2.resize( + depth.numpy(), (width, height), interpolation=cv2.INTER_CUBIC + ) + + return depth_resized + + +def write_depth(path, depth, bits=1): + """Write depth map to pfm and png file. + + Args: + path (str): filepath without extension + depth (array): depth + """ + write_pfm(path + ".pfm", depth.astype(np.float32)) + + depth_min = depth.min() + depth_max = depth.max() + + max_val = (2**(8*bits))-1 + + if depth_max - depth_min > np.finfo("float").eps: + out = max_val * (depth - depth_min) / (depth_max - depth_min) + else: + out = np.zeros(depth.shape, dtype=depth.type) + + if bits == 1: + cv2.imwrite(path + ".png", out.astype("uint8")) + elif bits == 2: + cv2.imwrite(path + ".png", out.astype("uint16")) + + return + + +def write_png(path, array, bits=2, absolute=True): + """Write array to png file. + + Args: + path (str): filepath without extension + array (array): array to be saved + """ + if absolute: + out = array + else: + array_min = np.min(array) + array_max = np.max(array) + + max_val = (2**(8*bits))-1 + + if array_max - array_min > np.finfo("float").eps: + out = max_val * (array - array_min) / (array_max - array_min) + else: + print(f"zero array not being saved at {path}") + return + + if bits == 1: + cv2.imwrite(path + ".png", out.astype("uint8"), [cv2.IMWRITE_PNG_COMPRESSION, 0]) + elif bits == 2: + cv2.imwrite(path + ".png", out.astype("uint16"), [cv2.IMWRITE_PNG_COMPRESSION, 0]) + + return + + +def normalize_unit_range(data): + """Normalize data array to [0, 1] range. + + Args: + data (array): input array + + Returns: + array: normalized array + """ + if np.max(data) - np.min(data) > np.finfo("float").eps: + normalized = (data - np.min(data)) / (np.max(data) - np.min(data)) + else: + raise ValueError("cannot normalize array, max-min range is 0") + + return normalized \ No newline at end of file diff --git a/src/Baselines/radarcam-depth/networks.py b/src/Baselines/radarcam-depth/networks.py new file mode 100644 index 0000000000000000000000000000000000000000..cf7aa4bbce9c61a559406d243b21a198b59c1c5d --- /dev/null +++ b/src/Baselines/radarcam-depth/networks.py @@ -0,0 +1,1516 @@ +import torch +from utils import net_utils +import torchvision +from linear_attention import LocalFeatureTransformer + +''' +Encoders +''' + + +class ResNetEncoder(torch.nn.Module): + ''' + ResNet encoder with skip connections + Arg(s): + n_layer : int + architecture type based on layers: 18, 34, 50 + input_channels : int + number of channels in input data + n_filters : list + number of filters to use for each block + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + def __init__(self, + n_layer, + input_channels=3, + n_filters=[32, 64, 128, 256, 256], + weight_initializer='kaiming_uniform', + activation_func='leaky_relu', + use_batch_norm=False): + super(ResNetEncoder, self).__init__() + + if n_layer == 18: + n_blocks = [2, 2, 2, 2] + resnet_block = net_utils.ResNetBlock + elif n_layer == 34: + n_blocks = [3, 4, 6, 3] + resnet_block = net_utils.ResNetBlock + else: + raise ValueError('Only supports 18, 34 layer architecture') + + for n in range(len(n_filters) - len(n_blocks) - 1): + n_blocks = n_blocks + [n_blocks[-1]] + + network_depth = len(n_filters) + + assert network_depth < 8, 'Does not support network depth of 8 or more' + assert network_depth == len(n_blocks) + 1 + + # Keep track on current block + block_idx = 0 + filter_idx = 0 + + activation_func = net_utils.activation_func(activation_func) + + in_channels, out_channels = [input_channels, n_filters[filter_idx]] + + # Resolution 1/1 -> 1/2 + self.conv1 = net_utils.Conv2d( + in_channels, + out_channels, + kernel_size=7, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + # Resolution 1/2 -> 1/4 + self.max_pool = torch.nn.MaxPool2d( + kernel_size=3, + stride=2, + padding=1) + + filter_idx = filter_idx + 1 + + in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]] + + self.blocks2 = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels=in_channels, + out_channels=out_channels, + stride=1, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + # Resolution 1/4 -> 1/8 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]] + + self.blocks3 = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels=in_channels, + out_channels=out_channels, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + # Resolution 1/8 -> 1/16 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]] + + self.blocks4 = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels=in_channels, + out_channels=out_channels, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + # Resolution 1/16 -> 1/32 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]] + + self.blocks5 = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels=in_channels, + out_channels=out_channels, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + # Resolution 1/32 -> 1/64 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + if filter_idx < len(n_filters): + + in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]] + + self.blocks6 = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels=in_channels, + out_channels=out_channels, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + else: + self.blocks6 = None + + # Resolution 1/64 -> 1/128 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + if filter_idx < len(n_filters): + + in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]] + + self.blocks7 = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels=in_channels, + out_channels=out_channels, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + else: + self.blocks7 = None + + def _make_layer(self, + network_block, + n_block, + in_channels, + out_channels, + stride, + weight_initializer, + activation_func, + use_batch_norm): + ''' + Creates a layer + Arg(s): + network_block : Object + block type + n_block : int + number of blocks to use in layer + in_channels : int + number of channels + out_channels : int + number of output channels + stride : int + stride of convolution + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + blocks = [] + + for n in range(n_block): + + if n == 0: + stride = stride + else: + in_channels = out_channels + stride = 1 + + block = network_block( + in_channels=in_channels, + out_channels=out_channels, + stride=stride, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + blocks.append(block) + + blocks = torch.nn.Sequential(*blocks) + + return blocks + + def forward(self, x): + ''' + Forward input x through the ResNet model + Arg(s): + x : torch.Tensor + Returns: + torch.Tensor[float32] : latent vector + list[torch.Tensor[float32]] : skip connections + ''' + + layers = [x] + + # Resolution 1/1 -> 1/2 + layers.append(self.conv1(layers[-1])) + + # Resolution 1/2 -> 1/4 + max_pool = self.max_pool(layers[-1]) + layers.append(self.blocks2(max_pool)) + + # Resolution 1/4 -> 1/8 + layers.append(self.blocks3(layers[-1])) + + # Resolution 1/8 -> 1/16 + layers.append(self.blocks4(layers[-1])) + + # Resolution 1/16 -> 1/32 + layers.append(self.blocks5(layers[-1])) + + # Resolution 1/32 -> 1/64 + if self.blocks6 is not None: + layers.append(self.blocks6(layers[-1])) + + # Resolution 1/64 -> 1/128 + if self.blocks7 is not None: + layers.append(self.blocks7(layers[-1])) + + return layers[-1], layers[1:-1] + + +class FullyConnectedEncoder(torch.nn.Module): + ''' + Fully connected encoder + Arg(s): + input_channels : int + number of input channels + n_neurons : list[int] + number of filters to use per layer + latent_size : int + number of output neuron + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : str + activation function after convolution + ''' + + def __init__(self, + input_channels=3, + n_neurons=[32, 64, 96, 128, 256], + latent_size=29 * 10, + weight_initializer='kaiming_uniform', + activation_func='leaky_relu'): + super(FullyConnectedEncoder, self).__init__() + + activation_func = net_utils.activation_func(activation_func) + + self.mlp = torch.nn.Sequential( + net_utils.FullyConnected( + in_features=input_channels, + out_features=n_neurons[0], + weight_initializer=weight_initializer, + activation_func=activation_func), + net_utils.FullyConnected( + in_features=n_neurons[0], + out_features=n_neurons[1], + weight_initializer=weight_initializer, + activation_func=activation_func), + net_utils.FullyConnected( + in_features=n_neurons[1], + out_features=n_neurons[2], + weight_initializer=weight_initializer, + activation_func=activation_func), + net_utils.FullyConnected( + in_features=n_neurons[2], + out_features=n_neurons[3], + weight_initializer=weight_initializer, + activation_func=activation_func), + net_utils.FullyConnected( + in_features=n_neurons[3], + out_features=n_neurons[4], + weight_initializer=weight_initializer, + activation_func=activation_func), + net_utils.FullyConnected( + in_features=n_neurons[4], + out_features=latent_size, + weight_initializer=weight_initializer, + activation_func=activation_func)) + + def forward(self, x): + return self.mlp(x) + + +class FusionNetEncoder(torch.nn.Module): + ''' + FusionNet encoder with skip connections + Arg(s): + n_layer : int + number of layer for encoder + input_channels_image : int + number of channels in input data + input_channels_depth : int + number of channels in input data + n_filters_per_block : list[int] + number of filters to use for each block + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + fusion_type : str + add, weight + ''' + + def __init__(self, + n_layer=18, + input_channels_image=3, + input_channels_depth=3, + n_filters_encoder_image=[32, 64, 128, 256, 256], + n_filters_encoder_depth=[32, 64, 128, 256, 256], + weight_initializer='kaiming_uniform', + activation_func='leaky_relu', + use_batch_norm=False, + fusion_type='add'): + super(FusionNetEncoder, self).__init__() + + self.fusion_type = fusion_type + + if n_layer == 18: + n_blocks = [2, 2, 2, 2] + elif n_layer == 34: + n_blocks = [3, 4, 6, 3] + else: + raise ValueError('Only supports 18, 34 layer architecture') + + resnet_block = net_utils.ResNetBlock + + assert len(n_filters_encoder_image) == len(n_filters_encoder_depth) + + for n in range(len(n_filters_encoder_image) - len(n_blocks) - 1): + n_blocks = n_blocks + [n_blocks[-1]] + + network_depth = len(n_filters_encoder_image) + + assert network_depth < 8, 'Does not support network depth of 8 or more' + assert network_depth == len(n_blocks) + 1 + + # Keep track on current block + block_idx = 0 + filter_idx = 0 + + activation_func = net_utils.activation_func(activation_func) + + # Resolution 1/1 -> 1/2 + self.conv1_image = net_utils.Conv2d( + input_channels_image, + n_filters_encoder_image[filter_idx], + kernel_size=7, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + self.conv1_depth = net_utils.Conv2d( + input_channels_depth, + n_filters_encoder_depth[filter_idx], + kernel_size=7, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + if fusion_type == 'add': + self.conv1_project = net_utils.Conv2d( + n_filters_encoder_depth[filter_idx], + n_filters_encoder_image[filter_idx], + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight': + + self.conv1_weight = net_utils.Conv2d( + n_filters_encoder_depth[filter_idx], + n_filters_encoder_depth[filter_idx], + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight_and_project': + + self.conv1_weight = net_utils.Conv2d( + n_filters_encoder_depth[filter_idx], + n_filters_encoder_image[filter_idx], + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + self.conv1_project = net_utils.Conv2d( + n_filters_encoder_depth[filter_idx], + n_filters_encoder_image[filter_idx], + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + # Resolution 1/2 -> 1/4 + self.max_pool = torch.nn.MaxPool2d( + kernel_size=3, + stride=2, + padding=1) + + filter_idx = filter_idx + 1 + + in_channels_image, out_channels_image = [ + n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx] + ] + + in_channels_depth, out_channels_depth = [ + n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx] + ] + + self.blocks2_image, self.blocks2_depth = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels_image=in_channels_image, + in_channels_depth=in_channels_depth, + out_channels_image=out_channels_image, + out_channels_depth=out_channels_depth, + stride=1, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + if fusion_type == 'add': + self.conv2_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight': + + self.conv2_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_depth, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight_and_project': + + self.conv2_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + self.conv2_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + # Resolution 1/4 -> 1/8 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + in_channels_image, out_channels_image = [ + n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx] + ] + + in_channels_depth, out_channels_depth = [ + n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx] + ] + + self.blocks3_image, self.blocks3_depth = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels_image=in_channels_image, + in_channels_depth=in_channels_depth, + out_channels_image=out_channels_image, + out_channels_depth=out_channels_depth, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + if fusion_type == 'add': + self.conv3_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight': + + self.conv3_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_depth, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight_and_project': + + self.conv3_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + self.conv3_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + # Resolution 1/8 -> 1/16 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + in_channels_image, out_channels_image = [ + n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx] + ] + + in_channels_depth, out_channels_depth = [ + n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx] + ] + + self.blocks4_image, self.blocks4_depth = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels_image=in_channels_image, + in_channels_depth=in_channels_depth, + out_channels_image=out_channels_image, + out_channels_depth=out_channels_depth, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + if fusion_type == 'add': + self.conv4_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight': + + self.conv4_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_depth, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight_and_project': + + self.conv4_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + self.conv4_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + # Resolution 1/16 -> 1/32 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + in_channels_image, out_channels_image = [ + n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx] + ] + + in_channels_depth, out_channels_depth = [ + n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx] + ] + + self.blocks5_image, self.blocks5_depth = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels_image=in_channels_image, + in_channels_depth=in_channels_depth, + out_channels_image=out_channels_image, + out_channels_depth=out_channels_depth, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + if fusion_type == 'add': + self.conv5_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + elif fusion_type == 'weight': + + self.conv5_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_depth, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + if fusion_type == 'weight_and_project': + self.conv5_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + self.conv5_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + # Resolution 1/32 -> 1/64 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + if filter_idx < len(n_filters_encoder_image): + + in_channels_image, out_channels_image = [ + n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx] + ] + + in_channels_depth, out_channels_depth = [ + n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx] + ] + + self.blocks6_image, self.blocks6_depth = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels_image=in_channels_image, + in_channels_depth=in_channels_depth, + out_channels_image=out_channels_image, + out_channels_depth=out_channels_depth, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + if fusion_type == 'add': + self.conv6_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + + if fusion_type == 'weight_and_project': + self.conv6_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + self.conv6_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + else: + self.blocks6_image = None + self.blocks6_depth = None + self.conv6_weight = None + self.conv6_project = None + + # Resolution 1/64 -> 1/128 + block_idx = block_idx + 1 + filter_idx = filter_idx + 1 + + if filter_idx < len(n_filters_encoder_image): + + in_channels_image, out_channels_image = [ + n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx] + ] + + in_channels_depth, out_channels_depth = [ + n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx] + ] + + self.blocks7_image, self.blocks7_depth = self._make_layer( + network_block=resnet_block, + n_block=n_blocks[block_idx], + in_channels_image=in_channels_image, + in_channels_depth=in_channels_depth, + out_channels_image=out_channels_image, + out_channels_depth=out_channels_depth, + stride=2, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + if fusion_type == 'weight_and_project': + self.conv7_weight = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('sigmoid'), + use_batch_norm=use_batch_norm) + + self.conv7_project = net_utils.Conv2d( + out_channels_depth, + out_channels_image, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=net_utils.activation_func('linear'), + use_batch_norm=use_batch_norm) + else: + self.blocks7_image = None + self.blocks7_depth = None + self.conv7_weight = None + self.conv7_project = None + + def _make_layer(self, + network_block, + n_block, + in_channels_image, + in_channels_depth, + out_channels_image, + out_channels_depth, + stride, + weight_initializer, + activation_func, + use_batch_norm): + ''' + Creates a layer + Arg(s): + network_block : Object + block type + n_block : int + number of blocks to use in layer + in_channels_image : int + number of channels in image branch + in_channels_depth : int + number of channels in depth branch + out_channels_image : int + number of output channels in image branch + out_channels_depth : int + number of output channels in depth branch + stride : int + stride of convolution + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + blocks_image = [] + blocks_depth = [] + + for n in range(n_block): + + if n == 0: + stride = stride + else: + in_channels_image = out_channels_image + in_channels_depth = out_channels_depth + stride = 1 + + block_image = network_block( + in_channels=in_channels_image, + out_channels=out_channels_image, + stride=stride, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + blocks_image.append(block_image) + + block_depth = network_block( + in_channels=in_channels_depth, + out_channels=out_channels_depth, + stride=stride, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + blocks_depth.append(block_depth) + + blocks_image = torch.nn.Sequential(*blocks_image) + blocks_depth = torch.nn.Sequential(*blocks_depth) + + return blocks_image, blocks_depth + + def forward(self, image, depth): + ''' + Forward input x through the ResNet model + Arg(s): + image : torch.Tensor + depth : torch.Tensor + Returns: + torch.Tensor[float32] : latent vector + list[torch.Tensor[float32]] : skip connections + ''' + + layers = [] + + # Resolution 1/1 -> 1/2 + conv1_image = self.conv1_image(image) + conv1_depth = self.conv1_depth(depth) + + if self.fusion_type == 'add': + conv1_project = self.conv1_project(conv1_depth) + conv1 = conv1_project + conv1_image + elif self.fusion_type == 'weight': + conv1_weight = self.conv1_weight(conv1_depth) + conv1 = conv1_weight * conv1_depth + conv1_image + elif self.fusion_type == 'weight_and_project': + conv1_weight = self.conv1_weight(conv1_depth) + conv1_project = self.conv1_project(conv1_depth) + conv1 = conv1_weight * conv1_project + conv1_image + elif self.fusion_type == 'concat': + conv1 = torch.cat([conv1_depth, conv1_image], dim=1) + else: + raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type)) + + layers.append(conv1) + + # Resolution 1/2 -> 1/4 + max_pool_image = self.max_pool(conv1_image) + max_pool_depth = self.max_pool(conv1_depth) + + blocks2_image = self.blocks2_image(max_pool_image) + blocks2_depth = self.blocks2_depth(max_pool_depth) + + if self.fusion_type == 'add': + conv2_project = self.conv2_project(blocks2_depth) + blocks2 = conv2_project + blocks2_image + elif self.fusion_type == 'weight': + conv2_weight = self.conv2_weight(blocks2_depth) + blocks2 = conv2_weight * blocks2_depth + blocks2_image + elif self.fusion_type == 'weight_and_project': + conv2_weight = self.conv2_weight(blocks2_depth) + conv2_project = self.conv2_project(blocks2_depth) + blocks2 = conv2_weight * conv2_project + blocks2_image + elif self.fusion_type == 'concat': + blocks2 = torch.cat([blocks2_image, blocks2_depth], dim=1) + else: + raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type)) + + layers.append(blocks2) + + # Resolution 1/4 -> 1/8 + blocks3_image = self.blocks3_image(blocks2_image) + blocks3_depth = self.blocks3_depth(blocks2_depth) + + if self.fusion_type == 'add': + conv3_project = self.conv3_project(blocks3_depth) + blocks3 = conv3_project + blocks3_image + elif self.fusion_type == 'weight': + conv3_weight = self.conv3_weight(blocks3_depth) + blocks3 = conv3_weight * blocks3_depth + blocks3_image + elif self.fusion_type == 'weight_and_project': + conv3_weight = self.conv3_weight(blocks3_depth) + conv3_project = self.conv3_project(blocks3_depth) + blocks3 = conv3_weight * conv3_project + blocks3_image + elif self.fusion_type == 'concat': + blocks3 = torch.cat([blocks3_image, blocks3_depth], dim=1) + else: + raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type)) + + layers.append(blocks3) + + # Resolution 1/8 -> 1/16 + blocks4_image = self.blocks4_image(blocks3_image) + blocks4_depth = self.blocks4_depth(blocks3_depth) + + if self.fusion_type == 'add': + conv4_project = self.conv4_project(blocks4_depth) + blocks4 = conv4_project + blocks4_image + elif self.fusion_type == 'weight': + conv4_weight = self.conv4_weight(blocks4_depth) + blocks4 = conv4_weight * blocks4_depth + blocks4_image + elif self.fusion_type == 'weight_and_project': + conv4_weight = self.conv4_weight(blocks4_depth) + conv4_project = self.conv4_project(blocks4_depth) + blocks4 = conv4_weight * conv4_project + blocks4_image + elif self.fusion_type == 'concat': + blocks4 = torch.cat([blocks4_image, blocks4_depth], dim=1) + else: + raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type)) + + layers.append(blocks4) + + # Resolution 1/16 -> 1/32 + blocks5_image = self.blocks5_image(blocks4_image) + blocks5_depth = self.blocks5_depth(blocks4_depth) + + if self.fusion_type == 'add': + conv5_project = self.conv5_project(blocks5_depth) + blocks5 = conv5_project + blocks5_image + elif self.fusion_type == 'weight': + conv5_weight = self.conv5_weight(blocks5_depth) + blocks5 = conv5_weight * blocks5_depth + blocks5_image + elif self.fusion_type == 'weight_and_project': + conv5_weight = self.conv5_weight(blocks5_depth) + conv5_project = self.conv5_project(blocks5_depth) + blocks5 = conv5_weight * conv5_project + blocks5_image + elif self.fusion_type == 'concat': + blocks5 = torch.cat([blocks5_image, blocks5_depth], dim=1) + else: + raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type)) + + layers.append(blocks5) + + # Resolution 1/32 -> 1/64 + if self.blocks6_image is not None and self.blocks6_depth is not None: + blocks6_image = self.blocks6_image(blocks5_image) + blocks6_depth = self.blocks6_depth(blocks5_depth) + + if self.fusion_type == 'add': + conv6_project = self.conv6_project(blocks6_depth) + blocks6 = conv6_project + blocks6_image + elif self.fusion_type == 'weight': + conv6_weight = self.conv6_weight(blocks6_depth) + blocks6 = conv6_weight * blocks6_depth + blocks6_image + elif self.fusion_type == 'weight_and_project': + conv6_weight = self.conv6_weight(blocks6_depth) + conv6_project = self.conv6_project(blocks6_depth) + blocks6 = conv6_weight * conv6_project + blocks6_image + elif self.fusion_type == 'concat': + blocks6 = torch.cat([blocks6_image, blocks6_depth], dim=1) + else: + raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type)) + + layers.append(blocks6) + + # Resolution 1/64 -> 1/128 + if self.blocks7_image is not None and self.blocks7_depth is not None: + blocks7_image = self.blocks7_image(blocks6_image) + blocks7_depth = self.blocks7_depth(blocks6_depth) + + if self.fusion_type == 'add': + conv7_project = self.conv7_project(blocks7_depth) + blocks7 = conv7_project + blocks7_image + elif self.fusion_type == 'weight': + conv7_weight = self.conv7_weight(blocks7_depth) + blocks7 = conv7_weight * blocks7_depth + blocks7_image + elif self.fusion_type == 'weight_and_project': + conv7_weight = self.conv7_weight(blocks7_depth) + conv7_project = self.conv7_project(blocks7_depth) + blocks7 = conv7_weight * conv7_project + blocks7_image + elif self.fusion_type == 'concat': + blocks7 = torch.cat([blocks7_image, blocks7_depth], dim=1) + else: + raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type)) + + layers.append(blocks7) + + return layers[-1], layers[:-1] + + +class RCNetEncoder(torch.nn.Module): + ''' + Radar association network + Arg(s): + in_channels_image : int + number of input channels for image (RGB) branch + in_channels_depth : int + number of input channels for depth branch + n_filters_encoder_image : int + number of filters for image (RGB) branch + n_neurons_encoder_depth : int + number of neurons for depth (radar) branch + latent_size_depth : int + size of latent vector + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + def __init__(self, + input_channels_image=3, + input_channels_depth=3, + input_patch_size_image=(900, 288), + n_filters_encoder_image=[32, 64, 128, 128, 128], + n_neurons_encoder_depth=[32, 64, 128, 128, 128], + latent_size_depth=128 * 29 * 10, + weight_initializer='kaiming_uniform', + activation_func='leaky_relu', + use_batch_norm=False): + super(RCNetEncoder, self).__init__() + + self.n_neuron_latent_depth = n_neurons_encoder_depth[-1] + + self.encoder_image = ResNetEncoder( + n_layer=18, + input_channels=input_channels_image, + n_filters=n_filters_encoder_image, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + self.attention = LocalFeatureTransformer(['self','cross'], n_layers=4, d_model=self.n_neuron_latent_depth) + + self.encoder_depth = FullyConnectedEncoder( + input_channels=input_channels_depth, + n_neurons=n_neurons_encoder_depth, + latent_size=latent_size_depth, + weight_initializer=weight_initializer, + activation_func=activation_func) + + self.input_patch_size_image =input_patch_size_image + + def forward(self, image, points, b_boxes): + # Image shape: (B, C, H, W) # Should be (B, 3, 768, 288) + # points shape: (B*K, X) + # b_boxes: [(K, 4) * B], this should be a list with B elements, and each element is (K, 4) size + # K is the number of radar points per image + # X is the radar dimension + + + # Define dimensions + shape = self.input_patch_size_image + latent_height = int(shape[-2] // 32.0) + latent_width = int(shape[-1] // 32.0) + batch_size = image.shape[0] + + # Define scales and feature sizes + skip_scales = [ 1 /2.0, 1/ 4.0, 1 / 8.0, 1 / 16.0, 1 / 32.0, 1 / 64.0, 1 / 128.0] + skip_feature_sizes = [ + (int(shape[-2] * skip_scale), + int(shape[-1] * skip_scale)) + for skip_scale in skip_scales + ] # Should be [(384, 144), (192, 72), (96, 36), (48, 18)] + + latent_scale = 1 / 32.0 + latent_feature_size = (latent_height, latent_width) # Should be (24, 9) + + # Forward the entire image + latent_image, skips_image = self.encoder_image(image) + + # ROI pooling on latent images + latent_image_pooled = torchvision.ops.roi_pool( + latent_image, b_boxes, + spatial_scale=latent_scale, + output_size=latent_feature_size + ) # (N*K, C, H, W) + + # ROI pooling on the skips + skips_image_pooled = [] + for skip_image_idx in range(len(skips_image)): + skips_image_pooled.append( + torchvision.ops.roi_pool( + skips_image[skip_image_idx], b_boxes, + spatial_scale=skip_scales[skip_image_idx], + output_size=skip_feature_sizes[skip_image_idx] + ) # (N*K, C, H, W) + ) + + # Radar points size: (bath_size * total_points_sampled, 3) + # latent_depth size: (batch_size * total_points_sampled, n_neuron_latent_depth, patch_w//32, patch_h//32) + # latent_image_pooled size = latent_depth size + latent_depth = self.encoder_depth(points) + latent_depth = latent_depth.view(points.shape[0], self.n_neuron_latent_depth, -1, latent_width) + + latent_depth_reshape = latent_depth.view(latent_depth.shape[0], latent_depth.shape[1], -1).permute(0, 2, 1) + latent_image_pooled_reshape = latent_image_pooled.view(latent_image_pooled.shape[0], + latent_image_pooled.shape[1], -1).permute(0, 2, 1) + latent_depth_tf, latent_image_pooled_tf = self.attention(latent_depth_reshape, latent_image_pooled_reshape) + latent_depth_tf = latent_depth_tf.permute(0, 2, 1).view(latent_depth.shape) + latent_image_pooled_tf = latent_image_pooled_tf.permute(0, 2, 1).view(latent_image_pooled.shape) + + # Concatenate the features + # latent = torch.cat([latent_image_pooled, latent_depth], dim=1) + latent = torch.cat([latent_image_pooled_tf, latent_depth_tf], dim=1) + return latent, skips_image_pooled + + +''' +Decoder +''' + + +class MultiScaleDecoder(torch.nn.Module): + ''' + Multi-scale decoder with skip connections + Arg(s): + input_channels : int + number of channels in input latent vector + output_channels : int + number of channels or classes in output + n_resolution : int + number of output resolutions (scales) for multi-scale prediction + n_filters : int list + number of filters to use at each decoder block + n_skips : int list + number of filters from skip connections + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + output_func : func + activation function for output + use_batch_norm : bool + if set, then applied batch normalization + deconv_type : str + deconvolution types available: transpose, up + ''' + + def __init__(self, + input_channels=256, + output_channels=1, + n_resolution=1, + n_filters=[256, 128, 64, 32, 16], + n_skips=[256, 128, 64, 32, 0], + weight_initializer='kaiming_uniform', + activation_func='leaky_relu', + output_func='linear', + use_batch_norm=False, + deconv_type='up'): + super(MultiScaleDecoder, self).__init__() + + network_depth = len(n_filters) + + assert network_depth < 8, 'Does not support network depth of 8 or more' + assert n_resolution > 0 and n_resolution < network_depth + + self.n_resolution = n_resolution + self.output_func = output_func + + activation_func = net_utils.activation_func(activation_func) + output_func = net_utils.activation_func(output_func) + + # Upsampling from lower to full resolution requires multi-scale + if 'upsample' in self.output_func and self.n_resolution < 2: + self.n_resolution = 2 + + filter_idx = 0 + + in_channels, skip_channels, out_channels = [ + input_channels, n_skips[filter_idx], n_filters[filter_idx] + ] + + # Resolution 1/128 -> 1/64 + if network_depth > 6: + self.deconv6 = net_utils.DecoderBlock( + in_channels, + skip_channels, + out_channels, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm, + deconv_type=deconv_type) + + filter_idx = filter_idx + 1 + + in_channels, skip_channels, out_channels = [ + n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx] + ] + else: + self.deconv6 = None + + # Resolution 1/64 -> 1/32 + if network_depth > 5: + self.deconv5 = net_utils.DecoderBlock( + in_channels, + skip_channels, + out_channels, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm, + deconv_type=deconv_type) + + filter_idx = filter_idx + 1 + + in_channels, skip_channels, out_channels = [ + n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx] + ] + else: + self.deconv5 = None + + # Resolution 1/32 -> 1/16 + self.deconv4 = net_utils.DecoderBlock( + in_channels, + skip_channels, + out_channels, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm, + deconv_type=deconv_type) + + # Resolution 1/16 -> 1/8 + filter_idx = filter_idx + 1 + + in_channels, skip_channels, out_channels = [ + n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx] + ] + + self.deconv3 = net_utils.DecoderBlock( + in_channels, + skip_channels, + out_channels, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm, + deconv_type=deconv_type) + + if self.n_resolution > 3: + self.output3 = net_utils.Conv2d( + out_channels, + output_channels, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=output_func, + use_batch_norm=False) + + # Resolution 1/8 -> 1/4 + filter_idx = filter_idx + 1 + + in_channels, skip_channels, out_channels = [ + n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx] + ] + + if self.n_resolution > 3: + skip_channels = skip_channels + output_channels + + self.deconv2 = net_utils.DecoderBlock( + in_channels, + skip_channels, + out_channels, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm, + deconv_type=deconv_type) + + if self.n_resolution > 2: + self.output2 = net_utils.Conv2d( + out_channels, + output_channels, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=output_func, + use_batch_norm=False) + + # Resolution 1/4 -> 1/2 + filter_idx = filter_idx + 1 + + in_channels, skip_channels, out_channels = [ + n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx] + ] + + if self.n_resolution > 2: + skip_channels = skip_channels + output_channels + + self.deconv1 = net_utils.DecoderBlock( + in_channels, + skip_channels, + out_channels, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm, + deconv_type=deconv_type) + + if self.n_resolution > 1: + self.output1 = net_utils.Conv2d( + out_channels, + output_channels, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=output_func, + use_batch_norm=False) + + # Resolution 1/2 -> 1/1 + filter_idx = filter_idx + 1 + + in_channels, skip_channels, out_channels = [ + n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx] + ] + + if self.n_resolution > 1: + skip_channels = skip_channels + output_channels + + self.deconv0 = net_utils.DecoderBlock( + in_channels, + skip_channels, + out_channels, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm, + deconv_type=deconv_type) + + self.output0 = net_utils.Conv2d( + out_channels, + output_channels, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=output_func, + use_batch_norm=False) + + def forward(self, x, skips, shape=None): + ''' + Forward latent vector x through decoder network + Arg(s): + x : torch.Tensor[float32] + latent vector + skips : list[torch.Tensor[float32]] + list of skip connection tensors (earlier are larger resolution) + shape : tuple[int] + (height, width) tuple denoting output size + Returns: + list[torch.Tensor[float32]] : list of outputs at multiple scales + ''' + + layers = [x] + outputs = [] + + # Start at the end and walk backwards through skip connections + n = len(skips) - 1 + + # Resolution 1/128 -> 1/64 + if self.deconv6 is not None: + layers.append(self.deconv6(layers[-1], skips[n])) + n = n - 1 + + # Resolution 1/64 -> 1/32 + if self.deconv5 is not None: + layers.append(self.deconv5(layers[-1], skips[n])) + n = n - 1 + + # Resolution 1/32 -> 1/16 + layers.append(self.deconv4(layers[-1], skips[n])) + + # Resolution 1/16 -> 1/8 + n = n - 1 + + layers.append(self.deconv3(layers[-1], skips[n])) + + if self.n_resolution > 3: + output3 = self.output3(layers[-1]) + outputs.append(output3) + + upsample_output3 = torch.nn.functional.interpolate( + input=outputs[-1], + scale_factor=2, + mode='bilinear', + align_corners=True) + + # Resolution 1/8 -> 1/4 + n = n - 1 + + skip = torch.cat([skips[n], upsample_output3], dim=1) if self.n_resolution > 3 else skips[n] + layers.append(self.deconv2(layers[-1], skip)) + + if self.n_resolution > 2: + output2 = self.output2(layers[-1]) + outputs.append(output2) + + upsample_output2 = torch.nn.functional.interpolate( + input=outputs[-1], + scale_factor=2, + mode='bilinear', + align_corners=True) + + # Resolution 1/4 -> 1/2 + n = n - 1 + + skip = torch.cat([skips[n], upsample_output2], dim=1) if self.n_resolution > 2 else skips[n] + layers.append(self.deconv1(layers[-1], skip)) + + if self.n_resolution > 1: + output1 = self.output1(layers[-1]) + outputs.append(output1) + + upsample_output1 = torch.nn.functional.interpolate( + input=outputs[-1], + scale_factor=2, + mode='bilinear', + align_corners=True) + + # Resolution 1/2 -> 1/1 + n = n - 1 + + if 'upsample' in self.output_func: + output0 = upsample_output1 + else: + if self.n_resolution > 1: + # If there is skip connection at layer 0 + skip = torch.cat([skips[n], upsample_output1], dim=1) if n == 0 else upsample_output1 + layers.append(self.deconv0(layers[-1], skip)) + else: + + if n == 0: + layers.append(self.deconv0(layers[-1], skips[n])) + else: + layers.append(self.deconv0(layers[-1], shape=shape[-2:])) + + output0 = self.output0(layers[-1]) + + outputs.append(output0) + + return outputs \ No newline at end of file diff --git a/src/Baselines/radarcam-depth/rcnet_inference.py b/src/Baselines/radarcam-depth/rcnet_inference.py new file mode 100644 index 0000000000000000000000000000000000000000..c7231570631ef42dbe4890bc6073493699b1fc2a --- /dev/null +++ b/src/Baselines/radarcam-depth/rcnet_inference.py @@ -0,0 +1,351 @@ +"""RC-Net inference: quasi-dense depth from a trained RC-Net. + +The run(), forward() and log_network_settings() functions below are copied +verbatim from RadarCam-Depth/RCNet/rcnet_main.py (the baseline's train() half +is replaced by rcnet_train_rice.py and is not vendored, which also drops the +tensorboard dependency). +""" + +import os + +import numpy as np +import torch +import torch.utils.data +import torchvision +from tqdm.auto import tqdm + +from data import data_utils, datasets +from rcnet_model import RCNetModel +from rcnet_transforms import Transforms +from utils.log_utils import log + + +def run(save_root, + + image_paths, + radar_paths, + gt_paths, + + restore_path, + patch_size, + normalized_image_range, + + encoder_type, + n_filters_encoder_image, + n_neurons_encoder_depth, + decoder_type, + n_filters_decoder, + weight_initializer, + activation_func, + response_thr=0.5): + + # Set up device + device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') + + ''' + Read input paths + ''' + + n_sample = len(image_paths) + + assert n_sample == len(radar_paths) + assert n_sample == len(gt_paths) + + ''' + Set up inputs and outputs + ''' + depth_predicted_paths = [] + response_predicted_paths = [] + depth_predicted_color_paths = [] + + inputs_outputs = [ + [ + 'training', + image_paths, + radar_paths, + gt_paths, + depth_predicted_paths, + depth_predicted_color_paths, + response_predicted_paths, + ] + ] + + ''' + Set up the model + ''' + # Build network + rcnet_model = RCNetModel( + input_channels_image=3, + input_channels_depth=3, + input_patch_size_image=patch_size, + encoder_type=encoder_type, + n_filters_encoder_image=n_filters_encoder_image, + n_neurons_encoder_depth=n_neurons_encoder_depth, + decoder_type=decoder_type, + n_filters_decoder=n_filters_decoder, + weight_initializer=weight_initializer, + activation_func=activation_func, + device=device) + + rcnet_model.eval() + rcnet_model.to(device) + rcnet_model.data_parallel() + + parameters_rcnet_model = rcnet_model.parameters() + + step, _ = rcnet_model.restore_model(restore_path) + + log('Restoring checkpoint from: \n{}\n'.format(restore_path)) + + log_network_settings( + log_path=None, + # Network settings + encoder_type=encoder_type, + n_filters_encoder_image=n_filters_encoder_image, + n_neurons_encoder_depth=n_neurons_encoder_depth, + decoder_type=decoder_type, + n_filters_decoder=n_filters_decoder, + # Weight settings + weight_initializer=weight_initializer, + activation_func=activation_func, + parameters_model=parameters_rcnet_model) + + ''' + Process each set of input and outputs + ''' + for paths in inputs_outputs: + # Unpack inputs and outputs + tag, \ + image_paths, \ + radar_paths, \ + ground_truth_paths, \ + depth_predicted_paths, \ + depth_predicted_color_paths, \ + response_predicted_paths, = paths + + # Create output paths for depth and response + for radar_path in radar_paths: + # Create path and store + file_id = os.path.basename(radar_path).split('.')[0] + depth_predicted_path = os.path.join(save_root, 'depth_predicted', file_id + '.png') + depth_predicted_paths.append(depth_predicted_path) + + depth_predicted_color_path = os.path.join(save_root, 'depth_predicted_colors', file_id + '.png') + depth_predicted_color_paths.append(depth_predicted_color_path) + + response_predicted_path = os.path.join(save_root, 'response_predicted', file_id + '.png') + response_predicted_paths.append(response_predicted_path) + + # Create directories + depth_predicted_dirpaths = np.unique([os.path.dirname(path) for path in depth_predicted_paths]) + depth_predicted_color_dirpaths = np.unique([os.path.dirname(path) for path in depth_predicted_color_paths]) + response_predicted_dirpaths = np.unique([os.path.dirname(path) for path in response_predicted_paths]) + for dirpaths in [depth_predicted_dirpaths, depth_predicted_color_dirpaths, response_predicted_dirpaths]: + for dirpath in dirpaths: + os.makedirs(dirpath, exist_ok=True) + + # Set up dataloader + dataloader = torch.utils.data.DataLoader( + datasets.RCNetInferenceDataset( + image_paths=image_paths, + radar_paths=radar_paths, + ground_truth_paths=ground_truth_paths), + batch_size=1, + shuffle=False, + num_workers=1, + drop_last=False) + + transforms = Transforms( + normalized_image_range=normalized_image_range) + + n_sample = len(dataloader) + + print('Processing {} samples...'.format(n_sample)) + + # Iterate through data loader + progress = tqdm(dataloader, total=n_sample, desc=f'{tag} inference', unit='frame') + for sample_idx, data in enumerate(progress): + with torch.no_grad(): + data = [ + datum.to(device) for datum in data + ] + + image, radar_points, ground_truth = data + bounding_boxes_list = [] + + pad_size_x = patch_size[1] // 2 + radar_points = radar_points.squeeze(dim=0) + + if radar_points.ndim == 1: + # Expand to 1 x 3 + radar_points = np.expand_dims(radar_points, axis=0) + + # get the shifted radar points after padding + for radar_point_idx in range(0, radar_points.shape[0]): + # Set radar point to the center of the patch + radar_points[radar_point_idx, 0] = radar_points[radar_point_idx, 0] + pad_size_x + bounding_box = torch.zeros(4) + bounding_box[0] = radar_points[radar_point_idx, 0] - pad_size_x + bounding_box[1] = 0 + bounding_box[2] = radar_points[radar_point_idx, 0] + pad_size_x + bounding_box[3] = image.shape[-2] + bounding_boxes_list.append(bounding_box) + + bounding_boxes_list = [torch.stack(bounding_boxes_list, dim=0)] + + [image], [radar_points], [bounding_boxes_list] = transforms.transform( + images_arr=[image], + points_arr=[radar_points], + bounding_boxes_arr=[bounding_boxes_list], + random_transform_probability=0.0) + + output_depth, output_response, _, inference_failed = forward_with_fallback( + model=rcnet_model, + image=image, + radar_points=radar_points, + bounding_boxes_list=bounding_boxes_list, + response_thr=response_thr, + device=device) + + output_depth = np.squeeze(output_depth.cpu().numpy()) + output_response = np.squeeze(output_response.cpu().numpy()) + + if inference_failed: + tqdm.write( + 'ERROR: RCNet inference failed for {}: ' + 'no positive response at threshold 0.0'.format( + os.path.basename(radar_paths[sample_idx]))) + + ''' + Save outputs + ''' + data_utils.save_depth(output_depth, depth_predicted_paths[sample_idx]) + data_utils.save_color_depth(output_depth, depth_predicted_color_paths[sample_idx]) + data_utils.save_response(output_response, response_predicted_paths[sample_idx]) +def forward_with_fallback(model, + image, + radar_points, + bounding_boxes_list, + response_thr=0.5, + device=torch.device('cuda')): + '''Run inference, lowering the response threshold only when necessary.''' + output_depth, output_response = forward( + model=model, + image=image, + radar_points=radar_points, + bounding_boxes_list=bounding_boxes_list, + response_thr=response_thr, + device=device) + + threshold = float(response_thr) + while not torch.any(output_depth > 0).item() and threshold > 0.0: + threshold = max(0.0, threshold - 0.05) + output_depth, output_response = forward( + model=model, + image=image, + radar_points=radar_points, + bounding_boxes_list=bounding_boxes_list, + response_thr=threshold, + device=device) + + inference_failed = not torch.any(output_depth > 0).item() + return output_depth, output_response, threshold, inference_failed + + +def forward(model, image, radar_points, bounding_boxes_list, response_thr=0.5, device=torch.device('cuda')): + # Determine crop size for possible radar correspondence + patch_size = model.input_patch_size_image + pad_size = patch_size[1] // 2 + + image = torchvision.transforms.functional.pad( + image, + (pad_size, 0, pad_size, 0), + padding_mode='edge') + start_y = image.shape[-2] - patch_size[0] + + output_tiles = [] + if radar_points.dim() == 3: + # Convert to 1 x N x 3 to N x 3 + radar_points = torch.squeeze(radar_points, dim=0) + + x_shifts = radar_points[:, 0].clone() + + height = image.shape[-2] + crop_height = height - patch_size[0] + + output_crops = model.forward( + image=image, + point=radar_points, + bounding_boxes=bounding_boxes_list, + return_logits=False) + for output_crop, x in zip(output_crops, x_shifts): + output = torch.zeros([1, image.shape[-2], image.shape[-1]], device=device) + + output_crop = torch.where(output_crop < response_thr, torch.zeros_like(output_crop), output_crop) + # Add crop to output + output[:, crop_height:, int(x) - pad_size:int(x) + pad_size] = output_crop + output_tiles.append(output) + + output_tiles = torch.cat(output_tiles, dim=0) + output_tiles = output_tiles[:, :, pad_size:-pad_size] + + # Find the max response over all tiles + output_response, output_indices = torch.max(output_tiles, dim=0, keepdim=True) + # Fill a floating-point depth map based on the selected radar point. + output_depth = torch.zeros_like(output_response) + for point_idx in range(radar_points.shape[0]): + output_depth = torch.where( + output_indices == point_idx, + torch.full_like( + output_response, + fill_value=float(radar_points[point_idx, 2]), + ), + output_depth) + + # Leave as 0s if we did not predict + output_depth = torch.where( + output_response == 0, + torch.zeros_like(output_depth), + output_depth) + + return output_depth, output_response + + +def log_network_settings(log_path, + # Network settings + encoder_type, + n_filters_encoder_image, + n_neurons_encoder_depth, + decoder_type, + n_filters_decoder, + # Weight settings + weight_initializer, + activation_func, + parameters_model=[]): + # Computer number of parameters + n_parameter = sum(p.numel() for p in parameters_model) + + n_parameter_text = 'n_parameter={}'.format(n_parameter) + n_parameter_vars = [] + + log('Network settings:', log_path) + log('encoder_type={}'.format(encoder_type), + log_path) + log('n_filters_encoder_image={}'.format(n_filters_encoder_image), + log_path) + log('n_neurons_encoder_depth={}'.format(n_neurons_encoder_depth), + log_path) + log('decoder_type={}'.format(decoder_type), + log_path) + log('n_filters_decoder={}'.format( + n_filters_decoder), + log_path) + log('', log_path) + + log('Weight settings:', log_path) + log(n_parameter_text.format(*n_parameter_vars), + log_path) + log('weight_initializer={} activation_func={}'.format( + weight_initializer, activation_func), + log_path) + log('', log_path) diff --git a/src/Baselines/radarcam-depth/rcnet_model.py b/src/Baselines/radarcam-depth/rcnet_model.py new file mode 100644 index 0000000000000000000000000000000000000000..2adce6f0e69232cbd00a8068a190057457d2af38 --- /dev/null +++ b/src/Baselines/radarcam-depth/rcnet_model.py @@ -0,0 +1,407 @@ +import torch, torchvision +from safetensors.torch import load_file +from utils import log_utils +import networks + + +class RCNetModel(object): + ''' + Image radar fusion to determine correspondence of radar to image + + Arg(s): + input_channels_image : int + number of channels in the image + input_channels_depth : int + number of channels in depth map + input_patch_size_image : int + patch of image to consider for radar point + encoder_type : str + encoder type + n_filters_encoder_image : list[int] + list of filters for each layer in image encoder + n_neurons_encoder_image : list[int] + list of neurons for each layer in depth encoder + decoder_type : str + decoder type + n_filters_decoder : list[int] + list of filters for each layer in decoder + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : str + activation function for network + device : torch.device + device for running model + ''' + + def __init__(self, + input_channels_image, + input_channels_depth, + input_patch_size_image, + encoder_type, + n_filters_encoder_image, + n_neurons_encoder_depth, + decoder_type, + n_filters_decoder, + weight_initializer='kaiming_uniform', + activation_func='leaky_relu', + device=torch.device('cuda')): + + self.input_patch_size_image = input_patch_size_image + self.device = device + + # height, width = input_patch_size_image + # latent_height = np.ceil(height / 32.0).astype(int) + # latent_width = np.ceil(width / 32.0).astype(int) + + height, width = input_patch_size_image + latent_height = int((height // 32.0)) + latent_width = int((width // 32.0)) + + latent_size_depth = latent_height * latent_width * n_neurons_encoder_depth[-1] + + # Build encoder + if 'rcnet' in encoder_type: + self.encoder = networks.RCNetEncoder( + input_channels_image=input_channels_image, + input_channels_depth=input_channels_depth, + input_patch_size_image=input_patch_size_image, + n_filters_encoder_image=n_filters_encoder_image, + n_neurons_encoder_depth=n_neurons_encoder_depth, + latent_size_depth=latent_size_depth, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm='batch_norm' in encoder_type) + else: + raise ValueError('Encoder type {} not supported.'.format(encoder_type)) + + # Calculate number of channels for latent and skip connections combining image + depth + n_skips = n_filters_encoder_image[:-1] + n_skips = n_skips[::-1] + [0] + + latent_channels = n_filters_encoder_image[-1] + n_neurons_encoder_depth[-1] + + # Build decoder + if 'multiscale' in decoder_type: + self.decoder = networks.MultiScaleDecoder( + input_channels=latent_channels, + output_channels=1, + n_resolution=1, + n_filters=n_filters_decoder, + n_skips=n_skips, + weight_initializer=weight_initializer, + activation_func=activation_func, + output_func='linear', + use_batch_norm='batch_norm' in decoder_type, + deconv_type='up') + else: + raise ValueError('Decoder type {} not supported.'.format(decoder_type)) + + # Move to device + self.to(self.device) + + def forward(self, image, point, bounding_boxes, return_logits=True): + ''' + Forwards the inputs through the network + + Arg(s): + image : torch.Tensor[float32] + N x 3 x H x W image + point : torch.Tensor[float32] + N x 3 input point + return_logits : bool + if set, then return logits otherwise sigmoid + Returns: + torch.Tensor[float32] : N x 1 x H x W logits (correspondence map) + ''' + + latent, skips = self.encoder(image, point, bounding_boxes) + + logits = self.decoder(x=latent, skips=skips, shape=self.input_patch_size_image)[-1] + + if return_logits: + return logits + else: + return torch.sigmoid(logits) + + def parameters(self): + ''' + Returns the list of parameters in the model + + Returns: + list[torch.Tensor[float32]] : list of parameters + ''' + + parameters = \ + list(self.encoder.parameters()) + \ + list(self.decoder.parameters()) + + return parameters + + def eval(self): + ''' + Sets model to evaluation mode + ''' + + self.encoder.eval() + self.decoder.eval() + + def to(self, device): + ''' + Moves model to specified device + + Arg(s): + device : torch.device + device for running model + ''' + + # Move to device + self.encoder.to(device) + self.decoder.to(device) + + def restore_model(self, checkpoint_path): + ''' + Restore weights of the model + + Arg(s): + checkpoint_path : str + path to checkpoint + Returns: + int : current step in optimization + None : retained for compatibility with the baseline inference call + ''' + + checkpoint = load_file(checkpoint_path, device='cpu') + encoder_prefix = 'radarnet_encoder.' + decoder_prefix = 'radarnet_decoder.' + encoder_state = { + key[len(encoder_prefix):]: value + for key, value in checkpoint.items() + if key.startswith(encoder_prefix) + } + decoder_state = { + key[len(decoder_prefix):]: value + for key, value in checkpoint.items() + if key.startswith(decoder_prefix) + } + self.encoder.load_state_dict(encoder_state) + self.decoder.load_state_dict(decoder_state) + return 0, None + + def data_parallel(self): + ''' + Allows multi-gpu split along batch + ''' + + self.encoder = torch.nn.DataParallel(self.encoder) + self.decoder = torch.nn.DataParallel(self.decoder) + + def log_summary(self, + summary_writer, + tag, + step, + image=None, + output_response=None, + output_label=None, + output_depth=None, + validity_map=None, + ground_truth_label=None, + ground_truth_depth=None, + scalars={}, + n_display=4): + ''' + Logs summary to Tensorboard + + Arg(s): + summary_writer : SummaryWriter + Tensorboard summary writer + tag : str + tag that prefixes names to log + step : int + current step in training + image : torch.Tensor[float32] + N x 3 x H x W image + output_response : torch.Tensor[float32] + N x 1 x H x W soft correspondence map + output_label : torch.Tensor[float32] + N x 1 x H x W binary correspondence map + output_depth : torch.Tensor[float32] + N x 1 x H x W depth map + validity_map : torch.Tensor[float32] + N x 1 x H x W validity map + ground_truth_label : torch.Tensor[float32] + N x 1 x H x W ground truth label + ground_truth_depth : torch.Tensor[float32] + N x 1 x H x W ground truth depth map + scalars : dict[str, float] + dictionary of scalars to log + n_display : int + number of images to display + ''' + + with torch.no_grad(): + + display_summary_image = [] + + display_summary_image_text = tag + + if image is not None: + image_summary = image[0:n_display, ...] + + display_summary_image_text += '-image' + + # Add to list of images to log + display_summary_image.append(image_summary.cpu()) + + if output_response is not None: + output_response_summary = output_response[0:n_display, ...] + + display_summary_image_text += '-output_response' + + # Add to list of images to log + display_summary_image.append( + log_utils.colorize( + output_response_summary.cpu(), + colormap='inferno')) + + # Log distribution of output response + summary_writer.add_histogram( + tag + '-output_response_distro', + output_response_summary, + global_step=step) + + if output_label is not None: + output_label_summary = output_label[0:n_display, ...] + + display_summary_image_text += '-output_label' + + # Add to list of images to log + display_summary_image.append( + log_utils.colorize( + output_label_summary.cpu(), + colormap='inferno')) + + # Log distribution of output and label + summary_writer.add_histogram( + tag + '-output_label_distro', + output_label_summary, + global_step=step) + + if ground_truth_label is not None: + ground_truth_label_summary = ground_truth_label[0:n_display, ...] + + validity_map_label_summary = torch.where( + ground_truth_label_summary > 0, + torch.ones_like(ground_truth_label_summary), + torch.zeros_like(ground_truth_label_summary)) + + display_summary_image_text += '-ground_truth_label' + + if output_label is not None: + display_summary_image_text += '-error' + + # Compute output error w.r.t. ground truth + ground_truth_label_error_summary = \ + torch.abs(output_label_summary - ground_truth_label_summary) + + ground_truth_label_error_summary = torch.where( + validity_map_label_summary == 1.0, + (ground_truth_label_error_summary + 1e-8) / (ground_truth_label_summary + 1e-8), + validity_map_label_summary) + + display_summary_image.append( + log_utils.colorize( + ground_truth_label_error_summary.cpu(), + colormap='inferno')) + + # Add to list of images to log + display_summary_image.append( + log_utils.colorize( + ground_truth_label_summary.cpu(), + colormap='inferno')) + + # Log distribution of ground truth + summary_writer.add_histogram( + tag + '_ground_truth_label_distro', + ground_truth_label, + global_step=step) + + if validity_map is not None: + validity_map_summary = validity_map[0:n_display, ...] + + display_summary_image_text += '-validity_map' + + # Add to list of images to log + display_summary_image.append( + log_utils.colorize( + validity_map_summary.cpu(), + colormap='inferno')) + + if output_depth is not None: + output_depth_summary = output_depth[0:n_display, ...] + + display_summary_image_text += '-output_depth' + + # Add to list of images to log + display_summary_image.append( + log_utils.colorize( + (output_depth_summary / 100.0).cpu(), + colormap='viridis')) + + # Log distribution of output depth + summary_writer.add_histogram( + tag + '-output_depth_distro', + output_depth, + global_step=step) + + if ground_truth_depth is not None: + ground_truth_depth = torch.unsqueeze(ground_truth_depth[:, 0, :, :], dim=1) + ground_truth_depth_summary = ground_truth_depth[0:n_display, ...] + + validity_map_summary = torch.where( + ground_truth_depth_summary > 0, + torch.ones_like(ground_truth_depth_summary), + torch.zeros_like(ground_truth_depth_summary)) + + display_summary_image_text += '-ground_truth_label' + + if output_depth is not None: + display_summary_image_text += '-error' + + # Compute output error w.r.t. ground truth + ground_truth_depth_error_summary = \ + torch.abs(output_depth_summary - ground_truth_depth_summary) + + ground_truth_depth_error_summary = torch.where( + validity_map_summary == 1.0, + (ground_truth_depth_error_summary + 1e-8) / (ground_truth_depth_summary + 1e-8), + validity_map_summary) + + display_summary_image.append( + log_utils.colorize( + (ground_truth_depth_error_summary / 0.05).cpu(), + colormap='inferno')) + + # Add to list of images to log + display_summary_image.append( + log_utils.colorize( + (ground_truth_depth_summary / 100.0).cpu(), + colormap='viridis')) + + # Log distribution of ground truth + summary_writer.add_histogram( + tag + '-ground_truth_distro', + ground_truth_depth, + global_step=step) + + # Log scalars to tensorboard + for (name, value) in scalars.items(): + summary_writer.add_scalar(tag + '-' + name, value, global_step=step) + + # Log image summaries to tensorboard + if len(display_summary_image) > 1: + display_summary_image = torch.cat(display_summary_image, dim=2) + + summary_writer.add_image( + display_summary_image_text, + torchvision.utils.make_grid(display_summary_image, nrow=n_display), + global_step=step) diff --git a/src/Baselines/radarcam-depth/rcnet_transforms.py b/src/Baselines/radarcam-depth/rcnet_transforms.py new file mode 100644 index 0000000000000000000000000000000000000000..da80d46f54f00ff1f5b6b7ec7c1558b99642ad22 --- /dev/null +++ b/src/Baselines/radarcam-depth/rcnet_transforms.py @@ -0,0 +1,432 @@ +import torch +import torchvision.transforms.functional as functional + + +class Transforms(object): + + def __init__(self, + normalized_image_range=[0, 255], + random_brightness=[-1], + random_contrast=[-1], + random_saturation=[-1], + random_noise_type='none', + random_noise_spread=-1, + random_flip_type=['none']): + ''' + Transforms and augmentation class + Note: brightness, contrast, gamma, hue, saturation augmentations expect + either type int in [0, 255] or float in [0, 1] + + Arg(s): + normalized_image_range : list[float] + intensity range after normalizing images + random_brightness : list[float] + brightness adjustment [0, B], from 0 (black image) to B factor increase + random_contrast : list[float] + contrast adjustment [0, C], from 0 (gray image) to C factor increase + random_saturation : list[float] + saturation adjustment [0, S], from 0 (black image) to S factor increase + random_noise_type : str + type of noise to add: gaussian, uniform + random_noise_spread : float + if gaussian, then standard deviation; if uniform, then min-max range + random_flip_type : list[str] + none, horizontal, vertical + ''' + + # Image normalization + self.normalized_image_range = normalized_image_range + + # Photometric augmentations + self.do_random_brightness = True if -1 not in random_brightness else False + self.random_brightness = random_brightness + self.do_random_contrast = True if -1 not in random_contrast else False + self.random_contrast = random_contrast + self.do_random_saturation = True if -1 not in random_saturation else False + self.random_saturation = random_saturation + + self.do_random_noise = \ + True if (random_noise_type != 'none' and random_noise_spread > -1) else False + + self.random_noise_type = random_noise_type + self.random_noise_spread = random_noise_spread + + # Geometric augmentations + self.do_random_horizontal_flip = True if 'horizontal' in random_flip_type else False + self.do_random_vertical_flip = True if 'vertical' in random_flip_type else False + + def transform(self, + images_arr, + labels_arr=[], + points_arr=[], + bounding_boxes_arr=[], + random_transform_probability=0.00): + ''' + Applies transform to images and ground truth + + Arg(s): + images_arr : list[torch.Tensor] + list of N x C x H x W tensors + labels_arr : list[torch.Tensor] + list of N x c x H x W tensors + points_arr : list[torch.Tensor] + list of N x 3 tensors + bounding_boxes_arr : list[torch.Tensor] + list of N x 4 tensors + random_transform_probability : float + probability to perform transform + Returns: + list[torch.Tensor] : list of transformed N x C x H x W image tensors + list[torch.Tensor] : list of transformed N x c x H x W label tensors + list[torch.Tensor] : list of transformed N x 3 points tensors + list[torch.Tensor] : list of transformed N x 4 bounding box tensors + ''' + + device = images_arr[0].device + + n_dim = images_arr[0].ndim + + if n_dim == 4: + n_batch, _, n_height, n_width = images_arr[0].shape + else: + raise ValueError('Unsupported number of dimensions: {}'.format(n_dim)) + + do_random_transform = \ + torch.rand(n_batch, device=device) <= random_transform_probability + + ''' + Photometric Transformations (applied only to images) + ''' + for idx, images in enumerate(images_arr): + # In case user pass in [0, 255] range image as float type + if torch.max(images) > 1.0: + images_arr[idx] = images.int() + + if self.do_random_brightness: + + do_brightness = torch.logical_and( + do_random_transform, + torch.rand(n_batch, device=device) <= 0.50) + + values = torch.rand(n_batch, device=device) + + brightness_min, brightness_max = self.random_brightness + factors = (brightness_max - brightness_min) * values + brightness_min + + images_arr = self.adjust_brightness(images_arr, do_brightness, factors) + + if self.do_random_contrast: + + do_contrast = torch.logical_and( + do_random_transform, + torch.rand(n_batch, device=device) <= 0.50) + + values = torch.rand(n_batch, device=device) + + contrast_min, contrast_max = self.random_contrast + factors = (contrast_max - contrast_min) * values + contrast_min + + images_arr = self.adjust_contrast(images_arr, do_contrast, factors) + + if self.do_random_saturation: + + do_saturation = torch.logical_and( + do_random_transform, + torch.rand(n_batch, device=device) <= 0.50) + + values = torch.rand(n_batch, device=device) + + saturation_min, saturation_max = self.random_saturation + factors = (saturation_max - saturation_min) * values + saturation_min + + images_arr = self.adjust_saturation(images_arr, do_saturation, factors) + + ''' + Convert all images to float and normalize + ''' + images_arr = [ + images.float() for images in images_arr + ] + + # Normalize images to a given range + images_arr = self.normalize_images( + images_arr, + normalized_image_range=self.normalized_image_range) + + ''' + Points augmentation + ''' + if self.do_random_noise: + + do_add_noise = torch.logical_and( + do_random_transform, + torch.rand(n_batch, device=device) <= 0.50) + + points_arr = self.add_noise( + points_arr, + do_add_noise=do_add_noise, + noise_type=self.random_noise_type, + noise_spread=self.random_noise_spread) + + ''' + Geometric transformations (applied to both images and labels) + ''' + if self.do_random_horizontal_flip: + + do_horizontal_flip = torch.logical_and( + do_random_transform, + torch.rand(n_batch, device=device) <= 0.50) + + images_arr = self.horizontal_flip( + images_arr, + do_horizontal_flip) + + labels_arr = self.horizontal_flip( + labels_arr, + do_horizontal_flip) + + + + for bounding_boxes in bounding_boxes_arr: + for bounding_box_idx in range(0,bounding_boxes.shape[0]): + do_hflip = do_horizontal_flip[bounding_box_idx] + for box_idx in range(0,bounding_boxes.shape[1]): + if do_hflip: + temp = bounding_boxes[bounding_box_idx, box_idx, 0].clone() + bounding_boxes[bounding_box_idx, box_idx, 0] = n_width - bounding_boxes[bounding_box_idx, box_idx, 2] + bounding_boxes[bounding_box_idx, box_idx, 2] = n_width - temp + + + if self.do_random_vertical_flip: + + do_vertical_flip = torch.logical_and( + do_random_transform, + torch.rand(n_batch, device=device) <= 0.50) + + images_arr = self.vertical_flip( + images_arr, + do_vertical_flip) + + labels_arr = self.vertical_flip( + labels_arr, + do_vertical_flip) + + for bounding_boxes in bounding_boxes_arr: + for bounding_box_idx in range(0,bounding_boxes.shape[0]): + do_vflip = do_vertical_flip[bounding_box_idx] + if do_vflip: + temp = bounding_boxes[bounding_box_idx, 1].clone() + bounding_boxes[bounding_box_idx, 1] = n_height - bounding_boxes[bounding_box_idx, 3] + bounding_boxes[bounding_box_idx, 3] = n_height - temp + + # Return the transformed inputs + outputs = [] + + if len(images_arr) > 0: + outputs.append(images_arr) + + if len(labels_arr) > 0: + outputs.append(labels_arr) + + if len(points_arr) > 0: + outputs.append(points_arr) + + if len(bounding_boxes_arr) > 0: + outputs.append(bounding_boxes_arr) + + if len(outputs) == 1: + return outputs[0] + else: + return outputs + + ''' + Photometric transforms + ''' + def normalize_images(self, images_arr, normalized_image_range=[0, 1]): + ''' + Normalize image to a given range + + Arg(s): + images_arr : list[torch.Tensor[float32]] + list of N x C x H x W tensors + normalized_image_range : list[float] + intensity range after normalizing images + Returns: + images_arr[torch.Tensor[float32]] : list of normalized N x C x H x W tensors + ''' + + if normalized_image_range == [0, 1]: + images_arr = [ + images / 255.0 for images in images_arr + ] + elif normalized_image_range == [-1, 1]: + images_arr = [ + 2.0 * (images / 255.0) - 1.0 for images in images_arr + ] + elif normalized_image_range == [0, 255]: + pass + else: + raise ValueError('Unsupported normalization range: {}'.format( + normalized_image_range)) + + return images_arr + + def adjust_brightness(self, images_arr, do_brightness, factors): + ''' + Adjust brightness on each sample + + Arg(s): + images_arr : list[torch.Tensor] + list of N x C x H x W tensors + do_brightness : bool + N booleans to determine if brightness is adjusted on each sample + factors : float + N floats to determine how much to adjust + Returns: + list[torch.Tensor] : list of transformed N x C x H x W image tensors + ''' + + for i, images in enumerate(images_arr): + + for b, image in enumerate(images): + if do_brightness[b]: + images[b, ...] = functional.adjust_brightness(image, factors[b]) + + images_arr[i] = images + + return images_arr + + def adjust_contrast(self, images_arr, do_contrast, factors): + ''' + Adjust contrast on each sample + + Arg(s): + images_arr : list[torch.Tensor] + list of N x C x H x W tensors + do_contrast : bool + N booleans to determine if contrast is adjusted on each sample + factors : float + N floats to determine how much to adjust + Returns: + list[torch.Tensor] : list of transformed N x C x H x W image tensors + ''' + + for i, images in enumerate(images_arr): + + for b, image in enumerate(images): + if do_contrast[b]: + images[b, ...] = functional.adjust_contrast(image, factors[b]) + + images_arr[i] = images + + return images_arr + + def adjust_saturation(self, images_arr, do_saturation, factors): + ''' + Adjust saturation on each sample + + Arg(s): + images_arr : list[torch.Tensor] + list of N x C x H x W tensors + do_saturation : bool + N booleans to determine if saturation is adjusted on each sample + gammas : float + N floats to determine how much to adjust + Returns: + list[torch.Tensor] : list of transformed N x C x H x W image tensors + ''' + + for i, images in enumerate(images_arr): + + for b, image in enumerate(images): + if do_saturation[b]: + images[b, ...] = functional.adjust_saturation(image, factors[b]) + + images_arr[i] = images + + return images_arr + + ''' + Geometric transforms + ''' + def horizontal_flip(self, images_arr, do_horizontal_flip): + ''' + Perform horizontal flip on each sample + + Arg(s): + images_arr : list[torch.Tensor[float32]] + list of N x C x H x W tensors + do_horizontal_flip : bool + N booleans to determine if horizontal flip is performed on each sample + Returns: + list[torch.Tensor[float32]] : list of transformed N x C x H x W image tensors + ''' + + for i, images in enumerate(images_arr): + + for b, image in enumerate(images): + if do_horizontal_flip[b]: + images[b, ...] = torch.flip(image, dims=[-1]) + + images_arr[i] = images + + return images_arr + + def vertical_flip(self, images_arr, do_vertical_flip): + ''' + Perform vertical flip on each sample + + Arg(s): + images_arr : list[torch.Tensor[float32]] + list of N x C x H x W tensors + do_vertical_flip : bool + N booleans to determine if vertical flip is performed on each sample + Returns: + list[torch.Tensor[float32]] : list of transformed N x C x H x W image tensors + ''' + + for i, images in enumerate(images_arr): + + for b, image in enumerate(images): + if do_vertical_flip[b]: + images[b, ...] = torch.flip(image, dims=[-2]) + + images_arr[i] = images + + return images_arr + + def add_noise(self, images_arr, do_add_noise, noise_type, noise_spread): + ''' + Add noise to images + + Arg(s): + images_arr : list[torch.Tensor] + list of N x C x H x W tensors + do_add_noise : bool + N booleans to determine if noise will be added + noise_type : str + gaussian, uniform + noise_spread : float + if gaussian, then standard deviation; if uniform, then min-max range + ''' + + for i, images in enumerate(images_arr): + device = images.device + + for b, image in enumerate(images): + if do_add_noise[b]: + + shape = image.shape + + if noise_type == 'gaussian': + image = image + noise_spread * torch.randn(*shape, device=device) + elif noise_type == 'uniform': + image = image + noise_spread * (torch.rand(*shape, device=device) - 0.5) + else: + raise ValueError('Unsupported noise type: {}'.format(noise_type)) + + images[b, ...] = image + + images_arr[i] = images + + return images_arr diff --git a/src/Baselines/radarcam-depth/rice_config.py b/src/Baselines/radarcam-depth/rice_config.py new file mode 100644 index 0000000000000000000000000000000000000000..431a6a59a6ae8c02bb6cd6bbd589e03bba8ba657 --- /dev/null +++ b/src/Baselines/radarcam-depth/rice_config.py @@ -0,0 +1,190 @@ +"""Nested YAML config loading, in the rcd_rice style. + +One config.yaml with sections (data / depth / rcnet / sml / mono / +global_alignment / runtime / wandb) deep-merged over DEFAULT_CONFIG, with +attribute access (cfg.rcnet.epochs) and data paths resolved relative to the +config file. The same module is shared by radarcam-depth and dataset_prep; +each folder's config.yaml only sets the sections it uses. + +Environment overrides (applied after the YAML merge): RICE_DATA_ROOT -> +data.output_root, RICE_RAW_DIR -> data.raw_dir. +""" + +import copy +import os +from typing import Any, Dict + +import yaml + +DEFAULT_CONFIG: Dict[str, Any] = { + "data": { + "raw_dir": "", + "split_json": "split.json", + "output_root": "rice_data", + "smoke_eval_root": "", + "limit": None, + }, + "depth": { + "max_radar_depth_m": 11.2, + "min_radar_depth_m": 0.05, + "min_pred_depth_m": 0.1, + "max_pred_depth_m": 20.0, # ZED GT can exceed the radar range + "min_eval_depth_m": 0.0, + "max_eval_depth_m": 11.2, + }, + "rcnet": { + # Shared-data convention: 288x512 images (uniform 0.4 scale of + # 1280x720), patch 288x96 -> latent (9, 3) as in the ZJU config. + "input_height": 288, + "input_width": 512, + "patch_size": [288, 96], + "total_points_sampled": 40, + "sample_probability_of_lidar": 0.10, + "normalized_image_range": [0, 1], + # Network (baseline architecture, unchanged) + "encoder_type": ["rcnet", "batch_norm"], + "n_filters_encoder_image": [32, 64, 128, 128, 128], + "n_neurons_encoder_depth": [32, 64, 128, 128, 128], + "decoder_type": ["multiscale", "batch_norm"], + "n_filters_decoder": [256, 128, 64, 32, 16], + "weight_initializer": "kaiming_uniform", + "activation_func": "leaky_relu", + # Augmentation (baseline defaults) + "augmentation_probability": 1.0, + "augmentation_random_brightness": [0.80, 1.20], + "augmentation_random_contrast": [0.80, 1.20], + "augmentation_random_saturation": [0.80, 1.20], + "augmentation_random_flip_type": ["horizontal"], + # Loss + "w_positive_class": 2.5, + "max_distance_correspondence": 0.5, + "set_invalid_to_negative_class": False, + # Training (paper Sec. IV-B: 50 epochs at lr 2e-4) + "batch_size": 6, # per GPU + "epochs": 50, + "learning_rate": 2e-4, + "weight_decay": 0.0, + "lr_milestones": [], + "lr_gamma": 0.5, + # Runtime + "num_workers": 0, + "log_freq": 50, + "save_dir": "checkpoints_rcnet", + "checkpoint_path": "", + "response_thr": 0.5, + }, + "sml": { + "mono_tag": "dpt_hybrid", + # Loss (w_lidar_loss MUST stay 0 for rice -- dense ZED gt) + "loss_func": "smoothl1", + "w_smoothness": 0.0, + "loss_smoothness_kernel_size": -1, + "w_lidar_loss": 0.0, + # Training (paper: lr 2e-4 -> 5e-5 after 20 of 40 epochs) + "batch_size": 8, # per GPU + "epochs": 40, + "learning_rate": 2e-4, + "lr_milestones": [20], + "lr_gamma": 0.25, + "weight_decay": 0.0, + # Runtime + "num_workers": 0, + "log_freq": 50, + "save_dir": "checkpoints_sml", + "checkpoint_path": "", + "save_visualizations": True, + "num_visualizations": 4, + "visualization_dir": "visualizations_sml", + }, + "mono": { + "model_type": "DPT_Hybrid", + "tag": "dpt_hybrid", + }, + "global_alignment": { + "mono_tag": "dpt_hybrid", + "min_points": 5, + }, + "runtime": { + "seed": 42, + "cpu": False, + "mixed_precision": "fp16", + }, + "wandb": { + "project": "radarcam-rice", + "entity": None, + "rcnet_run_name": "rcnet-run-1", + "sml_run_name": "sml-run-1", + "api_key": "", # empty -> use `wandb login` / WANDB_API_KEY env + }, +} + + +class ConfigNode(dict): + """Dictionary with attribute access for YAML config sections.""" + + def __getattr__(self, name: str) -> Any: + try: + return self[name] + except KeyError as exc: + raise AttributeError(name) from exc + + def __setattr__(self, name: str, value: Any) -> None: + self[name] = value + + +def _deep_update(base: Dict[str, Any], overrides: Dict[str, Any]) -> Dict[str, Any]: + for key, value in overrides.items(): + if isinstance(value, dict) and isinstance(base.get(key), dict): + _deep_update(base[key], value) + else: + base[key] = value + return base + + +def _to_node(value: Any) -> Any: + if isinstance(value, dict): + return ConfigNode({k: _to_node(v) for k, v in value.items()}) + if isinstance(value, list): + return [_to_node(v) for v in value] + return value + + +def _resolve_paths(cfg: Dict[str, Any], config_path: str) -> None: + config_dir = os.path.dirname(os.path.abspath(config_path)) + for key in ("raw_dir", "split_json", "output_root", "smoke_eval_root"): + value = cfg["data"][key] + if value and not os.path.isabs(value): + cfg["data"][key] = os.path.abspath(os.path.join(config_dir, value)) + + +def load_config(config_path: str) -> ConfigNode: + cfg = copy.deepcopy(DEFAULT_CONFIG) + with open(config_path, "r") as f: + payload = yaml.safe_load(f) or {} + if not isinstance(payload, dict): + raise ValueError("YAML config must contain a mapping at the top level.") + _deep_update(cfg, payload) + _resolve_paths(cfg, config_path) + + if os.environ.get("RICE_DATA_ROOT"): + cfg["data"]["output_root"] = os.path.abspath(os.environ["RICE_DATA_ROOT"]) + if os.environ.get("RICE_RAW_DIR"): + cfg["data"]["raw_dir"] = os.path.abspath(os.environ["RICE_RAW_DIR"]) + + return _to_node(cfg) + + +def to_plain_dict(value: Any) -> Any: + if isinstance(value, dict): + return {k: to_plain_dict(v) for k, v in value.items()} + if isinstance(value, list): + return [to_plain_dict(v) for v in value] + return value + + +def loggable_config(cfg: Any) -> Dict[str, Any]: + """Plain dict for wandb/checkpoint payloads, with the API key masked.""" + payload = to_plain_dict(cfg) + if payload.get("wandb", {}).get("api_key"): + payload["wandb"]["api_key"] = "***" + return payload diff --git a/src/Baselines/radarcam-depth/rice_paths.py b/src/Baselines/radarcam-depth/rice_paths.py new file mode 100644 index 0000000000000000000000000000000000000000..c748ac6cfa673a8486333cdf866f167f986a4cbe --- /dev/null +++ b/src/Baselines/radarcam-depth/rice_paths.py @@ -0,0 +1,48 @@ +"""ZJU-4DRadarCam-style directory layout under a data root. + +All scripts build a Layout from cfg.data.output_root (see rice_config.py) so +the on-disk structure is defined in exactly one place. This module is shared +by radarcam-depth and dataset_prep. +""" + +import os + +HERE = os.path.dirname(os.path.abspath(__file__)) + + +class Layout(object): + def __init__(self, data_root): + self.data_root = data_root + self.data_dir = os.path.join(data_root, "data") + self.result_dir = os.path.join(data_root, "result") + self.log_dir = os.path.join(data_root, "log") + + self.image_dir = os.path.join(self.data_dir, "image") + self.radar_npy_dir = os.path.join(self.data_dir, "radar") + self.radar_png_dir = os.path.join(self.data_dir, "radar_png") + self.gt_dir = os.path.join(self.data_dir, "gt") + self.gt_interp_dir = os.path.join(self.data_dir, "gt_interp") + + self.train_list = os.path.join(self.data_dir, "train.txt") + self.test_list = os.path.join(self.data_dir, "test.txt") + self.full_list = os.path.join(self.data_dir, "full.txt") + + self.mono_pred_dir = os.path.join(self.result_dir, "mono_pred") + self.ga_mono_dir = os.path.join(self.result_dir, "global_aligned_mono") + self.rcnet_result_dir = os.path.join(self.result_dir, "rcnet") + + def ensure_data_dirs(self): + for d in ( + self.image_dir, + self.radar_npy_dir, + self.radar_png_dir, + self.gt_dir, + self.gt_interp_dir, + self.result_dir, + self.log_dir, + ): + os.makedirs(d, exist_ok=True) + + +def build_layout(data_root): + return Layout(data_root) diff --git a/src/Baselines/radarcam-depth/sml_inference.py b/src/Baselines/radarcam-depth/sml_inference.py new file mode 100644 index 0000000000000000000000000000000000000000..3312925b829760afb0251da6864d90e14a3ad00c --- /dev/null +++ b/src/Baselines/radarcam-depth/sml_inference.py @@ -0,0 +1,266 @@ +"""SML evaluation: dense metric depth from a trained Scale Map Learner. + +The validate() and log_evaluation_results() functions below are copied +verbatim from RadarCam-Depth/SML/sml_main.py (the baseline's train() half is +replaced by sml_train_rice.py and is not vendored, which also drops the +tensorboard dependency). +""" + +import os +import time + +import numpy as np +import torch +import torch.utils.data + +import data.data_utils as data_utils +import data.SML_dataset as UTV +import modules.midas.transforms as transforms +import modules.midas.utils as utils +import utils.eval_utils as eval_utils +from utils.log_utils import log + + +def validate( + image_paths, + radar_paths, + gt_paths, + sparse_gt_paths, + rcnet_paths, + + best_results, + ScaleMapLearner, + step, + min_radar_valid_depth, + max_radar_valid_depth, + min_eval_depth, + max_eval_depth, + output_path, + + mono_pred_paths = None, + mono_ga_paths = None, + + save_output = False, + random_sample = False, + random_sample_size = 1000, + log_path = None, + depth_predictor = 'dpt_hybrid', + ): + + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + if random_sample: + random_sample_idx = np.random.choice(len(image_paths), random_sample_size, replace=False) + image_paths = [image_paths[idx] for idx in random_sample_idx] + radar_paths = [radar_paths[idx] for idx in random_sample_idx] + gt_paths = [gt_paths[idx] for idx in random_sample_idx] + sparse_gt_paths = [sparse_gt_paths[idx] for idx in random_sample_idx] + rcnet_paths = [rcnet_paths[idx] for idx in random_sample_idx] + if mono_ga_paths is not None: + mono_pred_paths = [mono_pred_paths[idx] for idx in random_sample_idx] + if mono_ga_paths is not None: + mono_ga_paths = [mono_ga_paths[idx] for idx in random_sample_idx] + + val_dataloader = torch.utils.data.DataLoader( + UTV.SML_dataset( + image_paths = image_paths, + radar_paths = radar_paths, + gt_paths = gt_paths, + sparse_gt_paths = sparse_gt_paths, + rcnet_paths = rcnet_paths, + mono_pred_paths = mono_pred_paths, + mono_ga_paths = mono_ga_paths, + ), + batch_size=1, + shuffle=False, + num_workers=1) + + n_sample = len(val_dataloader) + mae = np.zeros(n_sample) + rmse = np.zeros(n_sample) + imae = np.zeros(n_sample) + irmse = np.zeros(n_sample) + abs_rel = np.zeros(n_sample) + sq_rel = np.zeros(n_sample) + delta1 = np.zeros(n_sample) + + save_file_name = os.path.join(output_path, 'RadarCam-Depth') + if save_output: + os.makedirs(save_file_name, exist_ok=True) + os.makedirs(os.path.join(save_file_name, 'sml_depth'), exist_ok=True) + os.makedirs(os.path.join(save_file_name, 'sml_depth_color'), exist_ok=True) + + time_start = time.time() + + for idx, inputs in enumerate(val_dataloader): + inputs = [in_.to(device) for in_ in inputs] + + image, _, _, _, sparse_gt, rcnet, mono_ga_pos = inputs + input_height, input_width = image.shape[1:3] + + # transform + ScaleMapLearner_transform = transforms.get_transforms(depth_predictor, 'void', '150') + + rcnet_valid = (rcnet < max_radar_valid_depth) * (rcnet > min_radar_valid_depth) + rcnet_valid = rcnet_valid.bool() + rcnet[~rcnet_valid] = np.inf # set invalid depth + rcnet = 1.0 / rcnet + mono_ga = 1.0 / mono_ga_pos + + rcnet = rcnet.squeeze().cpu().numpy() + rcnet_valid = rcnet_valid.squeeze().cpu().numpy() + mono_ga_pos = mono_ga_pos.squeeze().cpu().numpy() + int_depth = mono_ga.squeeze().cpu().numpy() + + int_scales = np.ones_like(int_depth) + int_scales[rcnet_valid] = rcnet[rcnet_valid] / int_depth[rcnet_valid] + int_scales = utils.normalize_unit_range(int_scales.astype(np.float32)) + + # transforms + sample = {'image': image.squeeze().cpu().numpy(), + 'int_depth': int_depth, + 'int_scales': int_scales, + 'int_depth_no_tf': int_depth} + + sample = ScaleMapLearner_transform(sample) + + x = torch.cat([sample['int_depth'], sample['int_scales']], 0) + x = x.to(device) + d = sample['int_depth_no_tf'].to(device) + + with torch.no_grad(): + sml_pred, sml_scales = ScaleMapLearner.forward(x.unsqueeze(0), d.unsqueeze(0)) + sml_pred = ( + torch.nn.functional.interpolate( + 1.0 / sml_pred, + size=(input_height, input_width), + mode="bicubic", + align_corners=False, + ) + .squeeze() + .cpu() + .numpy() + ) + + sparse_gt = np.squeeze(sparse_gt.cpu().numpy()) + validity_map = np.where(sparse_gt > 0, 1, 0) + + # Select valid regions to evaluate + validity_mask = np.where(validity_map > 0, 1, 0) + min_max_mask = np.logical_and( + sparse_gt > min_eval_depth, + sparse_gt < max_eval_depth) + mask = np.where(np.logical_and(validity_mask, min_max_mask) > 0) + output_depth = sml_pred[mask] + sparse_gt = sparse_gt[mask] + + # Compute validation metrics + mae[idx] = eval_utils.mean_abs_err(1000.0 * output_depth, 1000.0 * sparse_gt) + rmse[idx] = eval_utils.root_mean_sq_err(1000.0 * output_depth, 1000.0 * sparse_gt) + imae[idx] = eval_utils.inv_mean_abs_err(0.001 * output_depth, 0.001 * sparse_gt) + irmse[idx] = eval_utils.inv_root_mean_sq_err(0.001 * output_depth, 0.001 * sparse_gt) + abs_rel[idx] = eval_utils.mean_abs_rel_err(1000.0 * output_depth, 1000.0 * sparse_gt) + sq_rel[idx] = eval_utils.mean_sq_rel_err(1000.0 * output_depth, 1000.0 * sparse_gt) + delta1[idx] = eval_utils.thr_acc(output_depth, sparse_gt) + print(mae[idx], rmse[idx], imae[idx], irmse[idx], abs_rel[idx], sq_rel[idx], delta1[idx]) + + if save_output: + basename = os.path.basename(image_paths[idx]).split('.')[0] + '.png' + print('Saving output {}'.format(basename)) + sky_mask = mono_ga_pos >= 200 + sml_pred[sky_mask] = mono_ga_pos[sky_mask] + data_utils.save_depth(sml_pred, os.path.join(save_file_name, 'sml_depth', basename)) + data_utils.save_color_depth(sml_pred, os.path.join(save_file_name, 'sml_depth_color', basename)) + + time_end = time.time() + print('Time taken: {:.4f} seconds'.format(time_end - time_start)) + print('average time per sample: {:.4f} seconds'.format((time_end - time_start) / len(image_paths))) + + # Compute mean metrics + mae = np.mean(mae) + rmse = np.mean(rmse) + imae = np.mean(imae) + irmse = np.mean(irmse) + abs_rel = np.mean(abs_rel) + sq_rel = np.mean(sq_rel) + delta1 = np.mean(delta1) + + # Print validation results to console + log_evaluation_results( + title='Validation results', + mae=mae, + rmse=rmse, + imae=imae, + irmse=irmse, + abs_rel=abs_rel, + sq_rel=sq_rel, + delta1=delta1, + step=step, + log_path=log_path) + + n_improve = 0 + if np.round(mae, 4) < np.round(best_results['mae'], 4): + n_improve = n_improve + 1 + if np.round(rmse, 4) < np.round(best_results['rmse'], 4): + n_improve = n_improve + 1 + if np.round(imae, 4) < np.round(best_results['imae'], 4): + n_improve = n_improve + 1 + if np.round(irmse, 4) < np.round(best_results['irmse'], 4): + n_improve = n_improve + 1 + if np.round(abs_rel, 4) < np.round(best_results['abs_rel'], 4): + n_improve = n_improve + 1 + if np.round(sq_rel, 4) < np.round(best_results['sq_rel'], 4): + n_improve = n_improve + 1 + if np.round(delta1, 4) > np.round(best_results['delta1'], 4): + n_improve = n_improve + 1 + + if n_improve > 3: + best_results['step'] = step + best_results['mae'] = mae + best_results['rmse'] = rmse + best_results['imae'] = imae + best_results['irmse'] = irmse + best_results['abs_rel'] = abs_rel + best_results['sq_rel'] = sq_rel + best_results['delta1'] = delta1 + + log_evaluation_results( + title='Best results', + mae=best_results['mae'], + rmse=best_results['rmse'], + imae=best_results['imae'], + irmse=best_results['irmse'], + step=best_results['step'], + abs_rel=best_results['abs_rel'], + sq_rel=best_results['sq_rel'], + delta1=best_results['delta1'], + log_path=log_path) + + return best_results + + +def log_evaluation_results(title, + mae, + rmse, + imae, + irmse, + abs_rel=None, + sq_rel=None, + delta1=None, + step=-1, + log_path=None): + + log(title + ':', log_path) + log('{:>8} {:>8} {:>8} {:>8} {:>8} {:>8} {:>8} {:>8}'.format( + 'Step', 'MAE', 'RMSE', 'iMAE', 'iRMSE', 'Abs_Rel', 'Sq_Rel', 'Delta1'), + log_path) + log('{:8} {:8.3f} {:8.3f} {:8.3f} {:8.3f} {:8.3f} {:8.3f} {:8.3f}'.format( + step, + mae, + rmse, + imae, + irmse, + abs_rel, + sq_rel, + delta1), + log_path) diff --git a/src/Baselines/radarcam-depth/smoke_eval_inference.py b/src/Baselines/radarcam-depth/smoke_eval_inference.py new file mode 100644 index 0000000000000000000000000000000000000000..1d269c2af173642016c86bde5ccefb74abd9ab98 --- /dev/null +++ b/src/Baselines/radarcam-depth/smoke_eval_inference.py @@ -0,0 +1,865 @@ +#!/usr/bin/env python3 +"""Smoke-Eval inference: RC-Net -> SML metric depth, one sequence at a time. + +Runs the complete RadarCam-Depth inference chain (RC-Net quasi-dense depth -> +Scale Map Learner metric depth) over *every* frame of *every* Smoke-Eval +sequence and writes one file per sequence: + + /_pred.npy float32 [N, 1, H, W], depth in METERS + +Nothing has to be pointed at by hand: the script discovers + + * the prepared Smoke-Eval root (ZJU-style layout, see rice_paths.Layout), + * the weights-only RC-Net safetensors file, + * the weights-only SML safetensors file, + +and each of them can still be overridden from the command line. + +DDP and configured mixed precision come from HuggingFace Accelerate. Every +rank takes a disjoint stride of the current sequence's frames. Per-sequence +predictions are gathered on the main process, deduplicated by frame index and +checked for completeness -- a sequence file is written only when all of its +frames are present exactly once. + +The Smoke-Eval data must already be prepared into the ZJU layout with +../dataset_prep (image / radar / gt + global-aligned mono depth), exactly like +the training data: SML consumes the globally aligned monocular depth, which is +produced there and not by this script. + +Usage: + accelerate launch smoke_eval_inference.py + accelerate launch smoke_eval_inference.py --smoke_root /path/to/smoke_data + python smoke_eval_inference.py --config config.yaml --output_dir preds +""" + +import argparse +import contextlib +import os +import pickle +import re +from collections import OrderedDict + +import numpy as np +import torch +import torch.utils.data +from accelerate import Accelerator +from accelerate.utils import set_seed +from safetensors.torch import load_file +from tqdm.auto import tqdm + +import data.data_utils as data_utils +import modules.midas.transforms as sml_transforms +import modules.midas.utils as midas_utils +import rice_paths +from modules.midas.midas_net_custom import MidasNet_small_videpth +from rcnet_inference import forward_with_fallback +from rcnet_model import RCNetModel +from rcnet_transforms import Transforms +from rice_config import load_config + +# Environment overrides, in the spirit of rice_config's RICE_DATA_ROOT. +SMOKE_ROOT_ENV = "SMOKE_EVAL_ROOT" + +# We distribute weights-only safetensors files, not training checkpoints. +CHECKPOINT_PREFERENCE = ("model.safetensors",) + +# "_" (the dataset_prep naming) or "/". +# The greedy sequence group makes the *last* numeric field the frame index, +# so sequence names may themselves contain digits, '-' and '_'. +DEFAULT_NAME_PATTERN = r"^(?P.+)[_/](?P\d+)$" + +DEFAULT_OUTPUT_DIRNAME = "prediction_smoke_eval" + + +# --------------------------------------------------------------------------- +# Discovery: dataset root and checkpoints +# --------------------------------------------------------------------------- + + +def _mono_ga_dirpath(layout, mono_tag): + return os.path.join(layout.ga_mono_dir, mono_tag + "_ls") + + +def _is_prepared_root(path, mono_tag): + """A prepared ZJU-layout root has the inputs both stages need.""" + if not path or not os.path.isdir(path): + return False + layout = rice_paths.build_layout(path) + required = [ + layout.image_dir, + layout.radar_npy_dir, + _mono_ga_dirpath(layout, mono_tag), + ] + return all(os.path.isdir(d) for d in required) + + +def _search_bases(cfg): + """Directories that plausibly hold a prepared Smoke-Eval root.""" + output_root = cfg.data.output_root + bases = [ + rice_paths.HERE, + os.path.dirname(rice_paths.HERE), + output_root, + os.path.dirname(output_root), + os.path.dirname(os.path.dirname(output_root)), + os.getcwd(), + ] + unique = [] + for base in bases: + base = os.path.abspath(base) + if base not in unique: + unique.append(base) + return unique + + +def _auto_candidates(cfg): + """Candidate roots: any 'smoke'-named directory near the project/data.""" + candidates = [] + for base in _search_bases(cfg): + if not os.path.isdir(base): + continue + try: + children = sorted(os.listdir(base)) + except OSError: + continue + for child in children: + if "smoke" not in child.lower(): + continue + path = os.path.join(base, child) + if os.path.isdir(path) and path not in candidates: + candidates.append(path) + return candidates + + +def _resolve_candidate(path, mono_tag): + """Accept the candidate itself or a single prepared root nested in it.""" + if _is_prepared_root(path, mono_tag): + return os.path.abspath(path) + if os.path.isdir(path): + for child in sorted(os.listdir(path)): + nested = os.path.join(path, child) + if _is_prepared_root(nested, mono_tag): + return os.path.abspath(nested) + return None + + +def discover_smoke_root(cfg, explicit, mono_tag): + """Locate the prepared Smoke-Eval root. + + Priority: --smoke_root, $SMOKE_EVAL_ROOT, data.smoke_eval_root in the YAML, + then a scan for 'smoke'-named directories beside the project and the + prepared training data. + """ + explicit_sources = [ + (explicit, "--smoke_root"), + (os.environ.get(SMOKE_ROOT_ENV), "${}".format(SMOKE_ROOT_ENV)), + (cfg.data.get("smoke_eval_root"), "data.smoke_eval_root in the config"), + ] + for path, origin in explicit_sources: + if not path: + continue + resolved = _resolve_candidate(path, mono_tag) + if resolved is None: + raise FileNotFoundError( + "{} points at '{}', which is not a prepared Smoke-Eval root " + "(expected data/image, data/radar and " + "result/global_aligned_mono/{}_ls underneath it).".format( + origin, path, mono_tag + ) + ) + return resolved + + for candidate in _auto_candidates(cfg): + resolved = _resolve_candidate(candidate, mono_tag) + if resolved is not None: + return resolved + + raise FileNotFoundError( + "Could not find a prepared Smoke-Eval root. Looked for directories " + "with 'smoke' in their name under:\n {}\n" + "A prepared root contains data/image, data/radar and " + "result/global_aligned_mono/{}_ls (produce it with " + "../dataset_prep/prepare_all.py on the Smoke-Eval recordings). " + "Pass --smoke_root, set ${}, or add data.smoke_eval_root to the " + "config to point at it explicitly.".format( + "\n ".join(_search_bases(cfg)), mono_tag, SMOKE_ROOT_ENV + ) + ) + + +def discover_checkpoint(save_dir, explicit, label): + """Locate the best checkpoint for one stage. + + Priority: the explicit safetensors path, then the configured save directory. + """ + if explicit: + path = explicit if os.path.isabs(explicit) else os.path.join(rice_paths.HERE, explicit) + if os.path.isfile(path): + return os.path.abspath(path) + if os.path.isdir(path): + save_dir = path + else: + raise FileNotFoundError("{} checkpoint not found: {}".format(label, explicit)) + + if not os.path.isabs(save_dir): + save_dir = os.path.join(rice_paths.HERE, save_dir) + + if not os.path.isdir(save_dir): + raise FileNotFoundError( + "{} checkpoint directory not found: {}".format(label, save_dir) + ) + + for name in CHECKPOINT_PREFERENCE: + path = os.path.join(save_dir, name) + if os.path.isfile(path): + return os.path.abspath(path) + + available = sorted( + f for f in os.listdir(save_dir) if f.endswith(".safetensors") + ) + raise FileNotFoundError( + "No usable {} checkpoint in {} (looked for {}; found {}).".format( + label, save_dir, ", ".join(CHECKPOINT_PREFERENCE), available or "none" + ) + ) + + +# --------------------------------------------------------------------------- +# Frame bookkeeping +# --------------------------------------------------------------------------- + + +def load_names(layout, explicit_list): + """Every frame name of the Smoke-Eval root, in file order, deduplicated.""" + if explicit_list: + list_path = explicit_list + else: + # full.txt covers every prepared frame; test.txt is the fallback when + # the prep wrote only a split. Otherwise read the image directory. + list_path = None + for candidate in (layout.full_list, layout.test_list, layout.train_list): + if os.path.isfile(candidate): + list_path = candidate + break + + if list_path is not None: + if not os.path.isfile(list_path): + raise FileNotFoundError("Frame list not found: {}".format(list_path)) + with open(list_path, "r") as f: + names = [line.strip() for line in f if line.strip()] + else: + if not os.path.isdir(layout.image_dir): + raise FileNotFoundError("Image directory not found: {}".format(layout.image_dir)) + names = sorted( + os.path.splitext(f)[0] + for f in os.listdir(layout.image_dir) + if f.endswith(".png") + ) + + if not names: + raise ValueError("No Smoke-Eval frames found (source: {}).".format(list_path or layout.image_dir)) + + seen = set() + unique = [] + for name in names: + if name not in seen: + seen.add(name) + unique.append(name) + return unique, (list_path or layout.image_dir) + + +def group_by_sequence(names, pattern): + """Map frame names to {sequence: [(frame_idx, name), ...]} sorted by frame.""" + regex = re.compile(pattern) + grouped = OrderedDict() + unmatched = [] + + for name in names: + match = regex.match(name) + if match is None: + unmatched.append(name) + continue + seq = match.group("seq") + frame_idx = int(match.group("frame")) + grouped.setdefault(seq, {}) + # Duplicate frame ids inside a sequence keep the first occurrence. + grouped[seq].setdefault(frame_idx, name) + + if unmatched: + raise ValueError( + "{} frame name(s) do not match the sequence/frame pattern '{}', " + "e.g. {}. Pass --name_pattern with named groups 'seq' and " + "'frame'.".format(len(unmatched), pattern, unmatched[:5]) + ) + + ordered = OrderedDict() + for seq in sorted(grouped): + ordered[seq] = [(idx, grouped[seq][idx]) for idx in sorted(grouped[seq])] + return ordered + + +def _safe_name(seq_name): + return seq_name.replace("/", "_").replace("\\", "_").lower() + + +# --------------------------------------------------------------------------- +# Dataset +# --------------------------------------------------------------------------- + + +class SmokeEvalFrames(torch.utils.data.Dataset): + """Per-frame inputs for the RC-Net + SML chain, from the prepared layout. + + Returns raw arrays; the geometry-dependent parts (bounding boxes, scale + map) are built in the inference loop, where the RC-Net output is known. + """ + + def __init__(self, entries, layout, mono_ga_dir, load_gt): + self.entries = entries + self.image_dir = layout.image_dir + self.radar_dir = layout.radar_npy_dir + self.gt_dir = layout.gt_dir if load_gt else None + self.mono_ga_dir = mono_ga_dir + + def __len__(self): + return len(self.entries) + + def __getitem__(self, index): + frame_idx, name = self.entries[index] + + # 0-255 HWC float; RC-Net's Transforms normalizes it, SML wants [0, 1]. + image = data_utils.load_image( + os.path.join(self.image_dir, name + ".png"), + normalize=False, + data_format="HWC", + ).astype(np.float32) + + radar_points = np.load(os.path.join(self.radar_dir, name + ".npy")) + radar_points = np.asarray(radar_points, dtype=np.float32) + if radar_points.ndim == 1: + radar_points = np.expand_dims(radar_points, axis=0) + if radar_points.size == 0: + radar_points = np.zeros((0, 3), dtype=np.float32) + + mono_ga = _load_metric_depth_png(os.path.join(self.mono_ga_dir, name + ".png")) + mono_ga = np.clip(mono_ga, 1e-3, None) + + sample = { + "frame_idx": int(frame_idx), + "name": name, + "image": image, + "radar_points": radar_points, + "mono_ga": mono_ga, + } + + if self.gt_dir is not None: + sample["gt"] = _load_metric_depth_png( + os.path.join(self.gt_dir, name + ".png") + ) + return sample + + +def _load_metric_depth_png(path): + """16-bit depth PNG -> float32 meters (the data_utils encoding).""" + return data_utils.load_depth(path, data_format="HW").astype(np.float32) + + +def _identity_collate(batch): + """Radar point counts vary per frame, so keep the batch as a list.""" + return batch + + +# --------------------------------------------------------------------------- +# Models +# --------------------------------------------------------------------------- + + +def _strip_module_prefix(state_dict): + return { + (k[len("module."):] if k.startswith("module.") else k): v + for k, v in state_dict.items() + } + + +def build_rcnet(cfg, checkpoint_path, device): + rc = cfg.rcnet + model = RCNetModel( + input_channels_image=3, + input_channels_depth=3, + input_patch_size_image=rc.patch_size, + encoder_type=rc.encoder_type, + n_filters_encoder_image=rc.n_filters_encoder_image, + n_neurons_encoder_depth=rc.n_neurons_encoder_depth, + decoder_type=rc.decoder_type, + n_filters_decoder=rc.n_filters_decoder, + weight_initializer=rc.weight_initializer, + activation_func=rc.activation_func, + device=device, + ) + + checkpoint = load_file(checkpoint_path, device="cpu") + encoder_prefix = "radarnet_encoder." + decoder_prefix = "radarnet_decoder." + encoder_state = { + key[len(encoder_prefix):]: value + for key, value in checkpoint.items() + if key.startswith(encoder_prefix) + } + decoder_state = { + key[len(decoder_prefix):]: value + for key, value in checkpoint.items() + if key.startswith(decoder_prefix) + } + if not encoder_state or not decoder_state: + raise ValueError( + "Unsupported RC-Net checkpoint format: {}".format(checkpoint_path) + ) + # The trainer stores 'module.'-prefixed keys for baseline compatibility; + # each DDP rank here runs an unwrapped copy, so strip the prefix. + model.encoder.load_state_dict( + _strip_module_prefix(encoder_state) + ) + model.decoder.load_state_dict( + _strip_module_prefix(decoder_state) + ) + model.eval() + model.to(device) + for parameter in model.parameters(): + parameter.requires_grad = False + return model + + +def build_sml(cfg, checkpoint_path, accelerator): + # The efficientnet backbone comes from torch.hub on first use; serialize + # the download so DDP ranks do not race for the cache. + with accelerator.main_process_first(): + model = MidasNet_small_videpth( + device="cpu", + min_pred=cfg.depth.min_pred_depth_m, + max_pred=cfg.depth.max_pred_depth_m, + ) + # BaseModel.load() understands the trainer's {"model": ...} payload. + model.load(checkpoint_path) + model.eval() + model.to(accelerator.device) + for parameter in model.parameters(): + parameter.requires_grad = False + return model + + +# --------------------------------------------------------------------------- +# Inference +# --------------------------------------------------------------------------- + + +def rcnet_quasi_dense(model, transforms, image_hwc, radar_points, patch_size, + response_thr, device): + """RC-Net quasi-dense depth for one frame -> float32 [H, W] in meters.""" + height, width = image_hwc.shape[:2] + + if radar_points.shape[0] == 0: + # No radar returns for this frame: SML still runs, on a scale map with + # no anchors. The frame is kept so the sequence stays complete. + return np.zeros((height, width), dtype=np.float32) + + image = torch.from_numpy(np.transpose(image_hwc, (2, 0, 1))).unsqueeze(0).to(device) + points = torch.from_numpy(radar_points).to(device) + + # Same boxes as rcnet_inference.run: full-height patches centered on each + # radar point in the horizontally padded image. Built on the model device + # because torchvision.ops.roi_pool does not move them for us. + pad_size_x = patch_size[1] // 2 + points[:, 0] = points[:, 0] + pad_size_x + x = points[:, 0] + bounding_boxes = torch.stack( + [ + x - pad_size_x, + torch.zeros_like(x), + x + pad_size_x, + torch.full_like(x, float(height)), + ], + dim=1, + ) + bounding_boxes_list = [bounding_boxes] + + [image], [points], [bounding_boxes_list] = transforms.transform( + images_arr=[image], + points_arr=[points], + bounding_boxes_arr=[bounding_boxes_list], + random_transform_probability=0.0, + ) + + output_depth, _, _, inference_failed = forward_with_fallback( + model=model, + image=image, + radar_points=points, + bounding_boxes_list=bounding_boxes_list, + response_thr=response_thr, + device=device, + ) + + if inference_failed: + # Keep going: an all-zero quasi-dense map degrades SML to the + # globally aligned mono depth for this frame rather than dropping it. + return np.zeros((height, width), dtype=np.float32) + + return np.squeeze(output_depth.float().cpu().numpy()).astype(np.float32) + + +def build_sml_sample(image_hwc, rcnet_depth, mono_ga, transform, depth_cfg): + """Scale-map inputs for one frame, matching sml_train_rice's math.""" + rcnet_valid = (rcnet_depth < depth_cfg.max_radar_depth_m) & ( + rcnet_depth > depth_cfg.min_radar_depth_m + ) + + int_depth = (1.0 / mono_ga).astype(np.float32) + int_scales = np.ones_like(int_depth) + int_scales[rcnet_valid] = (1.0 / rcnet_depth[rcnet_valid]) / int_depth[rcnet_valid] + if np.ptp(int_scales) > 0: + int_scales = midas_utils.normalize_unit_range(int_scales.astype(np.float32)) + else: + # No quasi-dense anchors in this frame: constant mid-range scale. + int_scales = np.full_like(int_depth, 0.5) + + sample = { + "image": (image_hwc / 255.0).astype(np.float32), + "int_depth": int_depth, + "int_scales": int_scales, + "int_depth_no_tf": int_depth, + } + sample = transform(sample) + x = torch.cat([sample["int_depth"], sample["int_scales"]], 0) + return x, sample["int_depth_no_tf"] + + +def run_sequence( + seq_name, + entries, + layout, + mono_ga_dir, + rcnet_model, + sml_model, + rcnet_transforms, + sml_transform, + cfg, + accelerator, + args, + output_dir, + gather_dir, +): + """Infer one sequence on all ranks and save _pred.npy on rank 0.""" + device = accelerator.device + rank = accelerator.process_index + world_size = accelerator.num_processes + + # Disjoint stride per rank: every frame is processed exactly once. + rank_entries = entries[rank::world_size] + + compute_metrics = (not args.no_metrics) and os.path.isdir(layout.gt_dir) + + # RC-Net pools ROIs with torchvision.ops.roi_pool, which only has an + # autocast kernel for CUDA (it casts feature map and boxes back to fp32 + # there, as during training). On CPU/MPS autocast would feed the kernel a + # half-precision feature map with fp32 boxes, so keep that stage in fp32. + rcnet_autocast = ( + accelerator.autocast if device.type == "cuda" else contextlib.nullcontext + ) + + loader = torch.utils.data.DataLoader( + SmokeEvalFrames(rank_entries, layout, mono_ga_dir, load_gt=compute_metrics), + batch_size=args.batch_size, + shuffle=False, + num_workers=args.num_workers, + pin_memory=False, + collate_fn=_identity_collate, + ) + + local_predictions = [] # (frame_idx, [1, H, W] float32) + abs_err_sum = 0.0 + sq_err_sum = 0.0 + n_valid_px = 0.0 + + progress = tqdm( + loader, + desc=" {}".format(seq_name), + disable=not accelerator.is_local_main_process, + leave=False, + ) + + for batch in progress: + batch_x = [] + batch_d = [] + + for sample in batch: + with rcnet_autocast(): + rcnet_depth = rcnet_quasi_dense( + model=rcnet_model, + transforms=rcnet_transforms, + image_hwc=sample["image"], + radar_points=sample["radar_points"], + patch_size=cfg.rcnet.patch_size, + response_thr=args.response_thr, + device=device, + ) + x, d = build_sml_sample( + image_hwc=sample["image"], + rcnet_depth=rcnet_depth, + mono_ga=sample["mono_ga"], + transform=sml_transform, + depth_cfg=cfg.depth, + ) + batch_x.append(x) + batch_d.append(d) + + x = torch.stack(batch_x, dim=0).to(device) + d = torch.stack(batch_d, dim=0).to(device) + + with accelerator.autocast(): + pred_inv, _ = sml_model(x, d) + + # Metric depth at the prepared resolution, as in sml_inference.validate. + height, width = batch[0]["image"].shape[:2] + pred_depth = torch.nn.functional.interpolate( + 1.0 / pred_inv.float(), + size=(height, width), + mode="bicubic", + align_corners=False, + ) + pred_np = pred_depth.detach().cpu().numpy().astype(np.float32) + + for i, sample in enumerate(batch): + frame_pred = pred_np[i] # [1, H, W], meters + local_predictions.append((sample["frame_idx"], frame_pred)) + + if compute_metrics: + gt = sample["gt"] + mask = ( + (gt > 0) + & (gt > cfg.depth.min_eval_depth_m) + & (gt < cfg.depth.max_eval_depth_m) + ) + if np.any(mask): + error = frame_pred[0][mask] - gt[mask] + abs_err_sum += float(np.abs(error).sum()) + sq_err_sum += float((error ** 2).sum()) + n_valid_px += float(mask.sum()) + + # Gather this sequence through per-rank files: bounded memory, and no + # object-collective large enough to matter. + accelerator.wait_for_everyone() + safe_seq = _safe_name(seq_name) + rank_file = os.path.join(gather_dir, "rank_{}_{}.pkl".format(rank, safe_seq)) + with open(rank_file, "wb") as f: + pickle.dump(local_predictions, f, protocol=pickle.HIGHEST_PROTOCOL) + + metrics_local = torch.tensor( + [abs_err_sum, sq_err_sum, n_valid_px], + device=device, + dtype=torch.float64, + ) + metrics_global = accelerator.reduce(metrics_local, reduction="sum") + accelerator.wait_for_everyone() + + if accelerator.is_main_process: + merged = [] + for r in range(world_size): + path = os.path.join(gather_dir, "rank_{}_{}.pkl".format(r, safe_seq)) + with open(path, "rb") as f: + merged.extend(pickle.load(f)) + os.remove(path) + + # Deduplicate by frame index (keep first) and order by frame. + by_frame = {} + for frame_idx, pred in merged: + if frame_idx not in by_frame: + by_frame[frame_idx] = pred + + expected = [frame_idx for frame_idx, _ in entries] + missing = [frame_idx for frame_idx in expected if frame_idx not in by_frame] + if missing: + raise RuntimeError( + "{}: missing {} prediction(s), first few frame indices: " + "{}".format(seq_name, len(missing), missing[:5]) + ) + + pred_stack = np.stack( + [by_frame[frame_idx] for frame_idx in sorted(by_frame)], axis=0 + ).astype(np.float32) + + if pred_stack.ndim != 4 or pred_stack.shape[1] != 1: + raise RuntimeError( + "{}: expected [N, 1, H, W], got {}".format(seq_name, pred_stack.shape) + ) + if pred_stack.shape[0] != len(expected): + raise RuntimeError( + "{}: expected {} frames, got {}".format( + seq_name, len(expected), pred_stack.shape[0] + ) + ) + if not np.isfinite(pred_stack).all(): + raise RuntimeError("{}: predictions contain NaN or Inf".format(seq_name)) + + out_path = os.path.join(output_dir, "{}_pred.npy".format(safe_seq)) + np.save(out_path, pred_stack) + + message = " saved {} frames shape={} range=[{:.2f}, {:.2f}] m -> {}".format( + pred_stack.shape[0], + tuple(pred_stack.shape), + float(pred_stack.min()), + float(pred_stack.max()), + out_path, + ) + if compute_metrics and float(metrics_global[2].item()) > 0: + count = float(metrics_global[2].item()) + mae = float(metrics_global[0].item()) / count + rmse = (float(metrics_global[1].item()) / count) ** 0.5 + message += "\n MAE={:.4f} m RMSE={:.4f} m".format(mae, rmse) + print(message, flush=True) + + accelerator.wait_for_everyone() + + +# --------------------------------------------------------------------------- +# Main +# --------------------------------------------------------------------------- + + +def parse_args(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--config", + default=os.path.join(rice_paths.HERE, "config.yaml"), + help="Path to the nested YAML config", + ) + parser.add_argument( + "--smoke_root", + default="", + help="Prepared Smoke-Eval root (default: auto-discovered)", + ) + parser.add_argument( + "--rcnet_checkpoint", + default="", + help="RC-Net .safetensors file", + ) + parser.add_argument( + "--sml_checkpoint", + default="", + help="SML .safetensors file", + ) + parser.add_argument( + "--output_dir", + default="", + help="Where to write _pred.npy (default: ./{})".format( + DEFAULT_OUTPUT_DIRNAME + ), + ) + parser.add_argument( + "--list", + default="", + help="Frame-name list (default: the root's full.txt, else test.txt, " + "else every image in data/image)", + ) + parser.add_argument( + "--name_pattern", + default=DEFAULT_NAME_PATTERN, + help="Regex with named groups 'seq' and 'frame' for frame names", + ) + parser.add_argument("--batch_size", type=int, default=None, help="Per-GPU SML batch") + parser.add_argument("--num_workers", type=int, default=None) + parser.add_argument("--response_thr", type=float, default=None) + parser.add_argument( + "--no_metrics", + action="store_true", + help="Skip the MAE/RMSE sanity check against the prepared ground truth", + ) + return parser.parse_args() + + +def main(): + args = parse_args() + cfg = load_config(args.config) + + if args.batch_size is None: + args.batch_size = cfg.sml.batch_size + if args.num_workers is None: + args.num_workers = cfg.sml.num_workers + if args.response_thr is None: + args.response_thr = cfg.rcnet.response_thr + mixed_precision = cfg.runtime.mixed_precision + + smoke_root = discover_smoke_root(cfg, args.smoke_root, cfg.sml.mono_tag) + layout = rice_paths.build_layout(smoke_root) + mono_ga_dir = _mono_ga_dirpath(layout, cfg.sml.mono_tag) + + rcnet_checkpoint = discover_checkpoint( + cfg.rcnet.save_dir, args.rcnet_checkpoint, "RC-Net" + ) + sml_checkpoint = discover_checkpoint(cfg.sml.save_dir, args.sml_checkpoint, "SML") + + output_dir = args.output_dir or os.path.join(rice_paths.HERE, DEFAULT_OUTPUT_DIRNAME) + output_dir = os.path.abspath(output_dir) + + names, names_source = load_names(layout, args.list) + sequences = group_by_sequence(names, args.name_pattern) + + set_seed(cfg.runtime.seed) + accelerator = Accelerator(mixed_precision=mixed_precision, cpu=cfg.runtime.cpu) + + gather_dir = os.path.join(output_dir, "_gather") + if accelerator.is_main_process: + os.makedirs(output_dir, exist_ok=True) + os.makedirs(gather_dir, exist_ok=True) + print("Smoke-Eval root : {}".format(smoke_root)) + print("Frame list : {}".format(names_source)) + print("RC-Net checkpoint : {}".format(rcnet_checkpoint)) + print("SML checkpoint : {}".format(sml_checkpoint)) + print("Output directory : {}".format(output_dir)) + print( + "Sequences: {} | frames: {} | processes: {} | mixed precision: {}".format( + len(sequences), len(names), accelerator.num_processes, mixed_precision + ) + ) + accelerator.wait_for_everyone() + + rcnet_model = build_rcnet(cfg, rcnet_checkpoint, accelerator.device) + sml_model = build_sml(cfg, sml_checkpoint, accelerator) + rcnet_transforms = Transforms( + normalized_image_range=cfg.rcnet.normalized_image_range + ) + sml_transform = sml_transforms.get_transforms(cfg.sml.mono_tag, "void", "150") + + with torch.no_grad(): + for seq_idx, (seq_name, entries) in enumerate(sequences.items()): + if accelerator.is_main_process: + print( + "\n[{}/{}] {} ({} frames)".format( + seq_idx + 1, len(sequences), seq_name, len(entries) + ), + flush=True, + ) + run_sequence( + seq_name=seq_name, + entries=entries, + layout=layout, + mono_ga_dir=mono_ga_dir, + rcnet_model=rcnet_model, + sml_model=sml_model, + rcnet_transforms=rcnet_transforms, + sml_transform=sml_transform, + cfg=cfg, + accelerator=accelerator, + args=args, + output_dir=output_dir, + gather_dir=gather_dir, + ) + + if accelerator.is_main_process: + if os.path.isdir(gather_dir) and not os.listdir(gather_dir): + os.rmdir(gather_dir) + print("\nInference complete: {} sequence file(s) in {}".format( + len(sequences), output_dir + )) + + +if __name__ == "__main__": + main() diff --git a/src/Baselines/radarcam-depth/utils/eval_utils.py b/src/Baselines/radarcam-depth/utils/eval_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..f48e56edbe680ffb87171a4a6d2dbfc694085e15 --- /dev/null +++ b/src/Baselines/radarcam-depth/utils/eval_utils.py @@ -0,0 +1,105 @@ +"""Metric utilities used by the RadarCam-Depth evaluation backend.""" +import numpy as np + + +def root_mean_sq_err(src, tgt): + ''' + Root mean squared error + Arg(s): + src : numpy[float32] + source array + tgt : numpy[float32] + target array + Returns: + float : root mean squared error + ''' + + return np.sqrt(np.mean((tgt - src) ** 2)) + +def mean_abs_err(src, tgt): + ''' + Mean absolute error + Arg(s): + src : numpy[float32] + source array + tgt : numpy[float32] + target array + Returns: + float : mean absolute error + ''' + + return np.mean(np.abs(tgt - src)) + +def inv_root_mean_sq_err(src, tgt): + ''' + Inverse root mean squared error + Arg(s): + src : numpy[float32] + source array + tgt : numpy[float32] + target array + Returns: + float : inverse root mean squared error + ''' + + return np.sqrt(np.mean(((1.0 / tgt) - (1.0 / src)) ** 2)) + +def inv_mean_abs_err(src, tgt): + ''' + Inverse mean absolute error + Arg(s): + src : numpy[float32] + source array + tgt : numpy[float32] + target array + Returns: + float : inverse mean absolute error + ''' + + return np.mean(np.abs((1.0 / tgt) - (1.0 / src))) + +def mean_abs_rel_err(src, tgt): + ''' + Mean absolute relative error (normalize absolute error) + Arg(s): + src : numpy[float32] + source array + tgt : numpy[float32] + target array + Returns: + float : mean absolute relative error between source and target + ''' + + return np.mean(np.abs(src - tgt) / tgt) + + +def mean_sq_rel_err(src, tgt): + ''' + Mean squared relative error (normalize squared error) + Arg(s): + src : numpy[float32] + source array + tgt : numpy[float32] + target array + Returns: + float : mean squared relative error between source and target + ''' + + return np.mean(((src - tgt) ** 2) / tgt) + + +def thr_acc(src, tgt, thr=1.25): + ''' + Threshold accuracy + Arg(s): + src : numpy[float32] + source array + tgt : numpy[float32] + target array + thr : float + threshold + Returns: + float : threshold accuracy + ''' + + return np.mean(np.maximum((tgt / src), (src / tgt)) < thr) diff --git a/src/Baselines/radarcam-depth/utils/log_utils.py b/src/Baselines/radarcam-depth/utils/log_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..7e906c4087d67bc372acf59ac26616aa42c50b67 --- /dev/null +++ b/src/Baselines/radarcam-depth/utils/log_utils.py @@ -0,0 +1,70 @@ +import os +import torch +import numpy as np +from matplotlib import pyplot as plt + + +def log(s, filepath=None, to_console=True): + ''' + Logs a string to either file or console + Arg(s): + s : str + string to log + filepath + output filepath for logging + to_console : bool + log to console + ''' + + if to_console: + print(s) + + if filepath is not None: + if not os.path.isdir(os.path.dirname(filepath)): + os.makedirs(os.path.dirname(filepath)) + with open(filepath, 'w+') as o: + o.write(s + '\n') + else: + with open(filepath, 'a+') as o: + o.write(s + '\n') + + +def colorize(T, colormap='magma', return_numpy=False): + ''' + Colorizes a 1-channel tensor with matplotlib colormaps + Arg(s): + T : torch.Tensor[float32] + 1-channel tensor + colormap : str + matplotlib colormap + ''' + + cm = plt.cm.get_cmap(colormap) + shape = T.shape + + # Convert to numpy array and transpose + if shape[0] > 1: + T = np.squeeze(np.transpose(T.cpu().numpy(), (0, 2, 3, 1))) + else: + T = np.squeeze(np.transpose(T.cpu().numpy(), (0, 2, 3, 1)), axis=-1) + + # Colorize using colormap + color = np.concatenate([ + np.expand_dims(cm(T[n, ...])[..., 0:3], 0) for n in range(T.shape[0])], + axis=0) + + if return_numpy: + return color + else: + # Transpose back to torch format + color = np.transpose(color, (0, 3, 1, 2)) + + # Convert back to tensor + return torch.from_numpy(color.astype(np.float32)) + + + +def log_params(log_path, params_dict): + with open(log_path, 'w') as log_file: + for param_name, param_value in params_dict.items(): + log_file.write(f"{param_name}: {param_value}\n") \ No newline at end of file diff --git a/src/Baselines/radarcam-depth/utils/net_utils.py b/src/Baselines/radarcam-depth/utils/net_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..36ef9960e54b074b1e0aeed16a0b31623e23a0cb --- /dev/null +++ b/src/Baselines/radarcam-depth/utils/net_utils.py @@ -0,0 +1,638 @@ +import torch + + +def activation_func(activation_fn): + ''' + Select activation function + Arg(s): + activation_fn : str + name of activation function + ''' + + if 'linear' in activation_fn: + return None + elif 'leaky_relu' in activation_fn: + return torch.nn.LeakyReLU(negative_slope=0.20, inplace=True) + elif 'relu' in activation_fn: + return torch.nn.ReLU() + elif 'elu' in activation_fn: + return torch.nn.ELU() + elif 'sigmoid' in activation_fn: + return torch.nn.Sigmoid() + else: + raise ValueError('Unsupported activation function: {}'.format(activation_fn)) + + +''' +Network layers +''' +class Conv2d(torch.nn.Module): + ''' + 2D convolution class + + Arg(s): + in_channels : int + number of input channels + out_channels : int + number of output channels + kernel_size : int + size of kernel + stride : int + stride of convolution + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + def __init__(self, + in_channels, + out_channels, + kernel_size=3, + stride=1, + weight_initializer='kaiming_uniform', + activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True), + use_batch_norm=False): + super(Conv2d, self).__init__() + + self.use_batch_norm = use_batch_norm + padding = kernel_size // 2 + + self.conv = torch.nn.Conv2d( + in_channels, + out_channels, + kernel_size=kernel_size, + stride=stride, + padding=padding, + bias=False) + + # Select the type of weight initialization, by default kaiming_uniform + if weight_initializer == 'kaiming_normal': + torch.nn.init.kaiming_normal_(self.conv.weight) + elif weight_initializer == 'xavier_normal': + torch.nn.init.xavier_normal_(self.conv.weight) + elif weight_initializer == 'xavier_uniform': + torch.nn.init.xavier_uniform_(self.conv.weight) + + self.activation_func = activation_func + + if self.use_batch_norm: + self.batch_norm = torch.nn.BatchNorm2d(out_channels) + + def forward(self, x): + conv = self.conv(x) + conv = self.batch_norm(conv) if self.use_batch_norm else conv + + if self.activation_func is not None: + return self.activation_func(conv) + else: + return conv + + +class TransposeConv2d(torch.nn.Module): + ''' + Transpose convolution class + + Arg(s): + in_channels : int + number of input channels + out_channels : int + number of output channels + kernel_size : int + size of kernel (k x k) + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + def __init__(self, + in_channels, + out_channels, + kernel_size=3, + weight_initializer='kaiming_uniform', + activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True), + use_batch_norm=False): + super(TransposeConv2d, self).__init__() + + self.use_batch_norm = use_batch_norm + padding = kernel_size // 2 + + self.deconv = torch.nn.ConvTranspose2d( + in_channels, + out_channels, + kernel_size=kernel_size, + stride=2, + padding=padding, + output_padding=1, + bias=False) + + # Select the type of weight initialization, by default kaiming_uniform + if weight_initializer == 'kaiming_normal': + torch.nn.init.kaiming_normal_(self.conv.weight) + elif weight_initializer == 'xavier_normal': + torch.nn.init.xavier_normal_(self.conv.weight) + elif weight_initializer == 'xavier_uniform': + torch.nn.init.xavier_uniform_(self.conv.weight) + + self.activation_func = activation_func + + if self.use_batch_norm: + self.batch_norm = torch.nn.BatchNorm2d(out_channels) + + def forward(self, x): + deconv = self.deconv(x) + deconv = self.batch_norm(deconv) if self.use_batch_norm else deconv + if self.activation_func is not None: + return self.activation_func(deconv) + else: + return deconv + + +class UpConv2d(torch.nn.Module): + ''' + Up-convolution (upsample + convolution) block class + + Arg(s): + in_channels : int + number of input channels + out_channels : int + number of output channels + shape : list[int] + two element tuple of ints (height, width) + kernel_size : int + size of kernel (k x k) + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + def __init__(self, + in_channels, + out_channels, + kernel_size=3, + weight_initializer='kaiming_uniform', + activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True), + use_batch_norm=False): + super(UpConv2d, self).__init__() + + self.conv = Conv2d( + in_channels, + out_channels, + kernel_size=kernel_size, + stride=1, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + def forward(self, x, shape): + upsample = torch.nn.functional.interpolate(x, size=shape) + conv = self.conv(upsample) + return conv + + +class FullyConnected(torch.nn.Module): + ''' + Fully connected layer + + Arg(s): + in_channels : int + number of input neurons + out_channels : int + number of output neurons + dropout_rate : float + probability to use dropout + ''' + + def __init__(self, + in_features, + out_features, + weight_initializer='kaiming_uniform', + activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True), + dropout_rate=0.00): + super(FullyConnected, self).__init__() + + self.fully_connected = torch.nn.Linear(in_features, out_features) + + if weight_initializer == 'kaiming_normal': + torch.nn.init.kaiming_normal_(self.fully_connected.weight) + elif weight_initializer == 'xavier_normal': + torch.nn.init.xavier_normal_(self.fully_connected.weight) + elif weight_initializer == 'xavier_uniform': + torch.nn.init.xavier_uniform_(self.fully_connected.weight) + + self.activation_func = activation_func + + if dropout_rate > 0.00 and dropout_rate <= 1.00: + self.dropout = torch.nn.Dropout(p=dropout_rate) + else: + self.dropout = None + + def forward(self, x): + fully_connected = self.fully_connected(x) + + if self.activation_func is not None: + fully_connected = self.activation_func(fully_connected) + + if self.dropout is not None: + return self.dropout(fully_connected) + else: + return fully_connected + + +''' +Network encoder blocks +''' +class ResNetBlock(torch.nn.Module): + ''' + Basic ResNet block class + Arg(s): + in_channels : int + number of input channels + out_channels : int + number of output channels + stride : int + stride of convolution + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + def __init__(self, + in_channels, + out_channels, + stride=1, + weight_initializer='kaiming_uniform', + activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True), + use_batch_norm=False): + super(ResNetBlock, self).__init__() + + self.activation_func = activation_func + + self.conv1 = Conv2d( + in_channels, + out_channels, + kernel_size=3, + stride=stride, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + self.conv2 = Conv2d( + out_channels, + out_channels, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + self.projection = Conv2d( + in_channels, + out_channels, + kernel_size=1, + stride=stride, + weight_initializer=weight_initializer, + activation_func=None, + use_batch_norm=False) + + def forward(self, x): + # Perform 2 convolutions + conv1 = self.conv1(x) + conv2 = self.conv2(conv1) + + # Perform projection if (1) shape does not match (2) channels do not match + in_shape = list(x.shape) + out_shape = list(conv2.shape) + if in_shape[2:4] != out_shape[2:4] or in_shape[1] != out_shape[1]: + X = self.projection(x) + else: + X = x + + # f(x) + x + return self.activation_func(conv2 + X) + + +class ResNetBottleneckBlock(torch.nn.Module): + ''' + ResNet bottleneck block class + + Arg(s): + in_channels : int + number of input channels + out_channels : int + number of output channels + stride : int + stride of convolution + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + def __init__(self, + in_channels, + out_channels, + stride=1, + weight_initializer='kaiming_uniform', + activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True), + use_batch_norm=False): + super(ResNetBottleneckBlock, self).__init__() + + self.activation_func = activation_func + + self.conv1 = Conv2d( + in_channels, + out_channels, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + self.conv2 = Conv2d( + out_channels, + out_channels, + kernel_size=3, + stride=stride, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + self.conv3 = Conv2d( + out_channels, + 4 * out_channels, + kernel_size=1, + stride=1, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + self.projection = Conv2d( + in_channels, + 4 * out_channels, + kernel_size=1, + stride=stride, + weight_initializer=weight_initializer, + activation_func=None, + use_batch_norm=False) + + def forward(self, x): + # Perform 2 convolutions + conv1 = self.conv1(x) + conv2 = self.conv2(conv1) + conv3 = self.conv3(conv2) + + # Perform projection if (1) shape does not match (2) channels do not match + in_shape = list(x.shape) + out_shape = list(conv2.shape) + if in_shape[2:4] != out_shape[2:4] or in_shape[1] != out_shape[1]: + X = self.projection(x) + else: + X = x + + # f(x) + x + return self.activation_func(conv3 + X) + + +class VGGNetBlock(torch.nn.Module): + ''' + VGGNet block class + + Arg(s): + in_channels : int + number of input channels + out_channels : int + number of output channels + n_conv : int + number of convolution layers + stride : int + stride of convolution + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + ''' + + def __init__(self, + in_channels, + out_channels, + n_conv=1, + stride=1, + weight_initializer='kaiming_uniform', + activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True), + use_batch_norm=False): + super(VGGNetBlock, self).__init__() + + layers = [] + for n in range(n_conv - 1): + conv = Conv2d( + in_channels, + out_channels, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + layers.append(conv) + in_channels = out_channels + + conv = Conv2d( + in_channels, + out_channels, + kernel_size=3, + stride=stride, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + layers.append(conv) + + self.conv_block = torch.nn.Sequential(*layers) + + def forward(self, x): + return self.conv_block(x) + + +''' +Network decoder blocks +''' +class DecoderBlock(torch.nn.Module): + ''' + Decoder block with skip connection + + Arg(s): + in_channels : int + number of input channels + skip_channels : int + number of skip connection channels + out_channels : int + number of output channels + weight_initializer : str + kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform + activation_func : func + activation function after convolution + use_batch_norm : bool + if set, then applied batch normalization + deconv_type : str + deconvolution types: transpose, up + ''' + + def __init__(self, + in_channels, + skip_channels, + out_channels, + weight_initializer='kaiming_uniform', + activation_func=torch.nn.LeakyReLU(negative_slope=0.10, inplace=True), + use_batch_norm=False, + deconv_type='up'): + super(DecoderBlock, self).__init__() + + self.skip_channels = skip_channels + self.deconv_type = deconv_type + + if deconv_type == 'transpose': + self.deconv = TransposeConv2d( + in_channels, + out_channels, + kernel_size=3, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + elif deconv_type == 'up': + self.deconv = UpConv2d( + in_channels, + out_channels, + kernel_size=3, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + concat_channels = skip_channels + out_channels + + self.conv = Conv2d( + concat_channels, + out_channels, + kernel_size=3, + stride=1, + weight_initializer=weight_initializer, + activation_func=activation_func, + use_batch_norm=use_batch_norm) + + def forward(self, x, skip=None, shape=None): + ''' + Forward input x through a decoder block and fuse with skip connection + + Arg(s): + x : torch.Tensor[float32] + N x C x h x w input tensor + skip : torch.Tensor[float32] + N x F x h x w skip connection + shape : tuple[int] + height, width (H, W) tuple denoting output shape + Returns: + torch.Tensor[float32] : N x K x H x W output tensor + ''' + + if self.deconv_type == 'transpose': + deconv = self.deconv(x) + elif self.deconv_type == 'up': + + if skip is not None: + shape = skip.shape[2:4] + elif shape is not None: + pass + else: + n_height, n_width = x.shape[2:4] + shape = (int(2 * n_height), int(2 * n_width)) + + deconv = self.deconv(x, shape=shape) + + if self.skip_channels > 0: + concat = torch.cat([deconv, skip], dim=1) + else: + concat = deconv + + return self.conv(concat) + + +''' +Utility function to pre-process sparse depth and input depth +''' +class OutlierRemoval(object): + ''' + Class to perform outlier removal based on depth difference in local neighborhood + + Arg(s): + kernel_size : int + local neighborhood to consider + threshold : float + depth difference threshold + ''' + + def __init__(self, kernel_size=7, threshold=1.5): + + self.kernel_size = kernel_size + self.threshold = threshold + + def remove_outliers(self, depth): + ''' + Removes erroneous measurements from sparse depth + + Arg(s): + depth : torch.Tensor[float32] + N x 1 x H x W tensor sparse depth + Returns: + torch.Tensor[float32] : N x 1 x H x W depth + ''' + + # Get valid locations + validity_map = torch.where( + depth > 0.0, + torch.ones_like(depth), + depth) + + # Replace all zeros with large values + max_value = 10 * torch.max(depth) + depth_max_filled = torch.where( + validity_map <= 0, + torch.full_like(depth, fill_value=max_value), + depth) + + # For each neighborhood find the smallest value + padding = self.kernel_size // 2 + depth_max_filled = torch.nn.functional.pad( + input=depth_max_filled, + pad=(padding, padding, padding, padding), + mode='constant', + value=max_value) + + min_values = -torch.nn.functional.max_pool2d( + input=-depth_max_filled, + kernel_size=self.kernel_size, + stride=1, + padding=0) + + # If measurement differs a lot from minimum value then remove + validity_map_clean = torch.where( + min_values < depth - self.threshold, + torch.zeros_like(validity_map), + torch.ones_like(validity_map)) + + # Update depth map + depth_clean = depth * validity_map_clean + + return depth_clean diff --git a/src/GRADE/stage2_diffusion_refinement/dataloader.py b/src/GRADE/stage2_diffusion_refinement/dataloader.py new file mode 100644 index 0000000000000000000000000000000000000000..4c75232207fcdd71c3a250411303c1cd5056127b --- /dev/null +++ b/src/GRADE/stage2_diffusion_refinement/dataloader.py @@ -0,0 +1,178 @@ +import json +import re +from typing import Optional, List, Dict, Tuple, Literal +from pathlib import Path +from torch.utils.data import DataLoader +import torch + +from rice_dataset import RiceDataset + + +_HELD_OUT_BUILDING_PATTERN = re.compile(r"keck|dell", re.IGNORECASE) + + +def _load_split_config() -> Dict[str, List[str]]: + """Load split configuration from split.json.""" + split_file = Path(__file__).parent / "split.json" + if not split_file.exists(): + raise FileNotFoundError(f"Split configuration not found: {split_file}") + + with open(split_file, "r") as f: + split_config = json.load(f) + + return split_config + + +def create_dataloader( + rice_root_dir: Optional[str] = None, + smoke_eval_root_dir: Optional[str] = None, + split: Literal["train", "val", "test"] = "train", + batch_size: int = 4, + shuffle: bool = True, + num_workers: int = 0, + frame_skip: int = 1, + num_frames: int = 1, + drop_last: bool = False, + seed: int = 42, + # Processing parameters + scale_factor: float = 0.001, + max_depth_m: float = 11.2, + depth_resolution: Tuple[int, int] = (128, 256), + use_rgb: bool = True, + rgb_resolution: Tuple[int, int] = (128, 256), +) -> DataLoader: + """Create a dataloader for the Rice dataset. + + Split rules + ----------- + train : Rice sequences NOT in test-rice (split.json), with an additional + case-insensitive exclusion for every Keck or Dell sequence + val : Rice sequences listed in test-rice (split.json) + test : All sequences discovered under smoke_eval_root_dir + + Args: + rice_root_dir: Root directory for Rice dataset (required for train/val) + smoke_eval_root_dir: Root directory for Smoke-Eval test sequences. + Required when split="test". + split: "train", "val", or "test" + batch_size: Batch size + shuffle: Whether to shuffle (train/val only; test is always sequential) + num_workers: Number of data loading workers + frame_skip: Frame sampling stride + num_frames: Number of frames F in the causal sliding window. + drop_last: Drop the last incomplete batch (use True for DDP training). + seed: Shared shuffle seed across all DDP ranks. + scale_factor: Radar amplitude scale factor + max_depth_m: Maximum depth in meters + depth_resolution: (height, width) for depth resizing + use_rgb: Whether to include RGB + rgb_resolution: (height, width) for RGB resizing + + Returns: + DataLoader instance + """ + split_config = _load_split_config() + val_rice_seqs = split_config.get("test-rice", []) + val_rice_seq_keys = {seq.casefold() for seq in val_rice_seqs} + + proc_kwargs = { + "scale_factor": scale_factor, + "max_depth_m": max_depth_m, + "depth_resolution": depth_resolution, + "use_rgb": use_rgb, + "rgb_resolution": rgb_resolution, + } + + if split in ("train", "val"): + if rice_root_dir is None: + raise ValueError("rice_root_dir is required for split='train' or 'val'") + rice_path = Path(rice_root_dir) + if not rice_path.exists(): + raise ValueError(f"rice_root_dir does not exist: {rice_path}") + + all_rice_seqs = sorted( + d.name + for d in rice_path.iterdir() + if d.is_dir() + and (d / "radar.npy").exists() + and (d / "zed_depth.npy").exists() + ) + + if split == "val": + sequences = [s for s in all_rice_seqs if s.casefold() in val_rice_seq_keys] + else: + sequences = [ + s + for s in all_rice_seqs + if s.casefold() not in val_rice_seq_keys + and _HELD_OUT_BUILDING_PATTERN.search(s) is None + ] + + if not sequences: + raise ValueError(f"No valid Rice sequences found for split '{split}'") + + dataset = RiceDataset( + str(rice_path), + sequences=sequences, + frame_skip=frame_skip, + num_frames=num_frames, + **proc_kwargs, + ) + + else: # test + if smoke_eval_root_dir is None: + raise ValueError("smoke_eval_root_dir is required for split='test'") + smoke_path = Path(smoke_eval_root_dir) + if not smoke_path.exists(): + raise ValueError(f"smoke_eval_root_dir does not exist: {smoke_path}") + + smoke_seqs = sorted( + d.name + for d in smoke_path.iterdir() + if d.is_dir() and not d.name.startswith(".") + ) + if not smoke_seqs: + raise ValueError("No sequences found in smoke_eval_root_dir") + + dataset = RiceDataset( + str(smoke_path), + sequences=smoke_seqs, + frame_skip=frame_skip, + num_frames=num_frames, + **proc_kwargs, + ) + + use_shuffle = shuffle and split in ("train", "val") + generator = None + if use_shuffle: + generator = torch.Generator() + generator.manual_seed(seed) + + return DataLoader( + dataset, + batch_size=batch_size, + shuffle=use_shuffle, + generator=generator, + num_workers=num_workers, + drop_last=drop_last, + ) + + +# ── Convenience helpers ──────────────────────────────────────────────────────── + + +def create_train_dataloader(**kwargs) -> DataLoader: + """Training dataloader (Rice non-val sequences).""" + return create_dataloader(split="train", **kwargs) + + +def create_val_dataloader(**kwargs) -> DataLoader: + """Validation dataloader (Rice test-rice sequences from split.json).""" + return create_dataloader(split="val", **kwargs) + + +def create_test_dataloader(**kwargs) -> DataLoader: + """Test dataloader (all Smoke-Eval sequences).""" + return create_dataloader(split="test", shuffle=False, **kwargs) + + diff --git a/src/GRADE/stage2_diffusion_refinement/inference_diffusion.py b/src/GRADE/stage2_diffusion_refinement/inference_diffusion.py new file mode 100644 index 0000000000000000000000000000000000000000..f63f2e7afe81d2bf06a3411bf948ceb28524a4f4 --- /dev/null +++ b/src/GRADE/stage2_diffusion_refinement/inference_diffusion.py @@ -0,0 +1,15 @@ +#!/usr/bin/env python3 +"""Run diffusion inference with validation-matched RadarDepth conditioning.""" + +from inference import parse_stage_args, run_diffusion + + +def main() -> None: + args = parse_stage_args( + "Run ours_diffusion with the train_diffusion.validate() sampling setup." + ) + run_diffusion(args.config) + + +if __name__ == "__main__": + main() diff --git a/src/GRADE/stage2_diffusion_refinement/inference_radar.py b/src/GRADE/stage2_diffusion_refinement/inference_radar.py new file mode 100644 index 0000000000000000000000000000000000000000..3c81684ef40ee331e15eb2fdebeced9ffbc09dfa --- /dev/null +++ b/src/GRADE/stage2_diffusion_refinement/inference_radar.py @@ -0,0 +1,13 @@ +#!/usr/bin/env python3 +"""Run the radar-only Smoke-Eval inference pass.""" + +from inference import parse_stage_args, run_radar + + +def main() -> None: + args = parse_stage_args("Run ours_radar inference on all Smoke-Eval sequences.") + run_radar(args.config) + + +if __name__ == "__main__": + main() diff --git a/src/GRADE/stage2_diffusion_refinement/split.json b/src/GRADE/stage2_diffusion_refinement/split.json new file mode 100644 index 0000000000000000000000000000000000000000..5d7addca12b366ae7facdf3f377bb09bfeb5225d --- /dev/null +++ b/src/GRADE/stage2_diffusion_refinement/split.json @@ -0,0 +1,14 @@ +{ + "test-rice": [ + "Dell-1", + "Dell-2", + "Smoke-Dell-1", + "Smoke-Dell-2", + "Keck-1", + "Keck-2", + "Keck-3", + "Smoke-keck-1", + "Smoke-keck-2", + "Smoke-keck-3" + ] +} \ No newline at end of file diff --git a/src/GRADE/stage2_diffusion_refinement/weather_cpu.py b/src/GRADE/stage2_diffusion_refinement/weather_cpu.py new file mode 100644 index 0000000000000000000000000000000000000000..c7365063ee98fa5c6b86ccb66f9d9cbe163fab8a --- /dev/null +++ b/src/GRADE/stage2_diffusion_refinement/weather_cpu.py @@ -0,0 +1,287 @@ +""" +Weather Simulation Module (Functional API) + +Implements fog and cloud effects using atmospheric scattering model. +Refactored to pure functions without a class wrapper. + +Usage: + from weather_cpu import add_fog, add_cloud + + # Fog: intensity 0-10 + foggy, mask = add_fog(img, intensity=5, return_attenuation=True) + + # Cloud: precise control + cloudy, mask = add_cloud(img, num_clouds=5, cloud_size=0.3, intensity=8, + color="white", return_attenuation=True) +""" + +import os +import random +import numpy as np +import cv2 +from typing import Literal, Optional, Tuple, Union + +try: + from noise import pnoise3, pnoise2 + + HAS_NOISE = True +except ImportError: + HAS_NOISE = False + print("Warning: 'noise' package not installed. Using random noise fallback.") + + +def _gen_perlin_noise(shape: Tuple[int, int], scale: float = 100.0) -> np.ndarray: + """Generate Perlin noise for realistic fog density.""" + if not HAS_NOISE: + noise = np.random.randn(*shape) * 50 + 128 + noise = cv2.GaussianBlur(noise.astype(np.float32), (31, 31), 0) + return np.clip(noise, 0, 255) + + h, w = shape + # Downsample for performance then resize + d_h, d_w = min(h, 480), min(w, 480) + + noise = np.zeros((d_h, d_w), dtype=np.float32) + s = 1.0 / scale + + # Simple 3-octave fractal noise + for y in range(d_h): + for x in range(d_w): + v = pnoise2(x * s, y * s, octaves=4, persistence=0.5, lacunarity=2.0) + noise[y, x] = (v + 1) * 128.0 + + return cv2.resize(noise, (w, h), interpolation=cv2.INTER_CUBIC) + + +def add_fog( + image: np.ndarray, + intensity: float = 5.0, # 0 to 10 + return_attenuation: bool = False, + cam_height: float = 20, + fog_height: float = 100, + haze_height: float = 35, + airlight_color: Tuple[int, int, int] = (210, 210, 210), +) -> Union[np.ndarray, Tuple[np.ndarray, np.ndarray]]: + """ + Add fog effect to image. + + Args: + image: np.ndarray [H, W, 3] uint8 + intensity: float [0, 10], 0=No Fog, 10=Max Fog + return_attenuation: bool, return (image, mask) if True + + Returns: + Augmented Image (uint8) OR (Augmented Image, Attenuation Map [0,1]) + """ + if intensity <= 0: + if return_attenuation: + return image, np.ones(image.shape[:2], dtype=np.float32) + return image + + h, w = image.shape[:2] + image_float = image.astype(np.float32) + + # Map Intensity [0, 10] -> Visibility [High, Low] -> Beta [Low, High] + # We map intensity directly to extinction coefficient (beta) + # Intensity 10 -> visibility ~20m -> beta ~ 0.2 + # Intensity 1 -> visibility ~1000m -> beta ~ 0.004 + max_beta = 0.2 # at intensity 10 + beta = (intensity / 10.0) * max_beta + + # Airlight (fog color) - usually white/gray for fog + airlight_color = np.asarray(airlight_color, dtype=np.float32) + + # 1. Height-based density + elevation = np.ones((h, w), dtype=np.float32) * cam_height + c = 1 - elevation / (fog_height + 1e-5) + c = np.clip(c, 0, 1) + + # 2. Perlin Noise density variation + noise = _gen_perlin_noise((h, w), scale=100.0 if w > 500 else 50.0) + noise_factor = noise / 255.0 + + # Combined Extinction Coefficient + # beta_spatial = beta * c * noise_factor + # Simplified: beta * noise helps create pockets + # We keep the height term 'c' to make it ground-hugging if desired, + # but for general camera views, uniform depth + noise is often enough. + # Let's mix uniform and noisy for robust "intensity" feel. + beta_spatial = beta * (0.6 + 0.4 * noise_factor) + + # 3. Distance Map (Scene Depth) + # Without depth map, we assume a "corridor" or flat plane recession + # y-coordinate approximation: bottom of image is close, top is far/sky + # Normalized Y from 1.0 (bottom) to 4.0 (horizon/top) + y_grad = np.linspace(4.0, 1.0, h).astype(np.float32) # shape (h,) + # specific distance model for visual effect + depth_proxy = np.tile(y_grad[:, np.newaxis], (1, w)) + # Scale depth by fog height concept + distance = depth_proxy * 10.0 # arbitrary scale meters + + # 4. Beer-Lambert Transmittance + # T = exp(-beta * d) + attenuation = np.exp(-beta_spatial * distance) + + # Apply + # I = J * T + A * (1 - T) + attenuation_3c = attenuation[:, :, np.newaxis] + airlight_3c = np.ones_like(image_float) * airlight_color + + out = image_float * attenuation_3c + airlight_3c * (1 - attenuation_3c) + out = np.clip(out, 0, 255).astype(np.uint8) + + if return_attenuation: + return out, attenuation.astype(np.float32) + return out + + +def add_cloud( + image: np.ndarray, + num_clouds: int = 4, + cloud_size: Union[float, int] = 0.4, # If float < 2, treated as ratio of min_dim + intensity: float = 8.0, # 0-10 opacity/density + color: Literal["white", "gray"] = "white", + return_attenuation: bool = False, +) -> Union[np.ndarray, Tuple[np.ndarray, np.ndarray]]: + """ + Add synthetic clouds to image. + + Args: + image: np.ndarray [H, W, 3] + num_clouds: Number of cloud patches + cloud_size: Size of clouds. If < 2.0, treated as ratio of image min dimension. + If > 2.0, treated as pixel size. + intensity: [0, 10] cloud opacity/thickness + color: "white" or "gray" + return_attenuation: Return mask + """ + if num_clouds <= 0 or intensity <= 0: + if return_attenuation: + return image, np.ones(image.shape[:2], dtype=np.float32) + return image + + h, w = image.shape[:2] + min_dim = min(h, w) + image_float = image.astype(np.float32) + + # Determine absolute size + if cloud_size < 2.0: + base_size = int(min_dim * cloud_size) + else: + base_size = int(cloud_size) + + # Create Cloud Mask (Accumulator) + # Starts at 0 (Clear/Transparent) -> 1 (Opaque/Cloudy) + # We will invert to Transmittance at the end (1=Clear, 0=Blocked) + cloud_density_map = np.zeros((h, w), dtype=np.float32) + + # Color + if color == "white": + cloud_rgb = np.array([210, 210, 210], dtype=np.float32) + else: # gray + cloud_rgb = np.array([50, 50, 50], dtype=np.float32) + + # Pattern Directory (if available, else synthetic) + pattern_dir = os.path.join( + os.path.dirname(os.path.dirname(os.path.dirname(__file__))), + "AdverseWeatherSimulation", + "data", + "patterns", + ) + has_patterns = os.path.exists(pattern_dir) and len(os.listdir(pattern_dir)) > 0 + + for _ in range(num_clouds): + # Random location + cx = np.random.randint(0, w) + cy = np.random.randint( + 0, h // 2 + ) # Clouds usually in sky/top half? Let's allow full image for "foggy cloud" + # Actually user might want fog-like clouds anywhere. Let's do full range but bias top? + # User asked for "add_cloud", usually implies sky, but for overlay tests fully random is safer. + cy = np.random.randint(0, h) + + # Randomize size slightly + this_size = int(base_size * random.uniform(0.8, 1.2)) + if this_size < 10: + this_size = 10 + + # Generate/Load Patch + patch = None + if has_patterns: + try: + fname = random.choice( + [f for f in os.listdir(pattern_dir) if f.endswith(("png", "jpg"))] + ) + img_path = os.path.join(pattern_dir, fname) + patch = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) + if patch is not None: + patch = ( + cv2.resize(patch, (this_size, this_size)).astype(np.float32) + / 255.0 + ) + except: + pass + + if patch is None: + # Synthetic Blob + y, x = np.ogrid[ + -this_size // 2 : this_size // 2, -this_size // 2 : this_size // 2 + ] + dist = np.sqrt(x * x + y * y) + radius = this_size // 2 + patch = (1 - dist / radius).clip(0, 1) + # Add noise + noise = np.random.rand(*patch.shape) * 0.4 + patch = (patch + noise * patch).clip(0, 1) + patch = cv2.GaussianBlur(patch, (15, 15), 0) + + # Paste Patch + h_sh, w_sh = patch.shape + x1 = cx - w_sh // 2 + y1 = cy - h_sh // 2 + x2 = x1 + w_sh + y2 = y1 + h_sh + + # Crop to bounds + pad_x1 = max(0, -x1) + pad_y1 = max(0, -y1) + crop_x1 = max(0, x1) + crop_y1 = max(0, y1) + crop_x2 = min(w, x2) + crop_y2 = min(h, y2) + + patch_x1 = pad_x1 + patch_y1 = pad_y1 + patch_x2 = patch_x1 + (crop_x2 - crop_x1) + patch_y2 = patch_y1 + (crop_y2 - crop_y1) + + if patch_x2 > patch.shape[1] or patch_y2 > patch.shape[0]: + continue + + valid_patch = patch[patch_y1:patch_y2, patch_x1:patch_x2] + + # Accumulate density (max) + cloud_density_map[crop_y1:crop_y2, crop_x1:crop_x2] = np.maximum( + cloud_density_map[crop_y1:crop_y2, crop_x1:crop_x2], valid_patch + ) + + # Apply Intensity Scaling + # Intensity [0, 10] -> Max Opacity [0, 1.0] + max_opacity = min(1.0, intensity / 10.0) + final_density = cloud_density_map * max_opacity + + # Transmittance = 1 - Density + attenuation = 1.0 - final_density + attenuation = np.clip(attenuation, 0, 1) + + # Composite + attenuation_3c = attenuation[:, :, np.newaxis] + cloud_color_layer = np.ones_like(image_float) * cloud_rgb + + # J * T + C * (1 - T) + out = image_float * attenuation_3c + cloud_color_layer * (1 - attenuation_3c) + out = np.clip(out, 0, 255).astype(np.uint8) + + if return_attenuation: + return out, attenuation.astype(np.float32) + return out diff --git a/src/inference_runtime.py b/src/inference_runtime.py new file mode 100644 index 0000000000000000000000000000000000000000..3ae4be69cf329d474c78ead6b6ce7134f8af14e9 --- /dev/null +++ b/src/inference_runtime.py @@ -0,0 +1,85 @@ +"""Shared launcher for the model-specific inference entry points.""" + +import argparse +import os +import subprocess +import sys +from pathlib import Path +from typing import Iterable, Sequence + + +def _has_option(arguments: Sequence[str], option: str) -> bool: + return any(arg == option or arg.startswith(f"{option}=") for arg in arguments) + + +def _merge_defaults(arguments: Sequence[str], defaults: Iterable[str]) -> list[str]: + merged = list(arguments) + pairs = list(defaults) + if len(pairs) % 2: + raise ValueError("Default inference arguments must be option/value pairs") + for index in range(0, len(pairs), 2): + option, value = pairs[index], pairs[index + 1] + if not _has_option(merged, option): + merged.extend((option, value)) + return merged + + +def build_command( + backend: Path, + gpu_ids: Sequence[int], + arguments: Sequence[str], +) -> list[str]: + if not gpu_ids or any(gpu_id < 0 for gpu_id in gpu_ids): + raise ValueError("--gpuid requires one or more non-negative GPU IDs") + if len(set(gpu_ids)) != len(gpu_ids): + raise ValueError("--gpuid values must be unique") + if not backend.is_file(): + raise FileNotFoundError(f"Inference backend not found: {backend}") + + command = [ + sys.executable, + "-m", + "accelerate.commands.launch", + "--num_processes", + str(len(gpu_ids)), + "--num_machines", + "1", + "--mixed_precision", + "fp16", + "--dynamo_backend", + "no", + ] + if len(gpu_ids) > 1: + command.append("--multi_gpu") + command.extend((str(backend), *arguments)) + return command + + +def launch(backend: Path, defaults: Iterable[str] = ()) -> None: + parser = argparse.ArgumentParser( + description=( + "Launch this model with Hugging Face Accelerate FP16. " + "Unrecognized arguments are forwarded to the model backend." + ) + ) + parser.add_argument( + "--gpuid", + nargs="+", + type=int, + default=[0], + help="Physical GPU IDs. Default: 0; example: --gpuid 0 2", + ) + runtime, backend_args = parser.parse_known_args() + backend = backend.resolve() + backend_args = _merge_defaults(backend_args, defaults) + + environment = os.environ.copy() + environment["CUDA_VISIBLE_DEVICES"] = ",".join(map(str, runtime.gpuid)) + command = build_command(backend, runtime.gpuid, backend_args) + completed = subprocess.run( + command, + cwd=backend.parent, + env=environment, + check=False, + ) + raise SystemExit(completed.returncode) diff --git a/src/models/cafnet/config.yaml b/src/models/cafnet/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..332846b7b95760918beb881b63f87c88eed20b70 --- /dev/null +++ b/src/models/cafnet/config.yaml @@ -0,0 +1,22 @@ +# Artifact evaluation uses only the packaged Smoke-Eval dataset. +test_base_dir: ../../../evaluation_dataset/Smoke-Eval +test_split: train +test_split_json: null +input_height: 288 +input_width: 512 +radar_max_depth_m: 11.2 +max_dist_correspondence: 0.5 +patch_size: [64, 128] +encoder: resnet34_bts +encoder_radar: resnet18 +radar_input_channels: 1 +bts_size: 512 +max_depth: 11.2 +batch_size: 8 +num_workers: 0 +seed: 42 +cpu: false +# CaFNet overflows on this GPU under fp16; use the released fp32 weights. +mixed_precision: "no" +checkpoint_path: ../../../checkpoints/baselines/cafnet/cafnet.safetensors +prediction_dir: ../../../inference_results/cafnet diff --git a/src/models/cafnet/inference.py b/src/models/cafnet/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..5fad0df57c661705147dd86c1e7fa7dec65917b3 --- /dev/null +++ b/src/models/cafnet/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "Baselines/cafnet/inference.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/cafnet_no_smoke/config.yaml b/src/models/cafnet_no_smoke/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..bd55c29061af921096b09648b6ca90e8c0f00cae --- /dev/null +++ b/src/models/cafnet_no_smoke/config.yaml @@ -0,0 +1,22 @@ +# Artifact evaluation uses only the packaged Smoke-Eval dataset. +test_base_dir: ../../../evaluation_dataset/Smoke-Eval +test_split: train +test_split_json: null +input_height: 288 +input_width: 512 +radar_max_depth_m: 11.2 +max_dist_correspondence: 0.5 +patch_size: [64, 128] +encoder: resnet34_bts +encoder_radar: resnet18 +radar_input_channels: 1 +bts_size: 512 +max_depth: 11.2 +batch_size: 8 +num_workers: 0 +seed: 42 +cpu: false +# CaFNet overflows on this GPU under fp16; use the released fp32 weights. +mixed_precision: "no" +checkpoint_path: ../../../checkpoints/baselines/cafnet_no_smoke/cafnet_no_smoke.safetensors +prediction_dir: ../../../inference_results/cafnet_no_smoke diff --git a/src/models/cafnet_no_smoke/inference.py b/src/models/cafnet_no_smoke/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..1327eb7f065e96400817f454c5389298609b3182 --- /dev/null +++ b/src/models/cafnet_no_smoke/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "Baselines/cafnet_no_smoke/inference.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/da3/config.yaml b/src/models/da3/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..1ac8f64a91566a06ed2ef3dbb99a972eee358e5f --- /dev/null +++ b/src/models/da3/config.yaml @@ -0,0 +1,9 @@ +# Runtime metadata for the DA3 direct backend. The runner passes the data, +# output, and checkpoint paths explicitly because the upstream-style CLI +# requires them. +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval +runtime: + seed: 42 + mixed_precision: fp16 + model_name: da3metric-large diff --git a/src/models/da3/inference.py b/src/models/da3/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..dded26fffa90b54fc4ec68bfee8251d3adc007cb --- /dev/null +++ b/src/models/da3/inference.py @@ -0,0 +1,13 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +ROOT = HERE.parents[2] +launch( + SRC / "Baselines/da3/inference.py", + ("--data_root", str(ROOT / "evaluation_dataset/Smoke-Eval"), "--checkpoint", str(ROOT / "checkpoints/baselines/da3/da3metric-large.safetensors"), "--output_dir", str(ROOT / "inference_results/da3")), +) diff --git a/src/models/grade/config.yaml b/src/models/grade/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..48d63af8aff14729c304d19f91ee0873bdab64f7 --- /dev/null +++ b/src/models/grade/config.yaml @@ -0,0 +1,14 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + resolution: {height: 288, width: 512} + scale_factor: 0.001 + max_depth_m: 11.2 + num_frames: 1 + num_workers: 0 +pretrained: + radar_model: ../../../checkpoints/grade/radar.safetensors + unet: ../../../checkpoints/grade/diffusion.safetensors + controlnet: ../../../checkpoints/grade/control.safetensors +diffusion: {num_train_timesteps: 1000} +inference: {frame_skip: 1, batch_size: 1, num_workers: 0} diff --git a/src/models/grade/inference.py b/src/models/grade/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..c6b9a874db8bd58a41ddb29adcde7834752856b9 --- /dev/null +++ b/src/models/grade/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "GRADE/stage2_diffusion_refinement/inference_full.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/grt/config.yaml b/src/models/grt/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..8b5f7eb13afccf2b136924cb80cd3a19c9285a0d --- /dev/null +++ b/src/models/grt/config.yaml @@ -0,0 +1,7 @@ +paths: + data_root: ../../../evaluation_dataset/Smoke-Eval +training: + batch_size: 1 + num_workers: 0 + mixed_precision: fp16 + seed: 42 diff --git a/src/models/grt/inference.py b/src/models/grt/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..d0ef38dc798b0a776c74be6c51e543da35a72df4 --- /dev/null +++ b/src/models/grt/inference.py @@ -0,0 +1,12 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch( + SRC / "Baselines/grt/inference.py", + ("--config", str(HERE / "config.yaml"), "--checkpoint", str(HERE.parents[2] / "checkpoints/baselines/grt/grt.safetensors"), "--output_dir", str(HERE.parents[2] / "inference_results/grt")), +) diff --git a/src/models/grt_image/config.yaml b/src/models/grt_image/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..80bf6d3c8d370d00bb0b8eecb632a2a0896bc36d --- /dev/null +++ b/src/models/grt_image/config.yaml @@ -0,0 +1,11 @@ +paths: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval +training: + batch_size: 1 + mixed_precision: fp16 + seed: 42 +data: + image_height: 288 + image_width: 512 +model: + resnet18_pretrained: false diff --git a/src/models/grt_image/inference.py b/src/models/grt_image/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..0d64370e14bd910fe345ebd6965f1aa90eb8d4b0 --- /dev/null +++ b/src/models/grt_image/inference.py @@ -0,0 +1,12 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch( + SRC / "Baselines/grt_image/inference.py", + ("--config", str(HERE / "config.yaml"), "--checkpoint", str(HERE.parents[2] / "checkpoints/baselines/grt_image/grt_image.safetensors"), "--output_dir", str(HERE.parents[2] / "inference_results/grt_image")), +) diff --git a/src/models/grt_no_doppler/config.yaml b/src/models/grt_no_doppler/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..5b3a343409ad518ff8ad9f6517c7f65a2080ea75 --- /dev/null +++ b/src/models/grt_no_doppler/config.yaml @@ -0,0 +1,9 @@ +paths: + test_data_root: ../../../evaluation_dataset/Smoke-Eval +training: {batch_size: 1, num_workers: 0, mixed_precision: fp16, seed: 42} +inference: + checkpoint_path: ../../../checkpoints/ablations/grt_no_doppler/grt_no_doppler.safetensors + output_dir: ../../../inference_results/grt_no_doppler + batch_size: 1 + num_workers: 0 + frame_skip: 1 diff --git a/src/models/grt_no_doppler/inference.py b/src/models/grt_no_doppler/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..a85bf7546293b4635954f6c1638b468d640a8820 --- /dev/null +++ b/src/models/grt_no_doppler/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "Ablation/grt_no_doppler/inference.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/grt_refine_freeze/config.yaml b/src/models/grt_refine_freeze/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..b3dc34d6873fe1a63fffa00a5fb456efbe48b5dc --- /dev/null +++ b/src/models/grt_refine_freeze/config.yaml @@ -0,0 +1,20 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + resolution: {height: 288, width: 512} + scale_factor: 0.001 + max_depth_m: 11.2 + num_frames: 1 + num_workers: 0 + test_skip: 1 + camera_intrinsics: {fx: 208.632642, fy: 208.632642, cx: 254.533594, cy: 149.243860} +pretrained: + grt_model: ../../../checkpoints/baselines/grt/grt.safetensors +diffusion: {num_train_timesteps: 1000} +inference: + unet_checkpoint: ../../../checkpoints/grade/diffusion.safetensors + controlnet_checkpoint: ../../../checkpoints/grade/control.safetensors + output_dir: ../../../inference_results/grt_refine_freeze + frame_skip: 1 + batch_size: 1 + num_workers: 0 diff --git a/src/models/grt_refine_freeze/inference.py b/src/models/grt_refine_freeze/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..303c48d232298082c0c743926cee143bd5b1aae8 --- /dev/null +++ b/src/models/grt_refine_freeze/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "Ablation/grt_refine_freeze/inference_control.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/grt_refine_retrain/config.yaml b/src/models/grt_refine_retrain/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..9910d96b022944768270c53c19430db1fe71ecd6 --- /dev/null +++ b/src/models/grt_refine_retrain/config.yaml @@ -0,0 +1,20 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + resolution: {height: 288, width: 512} + scale_factor: 0.001 + max_depth_m: 11.2 + num_frames: 1 + num_workers: 0 + test_skip: 1 + camera_intrinsics: {fx: 208.632642, fy: 208.632642, cx: 254.533594, cy: 149.243860} +pretrained: + grt_model: ../../../checkpoints/ablations/grt_refine_retrain/grt.safetensors +diffusion: {num_train_timesteps: 1000} +inference: + unet_checkpoint: ../../../checkpoints/ablations/grt_refine_retrain/diffusion.safetensors + controlnet_checkpoint: ../../../checkpoints/ablations/grt_refine_retrain/control.safetensors + output_dir: ../../../inference_results/grt_refine_retrain + frame_skip: 1 + batch_size: 1 + num_workers: 0 diff --git a/src/models/grt_refine_retrain/inference.py b/src/models/grt_refine_retrain/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..b82a47d077dae53c812f220508e93854f4ce5a28 --- /dev/null +++ b/src/models/grt_refine_retrain/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "Ablation/grt_refine_retrain/inference_control.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/ours_diffusion/config.yaml b/src/models/ours_diffusion/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..653a15870837e9d68d9d73af2dc740baf1e802c1 --- /dev/null +++ b/src/models/ours_diffusion/config.yaml @@ -0,0 +1,13 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + resolution: {height: 288, width: 512} + scale_factor: 0.001 + max_depth_m: 11.2 + num_frames: 1 + num_workers: 0 +pretrained: + radar_model: ../../../checkpoints/grade/radar.safetensors + unet: ../../../checkpoints/grade/diffusion.safetensors +diffusion: {num_train_timesteps: 1000} +inference: {frame_skip: 1, batch_size: 1, num_workers: 0} diff --git a/src/models/ours_diffusion/inference.py b/src/models/ours_diffusion/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..0be577338275f0917158703881a54d1e32a6275d --- /dev/null +++ b/src/models/ours_diffusion/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "GRADE/stage2_diffusion_refinement/inference_diffusion.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/ours_full/config.yaml b/src/models/ours_full/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..48d63af8aff14729c304d19f91ee0873bdab64f7 --- /dev/null +++ b/src/models/ours_full/config.yaml @@ -0,0 +1,14 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + resolution: {height: 288, width: 512} + scale_factor: 0.001 + max_depth_m: 11.2 + num_frames: 1 + num_workers: 0 +pretrained: + radar_model: ../../../checkpoints/grade/radar.safetensors + unet: ../../../checkpoints/grade/diffusion.safetensors + controlnet: ../../../checkpoints/grade/control.safetensors +diffusion: {num_train_timesteps: 1000} +inference: {frame_skip: 1, batch_size: 1, num_workers: 0} diff --git a/src/models/ours_full/inference.py b/src/models/ours_full/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..c6b9a874db8bd58a41ddb29adcde7834752856b9 --- /dev/null +++ b/src/models/ours_full/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "GRADE/stage2_diffusion_refinement/inference_full.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/ours_full_no_3d/config.yaml b/src/models/ours_full_no_3d/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..35c9f185c7c897a932e5711b4f452d11854acd8d --- /dev/null +++ b/src/models/ours_full_no_3d/config.yaml @@ -0,0 +1,14 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + resolution: {height: 288, width: 512} + scale_factor: 0.001 + max_depth_m: 11.2 + num_frames: 1 + num_workers: 0 +pretrained: + radar_model: ../../../checkpoints/ablations/ours_full_no_3d/radar.safetensors + unet: ../../../checkpoints/ablations/ours_full_no_3d/diffusion.safetensors + controlnet: ../../../checkpoints/ablations/ours_full_no_3d/control.safetensors +diffusion: {num_train_timesteps: 1000} +inference: {frame_skip: 1, batch_size: 1, num_workers: 0} diff --git a/src/models/ours_full_no_3d/inference.py b/src/models/ours_full_no_3d/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..ad19b46ffc5a6083b59bd649ace6400347b52aef --- /dev/null +++ b/src/models/ours_full_no_3d/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "Ablation/ours_full_no_3d/inference.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/ours_radar/config.yaml b/src/models/ours_radar/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..1fdffa9592a9581177a23770a2854f708cbc93ab --- /dev/null +++ b/src/models/ours_radar/config.yaml @@ -0,0 +1,12 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + resolution: {height: 288, width: 512} + scale_factor: 0.001 + max_depth_m: 11.2 + num_frames: 1 + num_workers: 0 +pretrained: + radar_model: ../../../checkpoints/grade/radar.safetensors +diffusion: {num_train_timesteps: 1000} +inference: {frame_skip: 1, batch_size: 1, num_workers: 0} diff --git a/src/models/ours_radar/inference.py b/src/models/ours_radar/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..d5a4b9318d2f9dcff0e064bc9a9776041dc90237 --- /dev/null +++ b/src/models/ours_radar/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "GRADE/stage2_diffusion_refinement/inference_radar.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/ours_radar_no_doppler/config.yaml b/src/models/ours_radar_no_doppler/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..39d416b9bb8729097514257eb331643d54b312ea --- /dev/null +++ b/src/models/ours_radar_no_doppler/config.yaml @@ -0,0 +1,13 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + test_root: ../../../evaluation_dataset/Smoke-Eval + scale_factor: 0.001 + max_depth_m: 11.2 + depth_resolution: [128, 256] + num_workers: 0 +inference: + checkpoint_path: ../../../checkpoints/ablations/ours_radar_no_doppler/ours_radar_no_doppler.safetensors + output_dir: ../../../inference_results/ours_radar_no_doppler + batch_size: 1 + num_workers: 0 + frame_skip: 1 diff --git a/src/models/ours_radar_no_doppler/inference.py b/src/models/ours_radar_no_doppler/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..5224ceaea4a788dc317a94b9e7e1e85c3d69e6c0 --- /dev/null +++ b/src/models/ours_radar_no_doppler/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "Ablation/ours_radar_no_doppler/inference.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/ours_radar_no_grad/config.yaml b/src/models/ours_radar_no_grad/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..8990d863216172fba1191235c587d3e9c5c165df --- /dev/null +++ b/src/models/ours_radar_no_grad/config.yaml @@ -0,0 +1,12 @@ +training: {batch_size: 1, mixed_precision: fp16, seed: 42} +data: + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval + resolution: {height: 288, width: 512} + scale_factor: 0.001 + max_depth_m: 11.2 + num_frames: 1 + num_workers: 0 +pretrained: + radar_model: ../../../checkpoints/ablations/ours_radar_no_grad/ours_radar_no_grad.safetensors +diffusion: {num_train_timesteps: 1000} +inference: {frame_skip: 1, batch_size: 1, num_workers: 0} diff --git a/src/models/ours_radar_no_grad/inference.py b/src/models/ours_radar_no_grad/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..0e6e36dce7a199472935e16b72110e451134f630 --- /dev/null +++ b/src/models/ours_radar_no_grad/inference.py @@ -0,0 +1,9 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +launch(SRC / "Ablation/ours_radar_no_grad/inference.py", ("--config", str(HERE / "config.yaml"))) diff --git a/src/models/radarcam-depth/config.yaml b/src/models/radarcam-depth/config.yaml new file mode 100644 index 0000000000000000000000000000000000000000..f7be597cc137d6b330368af98ace5efd9d7981a8 --- /dev/null +++ b/src/models/radarcam-depth/config.yaml @@ -0,0 +1,25 @@ +data: + # RadarCam-Depth consumes the separately prepared ZJU-style evaluation set. + smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval-RadarCam-Depth +depth: + # Preserve the bounds used to train the released RadarCam-Depth weights. + # They also define which RC-Net depths become SML scale anchors. + max_radar_depth_m: 11.2 + min_radar_depth_m: 0.05 + min_pred_depth_m: 0.1 + max_pred_depth_m: 11.2 + min_eval_depth_m: 0.0 + max_eval_depth_m: 11.2 +rcnet: + input_height: 288 + input_width: 512 + patch_size: [288, 96] + response_thr: 0.5 +sml: + mono_tag: dpt_hybrid + batch_size: 8 + num_workers: 0 +runtime: + seed: 42 + cpu: false + mixed_precision: bf16 diff --git a/src/models/radarcam-depth/inference.py b/src/models/radarcam-depth/inference.py new file mode 100644 index 0000000000000000000000000000000000000000..098267d4f0b13f334634edf4c7faa24bd2931a6c --- /dev/null +++ b/src/models/radarcam-depth/inference.py @@ -0,0 +1,13 @@ +from pathlib import Path +import sys + +SRC = Path(__file__).resolve().parents[2] +sys.path.insert(0, str(SRC)) +from inference_runtime import launch + +HERE = Path(__file__).resolve().parent +ROOT = HERE.parents[2] +launch( + SRC / "Baselines/radarcam-depth/smoke_eval_inference.py", + ("--config", str(HERE / "config.yaml"), "--rcnet_checkpoint", str(ROOT / "checkpoints/baselines/radarcam-depth/radarcam-depth_rcnet.safetensors"), "--sml_checkpoint", str(ROOT / "checkpoints/baselines/radarcam-depth/radarcam-depth_sml.safetensors"), "--output_dir", str(ROOT / "inference_results/radarcam-depth")), +)