Bin-0815 commited on
Commit
f348660
·
verified ·
1 Parent(s): 0e150d6

Release all GRADE models, checkpoints, and reviewed evaluation code (part 2)

Browse files

Add all 15 model paths and 20 safetensors files; include E1/E2 reproduction code, setup instructions, fixes for reviewer findings, and checkpoint SHA-256 manifest.

This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. src/Ablation/ours_radar_no_doppler/inference.py +267 -0
  2. src/Ablation/ours_radar_no_doppler/iq1m_dataset.py +305 -0
  3. src/Ablation/ours_radar_no_doppler/radar_depth.py +406 -0
  4. src/Ablation/ours_radar_no_doppler/rice_dataset.py +159 -0
  5. src/Ablation/ours_radar_no_doppler/split.json +34 -0
  6. src/Ablation/ours_radar_no_grad/inference.py +16 -0
  7. src/Baselines/cafnet/collate_fn_helpers.py +404 -0
  8. src/Baselines/cafnet/dataloader.py +100 -0
  9. src/Baselines/cafnet/extract_pcd_from_depth.py +96 -0
  10. src/Baselines/cafnet/inference.py +224 -0
  11. src/Baselines/cafnet/inference_config.yaml +29 -0
  12. src/Baselines/cafnet/models/bts.py +367 -0
  13. src/Baselines/cafnet/models/model.py +28 -0
  14. src/Baselines/cafnet/models/radar.py +212 -0
  15. src/Baselines/cafnet/rice_dataset.py +121 -0
  16. src/Baselines/cafnet/split.json +14 -0
  17. src/Baselines/cafnet_no_smoke/collate_fn_helpers.py +404 -0
  18. src/Baselines/cafnet_no_smoke/dataloader.py +100 -0
  19. src/Baselines/cafnet_no_smoke/extract_pcd_from_depth.py +96 -0
  20. src/Baselines/cafnet_no_smoke/inference.py +224 -0
  21. src/Baselines/cafnet_no_smoke/inference_config.yaml +29 -0
  22. src/Baselines/cafnet_no_smoke/models/bts.py +367 -0
  23. src/Baselines/cafnet_no_smoke/models/model.py +28 -0
  24. src/Baselines/cafnet_no_smoke/models/radar.py +212 -0
  25. src/Baselines/cafnet_no_smoke/rice_dataset.py +123 -0
  26. src/Baselines/cafnet_no_smoke/split.json +14 -0
  27. src/Baselines/da3/inference.py +179 -0
  28. src/Baselines/grt/augmentations.py +193 -0
  29. src/Baselines/grt/dataloader.py +330 -0
  30. src/Baselines/grt/grt_model.py +585 -0
  31. src/Baselines/grt/inference.py +220 -0
  32. src/Baselines/grt/split.json +16 -0
  33. src/Baselines/grt_image/augmentations.py +193 -0
  34. src/Baselines/grt_image/dataloader.py +344 -0
  35. src/Baselines/grt_image/grt_image_resnet_inference.example.yaml +16 -0
  36. src/Baselines/grt_image/grt_model.py +799 -0
  37. src/Baselines/grt_image/inference.py +224 -0
  38. src/Baselines/grt_image/split.json +16 -0
  39. src/Baselines/radarcam-depth/data/SML_dataset.py +83 -0
  40. src/Baselines/radarcam-depth/data/data_utils.py +326 -0
  41. src/Baselines/radarcam-depth/data/datasets.py +392 -0
  42. src/Baselines/radarcam-depth/linear_attention.py +184 -0
  43. src/Baselines/radarcam-depth/modules/estimator.py +188 -0
  44. src/Baselines/radarcam-depth/modules/midas/base_model.py +12 -0
  45. src/Baselines/radarcam-depth/modules/midas/blocks.py +197 -0
  46. src/Baselines/radarcam-depth/modules/midas/midas_net_custom.py +138 -0
  47. src/Baselines/radarcam-depth/modules/midas/normalization.py +109 -0
  48. src/Baselines/radarcam-depth/modules/midas/transforms.py +263 -0
  49. src/Baselines/radarcam-depth/modules/midas/utils.py +237 -0
  50. src/Baselines/radarcam-depth/networks.py +1516 -0
src/Ablation/ours_radar_no_doppler/inference.py ADDED
@@ -0,0 +1,267 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """Sequence-by-sequence RadarDepth no-Doppler inference on Smoke-Eval.
3
+
4
+ Launch with ``accelerate launch inference.py --config <config.yaml>``.
5
+ Each output is ``<sequence>_pred.npy`` with float32 shape ``[N, 1, H, W]``
6
+ and normalized depth clipped to ``[0, 1]``.
7
+ """
8
+
9
+ import argparse
10
+ import os
11
+ import pickle
12
+ from typing import Dict, List, Tuple
13
+
14
+ import numpy as np
15
+ import torch
16
+ import yaml
17
+ from accelerate import Accelerator
18
+ from accelerate.utils import set_seed
19
+ from safetensors.torch import load_file
20
+ from torch.utils.data import DataLoader
21
+ from tqdm import tqdm
22
+
23
+ from radar_depth import RadarDepth
24
+ from rice_dataset import RiceDataset
25
+
26
+
27
+ def _resolve_path(config_path: str, value: str) -> str:
28
+ if os.path.isabs(value):
29
+ return value
30
+ return os.path.normpath(
31
+ os.path.join(os.path.dirname(os.path.abspath(config_path)), value)
32
+ )
33
+
34
+
35
+ def _validate_prediction_array(predictions: np.ndarray, sequence: str) -> None:
36
+ if predictions.ndim != 4 or predictions.shape[1] != 1:
37
+ raise RuntimeError(
38
+ f"{sequence}: expected prediction shape [N, 1, H, W], "
39
+ f"got {predictions.shape}"
40
+ )
41
+ if not np.isfinite(predictions).all():
42
+ raise RuntimeError(f"{sequence}: predictions contain NaN or Inf")
43
+ if predictions.min() < 0.0 or predictions.max() > 1.0:
44
+ raise RuntimeError(
45
+ f"{sequence}: normalized predictions are outside [0, 1]: "
46
+ f"[{predictions.min()}, {predictions.max()}]"
47
+ )
48
+
49
+
50
+ def _merge_rank_results(
51
+ gather_dir: str,
52
+ sequence: str,
53
+ num_processes: int,
54
+ ) -> Dict[int, np.ndarray]:
55
+ safe_sequence = sequence.replace("/", "_").replace("\\", "_").lower()
56
+ merged: Dict[int, np.ndarray] = {}
57
+ for rank in range(num_processes):
58
+ rank_path = os.path.join(gather_dir, f"rank_{rank}_{safe_sequence}.pkl")
59
+ with open(rank_path, "rb") as handle:
60
+ rank_results = pickle.load(handle)
61
+ for frame_idx, prediction in rank_results:
62
+ merged.setdefault(int(frame_idx), prediction)
63
+ os.remove(rank_path)
64
+ return merged
65
+
66
+
67
+ def _save_sequence(
68
+ output_dir: str,
69
+ sequence: str,
70
+ predictions: Dict[int, np.ndarray],
71
+ expected_frames: List[int],
72
+ debug: bool,
73
+ ) -> np.ndarray:
74
+ if not predictions:
75
+ raise RuntimeError(f"{sequence}: inference produced no predictions")
76
+
77
+ if not debug:
78
+ missing = [frame for frame in expected_frames if frame not in predictions]
79
+ if missing:
80
+ raise RuntimeError(
81
+ f"{sequence}: missing {len(missing)} predictions "
82
+ f"(first few frame indices: {missing[:5]})"
83
+ )
84
+ ordered_frames = expected_frames
85
+ else:
86
+ ordered_frames = sorted(predictions)
87
+
88
+ prediction_array = np.stack(
89
+ [predictions[frame] for frame in ordered_frames], axis=0
90
+ ).astype(np.float32, copy=False)
91
+ _validate_prediction_array(prediction_array, sequence)
92
+
93
+ safe_sequence = sequence.replace("/", "_").replace("\\", "_").lower()
94
+ np.save(
95
+ os.path.join(output_dir, f"{safe_sequence}_pred.npy"),
96
+ prediction_array,
97
+ )
98
+ return prediction_array
99
+
100
+
101
+ def parse_args() -> argparse.Namespace:
102
+ parser = argparse.ArgumentParser(
103
+ description="Run RadarDepth no-Doppler inference on Smoke-Eval."
104
+ )
105
+ parser.add_argument(
106
+ "--config",
107
+ default="config_stage1_iq1m.yaml",
108
+ help="YAML config path",
109
+ )
110
+ parser.add_argument("--checkpoint", default=None, help="Override checkpoint path")
111
+ parser.add_argument("--output_dir", default=None, help="Override output directory")
112
+ parser.add_argument("--debug", action="store_true", help="Process one batch per sequence")
113
+ return parser.parse_args()
114
+
115
+
116
+ def main() -> None:
117
+ cli = parse_args()
118
+ with open(cli.config, "r") as handle:
119
+ config = yaml.safe_load(handle) or {}
120
+
121
+ training_config = config.get("training", {})
122
+ data_config = config.get("data", {})
123
+ inference_config = config.get("inference", {})
124
+
125
+ test_root_value = data_config.get("test_root")
126
+ if not test_root_value:
127
+ raise ValueError("config['data']['test_root'] is required")
128
+ test_root = _resolve_path(cli.config, str(test_root_value))
129
+
130
+ checkpoint_value = cli.checkpoint or inference_config.get("checkpoint_path")
131
+ if not checkpoint_value:
132
+ raise ValueError(
133
+ "Set config['inference']['checkpoint_path'] or pass --checkpoint"
134
+ )
135
+ checkpoint_path = _resolve_path(cli.config, str(checkpoint_value))
136
+ if not os.path.isfile(checkpoint_path):
137
+ raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
138
+
139
+ output_value = cli.output_dir or inference_config.get(
140
+ "output_dir", "inference_results"
141
+ )
142
+ output_dir = _resolve_path(cli.config, str(output_value))
143
+ batch_size = int(
144
+ inference_config.get("batch_size", training_config.get("batch_size", 1))
145
+ )
146
+ num_workers = int(
147
+ inference_config.get("num_workers", data_config.get("num_workers", 0))
148
+ )
149
+ frame_skip = int(inference_config.get("frame_skip", 1))
150
+ mixed_precision = "fp16"
151
+ scale_factor = float(data_config.get("scale_factor", 0.001))
152
+ max_depth_m = float(data_config.get("max_depth_m", 11.2))
153
+ depth_resolution = tuple(data_config.get("depth_resolution", [128, 256]))
154
+
155
+ accelerator = Accelerator(mixed_precision=mixed_precision)
156
+ set_seed(int(training_config.get("seed", 42)))
157
+
158
+ discovery_dataset = RiceDataset(
159
+ root_dir=test_root,
160
+ sequences=None,
161
+ frame_skip=frame_skip,
162
+ scale_factor=scale_factor,
163
+ max_depth_m=max_depth_m,
164
+ depth_resolution=depth_resolution,
165
+ use_rgb=False,
166
+ )
167
+ sequences = discovery_dataset.sequences
168
+ if not sequences:
169
+ raise ValueError(f"No valid Smoke-Eval sequences found under {test_root}")
170
+ del discovery_dataset
171
+
172
+ if accelerator.is_main_process:
173
+ os.makedirs(output_dir, exist_ok=True)
174
+ print(f"Smoke-Eval: {test_root} ({len(sequences)} sequences)")
175
+ print(f"Checkpoint: {checkpoint_path}")
176
+ print(f"Output: {output_dir}")
177
+ print(
178
+ f"Mixed precision: {mixed_precision} | "
179
+ f"processes: {accelerator.num_processes}"
180
+ )
181
+ accelerator.wait_for_everyone()
182
+
183
+ gather_dir = os.path.join(output_dir, "_gather")
184
+ os.makedirs(gather_dir, exist_ok=True)
185
+
186
+ model = RadarDepth(
187
+ output_height=int(depth_resolution[0]),
188
+ output_width=int(depth_resolution[1]),
189
+ )
190
+ model.load_state_dict(load_file(checkpoint_path, device="cpu"), strict=True)
191
+ model.eval()
192
+ model = accelerator.prepare(model)
193
+
194
+ for sequence_index, sequence in enumerate(sequences):
195
+ dataset = RiceDataset(
196
+ root_dir=test_root,
197
+ sequences=[sequence],
198
+ frame_skip=frame_skip,
199
+ scale_factor=scale_factor,
200
+ max_depth_m=max_depth_m,
201
+ depth_resolution=depth_resolution,
202
+ use_rgb=False,
203
+ )
204
+ expected_frames = [int(frame_idx) for _, frame_idx in dataset.index_map]
205
+ loader = DataLoader(
206
+ dataset,
207
+ batch_size=batch_size,
208
+ shuffle=False,
209
+ num_workers=num_workers,
210
+ pin_memory=(accelerator.device.type == "cuda"),
211
+ drop_last=False,
212
+ )
213
+ loader = accelerator.prepare(loader)
214
+
215
+ local_results: List[Tuple[int, np.ndarray]] = []
216
+ with torch.no_grad():
217
+ progress = tqdm(
218
+ loader,
219
+ desc=f"[{sequence_index + 1}/{len(sequences)}] {sequence}",
220
+ disable=not accelerator.is_local_main_process,
221
+ dynamic_ncols=True,
222
+ leave=False,
223
+ )
224
+ for batch in progress:
225
+ with accelerator.autocast():
226
+ prediction = model(batch["radar"]).clamp_(0.0, 1.0)
227
+ prediction_np = prediction.detach().float().cpu().numpy()
228
+ frame_indices = batch["frame_idx"].detach().cpu().tolist()
229
+ local_results.extend(
230
+ (int(frame_idx), prediction_np[index])
231
+ for index, frame_idx in enumerate(frame_indices)
232
+ )
233
+ if cli.debug:
234
+ break
235
+
236
+ accelerator.wait_for_everyone()
237
+ safe_sequence = sequence.replace("/", "_").replace("\\", "_").lower()
238
+ rank_path = os.path.join(
239
+ gather_dir,
240
+ f"rank_{accelerator.process_index}_{safe_sequence}.pkl",
241
+ )
242
+ with open(rank_path, "wb") as handle:
243
+ pickle.dump(local_results, handle, protocol=pickle.HIGHEST_PROTOCOL)
244
+ accelerator.wait_for_everyone()
245
+
246
+ if accelerator.is_main_process:
247
+ merged = _merge_rank_results(
248
+ gather_dir, sequence, accelerator.num_processes
249
+ )
250
+ prediction_array = _save_sequence(
251
+ output_dir,
252
+ sequence,
253
+ merged,
254
+ expected_frames,
255
+ cli.debug,
256
+ )
257
+ print(f"{sequence}: saved {prediction_array.shape}")
258
+ accelerator.wait_for_everyone()
259
+
260
+ if accelerator.is_main_process:
261
+ if os.path.isdir(gather_dir) and not os.listdir(gather_dir):
262
+ os.rmdir(gather_dir)
263
+ print(f"Saved {len(sequences)} sequence predictions to: {output_dir}")
264
+
265
+
266
+ if __name__ == "__main__":
267
+ main()
src/Ablation/ours_radar_no_doppler/iq1m_dataset.py ADDED
@@ -0,0 +1,305 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ from pathlib import Path
3
+ from typing import Dict, List, Optional, Tuple, Any, Union
4
+ import cv2
5
+ import numpy as np
6
+ import torch
7
+ from torch.utils.data import Dataset
8
+ from collate_fn_helpers import radar_collator, depth_collator, fisheye_rgb_collator
9
+
10
+
11
+ class IQ1MMultiModalDataset(Dataset):
12
+ """
13
+ Dataset for loading aligned lidar, radar, and video frames.
14
+
15
+ No-doppler ablation: radar is read from
16
+ root_dir/radar_no_doppler/<sequence>/amplitude.npy and phase.npy
17
+ (single doppler bin), and the doppler axis is repeated 64x so
18
+ downstream code sees the standard cube.
19
+
20
+ Args:
21
+ root_dir: Root directory containing 'lidar', 'radar_no_doppler',
22
+ 'video' folders
23
+ sequences: Optional list of sequence names to load. If None, loads all.
24
+ transform: Optional transform to apply to video frames
25
+ """
26
+
27
+ DOPPLER_BINS = 64
28
+
29
+ def __init__(
30
+ self,
31
+ root_dir: str,
32
+ sequences: Optional[List[str]] = None,
33
+ frame_skip: int = 1,
34
+ split_type: Optional[str] = None, # 'train', 'val', 'test', or None for all
35
+ # Processing parameters
36
+ scale_factor: float = 0.001,
37
+ max_depth_m: float = 11.2,
38
+ depth_resolution: Tuple[int, int] = (128, 256),
39
+ use_rgb: bool = True,
40
+ rgb_resolution: Tuple[int, int] = (128, 256),
41
+ ):
42
+ self.root_dir = Path(root_dir)
43
+ self.depth_dir = self.root_dir / "metric_depth"
44
+ self.radar_dir = self.root_dir / "radar_no_doppler"
45
+ self.video_dir = self.root_dir / "video"
46
+ self.frame_skip = max(1, frame_skip)
47
+ self.split_type = split_type
48
+
49
+ # Processing parameters
50
+ self.proc_params = {
51
+ "scale_factor": scale_factor,
52
+ "max_depth_m": max_depth_m,
53
+ "depth_res": depth_resolution,
54
+ "use_rgb": use_rgb,
55
+ "rgb_res": rgb_resolution,
56
+ }
57
+
58
+ # Load split configuration if split_type is specified
59
+ if split_type is not None:
60
+ split_config = self._load_split_config()
61
+ sequences = self._get_sequences_for_split(split_config, sequences)
62
+
63
+ # Discover sequences
64
+ self.sequences = self._discover_sequences(sequences)
65
+
66
+ # Build index mapping (global_idx -> (sequence_name, frame_idx))
67
+ self.index_map: List[Tuple[str, int]] = []
68
+ self.sequence_info: Dict[str, dict] = {}
69
+
70
+ # Memory-mapped numpy arrays for efficient loading
71
+ self._depth_mmap: Dict[str, np.memmap] = {}
72
+ self._radar_amplitude_mmap: Dict[str, np.memmap] = {}
73
+ self._radar_phase_mmap: Dict[str, np.memmap] = {}
74
+ self._video_captures: Dict[str, cv2.VideoCapture] = {}
75
+
76
+ self._build_index()
77
+
78
+ def _load_split_config(self) -> Dict:
79
+ """Load split configuration from iq1m_split.json"""
80
+ split_file = Path(__file__).parent / "iq1m_split.json"
81
+ if not split_file.exists():
82
+ raise FileNotFoundError(f"Split configuration not found: {split_file}")
83
+
84
+ with open(split_file, "r") as f:
85
+ split_config = json.load(f)
86
+
87
+ return split_config
88
+
89
+ def _get_sequences_for_split(
90
+ self, split_config: Dict, requested_sequences: Optional[List[str]] = None
91
+ ) -> Optional[List[str]]:
92
+ """Get sequences for the specified split type"""
93
+ if self.split_type == "test":
94
+ sequences = split_config.get("test", [])
95
+ elif self.split_type in ["train", "val"]:
96
+ # Get all available sequences
97
+ all_sequences = self._get_all_available_sequences()
98
+ test_sequences = set(split_config.get("test", []))
99
+ # Exclude test sequences
100
+ sequences = [s for s in all_sequences if s not in test_sequences]
101
+ else:
102
+ raise ValueError(
103
+ f"Invalid split_type: {self.split_type}. Must be 'train', 'val', 'test', or None"
104
+ )
105
+
106
+ # Filter by requested sequences if provided
107
+ if requested_sequences is not None:
108
+ sequences = [s for s in sequences if s in requested_sequences]
109
+
110
+ return sequences
111
+
112
+ def _get_all_available_sequences(self) -> List[str]:
113
+ """Get all available sequences from the dataset"""
114
+ depth_seqs = set(
115
+ d.name
116
+ for d in self.depth_dir.iterdir()
117
+ if d.is_dir() and not d.name.startswith(".")
118
+ )
119
+ radar_seqs = set(
120
+ d.name
121
+ for d in self.radar_dir.iterdir()
122
+ if d.is_dir() and not d.name.startswith(".")
123
+ )
124
+
125
+ # Only consider video if use_rgb is True
126
+ if self.proc_params["use_rgb"]:
127
+ video_seqs = set(
128
+ d.name
129
+ for d in self.video_dir.iterdir()
130
+ if d.is_dir() and not d.name.startswith(".")
131
+ )
132
+ # Find common sequences across all modalities
133
+ common_seqs = depth_seqs & radar_seqs & video_seqs
134
+ else:
135
+ # Only need depth and radar
136
+ common_seqs = depth_seqs & radar_seqs
137
+
138
+ return sorted(list(common_seqs))
139
+
140
+ def _discover_sequences(self, sequences: Optional[List[str]] = None) -> List[str]:
141
+ """Discover available sequences with required modalities."""
142
+ # Get sequences from each modality folder
143
+ depth_seqs = set(
144
+ d.name
145
+ for d in self.depth_dir.iterdir()
146
+ if d.is_dir() and not d.name.startswith(".")
147
+ )
148
+ radar_seqs = set(
149
+ d.name
150
+ for d in self.radar_dir.iterdir()
151
+ if d.is_dir() and not d.name.startswith(".")
152
+ )
153
+
154
+ # Only consider video if use_rgb is True
155
+ if self.proc_params["use_rgb"]:
156
+ video_seqs = set(
157
+ d.name
158
+ for d in self.video_dir.iterdir()
159
+ if d.is_dir() and not d.name.startswith(".")
160
+ )
161
+ # Find common sequences across all modalities
162
+ common_seqs = depth_seqs & radar_seqs & video_seqs
163
+ else:
164
+ # Only need depth and radar
165
+ common_seqs = depth_seqs & radar_seqs
166
+
167
+ if sequences is not None:
168
+ # Filter to requested sequences
169
+ common_seqs = common_seqs & set(sequences)
170
+
171
+ return sorted(list(common_seqs))
172
+
173
+ def _build_index(self):
174
+ """Build global index mapping and load metadata."""
175
+ for seq_name in self.sequences:
176
+ # Load metadata from radar_no_doppler (or radar/lidar if available)
177
+ metadata_path = self.radar_dir / seq_name / "metadata.json"
178
+ if not metadata_path.exists():
179
+ metadata_path = self.root_dir / "radar" / seq_name / "metadata.json"
180
+ if not metadata_path.exists():
181
+ metadata_path = self.root_dir / "lidar" / seq_name / "metadata.json"
182
+ with open(metadata_path, "r") as f:
183
+ metadata = json.load(f)
184
+
185
+ n_frames = metadata["n_frames"]
186
+ self.sequence_info[seq_name] = {
187
+ "n_frames": n_frames,
188
+ "metadata": metadata,
189
+ "start_idx": len(self.index_map),
190
+ }
191
+
192
+ # Add frames to index with skipping
193
+ # Range: 0, frame_skip, 2*frame_skip, ...
194
+ for frame_idx in range(0, n_frames, self.frame_skip):
195
+ self.index_map.append((seq_name, frame_idx))
196
+
197
+ self.sequence_info[seq_name]["end_idx"] = len(self.index_map)
198
+
199
+ def _get_depth_mmap(self, seq_name: str) -> np.memmap:
200
+ """Get or create memory-mapped metric depth array."""
201
+ if seq_name not in self._depth_mmap:
202
+ path = self.depth_dir / seq_name / "metric_depth.npy"
203
+ self._depth_mmap[seq_name] = np.load(path, mmap_mode="r")
204
+ return self._depth_mmap[seq_name]
205
+
206
+ def _get_radar_mmap(self, seq_name: str) -> Tuple[np.memmap, np.memmap]:
207
+ """Get or create memory-mapped radar arrays."""
208
+ if seq_name not in self._radar_amplitude_mmap:
209
+ amp_path = self.radar_dir / seq_name / "amplitude.npy"
210
+ phase_path = self.radar_dir / seq_name / "phase.npy"
211
+ self._radar_amplitude_mmap[seq_name] = np.load(amp_path, mmap_mode="r")
212
+ self._radar_phase_mmap[seq_name] = np.load(phase_path, mmap_mode="r")
213
+ return self._radar_amplitude_mmap[seq_name], self._radar_phase_mmap[seq_name]
214
+
215
+ def _get_video_capture(self, seq_name: str) -> cv2.VideoCapture:
216
+ """Get or create video capture object."""
217
+ if seq_name not in self._video_captures:
218
+ video_path = self.video_dir / seq_name / "video.avi"
219
+ cap = cv2.VideoCapture(str(video_path))
220
+ if not cap.isOpened():
221
+ raise RuntimeError(f"Failed to open video: {video_path}")
222
+ self._video_captures[seq_name] = cap
223
+ return self._video_captures[seq_name]
224
+
225
+ def _load_rgb_frame(self, seq_name: str, frame_idx: int) -> np.ndarray:
226
+ """Load a specific frame from video."""
227
+ cap = self._get_video_capture(seq_name)
228
+
229
+ # Seek to frame
230
+ cap.set(cv2.CAP_PROP_POS_FRAMES, frame_idx)
231
+ ret, frame = cap.read()
232
+
233
+ if not ret:
234
+ raise RuntimeError(f"Failed to read frame {frame_idx} from {seq_name}")
235
+
236
+ # Convert BGR to RGB
237
+ frame = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
238
+ return frame
239
+
240
+ def __len__(self) -> int:
241
+ return len(self.index_map)
242
+
243
+ def __getitem__(self, idx: int) -> Dict[str, Any]:
244
+ seq_name, frame_idx = self.index_map[idx]
245
+
246
+ # === Radar (always needed) ===
247
+ amp_mmap, phase_mmap = self._get_radar_mmap(seq_name)
248
+ radar_amp = torch.from_numpy(amp_mmap[frame_idx].copy()).float()
249
+ radar_phase = torch.from_numpy(phase_mmap[frame_idx].copy()).float()
250
+ # Single doppler bin -> repeat to the standard 64-bin cube so
251
+ # downstream code is unchanged
252
+ radar_amp = torch.repeat_interleave(radar_amp, self.DOPPLER_BINS, dim=0)
253
+ radar_phase = torch.repeat_interleave(radar_phase, self.DOPPLER_BINS, dim=0)
254
+
255
+ processed_radar = radar_collator(
256
+ radar_amp.unsqueeze(0),
257
+ radar_phase.unsqueeze(0),
258
+ scale_factor=self.proc_params["scale_factor"],
259
+ ).squeeze(0)
260
+
261
+ # === Depth (always needed) ===
262
+ depth_mmap = self._get_depth_mmap(seq_name)
263
+ depth = torch.from_numpy(depth_mmap[frame_idx].copy()).float().unsqueeze(0)
264
+
265
+ processed_depth = depth_collator(
266
+ depth.unsqueeze(0),
267
+ max_depth_m=self.proc_params["max_depth_m"],
268
+ target_size=self.proc_params["depth_res"],
269
+ ).squeeze(0)
270
+
271
+ out = {
272
+ "radar": processed_radar,
273
+ "depth": processed_depth,
274
+ "sequence": seq_name,
275
+ "frame_idx": frame_idx,
276
+ }
277
+
278
+ # === RGB (only if use_rgb is True) ===
279
+ if self.proc_params["use_rgb"]:
280
+ rgb = torch.from_numpy(
281
+ self._load_rgb_frame(seq_name, frame_idx)
282
+ ).float().permute(2, 0, 1) / 255.0
283
+
284
+ out["rgb"] = fisheye_rgb_collator(
285
+ rgb.unsqueeze(0),
286
+ target_size=self.proc_params["rgb_res"],
287
+ ).squeeze(0)
288
+
289
+ return out
290
+
291
+ def get_sequence_frames(self, seq_name: str) -> List[int]:
292
+ """Get global indices for all frames in a sequence."""
293
+ info = self.sequence_info[seq_name]
294
+ return list(range(info["start_idx"], info["end_idx"]))
295
+
296
+ def close(self):
297
+ """Release video capture resources."""
298
+ for cap in self._video_captures.values():
299
+ cap.release()
300
+ self._video_captures.clear()
301
+
302
+ def __del__(self):
303
+ self.close()
304
+
305
+
src/Ablation/ours_radar_no_doppler/radar_depth.py ADDED
@@ -0,0 +1,406 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ from typing import Tuple
4
+
5
+
6
+ class RadarPatchEmbed(nn.Module):
7
+ """
8
+ Radar Spectrum Patch Embedding Layer.
9
+
10
+ Takes 5D radar spectrum data and converts it into patch embeddings:
11
+ 1. Input: [B, 2, 256, 64, 8, 2] where channels are (magnitude, phase)
12
+ 2. Patchifies along range and doppler dimensions
13
+ 3. Outputs: [B, num_patches, embed_dim] where num_patches = 2048
14
+
15
+ Patch extraction:
16
+ - Range dimension (256): patch_size=4, stride=4 -> 64 patches
17
+ - Doppler dimension (64): patch_size=2, stride=2 -> 32 patches
18
+ - Total patches: 64 × 32 = 2048
19
+ - Each patch: [4 range × 2 doppler × 8 elevation × 2 azimuth] × 2 channels = 256 features
20
+ """
21
+
22
+ def __init__(
23
+ self,
24
+ input_shape: Tuple[int, int, int, int] = (
25
+ 256,
26
+ 64,
27
+ 8,
28
+ 2,
29
+ ), # (Range, Doppler, Elevation, Azimuth)
30
+ patch_size: Tuple[int, int, int, int] = (
31
+ 4,
32
+ 2,
33
+ 8,
34
+ 2,
35
+ ), # (Range, Doppler, Elevation, Azimuth)
36
+ stride: Tuple[int, int] = (4, 2), # (Range, Doppler)
37
+ embed_dim: int = 256,
38
+ in_channels: int = 2, # magnitude + phase
39
+ ):
40
+ super().__init__()
41
+
42
+ self.input_shape = input_shape
43
+ self.patch_size = patch_size
44
+ self.stride = stride
45
+ self.embed_dim = embed_dim
46
+ self.in_channels = in_channels
47
+
48
+ # Calculate number of patches
49
+ range_dim, doppler_dim, elev_dim, azim_dim = input_shape
50
+ patch_range, patch_doppler, patch_elev, patch_azim = patch_size
51
+ stride_range, stride_doppler = stride
52
+
53
+ self.num_patches_range = (range_dim - patch_range) // stride_range + 1 # 64
54
+ self.num_patches_doppler = (
55
+ doppler_dim - patch_doppler
56
+ ) // stride_doppler + 1 # 32
57
+ self.num_patches = self.num_patches_range * self.num_patches_doppler # 2048
58
+
59
+ # Each patch has: patch_range × patch_doppler × patch_elev × patch_azim features per channel
60
+ patch_volume = (
61
+ patch_range * patch_doppler * patch_elev * patch_azim
62
+ ) # 4×2×8×2 = 128
63
+ self.patch_features = patch_volume * in_channels # 128 × 2 = 256
64
+
65
+ # Linear projection from patch features to embedding dimension
66
+ self.proj = nn.Linear(self.patch_features, embed_dim)
67
+
68
+ print(f"Radar Patch Embedding Configuration:")
69
+ print(
70
+ f" Input shape: [B, {in_channels}, {range_dim}, {doppler_dim}, {elev_dim}, {azim_dim}]"
71
+ )
72
+ print(f" Patch size: {patch_size}")
73
+ print(f" Stride: {stride}")
74
+ print(
75
+ f" Number of patches (range × doppler): {self.num_patches_range} × {self.num_patches_doppler} = {self.num_patches}"
76
+ )
77
+ print(f" Patch features per channel: {patch_volume}")
78
+ print(f" Total patch features (mag+phase): {self.patch_features}")
79
+ print(f" Embedding dimension: {embed_dim}")
80
+
81
+ def extract_patches(self, x: torch.Tensor) -> torch.Tensor:
82
+ """
83
+ Extract patches from radar spectrum data.
84
+
85
+ Args:
86
+ x: [B, 2, 256, 64, 8, 2] (magnitude + phase channels)
87
+
88
+ Returns:
89
+ patches: [B, num_patches, patch_features]
90
+ """
91
+ batch_size = x.shape[0]
92
+ x_mag = x[:, 0] # [B, 256, 64, 8, 2]
93
+ x_phase = x[:, 1] # [B, 256, 64, 8, 2]
94
+
95
+ all_patches = []
96
+
97
+ # Extract patches with stride along range and doppler dimensions
98
+ for i in range(self.num_patches_range):
99
+ for j in range(self.num_patches_doppler):
100
+ start_range = i * self.stride[0]
101
+ end_range = start_range + self.patch_size[0]
102
+ start_doppler = j * self.stride[1]
103
+ end_doppler = start_doppler + self.patch_size[1]
104
+
105
+ # Extract patch from both channels
106
+ patch_mag = x_mag[
107
+ :, start_range:end_range, start_doppler:end_doppler, :, :
108
+ ]
109
+ patch_phase = x_phase[
110
+ :, start_range:end_range, start_doppler:end_doppler, :, :
111
+ ]
112
+
113
+ # Flatten patches
114
+ patch_mag_flat = patch_mag.flatten(1) # [B, 128]
115
+ patch_phase_flat = patch_phase.flatten(1) # [B, 128]
116
+
117
+ # Interleave magnitude and phase features
118
+ patch_interleaved = torch.stack(
119
+ [patch_mag_flat, patch_phase_flat], dim=-1
120
+ )
121
+ patch_interleaved = patch_interleaved.flatten(1, -1) # [B, 256]
122
+
123
+ all_patches.append(patch_interleaved)
124
+
125
+ # Stack all patches: [B, num_patches, patch_features]
126
+ all_patches = torch.stack(all_patches, dim=1)
127
+ return all_patches
128
+
129
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
130
+ """
131
+ Forward pass.
132
+
133
+ Args:
134
+ x: [B, 2, 256, 64, 8, 2]
135
+
136
+ Returns:
137
+ embeddings: [B, num_patches, embed_dim]
138
+ """
139
+ # Extract patches: [B, 2048, 256]
140
+ patches = self.extract_patches(x)
141
+
142
+ # Project to embedding dimension: [B, 2048, embed_dim]
143
+ embeddings = self.proj(patches)
144
+
145
+ return embeddings
146
+
147
+
148
+ class RadarEncoder(nn.Module):
149
+ """
150
+ Radar Vision Transformer (ViT) Encoder.
151
+ """
152
+
153
+ def __init__(
154
+ self,
155
+ input_shape: Tuple[int, int, int, int] = (256, 64, 8, 2),
156
+ patch_size: Tuple[int, int, int, int] = (4, 2, 8, 2),
157
+ stride: Tuple[int, int] = (4, 2),
158
+ embed_dim: int = 256,
159
+ num_heads: int = 8,
160
+ num_layers: int = 4,
161
+ mlp_ratio: float = 4.0,
162
+ dropout: float = 0.1,
163
+ ):
164
+ super().__init__()
165
+
166
+ self.embed_dim = embed_dim
167
+
168
+ # Patch embedding layer
169
+ self.patch_embed = RadarPatchEmbed(
170
+ input_shape=input_shape,
171
+ patch_size=patch_size,
172
+ stride=stride,
173
+ embed_dim=embed_dim,
174
+ in_channels=2,
175
+ )
176
+
177
+ self.num_patches = self.patch_embed.num_patches
178
+
179
+ # Learnable positional embeddings
180
+ self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches, embed_dim))
181
+
182
+ # Transformer encoder
183
+ encoder_layer = nn.TransformerEncoderLayer(
184
+ d_model=embed_dim,
185
+ nhead=num_heads,
186
+ dim_feedforward=int(embed_dim * mlp_ratio),
187
+ dropout=dropout,
188
+ activation="gelu",
189
+ batch_first=True,
190
+ norm_first=True,
191
+ )
192
+ self.transformer = nn.TransformerEncoder(
193
+ encoder_layer=encoder_layer,
194
+ num_layers=num_layers,
195
+ norm=nn.LayerNorm(embed_dim),
196
+ )
197
+
198
+ self._init_weights()
199
+
200
+ def _init_weights(self):
201
+ """Initialize weights."""
202
+ # Initialize positional embeddings
203
+ nn.init.trunc_normal_(self.pos_embed, std=0.02)
204
+
205
+ # Initialize patch embedding projection
206
+ if hasattr(self.patch_embed.proj, "weight"):
207
+ nn.init.xavier_uniform_(self.patch_embed.proj.weight)
208
+ if self.patch_embed.proj.bias is not None:
209
+ nn.init.zeros_(self.patch_embed.proj.bias)
210
+
211
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
212
+ x = self.patch_embed(x)
213
+ x = x + self.pos_embed
214
+ x = self.transformer(x)
215
+ return x
216
+
217
+
218
+ class TransformerDecoderBlock(nn.Module):
219
+ """Transformer decoder block with self-attention and feedforward"""
220
+
221
+ def __init__(self, embed_dim=384, num_heads=6, mlp_ratio=4.0, dropout=0.0):
222
+ super().__init__()
223
+ self.norm1 = nn.LayerNorm(embed_dim)
224
+ self.attn = nn.MultiheadAttention(
225
+ embed_dim, num_heads, dropout=dropout, batch_first=True
226
+ )
227
+ self.norm2 = nn.LayerNorm(embed_dim)
228
+ self.mlp = nn.Sequential(
229
+ nn.Linear(embed_dim, int(embed_dim * mlp_ratio)),
230
+ nn.GELU(),
231
+ nn.Dropout(dropout),
232
+ nn.Linear(int(embed_dim * mlp_ratio), embed_dim),
233
+ nn.Dropout(dropout),
234
+ )
235
+
236
+ def forward(self, x):
237
+ # Self-attention with residual
238
+ x = x + self.attn(self.norm1(x), self.norm1(x), self.norm1(x))[0]
239
+ # MLP with residual
240
+ x = x + self.mlp(self.norm2(x))
241
+ return x
242
+
243
+
244
+ class DepthDecoder(nn.Module):
245
+ """
246
+ Hybrid Transformer+CNN decoder for depth image generation.
247
+
248
+ Input: [batch_size, num_patches=2048, embed_dim=512]
249
+ Output: [batch_size, 1, height=128, width=256]
250
+
251
+ Architecture:
252
+ 1. Transformer decoder blocks (4 layers)
253
+ 2. Reshape to 2D feature map (64x32)
254
+ 3. CNN upsampling stages (64x32 -> 128x256)
255
+ """
256
+
257
+ def __init__(
258
+ self,
259
+ embed_dim=256,
260
+ num_patches=2048,
261
+ patch_grid_size=(64, 32), # Spatial structure from radar encoder
262
+ num_decoder_blocks=4,
263
+ num_heads=8,
264
+ mlp_ratio=4.0,
265
+ dropout=0.0,
266
+ output_height=128,
267
+ output_width=256,
268
+ output_channels=1,
269
+ ):
270
+ super().__init__()
271
+ self.embed_dim = embed_dim
272
+ self.num_patches = num_patches
273
+ self.patch_grid_size = patch_grid_size # (64, 32) spatial grid
274
+ self.output_height = output_height
275
+ self.output_width = output_width
276
+ self.output_channels = output_channels
277
+
278
+ # Transformer decoder blocks
279
+ self.decoder_blocks = nn.ModuleList(
280
+ [
281
+ TransformerDecoderBlock(embed_dim, num_heads, mlp_ratio, dropout)
282
+ for _ in range(num_decoder_blocks)
283
+ ]
284
+ )
285
+
286
+ self.norm = nn.LayerNorm(embed_dim)
287
+
288
+ # Projection to intermediate feature map
289
+ # From 64x32x256 to 64x32x128 (reduce dimension for upsampling)
290
+ self.feature_proj = nn.Conv2d(embed_dim, 128, kernel_size=1)
291
+
292
+ # Upsampling network: 64x32 -> 128x256
293
+ # Start from 64x32 (range x doppler), upsample to 128x256
294
+ self.upsample = nn.Sequential(
295
+ # Upsample doppler dimension: 64x32 -> 64x64
296
+ nn.Upsample(scale_factor=(1, 2), mode="bilinear", align_corners=False),
297
+ nn.Conv2d(128, 64, kernel_size=3, padding=1),
298
+ nn.BatchNorm2d(64),
299
+ nn.ReLU(inplace=True),
300
+ # Upsample both dimensions: 64x64 -> 128x128
301
+ nn.Upsample(scale_factor=(2, 2), mode="bilinear", align_corners=False),
302
+ nn.Conv2d(64, 32, kernel_size=3, padding=1),
303
+ nn.BatchNorm2d(32),
304
+ nn.ReLU(inplace=True),
305
+ # Upsample width dimension: 128x128 -> 128x256
306
+ nn.Upsample(scale_factor=(1, 2), mode="bilinear", align_corners=False),
307
+ nn.Conv2d(32, output_channels, kernel_size=3, padding=1),
308
+ nn.Sigmoid(), # Output in [0, 1] range
309
+ )
310
+
311
+ def forward(self, x):
312
+ """
313
+ Args:
314
+ x: [batch_size, num_patches, embed_dim]
315
+
316
+ Returns:
317
+ depth: [batch_size, output_channels, output_height, output_width]
318
+ """
319
+ batch_size = x.shape[0]
320
+
321
+ # Apply transformer decoder blocks
322
+ for block in self.decoder_blocks:
323
+ x = block(x)
324
+
325
+ x = self.norm(x)
326
+
327
+ # Reshape to spatial dimensions: [B, 2048, 512] -> [B, 64, 32, 512]
328
+ x = x.reshape(
329
+ batch_size,
330
+ self.patch_grid_size[0], # 64 (range)
331
+ self.patch_grid_size[1], # 32 (doppler)
332
+ self.embed_dim,
333
+ )
334
+
335
+ # Permute to channel-first: [B, H, W, C] -> [B, C, H, W]
336
+ x = x.permute(0, 3, 1, 2)
337
+ # Shape: [batch, 512, 64, 32]
338
+
339
+ # Project features
340
+ x = self.feature_proj(x)
341
+ # Shape: [batch, 128, 64, 32]
342
+
343
+ # Upsample to target resolution
344
+ depth = self.upsample(x)
345
+ # Shape: [batch, 1, 128, 256]
346
+
347
+ return depth
348
+
349
+
350
+ class RadarDepth(nn.Module):
351
+ """
352
+ End-to-end Radar to Depth model (Doppler-as-Channels).
353
+ """
354
+
355
+ def __init__(
356
+ self,
357
+ # Encoder args
358
+ input_shape: Tuple[int, int, int, int] = (256, 64, 8, 2),
359
+ patch_size: Tuple[int, int, int, int] = (4, 2, 8, 2),
360
+ stride: Tuple[int, int] = (4, 2),
361
+ embed_dim: int = 256,
362
+ encoder_num_heads: int = 8,
363
+ encoder_num_layers: int = 4,
364
+ encoder_mlp_ratio: float = 4.0,
365
+ encoder_dropout: float = 0.1,
366
+ # Decoder args
367
+ decoder_num_blocks: int = 4,
368
+ decoder_num_heads: int = 8,
369
+ output_height: int = 128,
370
+ output_width: int = 256,
371
+ ):
372
+ super().__init__()
373
+
374
+ self.encoder = RadarEncoder(
375
+ input_shape=input_shape,
376
+ patch_size=patch_size,
377
+ stride=stride,
378
+ embed_dim=embed_dim,
379
+ num_heads=encoder_num_heads,
380
+ num_layers=encoder_num_layers,
381
+ mlp_ratio=encoder_mlp_ratio,
382
+ dropout=encoder_dropout,
383
+ )
384
+
385
+ # Get patch info from encoder
386
+ num_patches = self.encoder.num_patches # 64
387
+
388
+ self.decoder = DepthDecoder(
389
+ embed_dim=embed_dim,
390
+ num_patches=num_patches,
391
+ num_decoder_blocks=decoder_num_blocks,
392
+ num_heads=decoder_num_heads,
393
+ output_height=output_height,
394
+ output_width=output_width,
395
+ )
396
+
397
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
398
+ x = self.encoder(x)
399
+ x = self.decoder(x)
400
+ return x
401
+
402
+
403
+ def create_radar_encoder(*args, **kwargs):
404
+ return RadarEncoder(*args, **kwargs)
405
+
406
+
src/Ablation/ours_radar_no_doppler/rice_dataset.py ADDED
@@ -0,0 +1,159 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ from pathlib import Path
4
+ from typing import Dict, List, Optional, Tuple
5
+ from collate_fn_helpers import dji_rgb_collator, radar_collator, depth_collator
6
+ from torch.utils.data import Dataset
7
+
8
+
9
+ class RiceDataset(Dataset):
10
+ """Dataset for radar, DJI RGB, and ZED depth.
11
+
12
+ No-doppler ablation: loads radar_no_doppler.npy (N, 1, elevation, azimuth,
13
+ range) and repeats the single doppler bin 64x so downstream code sees the
14
+ standard (64, elevation, azimuth, range) cube.
15
+ """
16
+
17
+ REQUIRED_FILES = ("radar_no_doppler.npy", "dji_rgb.npy", "zed_depth.npy")
18
+ DOPPLER_BINS = 64
19
+
20
+ def __init__(
21
+ self,
22
+ root_dir: str,
23
+ sequences: Optional[List[str]] = None,
24
+ frame_skip: int = 1,
25
+ depth_in_meters: bool = True,
26
+ rgb_normalize: bool = True,
27
+ # Processing parameters
28
+ scale_factor: float = 0.001,
29
+ max_depth_m: float = 11.2,
30
+ depth_resolution: Tuple[int, int] = (128, 256),
31
+ use_rgb: bool = True,
32
+ rgb_resolution: Tuple[int, int] = (128, 256),
33
+ ):
34
+ self.root_dir = Path(root_dir)
35
+ self.frame_skip = max(1, frame_skip)
36
+ self.depth_in_meters = depth_in_meters
37
+ self.rgb_normalize = rgb_normalize
38
+
39
+ # Processing parameters
40
+ self.proc_params = {
41
+ "scale_factor": scale_factor,
42
+ "max_depth_m": max_depth_m,
43
+ "depth_res": depth_resolution,
44
+ "use_rgb": use_rgb,
45
+ "rgb_res": rgb_resolution,
46
+ }
47
+
48
+ self.sequences = self._discover_sequences(sequences)
49
+ self.index_map: List[Tuple[str, int]] = []
50
+ self._seq_arrays: Dict[str, Dict] = {}
51
+
52
+ self._build_index()
53
+
54
+ def _discover_sequences(self, sequences: Optional[List[str]] = None) -> List[str]:
55
+ if not self.root_dir.is_dir():
56
+ raise FileNotFoundError(f"Root directory not found: {self.root_dir}")
57
+
58
+ all_seqs = sorted(
59
+ d.name
60
+ for d in self.root_dir.iterdir()
61
+ if d.is_dir() and not d.name.startswith(".")
62
+ )
63
+
64
+ # Required files always needed
65
+ required = ["radar_no_doppler.npy", "zed_depth.npy"]
66
+ # Add RGB if use_rgb is True
67
+ if self.proc_params["use_rgb"]:
68
+ required.append("dji_rgb.npy")
69
+
70
+ valid = [
71
+ name
72
+ for name in all_seqs
73
+ if all((self.root_dir / name / f).exists() for f in required)
74
+ ]
75
+ if sequences is not None:
76
+ valid = [s for s in valid if s in sequences]
77
+ return valid
78
+
79
+ def _build_index(self) -> None:
80
+ self.index_map.clear()
81
+ for seq_name in self.sequences:
82
+ radar = np.load(
83
+ self.root_dir / seq_name / "radar_no_doppler.npy", mmap_mode="r"
84
+ )
85
+ for i in range(0, radar.shape[0], self.frame_skip):
86
+ self.index_map.append((seq_name, i))
87
+
88
+ def _load_sequence_arrays(self, seq_name: str) -> Dict:
89
+ if seq_name not in self._seq_arrays:
90
+ seq_dir = self.root_dir / seq_name
91
+ arrays = {
92
+ "radar": np.load(seq_dir / "radar_no_doppler.npy", mmap_mode="r"),
93
+ "depth": np.load(seq_dir / "zed_depth.npy", mmap_mode="r"),
94
+ }
95
+ # Only load RGB if needed
96
+ if self.proc_params["use_rgb"]:
97
+ arrays["rgb"] = np.load(seq_dir / "dji_rgb.npy", mmap_mode="r")
98
+ self._seq_arrays[seq_name] = arrays
99
+ return self._seq_arrays[seq_name]
100
+
101
+ def __len__(self) -> int:
102
+ return len(self.index_map)
103
+
104
+ def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
105
+ seq_name, frame_idx = self.index_map[idx]
106
+ arrs = self._load_sequence_arrays(seq_name)
107
+
108
+ # === Radar (always needed) ===
109
+ # (1, elevation, azimuth, range) -> repeat single doppler bin to
110
+ # (64, elevation, azimuth, range) so downstream code is unchanged
111
+ radar = np.asarray(arrs["radar"][frame_idx].copy())
112
+ radar = np.repeat(radar, self.DOPPLER_BINS, axis=0)
113
+ radar_amp = torch.from_numpy(np.abs(radar).astype(np.float32))
114
+ radar_phase = torch.from_numpy((np.angle(radar) / np.pi).astype(np.float32))
115
+
116
+ processed_radar = radar_collator(
117
+ radar_amp.unsqueeze(0),
118
+ radar_phase.unsqueeze(0),
119
+ scale_factor=self.proc_params["scale_factor"],
120
+ ).squeeze(0)
121
+
122
+ # === Depth (always needed) ===
123
+ depth = np.asarray(arrs["depth"][frame_idx]).astype(np.float32)
124
+ if self.depth_in_meters:
125
+ depth = depth / 1000.0
126
+ invalid = ~(np.isfinite(depth) & (depth > 0))
127
+ depth[invalid] = 0.0
128
+ depth = depth[np.newaxis, ...]
129
+ depth_tensor = torch.from_numpy(depth).float()
130
+
131
+ processed_depth = depth_collator(
132
+ depth_tensor.unsqueeze(0),
133
+ max_depth_m=self.proc_params["max_depth_m"],
134
+ target_size=self.proc_params["depth_res"],
135
+ ).squeeze(0)
136
+
137
+ # === RGB (only if use_rgb is True) ===
138
+ out = {
139
+ "radar": processed_radar,
140
+ "depth": processed_depth,
141
+ "sequence": seq_name,
142
+ "frame_idx": frame_idx,
143
+ }
144
+
145
+ if self.proc_params["use_rgb"]:
146
+ rgb = np.asarray(arrs["rgb"][frame_idx])
147
+ rgb = np.transpose(rgb, (2, 0, 1))
148
+ if self.rgb_normalize:
149
+ rgb = rgb.astype(np.float32) / 255.0
150
+ rgb_tensor = torch.from_numpy(rgb)
151
+
152
+ out["rgb"] = dji_rgb_collator(
153
+ rgb_tensor.unsqueeze(0),
154
+ target_size=self.proc_params["rgb_res"],
155
+ ).squeeze(0)
156
+
157
+ return out
158
+
159
+
src/Ablation/ours_radar_no_doppler/split.json ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "test-rice": [
3
+ "Dell-1",
4
+ "Dell-2",
5
+ "Smoke-Dell-1",
6
+ "Smoke-Dell-2",
7
+ "Keck-1",
8
+ "Keck-2",
9
+ "Keck-3",
10
+ "Smoke-keck-1",
11
+ "Smoke-keck-2",
12
+ "Smoke-keck-3"
13
+ ],
14
+ "test-iq1m": [
15
+ "cfa.cfa.1.fwd",
16
+ "cfa.cfa.1.lat",
17
+ "cfa.cfa.3.fwd",
18
+ "cfa.cfa.3.lat",
19
+ "cfa.cfa.a.fwd",
20
+ "cfa.cfa.a.lat",
21
+ "morrison.morrison.1.fwd",
22
+ "morrison.morrison.1.lat",
23
+ "morrison.morrison.2.fwd",
24
+ "morrison.morrison.2.lat",
25
+ "posner.posner.1.fwd",
26
+ "posner.posner.1.lat",
27
+ "posner.posner.2.fwd",
28
+ "posner.posner.2.lat",
29
+ "posner.posner.3.fwd",
30
+ "posner.posner.3.lat",
31
+ "posner.posner.a.fwd",
32
+ "posner.posner.a.lat"
33
+ ]
34
+ }
src/Ablation/ours_radar_no_grad/inference.py ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Direct Accelerate backend for the no-gradient-loss Stage-1 ablation."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import sys
6
+ from pathlib import Path
7
+
8
+
9
+ STAGE2_DIR = Path(__file__).resolve().parents[2] / "GRADE" / "stage2_diffusion_refinement"
10
+ sys.path.insert(0, str(STAGE2_DIR))
11
+
12
+ from inference import main as stage2_main # noqa: E402
13
+
14
+
15
+ if __name__ == "__main__":
16
+ stage2_main()
src/Baselines/cafnet/collate_fn_helpers.py ADDED
@@ -0,0 +1,404 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import numpy as np
3
+ import torch
4
+ from functools import lru_cache
5
+ from typing import Callable, Dict, Sequence, Tuple, Union
6
+ from torchvision import transforms as T
7
+
8
+
9
+ IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
10
+ IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
11
+
12
+ # ZED intrinsics at 1280x720 reference resolution.
13
+ _K_ZED_REF = np.array(
14
+ [
15
+ [521.581604, 0.0, 636.33398438],
16
+ [0.0, 521.581604, 373.10964966],
17
+ [0.0, 0.0, 1.0],
18
+ ],
19
+ dtype=np.float64,
20
+ )
21
+ _ZED_REF_W = 1280
22
+ _ZED_REF_H = 720
23
+
24
+ # DJI calibration constants.
25
+ _CALIB_K_DJI = np.array(
26
+ [
27
+ [718.48555551, 0.0, 963.36465011],
28
+ [0.0, 720.25844189, 537.87569913],
29
+ [0.0, 0.0, 1.0],
30
+ ],
31
+ dtype=np.float64,
32
+ )
33
+ _CALIB_D_DJI = np.array(
34
+ [0.19022699, 0.03466753, 0.05858962, -0.07070669], dtype=np.float64
35
+ )
36
+ _CALIB_DEFISH_SHAPE = (1920, 1080)
37
+ _CALIB_DEFISH_BALANCE = 0.2
38
+ _CALIB_H_FULL = np.array(
39
+ [
40
+ [0.8274446551892256, -0.0742944198979625, 80.23797348979947],
41
+ [-0.014725864916652691, 0.8471179917075127, 28.27366063997317],
42
+ [-5.083573451500717e-05, -6.846079418201229e-05, 1.0],
43
+ ],
44
+ dtype=np.float64,
45
+ )
46
+ _CALIB_OUT_SIZE = (1918, 1105)
47
+ _CALIB_CROP = (115, 255, 1400, 760) # top, left, right, bottom
48
+
49
+
50
+ @lru_cache(maxsize=1)
51
+ def _get_dji_defish_maps() -> Tuple[np.ndarray, np.ndarray]:
52
+ r_defish = np.eye(3)
53
+ k_new_defish = cv2.fisheye.estimateNewCameraMatrixForUndistortRectify(
54
+ _CALIB_K_DJI,
55
+ _CALIB_D_DJI,
56
+ _CALIB_DEFISH_SHAPE,
57
+ r_defish,
58
+ balance=_CALIB_DEFISH_BALANCE,
59
+ fov_scale=1.0,
60
+ )
61
+ map1, map2 = cv2.fisheye.initUndistortRectifyMap(
62
+ _CALIB_K_DJI,
63
+ _CALIB_D_DJI,
64
+ r_defish,
65
+ k_new_defish,
66
+ _CALIB_DEFISH_SHAPE,
67
+ cv2.CV_16SC2,
68
+ )
69
+ return map1, map2
70
+
71
+
72
+ def resize_depth_mm(depth_mm: np.ndarray, target_size: Tuple[int, int]) -> np.ndarray:
73
+ target_h, target_w = target_size
74
+ if depth_mm.shape[:2] == (target_h, target_w):
75
+ return depth_mm
76
+ return cv2.resize(depth_mm, (target_w, target_h), interpolation=cv2.INTER_NEAREST)
77
+
78
+
79
+ def depth_collator(
80
+ depth: Union[torch.Tensor, np.ndarray],
81
+ max_depth_m: float = 11.2,
82
+ target_size: Tuple[int, int] = (128, 256),
83
+ ) -> Union[torch.Tensor, np.ndarray]:
84
+ """Clamp, normalize to [0, 1], and resize depth."""
85
+ is_numpy = isinstance(depth, np.ndarray)
86
+ if is_numpy:
87
+ depth = torch.from_numpy(depth)
88
+
89
+ depth = depth.float()
90
+ original_shape = depth.shape
91
+
92
+ if depth.dim() == 2:
93
+ depth = depth.unsqueeze(0)
94
+ elif depth.dim() == 3:
95
+ depth = depth.unsqueeze(1)
96
+
97
+ invalid_mask = ~(torch.isfinite(depth) & (depth >= 0))
98
+ depth[invalid_mask] = 0.0
99
+
100
+ depth = torch.clamp(depth, min=0.0, max=max_depth_m)
101
+ depth = depth / max_depth_m
102
+
103
+ invalid_mask = ~torch.isfinite(depth)
104
+ depth[invalid_mask] = 0.0
105
+
106
+ resized = T.Resize(
107
+ target_size, interpolation=T.InterpolationMode.BILINEAR, antialias=True
108
+ )(depth)
109
+
110
+ if len(original_shape) == 2:
111
+ resized = resized.squeeze(0)
112
+
113
+ return resized.numpy() if is_numpy else resized
114
+
115
+
116
+ def dji_rgb_collator(
117
+ image: torch.Tensor,
118
+ target_size: Tuple[int, int] = (128, 256),
119
+ ) -> torch.Tensor:
120
+ """Rectify and resize DJI RGB image batch.
121
+
122
+ Args:
123
+ image: Tensor with shape (B, C, H, W).
124
+ target_size: Target resolution as (height, width).
125
+
126
+ Returns:
127
+ Tensor in CHW format (B, C, H, W), float32 in [0, 1].
128
+ """
129
+ if not isinstance(image, torch.Tensor):
130
+ raise ValueError(f"Expected torch.Tensor, got {type(image)}")
131
+
132
+ if image.dim() != 4:
133
+ raise ValueError(
134
+ f"Expected 4D tensor (B, C, H, W), got {image.dim()}D tensor with shape {image.shape}"
135
+ )
136
+
137
+ map1_defish, map2_defish = _get_dji_defish_maps()
138
+ target_h, target_w = target_size
139
+
140
+ if image.max() <= 1.0:
141
+ img_batch = (image.permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8)
142
+ else:
143
+ img_batch = image.permute(0, 2, 3, 1).cpu().numpy().astype(np.uint8)
144
+
145
+ calibrated_images = []
146
+ for img in img_batch:
147
+ if img.shape[1] != 1920 or img.shape[0] != 1080:
148
+ img = cv2.resize(img, (1920, 1080), interpolation=cv2.INTER_LINEAR)
149
+
150
+ img = cv2.remap(img, map1_defish, map2_defish, interpolation=cv2.INTER_LINEAR)
151
+ img = cv2.warpPerspective(
152
+ img, _CALIB_H_FULL, _CALIB_OUT_SIZE, flags=cv2.INTER_LINEAR
153
+ )
154
+
155
+ top, left, right, bottom = _CALIB_CROP
156
+ img = img[top:bottom, left:right]
157
+ img = cv2.resize(img, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
158
+ calibrated_images.append(img)
159
+
160
+ out_batch = np.stack(calibrated_images, axis=0)
161
+ out_tensor = torch.from_numpy(out_batch).permute(0, 3, 1, 2).float() / 255.0
162
+ return out_tensor
163
+
164
+
165
+ def point_cloud_to_sparse_depth(
166
+ points_xyz: np.ndarray,
167
+ target_shape: Tuple[int, int],
168
+ max_depth_m: float,
169
+ ) -> np.ndarray:
170
+ """Project xyz radar points (meters) to a sparse depth image."""
171
+ target_h, target_w = target_shape
172
+ sparse_depth = np.zeros((target_h, target_w), dtype=np.float32)
173
+
174
+ if points_xyz.size == 0:
175
+ return sparse_depth
176
+
177
+ pts = np.asarray(points_xyz, dtype=np.float32)
178
+ if pts.ndim != 2 or pts.shape[1] != 3:
179
+ return sparse_depth
180
+
181
+ valid = np.isfinite(pts).all(axis=1)
182
+ valid &= pts[:, 2] > 0.0
183
+ valid &= pts[:, 2] <= float(max_depth_m)
184
+ pts = pts[valid]
185
+ if pts.shape[0] == 0:
186
+ return sparse_depth
187
+
188
+ sx = target_w / float(_ZED_REF_W)
189
+ sy = target_h / float(_ZED_REF_H)
190
+ fx = _K_ZED_REF[0, 0] * sx
191
+ fy = _K_ZED_REF[1, 1] * sy
192
+ cx = _K_ZED_REF[0, 2] * sx
193
+ cy = _K_ZED_REF[1, 2] * sy
194
+
195
+ z = pts[:, 2]
196
+ u = np.rint(pts[:, 0] * fx / z + cx).astype(np.int32)
197
+ v = np.rint(pts[:, 1] * fy / z + cy).astype(np.int32)
198
+
199
+ in_bounds = (u >= 0) & (u < target_w) & (v >= 0) & (v < target_h)
200
+ if not np.any(in_bounds):
201
+ return sparse_depth
202
+
203
+ u = u[in_bounds]
204
+ v = v[in_bounds]
205
+ z = z[in_bounds].astype(np.float32)
206
+
207
+ min_depth = np.full((target_h, target_w), np.inf, dtype=np.float32)
208
+ np.minimum.at(min_depth, (v, u), z)
209
+ min_depth[~np.isfinite(min_depth)] = 0.0
210
+ return min_depth
211
+
212
+
213
+ def build_radar_gt_map(
214
+ depth_m: np.ndarray,
215
+ sparse_depth: np.ndarray,
216
+ patch_size: Tuple[int, int],
217
+ max_dist_correspondence: float,
218
+ ) -> np.ndarray:
219
+ """Build confidence GT using local depth consistency around each radar pixel."""
220
+ h, w = depth_m.shape
221
+ radar_gt = np.zeros((h, w), dtype=np.float32)
222
+
223
+ ys, xs = np.where(sparse_depth > 0)
224
+ if len(ys) == 0:
225
+ return radar_gt
226
+
227
+ ext_h, ext_w = int(patch_size[0]), int(patch_size[1])
228
+ for y, x in zip(ys, xs):
229
+ radar_depth = sparse_depth[y, x]
230
+
231
+ delta_x1 = min(x, ext_w)
232
+ delta_y1 = min(y, ext_h)
233
+ delta_x2 = min(w - x, ext_w)
234
+ delta_y2 = min(h - y, ext_h)
235
+
236
+ x1 = x - delta_x1
237
+ y1 = y - delta_y1
238
+ x2 = x + delta_x2
239
+ y2 = y + delta_y2
240
+
241
+ distance = np.abs(depth_m[y1:y2, x1:x2] - radar_depth)
242
+ gt_label = (distance < float(max_dist_correspondence)).astype(np.float32)
243
+ radar_gt[y1:y2, x1:x2] = gt_label
244
+
245
+ return radar_gt
246
+
247
+
248
+ def make_rice_collate_fn(
249
+ input_height: int,
250
+ input_width: int,
251
+ radar_max_depth_m: float,
252
+ max_dist_correspondence: float,
253
+ patch_size: Tuple[int, int],
254
+ ) -> Callable[[Sequence[Dict[str, object]]], Tuple[torch.Tensor, ...]]:
255
+ """Create collate_fn for RiceDataset samples.
256
+
257
+ Each dataset sample should contain:
258
+ - sample_idx: int
259
+ - dji_rgb: (H, W, 3) uint8
260
+ - zed_depth_mm: (H, W) uint16
261
+ - radar_pcd_xyz: (N, 3) float32 in meters
262
+ """
263
+
264
+ mean = torch.tensor(IMAGENET_MEAN, dtype=torch.float32).view(1, 3, 1, 1)
265
+ std = torch.tensor(IMAGENET_STD, dtype=torch.float32).view(1, 3, 1, 1)
266
+
267
+ def _collate(batch: Sequence[Dict[str, object]]) -> Tuple[torch.Tensor, ...]:
268
+ if len(batch) == 0:
269
+ raise ValueError("Received empty batch in collate function")
270
+
271
+ sample_indices = []
272
+ rgb_batch = []
273
+ depth_batch = []
274
+ radar_batch = []
275
+ radar_gt_batch = []
276
+
277
+ for sample in batch:
278
+ sample_indices.append(int(sample["sample_idx"]))
279
+
280
+ rgb = np.asarray(sample["dji_rgb"]).copy()
281
+ if rgb.ndim != 3 or rgb.shape[2] != 3:
282
+ raise ValueError(f"Expected RGB shape (H, W, 3), got {rgb.shape}")
283
+ rgb_batch.append(torch.from_numpy(np.transpose(rgb, (2, 0, 1))))
284
+
285
+ depth_mm = np.asarray(sample["zed_depth_mm"]).copy()
286
+ depth_mm = resize_depth_mm(depth_mm, (input_height, input_width))
287
+ depth_m = depth_mm.astype(np.float32) / 1000.0
288
+ invalid = ~(np.isfinite(depth_m) & (depth_m > 0.0))
289
+ depth_m[invalid] = 0.0
290
+ depth_batch.append(depth_m)
291
+
292
+ radar_points = np.asarray(sample["radar_pcd_xyz"], dtype=np.float32)
293
+ if radar_points.ndim != 2 or radar_points.shape[1] != 3:
294
+ radar_points = np.zeros((0, 3), dtype=np.float32)
295
+
296
+ if radar_points.shape[0] == 0:
297
+ center_v = float(depth_m[input_height // 2, input_width // 2])
298
+ if not np.isfinite(center_v):
299
+ center_v = 0.0
300
+ radar_points = np.array([[0.0, 0.0, center_v]], dtype=np.float32)
301
+
302
+ sparse_depth = point_cloud_to_sparse_depth(
303
+ radar_points,
304
+ target_shape=(input_height, input_width),
305
+ max_depth_m=radar_max_depth_m,
306
+ )
307
+ radar_gt = build_radar_gt_map(
308
+ depth_m,
309
+ sparse_depth,
310
+ patch_size=patch_size,
311
+ max_dist_correspondence=max_dist_correspondence,
312
+ )
313
+ radar_batch.append(sparse_depth)
314
+ radar_gt_batch.append(radar_gt)
315
+
316
+ rgb_tensor = torch.stack(rgb_batch, dim=0).float()
317
+ rgb_tensor = dji_rgb_collator(rgb_tensor, target_size=(input_height, input_width))
318
+ rgb_tensor = (rgb_tensor - mean) / std
319
+
320
+ depth_tensor = torch.from_numpy(np.stack(depth_batch, axis=0)).float().unsqueeze(1)
321
+ radar_tensor = torch.from_numpy(np.stack(radar_batch, axis=0)).float().unsqueeze(1)
322
+ radar_gt_tensor = (
323
+ torch.from_numpy(np.stack(radar_gt_batch, axis=0)).float().unsqueeze(1)
324
+ )
325
+ idx_tensor = torch.tensor(sample_indices, dtype=torch.long)
326
+
327
+ return idx_tensor, rgb_tensor, depth_tensor, radar_tensor, radar_gt_tensor
328
+
329
+ return _collate
330
+
331
+
332
+ # Fisheye RGB Handler Functions ##
333
+ def fisheye_rgb_collator(
334
+ image: torch.Tensor,
335
+ target_size: Tuple[int, int] = (128, 256),
336
+ ) -> torch.Tensor:
337
+ """Calibrate and resize Fisheye RGB image batch.
338
+
339
+ Args:
340
+ image: Batch of Fisheye RGB images as torch tensor (B, C, H, W) in CHW format
341
+ target_size: Target resolution as (height, width)
342
+
343
+ Returns:
344
+ Batch of calibrated and resized torch tensors in CHW format (B, C, H, W)
345
+ """
346
+ IMAGE_WIDTH = 1920
347
+ IMAGE_HEIGHT = 1080
348
+ FOCAL_LENGTH_X = 0.613260
349
+ FOCAL_LENGTH_Y = 0.613260
350
+ CENTER_X = 0.5
351
+ CENTER_Y = 0.5
352
+ K1 = -0.120000
353
+ K2 = -0.015000
354
+
355
+ w, h = IMAGE_WIDTH, IMAGE_HEIGHT
356
+ x_out, y_out = np.meshgrid(np.arange(w), np.arange(h))
357
+ x_norm = (x_out - w * CENTER_X) / (w * FOCAL_LENGTH_X)
358
+ y_norm = (y_out - h * CENTER_Y) / (h * FOCAL_LENGTH_Y)
359
+ r = np.sqrt(x_norm**2 + y_norm**2)
360
+ r_distorted = r + K1 * r**2 + K2 * r**3
361
+ r_safe = np.where(r > 0, r, 1.0)
362
+ scale = np.where(r > 0, r_distorted / r_safe, 1.0)
363
+ x_norm_distorted = x_norm * scale
364
+ y_norm_distorted = y_norm * scale
365
+ map_x = (x_norm_distorted * (w * FOCAL_LENGTH_X) + w * CENTER_X).astype(np.float32)
366
+ map_y = (y_norm_distorted * (h * FOCAL_LENGTH_Y) + h * CENTER_Y).astype(np.float32)
367
+
368
+ if not isinstance(image, torch.Tensor):
369
+ raise ValueError(f"Expected torch.Tensor, got {type(image)}")
370
+
371
+ if image.dim() != 4:
372
+ raise ValueError(
373
+ f"Expected 4D tensor (B, C, H, W), got {image.dim()}D tensor with shape {image.shape}"
374
+ )
375
+
376
+ if image.max() <= 1.0:
377
+ img_batch = (image.permute(0, 2, 3, 1).cpu().numpy() * 255).astype(np.uint8)
378
+ else:
379
+ img_batch = image.permute(0, 2, 3, 1).cpu().numpy().astype(np.uint8)
380
+
381
+ calibrated_images = []
382
+ target_h, target_w = target_size
383
+
384
+ for img in img_batch:
385
+ if img.shape[1] != IMAGE_WIDTH or img.shape[0] != IMAGE_HEIGHT:
386
+ img = cv2.resize(
387
+ img, (IMAGE_WIDTH, IMAGE_HEIGHT), interpolation=cv2.INTER_LINEAR
388
+ )
389
+
390
+ img = cv2.remap(
391
+ img,
392
+ map_x,
393
+ map_y,
394
+ interpolation=cv2.INTER_LINEAR,
395
+ borderMode=cv2.BORDER_CONSTANT,
396
+ borderValue=(0, 0, 0),
397
+ )
398
+
399
+ img = cv2.resize(img, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
400
+ calibrated_images.append(img)
401
+
402
+ out_batch = np.stack(calibrated_images, axis=0)
403
+ out_tensor = torch.from_numpy(out_batch).permute(0, 3, 1, 2).float() / 255.0
404
+ return out_tensor
src/Baselines/cafnet/dataloader.py ADDED
@@ -0,0 +1,100 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Optional
2
+
3
+ from torch.utils.data import DataLoader
4
+
5
+ from collate_fn_helpers import make_rice_collate_fn
6
+ from rice_dataset import RiceDataset
7
+
8
+
9
+ def _build_dataset(
10
+ args,
11
+ split: str,
12
+ base_dir: Optional[str] = None,
13
+ split_json_path: Optional[str] = None,
14
+ ) -> RiceDataset:
15
+ return RiceDataset(
16
+ base_dir=base_dir or args.base_dir,
17
+ split_json_path=args.split_json if split_json_path is None else split_json_path,
18
+ split=split,
19
+ input_height=args.input_height,
20
+ input_width=args.input_width,
21
+ patch_size=args.patch_size,
22
+ )
23
+
24
+
25
+ def _build_loader(
26
+ args,
27
+ split: str,
28
+ batch_size: int,
29
+ shuffle: bool,
30
+ drop_last: bool,
31
+ pin_memory: bool,
32
+ base_dir: Optional[str] = None,
33
+ split_json_path: Optional[str] = None,
34
+ ):
35
+ dataset = _build_dataset(
36
+ args,
37
+ split=split,
38
+ base_dir=base_dir,
39
+ split_json_path=split_json_path,
40
+ )
41
+ collate_fn = make_rice_collate_fn(
42
+ input_height=args.input_height,
43
+ input_width=args.input_width,
44
+ radar_max_depth_m=args.radar_max_depth_m,
45
+ max_dist_correspondence=args.max_dist_correspondence,
46
+ patch_size=dataset.patch_size,
47
+ )
48
+ return DataLoader(
49
+ dataset,
50
+ batch_size=batch_size,
51
+ shuffle=shuffle,
52
+ num_workers=args.num_workers,
53
+ pin_memory=pin_memory,
54
+ drop_last=drop_last,
55
+ collate_fn=collate_fn,
56
+ )
57
+
58
+
59
+ def create_train_test_loaders(args, pin_memory: bool = False):
60
+ train_loader = _build_loader(
61
+ args,
62
+ split="train",
63
+ batch_size=args.batch_size,
64
+ shuffle=True,
65
+ drop_last=True,
66
+ pin_memory=pin_memory,
67
+ )
68
+ test_loader = _build_loader(
69
+ args,
70
+ split="test",
71
+ batch_size=args.batch_size,
72
+ shuffle=False,
73
+ drop_last=False,
74
+ pin_memory=pin_memory,
75
+ )
76
+ return train_loader, test_loader
77
+
78
+
79
+ def create_inference_loader(args, pin_memory: bool = False):
80
+ """Create the single packaged Smoke-Eval loader used for inference."""
81
+
82
+ test_base_dir = getattr(args, "test_base_dir", "")
83
+ if not test_base_dir:
84
+ raise ValueError("Config must define 'test_base_dir' for inference.")
85
+
86
+ test_split = getattr(args, "test_split", "train")
87
+ test_split_json = getattr(args, "test_split_json", None)
88
+ if not test_split_json:
89
+ test_split_json = None
90
+
91
+ return _build_loader(
92
+ args,
93
+ split=test_split,
94
+ batch_size=args.batch_size,
95
+ shuffle=False,
96
+ drop_last=False,
97
+ pin_memory=pin_memory,
98
+ base_dir=test_base_dir,
99
+ split_json_path=test_split_json,
100
+ )
src/Baselines/cafnet/extract_pcd_from_depth.py ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import numpy as np
3
+
4
+ # ZED intrinsics at reference resolution 1280x720 (same values as PointCloudConverter)
5
+ _K_ZED_REF = np.array(
6
+ [
7
+ [521.581604, 0.0, 636.33398438],
8
+ [0.0, 521.581604, 373.10964966],
9
+ [0.0, 0.0, 1.0],
10
+ ],
11
+ dtype=np.float64,
12
+ )
13
+ _ZED_REF_W = 1280
14
+ _ZED_REF_H = 720
15
+
16
+
17
+ def sample_depth_as_radar(
18
+ depth_mm: np.ndarray,
19
+ n_samples: int = 100,
20
+ target_shape: tuple = (300, 1280),
21
+ max_depth_m: float = 11.2,
22
+ seed: int | None = None,
23
+ ) -> tuple:
24
+ """
25
+ Randomly sample points from a ground truth ZED depth map and treat them as
26
+ radar points, mimicking the sparse depth input the model expects.
27
+
28
+ The input depth is resized from its native resolution (e.g. 896x504) to
29
+ target_shape using nearest-neighbor interpolation so raw mm values are
30
+ preserved. Camera intrinsics are scaled from the 1280x720 ZED reference to
31
+ match the target resolution.
32
+
33
+ Args:
34
+ depth_mm: Ground truth depth map, shape (H, W), dtype uint16, in mm.
35
+ n_samples: Number of points to randomly sample (default: 100).
36
+ target_shape: (target_H, target_W) to resize to before sampling.
37
+ Default (300, 1280) matches the model's required input.
38
+ max_depth_m: Maximum valid depth in meters — pixels beyond this are
39
+ treated as invalid (default: 11.2 m).
40
+ seed: Optional random seed for reproducibility.
41
+
42
+ Returns:
43
+ points (np.ndarray): (N, 3) float32 array of [X, Y, Z] in meters,
44
+ in camera coordinate frame. N <= n_samples.
45
+ sparse_depth (np.ndarray): (target_H, target_W) float32 sparse depth map
46
+ with only the N sampled pixels filled (meters),
47
+ zeros elsewhere.
48
+ """
49
+ target_h, target_w = target_shape
50
+
51
+ # --- 1. Resize depth map (nearest-neighbor preserves raw mm values) ---
52
+ in_h, in_w = depth_mm.shape
53
+ if (in_h, in_w) != (target_h, target_w):
54
+ depth_resized = cv2.resize(
55
+ depth_mm, (target_w, target_h), interpolation=cv2.INTER_NEAREST
56
+ )
57
+ else:
58
+ depth_resized = depth_mm.copy()
59
+
60
+ # --- 2. Scale intrinsics from 1280x720 reference to target resolution ---
61
+ sx = target_w / float(_ZED_REF_W)
62
+ sy = target_h / float(_ZED_REF_H)
63
+ fx = _K_ZED_REF[0, 0] * sx
64
+ fy = _K_ZED_REF[1, 1] * sy
65
+ cx = _K_ZED_REF[0, 2] * sx
66
+ cy = _K_ZED_REF[1, 2] * sy
67
+
68
+ # --- 3. Convert to float meters and find valid pixels ---
69
+ depth_m = depth_resized.astype(np.float32) / 1000.0
70
+ valid_mask = (depth_m > 0) & (depth_m <= max_depth_m)
71
+ valid_v, valid_u = np.where(valid_mask) # row (V), col (U)
72
+
73
+ if len(valid_v) == 0:
74
+ return (
75
+ np.zeros((0, 3), dtype=np.float32),
76
+ np.zeros((target_h, target_w), dtype=np.float32),
77
+ )
78
+
79
+ # --- 4. Randomly sample up to n_samples valid pixels ---
80
+ rng = np.random.default_rng(seed)
81
+ n = min(n_samples, len(valid_v))
82
+ indices = rng.choice(len(valid_v), size=n, replace=False)
83
+ sampled_v = valid_v[indices]
84
+ sampled_u = valid_u[indices]
85
+ sampled_z = depth_m[sampled_v, sampled_u]
86
+
87
+ # --- 5. Back-project to 3D camera coordinates (pinhole model) ---
88
+ X = (sampled_u - cx) * sampled_z / fx
89
+ Y = (sampled_v - cy) * sampled_z / fy
90
+ points = np.stack([X, Y, sampled_z], axis=1).astype(np.float32) # (N, 3)
91
+
92
+ # --- 6. Build sparse depth map ---
93
+ sparse_depth = np.zeros((target_h, target_w), dtype=np.float32)
94
+ sparse_depth[sampled_v, sampled_u] = sampled_z
95
+
96
+ return points, sparse_depth
src/Baselines/cafnet/inference.py ADDED
@@ -0,0 +1,224 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ from typing import Dict, List
4
+
5
+ import numpy as np
6
+ import torch
7
+ import torch.distributed as dist
8
+ import yaml
9
+ from accelerate import Accelerator
10
+ from accelerate.utils import DistributedDataParallelKwargs, set_seed
11
+ from safetensors.torch import load_file
12
+ from tqdm.auto import tqdm
13
+
14
+ from dataloader import create_inference_loader
15
+ from models.model import CaFNet
16
+
17
+
18
+ DEFAULT_CONFIG = {
19
+ # Packaged evaluation dataset.
20
+ "base_dir": "",
21
+ "split_json": None,
22
+ "test_base_dir": None,
23
+ "test_split": "train",
24
+ "test_split_json": None,
25
+ # Input and radar processing
26
+ "input_height": 288,
27
+ "input_width": 512,
28
+ "radar_max_depth_m": 11.2,
29
+ "max_dist_correspondence": 0.5,
30
+ "patch_size": None,
31
+ # Model
32
+ "encoder": "resnet34_bts",
33
+ "encoder_radar": "resnet18",
34
+ "radar_input_channels": 1,
35
+ "bts_size": 512,
36
+ "max_depth": 11.2,
37
+ # Runtime
38
+ "batch_size": 8,
39
+ # Windows uses spawn-based multiprocessing; keep the public evaluation
40
+ # entry point portable and deterministic by default.
41
+ "num_workers": 0,
42
+ "seed": 42,
43
+ "cpu": False,
44
+ "mixed_precision": "fp16",
45
+ "checkpoint_path": "checkpoints/cafnet.safetensors",
46
+ "prediction_dir": "prediction",
47
+ }
48
+
49
+
50
+ def parse_args():
51
+ parser = argparse.ArgumentParser(description="Run CaFNet inference on Smoke-Eval.")
52
+ parser.add_argument("--config", type=str, required=True, help="Path to YAML config")
53
+ return parser.parse_args()
54
+
55
+
56
+ def load_config(path):
57
+ with open(path, "r") as f:
58
+ cfg = yaml.safe_load(f) or {}
59
+ if not isinstance(cfg, dict):
60
+ raise ValueError("Config must be a YAML mapping (key-value pairs).")
61
+
62
+ merged = dict(DEFAULT_CONFIG)
63
+ merged.update(cfg)
64
+
65
+ if not merged["test_base_dir"]:
66
+ raise ValueError("Config must define 'test_base_dir'.")
67
+ if not merged["checkpoint_path"]:
68
+ raise ValueError("Config must define 'checkpoint_path'.")
69
+ if not os.path.isfile(merged["checkpoint_path"]):
70
+ raise FileNotFoundError(f"Checkpoint not found: {merged['checkpoint_path']}")
71
+ if merged.get("radar_input_channels", 1) != 1:
72
+ raise ValueError("radar_input_channels must be 1 for this setup.")
73
+
74
+ return argparse.Namespace(**merged)
75
+
76
+
77
+ def build_model_args(args):
78
+ return argparse.Namespace(
79
+ encoder=args.encoder,
80
+ encoder_radar=args.encoder_radar,
81
+ radar_input_channels=args.radar_input_channels,
82
+ input_height=args.input_height,
83
+ input_width=args.input_width,
84
+ max_depth=args.max_depth,
85
+ bts_size=args.bts_size,
86
+ )
87
+
88
+
89
+ def _extract_model_state(checkpoint):
90
+ if isinstance(checkpoint, dict) and isinstance(checkpoint.get("model"), dict):
91
+ return checkpoint["model"]
92
+ if isinstance(checkpoint, dict):
93
+ return checkpoint
94
+ raise ValueError("Unsupported checkpoint format.")
95
+
96
+
97
+ def _gather_objects(accelerator, obj):
98
+ if accelerator.num_processes == 1:
99
+ return [obj]
100
+ if not dist.is_available() or not dist.is_initialized():
101
+ return [obj]
102
+
103
+ gathered = [None for _ in range(accelerator.num_processes)]
104
+ dist.all_gather_object(gathered, obj)
105
+ return gathered
106
+
107
+
108
+ def _merge_predictions(all_rank_predictions):
109
+ merged: Dict[str, Dict[int, np.ndarray]] = {}
110
+ for rank_dict in all_rank_predictions:
111
+ if not rank_dict:
112
+ continue
113
+ for seq_name, frame_map in rank_dict.items():
114
+ seq_slot = merged.setdefault(seq_name, {})
115
+ for frame_idx, pred in frame_map.items():
116
+ frame_idx = int(frame_idx)
117
+ if frame_idx not in seq_slot:
118
+ seq_slot[frame_idx] = pred
119
+ return merged
120
+
121
+
122
+ def _save_sequence_predictions(predictions, out_dir):
123
+ os.makedirs(out_dir, exist_ok=True)
124
+ for seq_name in sorted(predictions.keys()):
125
+ frame_map = predictions[seq_name]
126
+ ordered_frames = sorted(frame_map.keys())
127
+ if not ordered_frames:
128
+ pred_stack = np.zeros((0,), dtype=np.float32)
129
+ else:
130
+ pred_stack = np.stack([frame_map[k] for k in ordered_frames], axis=0).astype(
131
+ np.float32,
132
+ copy=False,
133
+ )
134
+ np.save(os.path.join(out_dir, f"{seq_name.lower()}_pred.npy"), pred_stack)
135
+
136
+
137
+ def _run_loader_inference(accelerator, model, loader, samples, save_dir, desc):
138
+ model.eval()
139
+ local_preds: Dict[str, Dict[int, np.ndarray]] = {}
140
+
141
+ with torch.no_grad():
142
+ pbar = tqdm(
143
+ loader,
144
+ desc=desc,
145
+ disable=not accelerator.is_local_main_process,
146
+ dynamic_ncols=True,
147
+ leave=False,
148
+ )
149
+ for batch in pbar:
150
+ sample_idx, image, depth_gt, radar, radar_gt = batch
151
+
152
+ image = image.to(accelerator.device, non_blocking=True)
153
+ radar = radar.to(accelerator.device, non_blocking=True)
154
+ # Kept for parity with validation loop structure.
155
+ _ = depth_gt.to(accelerator.device, non_blocking=True)
156
+ _ = radar_gt.to(accelerator.device, non_blocking=True)
157
+
158
+ focal = torch.ones((image.size(0),), device=image.device)
159
+ _, _, _, _, depth_est, _, _ = model(image, radar, focal)
160
+
161
+ pred_np = depth_est.detach().float().cpu().numpy()
162
+ if pred_np.ndim == 4 and pred_np.shape[1] == 1:
163
+ pred_np = pred_np[:, 0]
164
+
165
+ if torch.is_tensor(sample_idx):
166
+ sample_idx_list = sample_idx.detach().cpu().tolist()
167
+ else:
168
+ sample_idx_list = list(sample_idx)
169
+
170
+ for local_i, sample_i in enumerate(sample_idx_list):
171
+ seq_name, frame_idx = samples[int(sample_i)]
172
+ seq_slot = local_preds.setdefault(seq_name, {})
173
+ frame_idx = int(frame_idx)
174
+ if frame_idx not in seq_slot:
175
+ seq_slot[frame_idx] = pred_np[local_i].astype(np.float32, copy=False)
176
+
177
+ gathered = _gather_objects(accelerator, local_preds)
178
+ if accelerator.is_main_process:
179
+ merged = _merge_predictions(gathered)
180
+ _save_sequence_predictions(merged, save_dir)
181
+
182
+ accelerator.wait_for_everyone()
183
+
184
+
185
+ def main():
186
+ cli = parse_args()
187
+ args = load_config(cli.config)
188
+
189
+ set_seed(args.seed)
190
+ ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
191
+ accelerator = Accelerator(
192
+ mixed_precision=None if args.mixed_precision in ("no", "none") else args.mixed_precision,
193
+ cpu=args.cpu,
194
+ kwargs_handlers=[ddp_kwargs],
195
+ )
196
+
197
+ test_loader = create_inference_loader(
198
+ args,
199
+ pin_memory=(accelerator.device.type == "cuda"),
200
+ )
201
+ test_samples: List = test_loader.dataset.samples
202
+
203
+ model = CaFNet(build_model_args(args))
204
+
205
+ model, test_loader = accelerator.prepare(model, test_loader)
206
+
207
+ state_dict = load_file(args.checkpoint_path, device="cpu")
208
+ accelerator.unwrap_model(model).load_state_dict(state_dict, strict=True)
209
+
210
+ _run_loader_inference(
211
+ accelerator=accelerator,
212
+ model=model,
213
+ loader=test_loader,
214
+ samples=test_samples,
215
+ save_dir=args.prediction_dir,
216
+ desc="Inference",
217
+ )
218
+
219
+ if accelerator.is_main_process:
220
+ print(f"Saved predictions to: {args.prediction_dir}")
221
+
222
+
223
+ if __name__ == "__main__":
224
+ main()
src/Baselines/cafnet/inference_config.yaml ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # CaFNet inference config for the packaged Smoke-Eval data.
2
+ test_base_dir: "../../../evaluation_dataset/Smoke-Eval"
3
+ test_split: "train"
4
+ test_split_json: null
5
+
6
+ # Input and radar preprocessing
7
+ input_height: 288
8
+ input_width: 512
9
+ radar_max_depth_m: 11.2
10
+ max_dist_correspondence: 0.5
11
+ patch_size: [64, 128]
12
+
13
+ # Model architecture
14
+ encoder: resnet34_bts
15
+ encoder_radar: resnet18
16
+ radar_input_channels: 1
17
+ bts_size: 512
18
+ max_depth: 11.2
19
+
20
+ # Runtime
21
+ batch_size: 32
22
+ num_workers: 0
23
+ seed: 42
24
+ cpu: false
25
+ mixed_precision: "fp16"
26
+
27
+ # Checkpoint and output root
28
+ checkpoint_path: "../../../checkpoints/baselines/cafnet/cafnet.safetensors"
29
+ prediction_dir: "prediction"
src/Baselines/cafnet/models/bts.py ADDED
@@ -0,0 +1,367 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (C) 2019 Jin Han Lee
2
+ #
3
+ # This file is a part of BTS.
4
+ # This program is free software: you can redistribute it and/or modify
5
+ # it under the terms of the GNU General Public License as published by
6
+ # the Free Software Foundation, either version 3 of the License, or
7
+ # (at your option) any later version.
8
+ #
9
+ # This program is distributed in the hope that it will be useful,
10
+ # but WITHOUT ANY WARRANTY; without even the implied warranty of
11
+ # MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
12
+ # GNU General Public License for more details.
13
+ #
14
+ # You should have received a copy of the GNU General Public License
15
+ # along with this program. If not, see <http://www.gnu.org/licenses/>
16
+
17
+ import torch
18
+ import torch.nn as nn
19
+ import torch.nn.functional as torch_nn_func
20
+ import math
21
+
22
+
23
+ def bn_init_as_tf(m):
24
+ if isinstance(m, nn.BatchNorm2d):
25
+ m.track_running_stats = True # These two lines enable using stats (moving mean and var) loaded from pretrained model
26
+ m.eval() # or zero mean and variance of one if the batch norm layer has no pretrained values
27
+ m.affine = True
28
+ m.requires_grad = True
29
+
30
+
31
+ def weights_init_xavier(m):
32
+ if isinstance(m, nn.Conv2d):
33
+ torch.nn.init.xavier_uniform_(m.weight)
34
+ if m.bias is not None:
35
+ torch.nn.init.zeros_(m.bias)
36
+
37
+
38
+ class atrous_conv(nn.Sequential):
39
+ def __init__(self, in_channels, out_channels, dilation, apply_bn_first=True):
40
+ super(atrous_conv, self).__init__()
41
+ self.atrous_conv = torch.nn.Sequential()
42
+ if apply_bn_first:
43
+ self.atrous_conv.add_module('first_bn', nn.BatchNorm2d(in_channels, momentum=0.01, affine=True, track_running_stats=True, eps=1.1e-5))
44
+
45
+ self.atrous_conv.add_module('aconv_sequence', nn.Sequential(nn.ReLU(),
46
+ nn.Conv2d(in_channels=in_channels, out_channels=out_channels*2, bias=False, kernel_size=1, stride=1, padding=0),
47
+ nn.BatchNorm2d(out_channels*2, momentum=0.01, affine=True, track_running_stats=True),
48
+ nn.ReLU(),
49
+ nn.Conv2d(in_channels=out_channels * 2, out_channels=out_channels, bias=False, kernel_size=3, stride=1,
50
+ padding=(dilation, dilation), dilation=dilation)))
51
+
52
+ def forward(self, x):
53
+ return self.atrous_conv.forward(x)
54
+
55
+ class upconv(nn.Module):
56
+ def __init__(self, in_channels, out_channels, ratio=2):
57
+ super(upconv, self).__init__()
58
+ self.elu = nn.ELU()
59
+ self.conv = nn.Conv2d(in_channels=in_channels, out_channels=out_channels, bias=False, kernel_size=3, stride=1, padding=1)
60
+ self.ratio = ratio
61
+
62
+ def forward(self, x):
63
+ up_x = torch_nn_func.interpolate(x, scale_factor=self.ratio, mode='nearest')
64
+ out = self.conv(up_x)
65
+ out = self.elu(out)
66
+ return out
67
+
68
+ class reduction_1x1(nn.Sequential):
69
+ def __init__(self, num_in_filters, num_out_filters, max_depth, is_final=False):
70
+ super(reduction_1x1, self).__init__()
71
+ self.max_depth = max_depth
72
+ self.is_final = is_final
73
+ self.sigmoid = nn.Sigmoid()
74
+ self.reduc = torch.nn.Sequential()
75
+
76
+ while num_out_filters >= 4:
77
+ if num_out_filters < 8:
78
+ if self.is_final:
79
+ self.reduc.add_module('final', torch.nn.Sequential(nn.Conv2d(num_in_filters, out_channels=1, bias=False,
80
+ kernel_size=1, stride=1, padding=0),
81
+ nn.Sigmoid()))
82
+ else:
83
+ self.reduc.add_module('plane_params', torch.nn.Conv2d(num_in_filters, out_channels=3, bias=False,
84
+ kernel_size=1, stride=1, padding=0))
85
+ break
86
+ else:
87
+ self.reduc.add_module('inter_{}_{}'.format(num_in_filters, num_out_filters),
88
+ torch.nn.Sequential(nn.Conv2d(in_channels=num_in_filters, out_channels=num_out_filters,
89
+ bias=False, kernel_size=1, stride=1, padding=0),
90
+ nn.ELU()))
91
+
92
+ num_in_filters = num_out_filters
93
+ num_out_filters = num_out_filters // 2
94
+
95
+ def forward(self, net):
96
+ net = self.reduc.forward(net)
97
+ if not self.is_final:
98
+ theta = self.sigmoid(net[:, 0, :, :]) * math.pi / 3
99
+ phi = self.sigmoid(net[:, 1, :, :]) * math.pi * 2
100
+ dist = self.sigmoid(net[:, 2, :, :]) * self.max_depth
101
+ n1 = torch.mul(torch.sin(theta), torch.cos(phi)).unsqueeze(1)
102
+ n2 = torch.mul(torch.sin(theta), torch.sin(phi)).unsqueeze(1)
103
+ n3 = torch.cos(theta).unsqueeze(1)
104
+ n4 = dist.unsqueeze(1)
105
+ net = torch.cat([n1, n2, n3, n4], dim=1)
106
+
107
+ return net
108
+
109
+ class local_planar_guidance(nn.Module):
110
+ def __init__(self, upratio):
111
+ super(local_planar_guidance, self).__init__()
112
+ self.upratio = upratio
113
+ self.u = torch.arange(self.upratio).reshape([1, 1, self.upratio]).float()
114
+ self.v = torch.arange(int(self.upratio)).reshape([1, self.upratio, 1]).float()
115
+ self.upratio = float(upratio)
116
+
117
+ def forward(self, plane_eq, focal):
118
+ plane_eq_expanded = torch.repeat_interleave(plane_eq, int(self.upratio), 2)
119
+ plane_eq_expanded = torch.repeat_interleave(plane_eq_expanded, int(self.upratio), 3)
120
+ n1 = plane_eq_expanded[:, 0, :, :]
121
+ n2 = plane_eq_expanded[:, 1, :, :]
122
+ n3 = plane_eq_expanded[:, 2, :, :]
123
+ n4 = plane_eq_expanded[:, 3, :, :]
124
+
125
+ u = self.u.repeat(plane_eq.size(0), plane_eq.size(2) * int(self.upratio), plane_eq.size(3)).cuda()
126
+ u = (u - (self.upratio - 1) * 0.5) / self.upratio
127
+
128
+ v = self.v.repeat(plane_eq.size(0), plane_eq.size(2), plane_eq.size(3) * int(self.upratio)).cuda()
129
+ v = (v - (self.upratio - 1) * 0.5) / self.upratio
130
+
131
+ return n4 / (n1 * u + n2 * v + n3)
132
+
133
+ class bts_gated_fuse(nn.Module):
134
+ def __init__(self, params, feat_out_channels, feat_out_channels_rad, num_features=512):
135
+ super(bts_gated_fuse, self).__init__()
136
+ self.params = params
137
+ self.weight5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[4], feat_out_channels[4], 1, 1, bias=False),
138
+ nn.Sigmoid())
139
+ self.project5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[4], feat_out_channels[4], 1, 1, bias=False),
140
+ nn.ReLU())
141
+ self.upconv5 = upconv(feat_out_channels[4], num_features)
142
+ self.bn5 = nn.BatchNorm2d(num_features, momentum=0.01, affine=True, eps=1.1e-5)
143
+
144
+ self.conv5 = torch.nn.Sequential(nn.Conv2d(num_features + feat_out_channels[3], num_features, 3, 1, 1, bias=False),
145
+ nn.ELU())
146
+
147
+ self.weight4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[3], num_features, 1, 1, bias=False),
148
+ nn.Sigmoid())
149
+ self.project4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[3], num_features, 1, 1, bias=False),
150
+ nn.ReLU())
151
+ self.upconv4 = upconv(num_features, num_features // 2)
152
+ self.bn4 = nn.BatchNorm2d(num_features // 2, momentum=0.01, affine=True, eps=1.1e-5)
153
+ self.conv4 = torch.nn.Sequential(nn.Conv2d(num_features // 2 + feat_out_channels[2], num_features // 2, 3, 1, 1, bias=False),
154
+ nn.ELU())
155
+ self.bn4_2 = nn.BatchNorm2d(num_features // 2, momentum=0.01, affine=True, eps=1.1e-5)
156
+
157
+ self.daspp_3 = atrous_conv(num_features // 2, num_features // 4, 3, apply_bn_first=False)
158
+ self.daspp_6 = atrous_conv(num_features // 2 + num_features // 4 + feat_out_channels[2], num_features // 4, 6)
159
+ self.daspp_12 = atrous_conv(num_features + feat_out_channels[2], num_features // 4, 12)
160
+ self.daspp_18 = atrous_conv(num_features + num_features // 4 + feat_out_channels[2], num_features // 4, 18)
161
+ self.daspp_24 = atrous_conv(num_features + num_features // 2 + feat_out_channels[2], num_features // 4, 24)
162
+ self.daspp_conv = torch.nn.Sequential(nn.Conv2d(num_features + num_features // 2 + num_features // 4, num_features // 4, 3, 1, 1, bias=False),
163
+ nn.ELU())
164
+ self.reduc8x8 = reduction_1x1(num_features // 4, num_features // 4, self.params.max_depth)
165
+ self.lpg8x8 = local_planar_guidance(8)
166
+
167
+ self.weight3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[2], num_features // 4, 1, 1, bias=False),
168
+ nn.Sigmoid())
169
+ self.project3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[2], num_features // 4, 1, 1, bias=False),
170
+ nn.ReLU())
171
+ self.upconv3 = upconv(num_features // 4, num_features // 4)
172
+ self.bn3 = nn.BatchNorm2d(num_features // 4, momentum=0.01, affine=True, eps=1.1e-5)
173
+ self.conv3 = torch.nn.Sequential(nn.Conv2d(num_features // 4 + feat_out_channels[1] + 1, num_features // 4, 3, 1, 1, bias=False),
174
+ nn.ELU())
175
+ self.reduc4x4 = reduction_1x1(num_features // 4, num_features // 8, self.params.max_depth)
176
+ self.lpg4x4 = local_planar_guidance(4)
177
+
178
+ self.weight2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[1], num_features // 4, 1, 1, bias=False),
179
+ nn.Sigmoid())
180
+ self.project2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[1], num_features // 4, 1, 1, bias=False),
181
+ nn.ReLU())
182
+ self.upconv2 = upconv(num_features // 4, num_features // 8)
183
+ self.bn2 = nn.BatchNorm2d(num_features // 8, momentum=0.01, affine=True, eps=1.1e-5)
184
+ self.conv2 = torch.nn.Sequential(nn.Conv2d(num_features // 8 + feat_out_channels[0] + 1, num_features // 8, 3, 1, 1, bias=False),
185
+ nn.ELU())
186
+
187
+ self.reduc2x2 = reduction_1x1(num_features // 8, num_features // 16, self.params.max_depth)
188
+ self.lpg2x2 = local_planar_guidance(2)
189
+
190
+ self.weight1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[0], num_features // 8, 1, 1, bias=False),
191
+ nn.Sigmoid())
192
+ self.project1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[0], num_features // 8, 1, 1, bias=False),
193
+ nn.ReLU())
194
+ self.upconv1 = upconv(num_features // 8, num_features // 16)
195
+ self.reduc1x1 = reduction_1x1(num_features // 16, num_features // 32, self.params.max_depth, is_final=True)
196
+ self.conv1 = torch.nn.Sequential(nn.Conv2d(num_features // 16 + 4, num_features // 16, 3, 1, 1, bias=False),
197
+ nn.ELU())
198
+ self.get_depth = torch.nn.Sequential(nn.Conv2d(num_features // 16, 1, 3, 1, 1, bias=False),
199
+ nn.Sigmoid())
200
+
201
+ self.pool5 = torch.nn.AvgPool2d(32, 32)
202
+ self.pool4 = torch.nn.AvgPool2d(16, 16)
203
+ self.pool3 = torch.nn.AvgPool2d(8, 8)
204
+ self.pool2 = torch.nn.AvgPool2d(4, 4)
205
+ self.pool1 = torch.nn.AvgPool2d(2, 2)
206
+
207
+ def forward(self, img_features, rad_features, focal, radar_confidence):
208
+ skip0, skip1, skip2, skip3 = img_features[0], img_features[1], img_features[2], img_features[3]
209
+ rad_skip0, rad_skip1, rad_skip2, rad_skip3 = rad_features[0], rad_features[1], rad_features[2], rad_features[3]
210
+
211
+ # prepare radar confidence
212
+ radar_confidence5 = self.pool5(radar_confidence)
213
+ radar_confidence4 = self.pool4(radar_confidence)
214
+ radar_confidence3 = self.pool3(radar_confidence)
215
+ radar_confidence2 = self.pool2(radar_confidence)
216
+ radar_confidence1 = self.pool1(radar_confidence)
217
+
218
+
219
+ rad_weight5 = self.weight5(rad_features[4])
220
+ rad_project5 = self.project5(rad_features[4])
221
+
222
+ dense_features = torch.nn.ReLU()(img_features[4])
223
+ dense_features = dense_features + rad_weight5*rad_project5*radar_confidence5
224
+ upconv5 = self.upconv5(dense_features) # H/16
225
+ upconv5 = self.bn5(upconv5)
226
+ concat5 = torch.cat([upconv5, skip3], dim=1)
227
+ iconv5 = self.conv5(concat5)
228
+
229
+ rad_weight4 = self.weight4(rad_skip3)
230
+ rad_project4 = self.project4(rad_skip3)
231
+
232
+ iconv5 = iconv5 + rad_weight4*rad_project4*radar_confidence4
233
+ upconv4 = self.upconv4(iconv5) # H/8
234
+ upconv4 = self.bn4(upconv4)
235
+ concat4 = torch.cat([upconv4, skip2], dim=1)
236
+ iconv4 = self.conv4(concat4)
237
+ iconv4 = self.bn4_2(iconv4)
238
+
239
+ daspp_3 = self.daspp_3(iconv4)
240
+ concat4_2 = torch.cat([concat4, daspp_3], dim=1)
241
+ daspp_6 = self.daspp_6(concat4_2)
242
+ concat4_3 = torch.cat([concat4_2, daspp_6], dim=1)
243
+ daspp_12 = self.daspp_12(concat4_3)
244
+ concat4_4 = torch.cat([concat4_3, daspp_12], dim=1)
245
+ daspp_18 = self.daspp_18(concat4_4)
246
+ concat4_5 = torch.cat([concat4_4, daspp_18], dim=1)
247
+ daspp_24 = self.daspp_24(concat4_5)
248
+ concat4_daspp = torch.cat([iconv4, daspp_3, daspp_6, daspp_12, daspp_18, daspp_24], dim=1)
249
+ daspp_feat = self.daspp_conv(concat4_daspp)
250
+ rad_weight3 = self.weight3(rad_skip2)
251
+ rad_project3 = self.project3(rad_skip2)
252
+ daspp_feat = daspp_feat + rad_weight3*rad_project3*radar_confidence3
253
+
254
+ reduc8x8 = self.reduc8x8(daspp_feat)
255
+ plane_normal_8x8 = reduc8x8[:, :3, :, :]
256
+ plane_normal_8x8 = torch_nn_func.normalize(plane_normal_8x8, 2, 1)
257
+ plane_dist_8x8 = reduc8x8[:, 3, :, :]
258
+ plane_eq_8x8 = torch.cat([plane_normal_8x8, plane_dist_8x8.unsqueeze(1)], 1)
259
+ depth_8x8 = self.lpg8x8(plane_eq_8x8, focal)
260
+ depth_8x8_scaled = depth_8x8.unsqueeze(1) / self.params.max_depth
261
+ depth_8x8_scaled_ds = torch_nn_func.interpolate(depth_8x8_scaled, scale_factor=0.25, mode='nearest')
262
+
263
+ upconv3 = self.upconv3(daspp_feat) # H/4
264
+ upconv3 = self.bn3(upconv3)
265
+ concat3 = torch.cat([upconv3, skip1, depth_8x8_scaled_ds], dim=1)
266
+ iconv3 = self.conv3(concat3)
267
+ rad_weight2 = self.weight2(rad_skip1)
268
+ rad_project2 = self.project2(rad_skip1)
269
+ iconv3 = iconv3 + rad_weight2*rad_project2*radar_confidence2
270
+
271
+ reduc4x4 = self.reduc4x4(iconv3)
272
+ plane_normal_4x4 = reduc4x4[:, :3, :, :]
273
+ plane_normal_4x4 = torch_nn_func.normalize(plane_normal_4x4, 2, 1)
274
+ plane_dist_4x4 = reduc4x4[:, 3, :, :]
275
+ plane_eq_4x4 = torch.cat([plane_normal_4x4, plane_dist_4x4.unsqueeze(1)], 1)
276
+ depth_4x4 = self.lpg4x4(plane_eq_4x4, focal)
277
+ depth_4x4_scaled = depth_4x4.unsqueeze(1) / self.params.max_depth
278
+ depth_4x4_scaled_ds = torch_nn_func.interpolate(depth_4x4_scaled, scale_factor=0.5, mode='nearest')
279
+
280
+ upconv2 = self.upconv2(iconv3) # H/2
281
+ upconv2 = self.bn2(upconv2)
282
+ concat2 = torch.cat([upconv2, skip0, depth_4x4_scaled_ds], dim=1)
283
+ iconv2 = self.conv2(concat2)
284
+ rad_weight1 = self.weight1(rad_skip0)
285
+ rad_project1 = self.project1(rad_skip0)
286
+ iconv2 = iconv2 + rad_weight1*rad_project1*radar_confidence1
287
+
288
+ reduc2x2 = self.reduc2x2(iconv2)
289
+ plane_normal_2x2 = reduc2x2[:, :3, :, :]
290
+ plane_normal_2x2 = torch_nn_func.normalize(plane_normal_2x2, 2, 1)
291
+ plane_dist_2x2 = reduc2x2[:, 3, :, :]
292
+ plane_eq_2x2 = torch.cat([plane_normal_2x2, plane_dist_2x2.unsqueeze(1)], 1)
293
+ depth_2x2 = self.lpg2x2(plane_eq_2x2, focal)
294
+ depth_2x2_scaled = depth_2x2.unsqueeze(1) / self.params.max_depth
295
+
296
+ rad_weight1 = self.weight1(rad_skip0)
297
+ rad_project1 = self.project1(rad_skip0)
298
+
299
+ upconv1 = self.upconv1(iconv2)
300
+ reduc1x1 = self.reduc1x1(upconv1)
301
+ concat1 = torch.cat([upconv1, reduc1x1, depth_2x2_scaled, depth_4x4_scaled, depth_8x8_scaled], dim=1)
302
+ iconv1 = self.conv1(concat1)
303
+ final_depth = self.params.max_depth * self.get_depth(iconv1)
304
+
305
+ return depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth
306
+
307
+ class encoder_image(nn.Module):
308
+ def __init__(self, params):
309
+ super(encoder_image, self).__init__()
310
+ self.params = params
311
+ import torchvision.models as models
312
+ if params.encoder == 'densenet121_bts':
313
+ self.base_model = models.densenet121(pretrained=False).features
314
+ self.feat_names = ['relu0', 'pool0', 'transition1', 'transition2', 'norm5']
315
+ self.feat_out_channels = [64, 64, 128, 256, 1024]
316
+ elif params.encoder == 'densenet161_bts':
317
+ self.base_model = models.densenet161(pretrained=False).features
318
+ self.feat_names = ['relu0', 'pool0', 'transition1', 'transition2', 'norm5']
319
+ self.feat_out_channels = [96, 96, 192, 384, 2208]
320
+ elif params.encoder == 'resnet50_bts':
321
+ self.base_model = models.resnet50(pretrained=False)
322
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
323
+ self.feat_out_channels = [64, 256, 512, 1024, 2048]
324
+ elif params.encoder == 'resnet34_bts':
325
+ self.base_model = models.resnet34(pretrained=False)
326
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
327
+ self.feat_out_channels = [64, 64, 128, 256, 512]
328
+ elif params.encoder == 'resnet18_bts':
329
+ self.base_model = models.resnet18(pretrained=False)
330
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
331
+ self.feat_out_channels = [64, 64, 128, 256, 512]
332
+ elif params.encoder == 'resnet101_bts':
333
+ self.base_model = models.resnet101(pretrained=False)
334
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
335
+ self.feat_out_channels = [64, 256, 512, 1024, 2048]
336
+ elif params.encoder == 'resnext50_bts':
337
+ self.base_model = models.resnext50_32x4d(pretrained=False)
338
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
339
+ self.feat_out_channels = [64, 256, 512, 1024, 2048]
340
+ elif params.encoder == 'resnext101_bts':
341
+ self.base_model = models.resnext101_32x8d(pretrained=False)
342
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
343
+ self.feat_out_channels = [64, 256, 512, 1024, 2048]
344
+ elif params.encoder == 'mobilenetv2_bts':
345
+ self.base_model = models.mobilenet_v2(pretrained=False).features
346
+ self.feat_inds = [2, 4, 7, 11, 19]
347
+ self.feat_out_channels = [16, 24, 32, 64, 1280]
348
+ self.feat_names = []
349
+ else:
350
+ print('Not supported encoder: {}'.format(params.encoder))
351
+
352
+ def forward(self, x):
353
+ feature = x
354
+ skip_feat = []
355
+ i = 1
356
+ for k, v in self.base_model._modules.items():
357
+ if 'fc' in k or 'avgpool' in k:
358
+ continue
359
+ feature = v(feature)
360
+ if self.params.encoder == 'mobilenetv2_bts':
361
+ if i == 2 or i == 4 or i == 7 or i == 11 or i == 19:
362
+ skip_feat.append(feature)
363
+ else:
364
+ if any(x in k for x in self.feat_names):
365
+ skip_feat.append(feature)
366
+ i = i + 1
367
+ return skip_feat
src/Baselines/cafnet/models/model.py ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ from models.bts import encoder_image, bts_gated_fuse
4
+ from models.radar import encoder_radar_sparse_conv, encoder_radar_sub, decoder_radar
5
+
6
+ class CaFNet(nn.Module):
7
+ def __init__(self, params, threshold=0.4):
8
+ super(CaFNet, self).__init__()
9
+ self.threshold = threshold
10
+ self.encoder = encoder_image(params)
11
+ self.encoder_radar1 = encoder_radar_sparse_conv(params)
12
+ self.decoder_radar = decoder_radar(params, self.encoder.feat_out_channels, self.encoder_radar1.feat_out_channels)
13
+ self.encoder_radar2 = encoder_radar_sub(params)
14
+ self.decoder = bts_gated_fuse(params, self.encoder.feat_out_channels, self.encoder_radar2.feat_out_channels, params.bts_size)
15
+
16
+
17
+ def forward(self, x, radar, focal):
18
+
19
+ skip_feat = self.encoder(x)
20
+ skip_feat_radar = self.encoder_radar1(radar)
21
+ rad_confidence, rad_depth = self.decoder_radar(skip_feat, skip_feat_radar)
22
+ mask = (rad_confidence > self.threshold).float()
23
+ radar_new_input = torch.cat([mask*rad_depth, radar], axis=1)
24
+ skip_feat_radar_new = self.encoder_radar2(radar_new_input)
25
+
26
+ depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth = self.decoder(skip_feat, skip_feat_radar_new, focal, rad_confidence)
27
+
28
+ return depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth, rad_confidence, rad_depth
src/Baselines/cafnet/models/radar.py ADDED
@@ -0,0 +1,212 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from models.bts import upconv
2
+ import torch
3
+ import torch.nn as nn
4
+ import torchvision.models as models
5
+
6
+ class encoder_radar_sparse_conv(nn.Module):
7
+ def __init__(self, params):
8
+ # radar encoder for the first stage
9
+ super(encoder_radar_sparse_conv, self).__init__()
10
+
11
+ self.params = params
12
+ self.sparse_conv1 = SparseConv(params.radar_input_channels, 16, 7, activation='elu')
13
+ self.sparse_conv2 = SparseConv(16, 16, 5, activation='elu')
14
+ self.sparse_conv3 = SparseConv(16, 16, 3, activation='elu')
15
+ self.sparse_conv4 = SparseConv(16, 3, 3, activation='elu')
16
+
17
+ if params.encoder_radar == 'resnet34':
18
+ self.base_model_radar = models.resnet34(pretrained=False)
19
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
20
+ self.feat_out_channels = [64, 64, 128, 256, 512]
21
+ elif params.encoder_radar == 'resnet18':
22
+ self.base_model_radar = models.resnet18(pretrained=False)
23
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
24
+ self.feat_out_channels = [64, 64, 128, 256, 512]
25
+ else:
26
+ print('Not supported encoder: {}'.format(params.encoder))
27
+
28
+ def forward(self, x):
29
+ mask = (x[:, 0] > 0).float().unsqueeze(1)
30
+ feature = x
31
+ feature, mask = self.sparse_conv1(feature, mask)
32
+ feature, mask = self.sparse_conv2(feature, mask)
33
+ feature, mask = self.sparse_conv3(feature, mask)
34
+ feature, mask = self.sparse_conv4(feature, mask)
35
+
36
+ skip_feat = []
37
+ i = 1
38
+ for k, v in self.base_model_radar._modules.items():
39
+ if 'fc' in k or 'avgpool' in k:
40
+ continue
41
+ feature = v(feature)
42
+ if any(x in k for x in self.feat_names):
43
+ skip_feat.append(feature)
44
+ i = i + 1
45
+ return skip_feat
46
+
47
+ class encoder_radar_sub(nn.Module):
48
+ def __init__(self, params):
49
+ # radar encoder for the second stage
50
+ super(encoder_radar_sub, self).__init__()
51
+
52
+ self.params = params
53
+ import torchvision.models as models
54
+ self.conv = torch.nn.Sequential(nn.Conv2d(params.radar_input_channels+1, 3, 3, 1, 1, bias=False),
55
+ nn.ELU())
56
+
57
+ if params.encoder_radar == 'resnet34':
58
+ self.base_model_radar = models.resnet34(pretrained=False)
59
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
60
+ self.feat_out_channels = [64, 64, 128, 256, 512]
61
+ elif params.encoder_radar == 'resnet18':
62
+ self.base_model_radar = models.resnet18(pretrained=False)
63
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
64
+ self.feat_out_channels = [64, 64, 128, 256, 512]
65
+ else:
66
+ print('Not supported encoder: {}'.format(params.encoder))
67
+ def forward(self, x):
68
+ feature = x
69
+ feature = self.conv(feature)
70
+ skip_feat = []
71
+ i = 1
72
+ for k, v in self.base_model_radar._modules.items():
73
+ if 'fc' in k or 'avgpool' in k:
74
+ continue
75
+ feature = v(feature)
76
+ if any(x in k for x in self.feat_names):
77
+ skip_feat.append(feature)
78
+ i = i + 1
79
+ return skip_feat
80
+
81
+
82
+ class decoder_radar(nn.Module):
83
+ def __init__(self, params, feat_out_channels_img, feat_out_channels_radar):
84
+ super(decoder_radar, self).__init__()
85
+ self.params = params
86
+ self.upconv5 = upconv(feat_out_channels_img[4]+feat_out_channels_radar[4], feat_out_channels_radar[4]//2)
87
+ self.bn5 = nn.BatchNorm2d(feat_out_channels_radar[4]//2, momentum=0.01, affine=True, eps=1.1e-5)
88
+ self.conv5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[4]//2, feat_out_channels_radar[4]//2, 3, 1, 1, bias=False),
89
+ nn.ELU())
90
+
91
+ 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)
92
+ self.bn4 = nn.BatchNorm2d(feat_out_channels_radar[3]//2, momentum=0.01, affine=True, eps=1.1e-5)
93
+ self.conv4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[3]//2, feat_out_channels_radar[3]//2, 3, 1, 1, bias=False),
94
+ nn.ELU())
95
+
96
+ 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)
97
+ self.bn3 = nn.BatchNorm2d(feat_out_channels_radar[2]//2, momentum=0.01, affine=True, eps=1.1e-5)
98
+ self.conv3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[2]//2, feat_out_channels_radar[2]//2, 3, 1, 1, bias=False),
99
+ nn.ELU())
100
+
101
+ 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)
102
+ self.bn2 = nn.BatchNorm2d(feat_out_channels_radar[1]//2, momentum=0.01, affine=True, eps=1.1e-5)
103
+ self.conv2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[1]//2, feat_out_channels_radar[1]//2, 3, 1, 1, bias=False),
104
+ nn.ELU())
105
+
106
+ 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)
107
+ self.bn1 = nn.BatchNorm2d(feat_out_channels_radar[0]//2, momentum=0.01, affine=True, eps=1.1e-5)
108
+ self.conv1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, feat_out_channels_radar[0]//2, 3, 1, 1, bias=False),
109
+ nn.ELU())
110
+
111
+ # self.get_depth = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, 1, 3, 1, 1, bias=False),
112
+ # nn.Sigmoid())
113
+
114
+ self.get_depth = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, 2, 3, 1, 1, bias=False),
115
+ nn.Sigmoid())
116
+
117
+ def forward(self, image_features, radar_features):
118
+ 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]
119
+ 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]
120
+ final = torch.cat([img_final, rad_final], axis=1)
121
+ upconv5 = self.upconv5(final)
122
+ upconv5 = self.bn5(upconv5)
123
+ upconv5 = self.conv5(upconv5)
124
+ upconv5 = torch.cat([img_skip3, rad_skip3, upconv5], axis=1)
125
+
126
+ upconv4 = self.upconv4(upconv5)
127
+ upconv4 = self.bn4(upconv4)
128
+ upconv4 = self.conv4(upconv4)
129
+ upconv4 = torch.cat([img_skip2, rad_skip2, upconv4], axis=1)
130
+
131
+ upconv3 = self.upconv3(upconv4)
132
+ upconv3 = self.bn3(upconv3)
133
+ upconv3 = self.conv3(upconv3)
134
+ upconv3 = torch.cat([img_skip1, rad_skip1, upconv3], axis=1)
135
+
136
+ upconv2 = self.upconv2(upconv3)
137
+ upconv2 = self.bn2(upconv2)
138
+ upconv2 = self.conv2(upconv2)
139
+ upconv2 = torch.cat([img_skip0, rad_skip0, upconv2], axis=1)
140
+
141
+ upconv1 = self.upconv1(upconv2)
142
+ upconv1 = self.bn1(upconv1)
143
+ upconv1 = self.conv1(upconv1)
144
+
145
+ # confidence = self.get_depth(upconv1)
146
+ # depth = self.params.max_depth * confidence
147
+ depth_conf = self.get_depth(upconv1)
148
+ depth = self.params.max_depth * depth_conf[:, 0:1]
149
+ confidence = depth_conf[:, 1:2]
150
+
151
+ return confidence, depth
152
+
153
+
154
+ class SparseConv(nn.Module):
155
+
156
+ def __init__(self,
157
+ in_channels,
158
+ out_channels,
159
+ kernel_size,
160
+ activation='relu'):
161
+ super().__init__()
162
+
163
+ padding = kernel_size//2
164
+
165
+ self.conv = nn.Conv2d(
166
+ in_channels,
167
+ out_channels,
168
+ kernel_size=kernel_size,
169
+ padding=padding,
170
+ bias=False)
171
+
172
+ self.bias = nn.Parameter(
173
+ torch.zeros(out_channels),
174
+ requires_grad=True)
175
+
176
+ self.sparsity = nn.Conv2d(
177
+ in_channels,
178
+ out_channels,
179
+ kernel_size=kernel_size,
180
+ padding=padding,
181
+ bias=False)
182
+
183
+ kernel = torch.FloatTensor(torch.ones([kernel_size, kernel_size])).unsqueeze(0).unsqueeze(0)
184
+
185
+ self.sparsity.weight = nn.Parameter(
186
+ data=kernel,
187
+ requires_grad=False)
188
+
189
+ if activation == 'relu':
190
+ self.act = nn.ReLU(inplace=False)
191
+ elif activation == 'sigmoid':
192
+ self.act = nn.Sigmoid()
193
+ elif activation == 'elu':
194
+ self.act = nn.ELU()
195
+
196
+ self.max_pool = nn.MaxPool2d(
197
+ kernel_size,
198
+ stride=1,
199
+ padding=padding)
200
+
201
+
202
+
203
+ def forward(self, x, mask):
204
+ x = x*mask
205
+ x = self.conv(x)
206
+ normalizer = 1/(self.sparsity(mask)+1e-8)
207
+ x = x * normalizer + self.bias.unsqueeze(0).unsqueeze(2).unsqueeze(3)
208
+ x = self.act(x)
209
+
210
+ mask = self.max_pool(mask)
211
+
212
+ return x, mask
src/Baselines/cafnet/rice_dataset.py ADDED
@@ -0,0 +1,121 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import os
3
+ from typing import Dict, List, Optional, Tuple
4
+
5
+ import numpy as np
6
+ from torch.utils.data import Dataset
7
+
8
+
9
+ class RiceDataset(Dataset):
10
+ """Raw Rice dataset reader for DJI RGB, ZED depth and radar point clouds.
11
+
12
+ This dataset returns raw per-frame arrays and leaves geometric processing to
13
+ `collate_fn_helpers.make_rice_collate_fn`.
14
+ """
15
+
16
+ def __init__(
17
+ self,
18
+ base_dir: str,
19
+ split_json_path: Optional[str] = None,
20
+ split: str = "train",
21
+ input_height: int = 288,
22
+ input_width: int = 512,
23
+ patch_size: Optional[Tuple[int, int]] = None,
24
+ ):
25
+ self.base_dir = base_dir
26
+ self.split = split
27
+ self.input_height = int(input_height)
28
+ self.input_width = int(input_width)
29
+ self.patch_size = self._resolve_patch_size(patch_size)
30
+
31
+ test_sequences = self._load_test_split(split_json_path)
32
+
33
+ all_sequences = sorted(
34
+ d
35
+ for d in os.listdir(base_dir)
36
+ if os.path.isdir(os.path.join(base_dir, d)) and not d.startswith(".")
37
+ )
38
+
39
+ self.sequences: List[str] = []
40
+ for seq in all_sequences:
41
+ if split == "train" and seq in test_sequences:
42
+ continue
43
+ if split == "test" and seq not in test_sequences:
44
+ continue
45
+ if self._is_valid_sequence(os.path.join(base_dir, seq)):
46
+ self.sequences.append(seq)
47
+
48
+ self.dji_rgb_mmaps: Dict[str, np.memmap] = {}
49
+ self.zed_depth_mmaps: Dict[str, np.memmap] = {}
50
+ self.samples: List[Tuple[str, int]] = []
51
+
52
+ for seq in self.sequences:
53
+ seq_dir = os.path.join(self.base_dir, seq)
54
+ dji_rgb_path = os.path.join(seq_dir, "dji_rgb.npy")
55
+ zed_depth_path = os.path.join(seq_dir, "zed_depth.npy")
56
+
57
+ self.dji_rgb_mmaps[seq] = np.load(dji_rgb_path, mmap_mode="r")
58
+ self.zed_depth_mmaps[seq] = np.load(zed_depth_path, mmap_mode="r")
59
+
60
+ n_frames = min(
61
+ len(self.dji_rgb_mmaps[seq]),
62
+ len(self.zed_depth_mmaps[seq]),
63
+ )
64
+ for frame_idx in range(n_frames):
65
+ self.samples.append((seq, frame_idx))
66
+
67
+ def _resolve_patch_size(
68
+ self, patch_size: Optional[Tuple[int, int]]
69
+ ) -> Tuple[int, int]:
70
+ if patch_size is not None:
71
+ return int(patch_size[0]), int(patch_size[1])
72
+
73
+ # Scale default CaFNet patch size (50, 150) from 352x704.
74
+ base_h, base_w = 352, 704
75
+ scale_h = self.input_height / float(base_h)
76
+ scale_w = self.input_width / float(base_w)
77
+ ext_h = max(1, int(round(50 * scale_h)))
78
+ ext_w = max(1, int(round(150 * scale_w)))
79
+ return ext_h, ext_w
80
+
81
+ def _load_test_split(self, split_json_path: Optional[str]) -> set:
82
+ if not split_json_path or not os.path.exists(split_json_path):
83
+ return set()
84
+ with open(split_json_path, "r") as f:
85
+ payload = json.load(f)
86
+ return set(payload.get("test", []))
87
+
88
+ def _is_valid_sequence(self, seq_dir: str) -> bool:
89
+ dji_rgb_path = os.path.join(seq_dir, "dji_rgb.npy")
90
+ zed_depth_path = os.path.join(seq_dir, "zed_depth.npy")
91
+ pcd_dir = os.path.join(seq_dir, "pcd")
92
+ return (
93
+ os.path.exists(dji_rgb_path)
94
+ and os.path.exists(zed_depth_path)
95
+ and os.path.isdir(pcd_dir)
96
+ )
97
+
98
+ def __len__(self) -> int:
99
+ return len(self.samples)
100
+
101
+ def __getitem__(self, idx: int) -> Dict[str, object]:
102
+ seq, frame_idx = self.samples[idx]
103
+ seq_dir = os.path.join(self.base_dir, seq)
104
+
105
+ dji_rgb = np.asarray(self.dji_rgb_mmaps[seq][frame_idx]).copy()
106
+ zed_depth_mm = np.asarray(self.zed_depth_mmaps[seq][frame_idx]).copy()
107
+
108
+ pcd_path = os.path.join(seq_dir, "pcd", f"pcd_{frame_idx}.npy")
109
+ if os.path.exists(pcd_path):
110
+ radar_pcd_xyz = np.asarray(np.load(pcd_path), dtype=np.float32)
111
+ else:
112
+ radar_pcd_xyz = np.zeros((0, 3), dtype=np.float32)
113
+
114
+ return {
115
+ "sample_idx": idx,
116
+ "sequence": seq,
117
+ "frame_idx": frame_idx,
118
+ "dji_rgb": dji_rgb,
119
+ "zed_depth_mm": zed_depth_mm,
120
+ "radar_pcd_xyz": radar_pcd_xyz,
121
+ }
src/Baselines/cafnet/split.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "test": [
3
+ "Dell-1",
4
+ "Dell-2",
5
+ "Smoke-Dell-1",
6
+ "Smoke-Dell-2",
7
+ "Keck-1",
8
+ "Keck-2",
9
+ "Keck-3",
10
+ "Smoke-keck-1",
11
+ "Smoke-keck-2",
12
+ "Smoke-keck-3"
13
+ ]
14
+ }
src/Baselines/cafnet_no_smoke/collate_fn_helpers.py ADDED
@@ -0,0 +1,404 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import numpy as np
3
+ import torch
4
+ from functools import lru_cache
5
+ from typing import Callable, Dict, Sequence, Tuple, Union
6
+ from torchvision import transforms as T
7
+
8
+
9
+ IMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)
10
+ IMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)
11
+
12
+ # ZED intrinsics at 1280x720 reference resolution.
13
+ _K_ZED_REF = np.array(
14
+ [
15
+ [521.581604, 0.0, 636.33398438],
16
+ [0.0, 521.581604, 373.10964966],
17
+ [0.0, 0.0, 1.0],
18
+ ],
19
+ dtype=np.float64,
20
+ )
21
+ _ZED_REF_W = 1280
22
+ _ZED_REF_H = 720
23
+
24
+ # DJI calibration constants.
25
+ _CALIB_K_DJI = np.array(
26
+ [
27
+ [718.48555551, 0.0, 963.36465011],
28
+ [0.0, 720.25844189, 537.87569913],
29
+ [0.0, 0.0, 1.0],
30
+ ],
31
+ dtype=np.float64,
32
+ )
33
+ _CALIB_D_DJI = np.array(
34
+ [0.19022699, 0.03466753, 0.05858962, -0.07070669], dtype=np.float64
35
+ )
36
+ _CALIB_DEFISH_SHAPE = (1920, 1080)
37
+ _CALIB_DEFISH_BALANCE = 0.2
38
+ _CALIB_H_FULL = np.array(
39
+ [
40
+ [0.8274446551892256, -0.0742944198979625, 80.23797348979947],
41
+ [-0.014725864916652691, 0.8471179917075127, 28.27366063997317],
42
+ [-5.083573451500717e-05, -6.846079418201229e-05, 1.0],
43
+ ],
44
+ dtype=np.float64,
45
+ )
46
+ _CALIB_OUT_SIZE = (1918, 1105)
47
+ _CALIB_CROP = (115, 255, 1400, 760) # top, left, right, bottom
48
+
49
+
50
+ @lru_cache(maxsize=1)
51
+ def _get_dji_defish_maps() -> Tuple[np.ndarray, np.ndarray]:
52
+ r_defish = np.eye(3)
53
+ k_new_defish = cv2.fisheye.estimateNewCameraMatrixForUndistortRectify(
54
+ _CALIB_K_DJI,
55
+ _CALIB_D_DJI,
56
+ _CALIB_DEFISH_SHAPE,
57
+ r_defish,
58
+ balance=_CALIB_DEFISH_BALANCE,
59
+ fov_scale=1.0,
60
+ )
61
+ map1, map2 = cv2.fisheye.initUndistortRectifyMap(
62
+ _CALIB_K_DJI,
63
+ _CALIB_D_DJI,
64
+ r_defish,
65
+ k_new_defish,
66
+ _CALIB_DEFISH_SHAPE,
67
+ cv2.CV_16SC2,
68
+ )
69
+ return map1, map2
70
+
71
+
72
+ def resize_depth_mm(depth_mm: np.ndarray, target_size: Tuple[int, int]) -> np.ndarray:
73
+ target_h, target_w = target_size
74
+ if depth_mm.shape[:2] == (target_h, target_w):
75
+ return depth_mm
76
+ return cv2.resize(depth_mm, (target_w, target_h), interpolation=cv2.INTER_NEAREST)
77
+
78
+
79
+ def depth_collator(
80
+ depth: Union[torch.Tensor, np.ndarray],
81
+ max_depth_m: float = 11.2,
82
+ target_size: Tuple[int, int] = (128, 256),
83
+ ) -> Union[torch.Tensor, np.ndarray]:
84
+ """Clamp, normalize to [0, 1], and resize depth."""
85
+ is_numpy = isinstance(depth, np.ndarray)
86
+ if is_numpy:
87
+ depth = torch.from_numpy(depth)
88
+
89
+ depth = depth.float()
90
+ original_shape = depth.shape
91
+
92
+ if depth.dim() == 2:
93
+ depth = depth.unsqueeze(0)
94
+ elif depth.dim() == 3:
95
+ depth = depth.unsqueeze(1)
96
+
97
+ invalid_mask = ~(torch.isfinite(depth) & (depth >= 0))
98
+ depth[invalid_mask] = 0.0
99
+
100
+ depth = torch.clamp(depth, min=0.0, max=max_depth_m)
101
+ depth = depth / max_depth_m
102
+
103
+ invalid_mask = ~torch.isfinite(depth)
104
+ depth[invalid_mask] = 0.0
105
+
106
+ resized = T.Resize(
107
+ target_size, interpolation=T.InterpolationMode.BILINEAR, antialias=True
108
+ )(depth)
109
+
110
+ if len(original_shape) == 2:
111
+ resized = resized.squeeze(0)
112
+
113
+ return resized.numpy() if is_numpy else resized
114
+
115
+
116
+ def dji_rgb_collator(
117
+ image: torch.Tensor,
118
+ target_size: Tuple[int, int] = (128, 256),
119
+ ) -> torch.Tensor:
120
+ """Rectify and resize DJI RGB image batch.
121
+
122
+ Args:
123
+ image: Tensor with shape (B, C, H, W).
124
+ target_size: Target resolution as (height, width).
125
+
126
+ Returns:
127
+ Tensor in CHW format (B, C, H, W), float32 in [0, 1].
128
+ """
129
+ if not isinstance(image, torch.Tensor):
130
+ raise ValueError(f"Expected torch.Tensor, got {type(image)}")
131
+
132
+ if image.dim() != 4:
133
+ raise ValueError(
134
+ f"Expected 4D tensor (B, C, H, W), got {image.dim()}D tensor with shape {image.shape}"
135
+ )
136
+
137
+ map1_defish, map2_defish = _get_dji_defish_maps()
138
+ target_h, target_w = target_size
139
+
140
+ if image.max() <= 1.0:
141
+ img_batch = (image.permute(0, 2, 3, 1).cpu().numpy() * 255.0).astype(np.uint8)
142
+ else:
143
+ img_batch = image.permute(0, 2, 3, 1).cpu().numpy().astype(np.uint8)
144
+
145
+ calibrated_images = []
146
+ for img in img_batch:
147
+ if img.shape[1] != 1920 or img.shape[0] != 1080:
148
+ img = cv2.resize(img, (1920, 1080), interpolation=cv2.INTER_LINEAR)
149
+
150
+ img = cv2.remap(img, map1_defish, map2_defish, interpolation=cv2.INTER_LINEAR)
151
+ img = cv2.warpPerspective(
152
+ img, _CALIB_H_FULL, _CALIB_OUT_SIZE, flags=cv2.INTER_LINEAR
153
+ )
154
+
155
+ top, left, right, bottom = _CALIB_CROP
156
+ img = img[top:bottom, left:right]
157
+ img = cv2.resize(img, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
158
+ calibrated_images.append(img)
159
+
160
+ out_batch = np.stack(calibrated_images, axis=0)
161
+ out_tensor = torch.from_numpy(out_batch).permute(0, 3, 1, 2).float() / 255.0
162
+ return out_tensor
163
+
164
+
165
+ def point_cloud_to_sparse_depth(
166
+ points_xyz: np.ndarray,
167
+ target_shape: Tuple[int, int],
168
+ max_depth_m: float,
169
+ ) -> np.ndarray:
170
+ """Project xyz radar points (meters) to a sparse depth image."""
171
+ target_h, target_w = target_shape
172
+ sparse_depth = np.zeros((target_h, target_w), dtype=np.float32)
173
+
174
+ if points_xyz.size == 0:
175
+ return sparse_depth
176
+
177
+ pts = np.asarray(points_xyz, dtype=np.float32)
178
+ if pts.ndim != 2 or pts.shape[1] != 3:
179
+ return sparse_depth
180
+
181
+ valid = np.isfinite(pts).all(axis=1)
182
+ valid &= pts[:, 2] > 0.0
183
+ valid &= pts[:, 2] <= float(max_depth_m)
184
+ pts = pts[valid]
185
+ if pts.shape[0] == 0:
186
+ return sparse_depth
187
+
188
+ sx = target_w / float(_ZED_REF_W)
189
+ sy = target_h / float(_ZED_REF_H)
190
+ fx = _K_ZED_REF[0, 0] * sx
191
+ fy = _K_ZED_REF[1, 1] * sy
192
+ cx = _K_ZED_REF[0, 2] * sx
193
+ cy = _K_ZED_REF[1, 2] * sy
194
+
195
+ z = pts[:, 2]
196
+ u = np.rint(pts[:, 0] * fx / z + cx).astype(np.int32)
197
+ v = np.rint(pts[:, 1] * fy / z + cy).astype(np.int32)
198
+
199
+ in_bounds = (u >= 0) & (u < target_w) & (v >= 0) & (v < target_h)
200
+ if not np.any(in_bounds):
201
+ return sparse_depth
202
+
203
+ u = u[in_bounds]
204
+ v = v[in_bounds]
205
+ z = z[in_bounds].astype(np.float32)
206
+
207
+ min_depth = np.full((target_h, target_w), np.inf, dtype=np.float32)
208
+ np.minimum.at(min_depth, (v, u), z)
209
+ min_depth[~np.isfinite(min_depth)] = 0.0
210
+ return min_depth
211
+
212
+
213
+ def build_radar_gt_map(
214
+ depth_m: np.ndarray,
215
+ sparse_depth: np.ndarray,
216
+ patch_size: Tuple[int, int],
217
+ max_dist_correspondence: float,
218
+ ) -> np.ndarray:
219
+ """Build confidence GT using local depth consistency around each radar pixel."""
220
+ h, w = depth_m.shape
221
+ radar_gt = np.zeros((h, w), dtype=np.float32)
222
+
223
+ ys, xs = np.where(sparse_depth > 0)
224
+ if len(ys) == 0:
225
+ return radar_gt
226
+
227
+ ext_h, ext_w = int(patch_size[0]), int(patch_size[1])
228
+ for y, x in zip(ys, xs):
229
+ radar_depth = sparse_depth[y, x]
230
+
231
+ delta_x1 = min(x, ext_w)
232
+ delta_y1 = min(y, ext_h)
233
+ delta_x2 = min(w - x, ext_w)
234
+ delta_y2 = min(h - y, ext_h)
235
+
236
+ x1 = x - delta_x1
237
+ y1 = y - delta_y1
238
+ x2 = x + delta_x2
239
+ y2 = y + delta_y2
240
+
241
+ distance = np.abs(depth_m[y1:y2, x1:x2] - radar_depth)
242
+ gt_label = (distance < float(max_dist_correspondence)).astype(np.float32)
243
+ radar_gt[y1:y2, x1:x2] = gt_label
244
+
245
+ return radar_gt
246
+
247
+
248
+ def make_rice_collate_fn(
249
+ input_height: int,
250
+ input_width: int,
251
+ radar_max_depth_m: float,
252
+ max_dist_correspondence: float,
253
+ patch_size: Tuple[int, int],
254
+ ) -> Callable[[Sequence[Dict[str, object]]], Tuple[torch.Tensor, ...]]:
255
+ """Create collate_fn for RiceDataset samples.
256
+
257
+ Each dataset sample should contain:
258
+ - sample_idx: int
259
+ - dji_rgb: (H, W, 3) uint8
260
+ - zed_depth_mm: (H, W) uint16
261
+ - radar_pcd_xyz: (N, 3) float32 in meters
262
+ """
263
+
264
+ mean = torch.tensor(IMAGENET_MEAN, dtype=torch.float32).view(1, 3, 1, 1)
265
+ std = torch.tensor(IMAGENET_STD, dtype=torch.float32).view(1, 3, 1, 1)
266
+
267
+ def _collate(batch: Sequence[Dict[str, object]]) -> Tuple[torch.Tensor, ...]:
268
+ if len(batch) == 0:
269
+ raise ValueError("Received empty batch in collate function")
270
+
271
+ sample_indices = []
272
+ rgb_batch = []
273
+ depth_batch = []
274
+ radar_batch = []
275
+ radar_gt_batch = []
276
+
277
+ for sample in batch:
278
+ sample_indices.append(int(sample["sample_idx"]))
279
+
280
+ rgb = np.asarray(sample["dji_rgb"]).copy()
281
+ if rgb.ndim != 3 or rgb.shape[2] != 3:
282
+ raise ValueError(f"Expected RGB shape (H, W, 3), got {rgb.shape}")
283
+ rgb_batch.append(torch.from_numpy(np.transpose(rgb, (2, 0, 1))))
284
+
285
+ depth_mm = np.asarray(sample["zed_depth_mm"]).copy()
286
+ depth_mm = resize_depth_mm(depth_mm, (input_height, input_width))
287
+ depth_m = depth_mm.astype(np.float32) / 1000.0
288
+ invalid = ~(np.isfinite(depth_m) & (depth_m > 0.0))
289
+ depth_m[invalid] = 0.0
290
+ depth_batch.append(depth_m)
291
+
292
+ radar_points = np.asarray(sample["radar_pcd_xyz"], dtype=np.float32)
293
+ if radar_points.ndim != 2 or radar_points.shape[1] != 3:
294
+ radar_points = np.zeros((0, 3), dtype=np.float32)
295
+
296
+ if radar_points.shape[0] == 0:
297
+ center_v = float(depth_m[input_height // 2, input_width // 2])
298
+ if not np.isfinite(center_v):
299
+ center_v = 0.0
300
+ radar_points = np.array([[0.0, 0.0, center_v]], dtype=np.float32)
301
+
302
+ sparse_depth = point_cloud_to_sparse_depth(
303
+ radar_points,
304
+ target_shape=(input_height, input_width),
305
+ max_depth_m=radar_max_depth_m,
306
+ )
307
+ radar_gt = build_radar_gt_map(
308
+ depth_m,
309
+ sparse_depth,
310
+ patch_size=patch_size,
311
+ max_dist_correspondence=max_dist_correspondence,
312
+ )
313
+ radar_batch.append(sparse_depth)
314
+ radar_gt_batch.append(radar_gt)
315
+
316
+ rgb_tensor = torch.stack(rgb_batch, dim=0).float()
317
+ rgb_tensor = dji_rgb_collator(rgb_tensor, target_size=(input_height, input_width))
318
+ rgb_tensor = (rgb_tensor - mean) / std
319
+
320
+ depth_tensor = torch.from_numpy(np.stack(depth_batch, axis=0)).float().unsqueeze(1)
321
+ radar_tensor = torch.from_numpy(np.stack(radar_batch, axis=0)).float().unsqueeze(1)
322
+ radar_gt_tensor = (
323
+ torch.from_numpy(np.stack(radar_gt_batch, axis=0)).float().unsqueeze(1)
324
+ )
325
+ idx_tensor = torch.tensor(sample_indices, dtype=torch.long)
326
+
327
+ return idx_tensor, rgb_tensor, depth_tensor, radar_tensor, radar_gt_tensor
328
+
329
+ return _collate
330
+
331
+
332
+ # Fisheye RGB Handler Functions ##
333
+ def fisheye_rgb_collator(
334
+ image: torch.Tensor,
335
+ target_size: Tuple[int, int] = (128, 256),
336
+ ) -> torch.Tensor:
337
+ """Calibrate and resize Fisheye RGB image batch.
338
+
339
+ Args:
340
+ image: Batch of Fisheye RGB images as torch tensor (B, C, H, W) in CHW format
341
+ target_size: Target resolution as (height, width)
342
+
343
+ Returns:
344
+ Batch of calibrated and resized torch tensors in CHW format (B, C, H, W)
345
+ """
346
+ IMAGE_WIDTH = 1920
347
+ IMAGE_HEIGHT = 1080
348
+ FOCAL_LENGTH_X = 0.613260
349
+ FOCAL_LENGTH_Y = 0.613260
350
+ CENTER_X = 0.5
351
+ CENTER_Y = 0.5
352
+ K1 = -0.120000
353
+ K2 = -0.015000
354
+
355
+ w, h = IMAGE_WIDTH, IMAGE_HEIGHT
356
+ x_out, y_out = np.meshgrid(np.arange(w), np.arange(h))
357
+ x_norm = (x_out - w * CENTER_X) / (w * FOCAL_LENGTH_X)
358
+ y_norm = (y_out - h * CENTER_Y) / (h * FOCAL_LENGTH_Y)
359
+ r = np.sqrt(x_norm**2 + y_norm**2)
360
+ r_distorted = r + K1 * r**2 + K2 * r**3
361
+ r_safe = np.where(r > 0, r, 1.0)
362
+ scale = np.where(r > 0, r_distorted / r_safe, 1.0)
363
+ x_norm_distorted = x_norm * scale
364
+ y_norm_distorted = y_norm * scale
365
+ map_x = (x_norm_distorted * (w * FOCAL_LENGTH_X) + w * CENTER_X).astype(np.float32)
366
+ map_y = (y_norm_distorted * (h * FOCAL_LENGTH_Y) + h * CENTER_Y).astype(np.float32)
367
+
368
+ if not isinstance(image, torch.Tensor):
369
+ raise ValueError(f"Expected torch.Tensor, got {type(image)}")
370
+
371
+ if image.dim() != 4:
372
+ raise ValueError(
373
+ f"Expected 4D tensor (B, C, H, W), got {image.dim()}D tensor with shape {image.shape}"
374
+ )
375
+
376
+ if image.max() <= 1.0:
377
+ img_batch = (image.permute(0, 2, 3, 1).cpu().numpy() * 255).astype(np.uint8)
378
+ else:
379
+ img_batch = image.permute(0, 2, 3, 1).cpu().numpy().astype(np.uint8)
380
+
381
+ calibrated_images = []
382
+ target_h, target_w = target_size
383
+
384
+ for img in img_batch:
385
+ if img.shape[1] != IMAGE_WIDTH or img.shape[0] != IMAGE_HEIGHT:
386
+ img = cv2.resize(
387
+ img, (IMAGE_WIDTH, IMAGE_HEIGHT), interpolation=cv2.INTER_LINEAR
388
+ )
389
+
390
+ img = cv2.remap(
391
+ img,
392
+ map_x,
393
+ map_y,
394
+ interpolation=cv2.INTER_LINEAR,
395
+ borderMode=cv2.BORDER_CONSTANT,
396
+ borderValue=(0, 0, 0),
397
+ )
398
+
399
+ img = cv2.resize(img, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
400
+ calibrated_images.append(img)
401
+
402
+ out_batch = np.stack(calibrated_images, axis=0)
403
+ out_tensor = torch.from_numpy(out_batch).permute(0, 3, 1, 2).float() / 255.0
404
+ return out_tensor
src/Baselines/cafnet_no_smoke/dataloader.py ADDED
@@ -0,0 +1,100 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Optional
2
+
3
+ from torch.utils.data import DataLoader
4
+
5
+ from collate_fn_helpers import make_rice_collate_fn
6
+ from rice_dataset import RiceDataset
7
+
8
+
9
+ def _build_dataset(
10
+ args,
11
+ split: str,
12
+ base_dir: Optional[str] = None,
13
+ split_json_path: Optional[str] = None,
14
+ ) -> RiceDataset:
15
+ return RiceDataset(
16
+ base_dir=base_dir or args.base_dir,
17
+ split_json_path=args.split_json if split_json_path is None else split_json_path,
18
+ split=split,
19
+ input_height=args.input_height,
20
+ input_width=args.input_width,
21
+ patch_size=args.patch_size,
22
+ )
23
+
24
+
25
+ def _build_loader(
26
+ args,
27
+ split: str,
28
+ batch_size: int,
29
+ shuffle: bool,
30
+ drop_last: bool,
31
+ pin_memory: bool,
32
+ base_dir: Optional[str] = None,
33
+ split_json_path: Optional[str] = None,
34
+ ):
35
+ dataset = _build_dataset(
36
+ args,
37
+ split=split,
38
+ base_dir=base_dir,
39
+ split_json_path=split_json_path,
40
+ )
41
+ collate_fn = make_rice_collate_fn(
42
+ input_height=args.input_height,
43
+ input_width=args.input_width,
44
+ radar_max_depth_m=args.radar_max_depth_m,
45
+ max_dist_correspondence=args.max_dist_correspondence,
46
+ patch_size=dataset.patch_size,
47
+ )
48
+ return DataLoader(
49
+ dataset,
50
+ batch_size=batch_size,
51
+ shuffle=shuffle,
52
+ num_workers=args.num_workers,
53
+ pin_memory=pin_memory,
54
+ drop_last=drop_last,
55
+ collate_fn=collate_fn,
56
+ )
57
+
58
+
59
+ def create_train_test_loaders(args, pin_memory: bool = False):
60
+ train_loader = _build_loader(
61
+ args,
62
+ split="train",
63
+ batch_size=args.batch_size,
64
+ shuffle=True,
65
+ drop_last=True,
66
+ pin_memory=pin_memory,
67
+ )
68
+ test_loader = _build_loader(
69
+ args,
70
+ split="test",
71
+ batch_size=args.batch_size,
72
+ shuffle=False,
73
+ drop_last=False,
74
+ pin_memory=pin_memory,
75
+ )
76
+ return train_loader, test_loader
77
+
78
+
79
+ def create_inference_loader(args, pin_memory: bool = False):
80
+ """Create the single packaged Smoke-Eval loader used for inference."""
81
+
82
+ test_base_dir = getattr(args, "test_base_dir", "")
83
+ if not test_base_dir:
84
+ raise ValueError("Config must define 'test_base_dir' for inference.")
85
+
86
+ test_split = getattr(args, "test_split", "train")
87
+ test_split_json = getattr(args, "test_split_json", None)
88
+ if not test_split_json:
89
+ test_split_json = None
90
+
91
+ return _build_loader(
92
+ args,
93
+ split=test_split,
94
+ batch_size=args.batch_size,
95
+ shuffle=False,
96
+ drop_last=False,
97
+ pin_memory=pin_memory,
98
+ base_dir=test_base_dir,
99
+ split_json_path=test_split_json,
100
+ )
src/Baselines/cafnet_no_smoke/extract_pcd_from_depth.py ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import numpy as np
3
+
4
+ # ZED intrinsics at reference resolution 1280x720 (same values as PointCloudConverter)
5
+ _K_ZED_REF = np.array(
6
+ [
7
+ [521.581604, 0.0, 636.33398438],
8
+ [0.0, 521.581604, 373.10964966],
9
+ [0.0, 0.0, 1.0],
10
+ ],
11
+ dtype=np.float64,
12
+ )
13
+ _ZED_REF_W = 1280
14
+ _ZED_REF_H = 720
15
+
16
+
17
+ def sample_depth_as_radar(
18
+ depth_mm: np.ndarray,
19
+ n_samples: int = 100,
20
+ target_shape: tuple = (300, 1280),
21
+ max_depth_m: float = 11.2,
22
+ seed: int | None = None,
23
+ ) -> tuple:
24
+ """
25
+ Randomly sample points from a ground truth ZED depth map and treat them as
26
+ radar points, mimicking the sparse depth input the model expects.
27
+
28
+ The input depth is resized from its native resolution (e.g. 896x504) to
29
+ target_shape using nearest-neighbor interpolation so raw mm values are
30
+ preserved. Camera intrinsics are scaled from the 1280x720 ZED reference to
31
+ match the target resolution.
32
+
33
+ Args:
34
+ depth_mm: Ground truth depth map, shape (H, W), dtype uint16, in mm.
35
+ n_samples: Number of points to randomly sample (default: 100).
36
+ target_shape: (target_H, target_W) to resize to before sampling.
37
+ Default (300, 1280) matches the model's required input.
38
+ max_depth_m: Maximum valid depth in meters — pixels beyond this are
39
+ treated as invalid (default: 11.2 m).
40
+ seed: Optional random seed for reproducibility.
41
+
42
+ Returns:
43
+ points (np.ndarray): (N, 3) float32 array of [X, Y, Z] in meters,
44
+ in camera coordinate frame. N <= n_samples.
45
+ sparse_depth (np.ndarray): (target_H, target_W) float32 sparse depth map
46
+ with only the N sampled pixels filled (meters),
47
+ zeros elsewhere.
48
+ """
49
+ target_h, target_w = target_shape
50
+
51
+ # --- 1. Resize depth map (nearest-neighbor preserves raw mm values) ---
52
+ in_h, in_w = depth_mm.shape
53
+ if (in_h, in_w) != (target_h, target_w):
54
+ depth_resized = cv2.resize(
55
+ depth_mm, (target_w, target_h), interpolation=cv2.INTER_NEAREST
56
+ )
57
+ else:
58
+ depth_resized = depth_mm.copy()
59
+
60
+ # --- 2. Scale intrinsics from 1280x720 reference to target resolution ---
61
+ sx = target_w / float(_ZED_REF_W)
62
+ sy = target_h / float(_ZED_REF_H)
63
+ fx = _K_ZED_REF[0, 0] * sx
64
+ fy = _K_ZED_REF[1, 1] * sy
65
+ cx = _K_ZED_REF[0, 2] * sx
66
+ cy = _K_ZED_REF[1, 2] * sy
67
+
68
+ # --- 3. Convert to float meters and find valid pixels ---
69
+ depth_m = depth_resized.astype(np.float32) / 1000.0
70
+ valid_mask = (depth_m > 0) & (depth_m <= max_depth_m)
71
+ valid_v, valid_u = np.where(valid_mask) # row (V), col (U)
72
+
73
+ if len(valid_v) == 0:
74
+ return (
75
+ np.zeros((0, 3), dtype=np.float32),
76
+ np.zeros((target_h, target_w), dtype=np.float32),
77
+ )
78
+
79
+ # --- 4. Randomly sample up to n_samples valid pixels ---
80
+ rng = np.random.default_rng(seed)
81
+ n = min(n_samples, len(valid_v))
82
+ indices = rng.choice(len(valid_v), size=n, replace=False)
83
+ sampled_v = valid_v[indices]
84
+ sampled_u = valid_u[indices]
85
+ sampled_z = depth_m[sampled_v, sampled_u]
86
+
87
+ # --- 5. Back-project to 3D camera coordinates (pinhole model) ---
88
+ X = (sampled_u - cx) * sampled_z / fx
89
+ Y = (sampled_v - cy) * sampled_z / fy
90
+ points = np.stack([X, Y, sampled_z], axis=1).astype(np.float32) # (N, 3)
91
+
92
+ # --- 6. Build sparse depth map ---
93
+ sparse_depth = np.zeros((target_h, target_w), dtype=np.float32)
94
+ sparse_depth[sampled_v, sampled_u] = sampled_z
95
+
96
+ return points, sparse_depth
src/Baselines/cafnet_no_smoke/inference.py ADDED
@@ -0,0 +1,224 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import argparse
2
+ import os
3
+ from typing import Dict, List
4
+
5
+ import numpy as np
6
+ import torch
7
+ import torch.distributed as dist
8
+ import yaml
9
+ from accelerate import Accelerator
10
+ from accelerate.utils import DistributedDataParallelKwargs, set_seed
11
+ from safetensors.torch import load_file
12
+ from tqdm.auto import tqdm
13
+
14
+ from dataloader import create_inference_loader
15
+ from models.model import CaFNet
16
+
17
+
18
+ DEFAULT_CONFIG = {
19
+ # Packaged evaluation dataset.
20
+ "base_dir": "",
21
+ "split_json": None,
22
+ "test_base_dir": None,
23
+ "test_split": "train",
24
+ "test_split_json": None,
25
+ # Input and radar processing
26
+ "input_height": 288,
27
+ "input_width": 512,
28
+ "radar_max_depth_m": 11.2,
29
+ "max_dist_correspondence": 0.5,
30
+ "patch_size": None,
31
+ # Model
32
+ "encoder": "resnet34_bts",
33
+ "encoder_radar": "resnet18",
34
+ "radar_input_channels": 1,
35
+ "bts_size": 512,
36
+ "max_depth": 11.2,
37
+ # Runtime
38
+ "batch_size": 8,
39
+ # Windows uses spawn-based multiprocessing; keep the public evaluation
40
+ # entry point portable and deterministic by default.
41
+ "num_workers": 0,
42
+ "seed": 42,
43
+ "cpu": False,
44
+ "mixed_precision": "fp16",
45
+ "checkpoint_path": "checkpoints/cafnet_no_smoke.safetensors",
46
+ "prediction_dir": "prediction",
47
+ }
48
+
49
+
50
+ def parse_args():
51
+ parser = argparse.ArgumentParser(description="Run CaFNet inference on Smoke-Eval.")
52
+ parser.add_argument("--config", type=str, required=True, help="Path to YAML config")
53
+ return parser.parse_args()
54
+
55
+
56
+ def load_config(path):
57
+ with open(path, "r") as f:
58
+ cfg = yaml.safe_load(f) or {}
59
+ if not isinstance(cfg, dict):
60
+ raise ValueError("Config must be a YAML mapping (key-value pairs).")
61
+
62
+ merged = dict(DEFAULT_CONFIG)
63
+ merged.update(cfg)
64
+
65
+ if not merged["test_base_dir"]:
66
+ raise ValueError("Config must define 'test_base_dir'.")
67
+ if not merged["checkpoint_path"]:
68
+ raise ValueError("Config must define 'checkpoint_path'.")
69
+ if not os.path.isfile(merged["checkpoint_path"]):
70
+ raise FileNotFoundError(f"Checkpoint not found: {merged['checkpoint_path']}")
71
+ if merged.get("radar_input_channels", 1) != 1:
72
+ raise ValueError("radar_input_channels must be 1 for this setup.")
73
+
74
+ return argparse.Namespace(**merged)
75
+
76
+
77
+ def build_model_args(args):
78
+ return argparse.Namespace(
79
+ encoder=args.encoder,
80
+ encoder_radar=args.encoder_radar,
81
+ radar_input_channels=args.radar_input_channels,
82
+ input_height=args.input_height,
83
+ input_width=args.input_width,
84
+ max_depth=args.max_depth,
85
+ bts_size=args.bts_size,
86
+ )
87
+
88
+
89
+ def _extract_model_state(checkpoint):
90
+ if isinstance(checkpoint, dict) and isinstance(checkpoint.get("model"), dict):
91
+ return checkpoint["model"]
92
+ if isinstance(checkpoint, dict):
93
+ return checkpoint
94
+ raise ValueError("Unsupported checkpoint format.")
95
+
96
+
97
+ def _gather_objects(accelerator, obj):
98
+ if accelerator.num_processes == 1:
99
+ return [obj]
100
+ if not dist.is_available() or not dist.is_initialized():
101
+ return [obj]
102
+
103
+ gathered = [None for _ in range(accelerator.num_processes)]
104
+ dist.all_gather_object(gathered, obj)
105
+ return gathered
106
+
107
+
108
+ def _merge_predictions(all_rank_predictions):
109
+ merged: Dict[str, Dict[int, np.ndarray]] = {}
110
+ for rank_dict in all_rank_predictions:
111
+ if not rank_dict:
112
+ continue
113
+ for seq_name, frame_map in rank_dict.items():
114
+ seq_slot = merged.setdefault(seq_name, {})
115
+ for frame_idx, pred in frame_map.items():
116
+ frame_idx = int(frame_idx)
117
+ if frame_idx not in seq_slot:
118
+ seq_slot[frame_idx] = pred
119
+ return merged
120
+
121
+
122
+ def _save_sequence_predictions(predictions, out_dir):
123
+ os.makedirs(out_dir, exist_ok=True)
124
+ for seq_name in sorted(predictions.keys()):
125
+ frame_map = predictions[seq_name]
126
+ ordered_frames = sorted(frame_map.keys())
127
+ if not ordered_frames:
128
+ pred_stack = np.zeros((0,), dtype=np.float32)
129
+ else:
130
+ pred_stack = np.stack([frame_map[k] for k in ordered_frames], axis=0).astype(
131
+ np.float32,
132
+ copy=False,
133
+ )
134
+ np.save(os.path.join(out_dir, f"{seq_name.lower()}_pred.npy"), pred_stack)
135
+
136
+
137
+ def _run_loader_inference(accelerator, model, loader, samples, save_dir, desc):
138
+ model.eval()
139
+ local_preds: Dict[str, Dict[int, np.ndarray]] = {}
140
+
141
+ with torch.no_grad():
142
+ pbar = tqdm(
143
+ loader,
144
+ desc=desc,
145
+ disable=not accelerator.is_local_main_process,
146
+ dynamic_ncols=True,
147
+ leave=False,
148
+ )
149
+ for batch in pbar:
150
+ sample_idx, image, depth_gt, radar, radar_gt = batch
151
+
152
+ image = image.to(accelerator.device, non_blocking=True)
153
+ radar = radar.to(accelerator.device, non_blocking=True)
154
+ # Kept for parity with validation loop structure.
155
+ _ = depth_gt.to(accelerator.device, non_blocking=True)
156
+ _ = radar_gt.to(accelerator.device, non_blocking=True)
157
+
158
+ focal = torch.ones((image.size(0),), device=image.device)
159
+ _, _, _, _, depth_est, _, _ = model(image, radar, focal)
160
+
161
+ pred_np = depth_est.detach().float().cpu().numpy()
162
+ if pred_np.ndim == 4 and pred_np.shape[1] == 1:
163
+ pred_np = pred_np[:, 0]
164
+
165
+ if torch.is_tensor(sample_idx):
166
+ sample_idx_list = sample_idx.detach().cpu().tolist()
167
+ else:
168
+ sample_idx_list = list(sample_idx)
169
+
170
+ for local_i, sample_i in enumerate(sample_idx_list):
171
+ seq_name, frame_idx = samples[int(sample_i)]
172
+ seq_slot = local_preds.setdefault(seq_name, {})
173
+ frame_idx = int(frame_idx)
174
+ if frame_idx not in seq_slot:
175
+ seq_slot[frame_idx] = pred_np[local_i].astype(np.float32, copy=False)
176
+
177
+ gathered = _gather_objects(accelerator, local_preds)
178
+ if accelerator.is_main_process:
179
+ merged = _merge_predictions(gathered)
180
+ _save_sequence_predictions(merged, save_dir)
181
+
182
+ accelerator.wait_for_everyone()
183
+
184
+
185
+ def main():
186
+ cli = parse_args()
187
+ args = load_config(cli.config)
188
+
189
+ set_seed(args.seed)
190
+ ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=True)
191
+ accelerator = Accelerator(
192
+ mixed_precision=None if args.mixed_precision in ("no", "none") else args.mixed_precision,
193
+ cpu=args.cpu,
194
+ kwargs_handlers=[ddp_kwargs],
195
+ )
196
+
197
+ test_loader = create_inference_loader(
198
+ args,
199
+ pin_memory=(accelerator.device.type == "cuda"),
200
+ )
201
+ test_samples: List = test_loader.dataset.samples
202
+
203
+ model = CaFNet(build_model_args(args))
204
+
205
+ model, test_loader = accelerator.prepare(model, test_loader)
206
+
207
+ state_dict = load_file(args.checkpoint_path, device="cpu")
208
+ accelerator.unwrap_model(model).load_state_dict(state_dict, strict=True)
209
+
210
+ _run_loader_inference(
211
+ accelerator=accelerator,
212
+ model=model,
213
+ loader=test_loader,
214
+ samples=test_samples,
215
+ save_dir=args.prediction_dir,
216
+ desc="Inference",
217
+ )
218
+
219
+ if accelerator.is_main_process:
220
+ print(f"Saved predictions to: {args.prediction_dir}")
221
+
222
+
223
+ if __name__ == "__main__":
224
+ main()
src/Baselines/cafnet_no_smoke/inference_config.yaml ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # CaFNet-no-smoke inference config for the packaged Smoke-Eval data.
2
+ test_base_dir: "../../../evaluation_dataset/Smoke-Eval"
3
+ test_split: "train"
4
+ test_split_json: null
5
+
6
+ # Input and radar preprocessing
7
+ input_height: 288
8
+ input_width: 512
9
+ radar_max_depth_m: 11.2
10
+ max_dist_correspondence: 0.5
11
+ patch_size: [64, 128]
12
+
13
+ # Model architecture
14
+ encoder: resnet34_bts
15
+ encoder_radar: resnet18
16
+ radar_input_channels: 1
17
+ bts_size: 512
18
+ max_depth: 11.2
19
+
20
+ # Runtime
21
+ batch_size: 32
22
+ num_workers: 0
23
+ seed: 42
24
+ cpu: false
25
+ mixed_precision: "fp16"
26
+
27
+ # Checkpoint and output root
28
+ checkpoint_path: "../../../checkpoints/baselines/cafnet_no_smoke/cafnet_no_smoke.safetensors"
29
+ prediction_dir: "prediction"
src/Baselines/cafnet_no_smoke/models/bts.py ADDED
@@ -0,0 +1,367 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright (C) 2019 Jin Han Lee
2
+ #
3
+ # This file is a part of BTS.
4
+ # This program is free software: you can redistribute it and/or modify
5
+ # it under the terms of the GNU General Public License as published by
6
+ # the Free Software Foundation, either version 3 of the License, or
7
+ # (at your option) any later version.
8
+ #
9
+ # This program is distributed in the hope that it will be useful,
10
+ # but WITHOUT ANY WARRANTY; without even the implied warranty of
11
+ # MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
12
+ # GNU General Public License for more details.
13
+ #
14
+ # You should have received a copy of the GNU General Public License
15
+ # along with this program. If not, see <http://www.gnu.org/licenses/>
16
+
17
+ import torch
18
+ import torch.nn as nn
19
+ import torch.nn.functional as torch_nn_func
20
+ import math
21
+
22
+
23
+ def bn_init_as_tf(m):
24
+ if isinstance(m, nn.BatchNorm2d):
25
+ m.track_running_stats = True # These two lines enable using stats (moving mean and var) loaded from pretrained model
26
+ m.eval() # or zero mean and variance of one if the batch norm layer has no pretrained values
27
+ m.affine = True
28
+ m.requires_grad = True
29
+
30
+
31
+ def weights_init_xavier(m):
32
+ if isinstance(m, nn.Conv2d):
33
+ torch.nn.init.xavier_uniform_(m.weight)
34
+ if m.bias is not None:
35
+ torch.nn.init.zeros_(m.bias)
36
+
37
+
38
+ class atrous_conv(nn.Sequential):
39
+ def __init__(self, in_channels, out_channels, dilation, apply_bn_first=True):
40
+ super(atrous_conv, self).__init__()
41
+ self.atrous_conv = torch.nn.Sequential()
42
+ if apply_bn_first:
43
+ self.atrous_conv.add_module('first_bn', nn.BatchNorm2d(in_channels, momentum=0.01, affine=True, track_running_stats=True, eps=1.1e-5))
44
+
45
+ self.atrous_conv.add_module('aconv_sequence', nn.Sequential(nn.ReLU(),
46
+ nn.Conv2d(in_channels=in_channels, out_channels=out_channels*2, bias=False, kernel_size=1, stride=1, padding=0),
47
+ nn.BatchNorm2d(out_channels*2, momentum=0.01, affine=True, track_running_stats=True),
48
+ nn.ReLU(),
49
+ nn.Conv2d(in_channels=out_channels * 2, out_channels=out_channels, bias=False, kernel_size=3, stride=1,
50
+ padding=(dilation, dilation), dilation=dilation)))
51
+
52
+ def forward(self, x):
53
+ return self.atrous_conv.forward(x)
54
+
55
+ class upconv(nn.Module):
56
+ def __init__(self, in_channels, out_channels, ratio=2):
57
+ super(upconv, self).__init__()
58
+ self.elu = nn.ELU()
59
+ self.conv = nn.Conv2d(in_channels=in_channels, out_channels=out_channels, bias=False, kernel_size=3, stride=1, padding=1)
60
+ self.ratio = ratio
61
+
62
+ def forward(self, x):
63
+ up_x = torch_nn_func.interpolate(x, scale_factor=self.ratio, mode='nearest')
64
+ out = self.conv(up_x)
65
+ out = self.elu(out)
66
+ return out
67
+
68
+ class reduction_1x1(nn.Sequential):
69
+ def __init__(self, num_in_filters, num_out_filters, max_depth, is_final=False):
70
+ super(reduction_1x1, self).__init__()
71
+ self.max_depth = max_depth
72
+ self.is_final = is_final
73
+ self.sigmoid = nn.Sigmoid()
74
+ self.reduc = torch.nn.Sequential()
75
+
76
+ while num_out_filters >= 4:
77
+ if num_out_filters < 8:
78
+ if self.is_final:
79
+ self.reduc.add_module('final', torch.nn.Sequential(nn.Conv2d(num_in_filters, out_channels=1, bias=False,
80
+ kernel_size=1, stride=1, padding=0),
81
+ nn.Sigmoid()))
82
+ else:
83
+ self.reduc.add_module('plane_params', torch.nn.Conv2d(num_in_filters, out_channels=3, bias=False,
84
+ kernel_size=1, stride=1, padding=0))
85
+ break
86
+ else:
87
+ self.reduc.add_module('inter_{}_{}'.format(num_in_filters, num_out_filters),
88
+ torch.nn.Sequential(nn.Conv2d(in_channels=num_in_filters, out_channels=num_out_filters,
89
+ bias=False, kernel_size=1, stride=1, padding=0),
90
+ nn.ELU()))
91
+
92
+ num_in_filters = num_out_filters
93
+ num_out_filters = num_out_filters // 2
94
+
95
+ def forward(self, net):
96
+ net = self.reduc.forward(net)
97
+ if not self.is_final:
98
+ theta = self.sigmoid(net[:, 0, :, :]) * math.pi / 3
99
+ phi = self.sigmoid(net[:, 1, :, :]) * math.pi * 2
100
+ dist = self.sigmoid(net[:, 2, :, :]) * self.max_depth
101
+ n1 = torch.mul(torch.sin(theta), torch.cos(phi)).unsqueeze(1)
102
+ n2 = torch.mul(torch.sin(theta), torch.sin(phi)).unsqueeze(1)
103
+ n3 = torch.cos(theta).unsqueeze(1)
104
+ n4 = dist.unsqueeze(1)
105
+ net = torch.cat([n1, n2, n3, n4], dim=1)
106
+
107
+ return net
108
+
109
+ class local_planar_guidance(nn.Module):
110
+ def __init__(self, upratio):
111
+ super(local_planar_guidance, self).__init__()
112
+ self.upratio = upratio
113
+ self.u = torch.arange(self.upratio).reshape([1, 1, self.upratio]).float()
114
+ self.v = torch.arange(int(self.upratio)).reshape([1, self.upratio, 1]).float()
115
+ self.upratio = float(upratio)
116
+
117
+ def forward(self, plane_eq, focal):
118
+ plane_eq_expanded = torch.repeat_interleave(plane_eq, int(self.upratio), 2)
119
+ plane_eq_expanded = torch.repeat_interleave(plane_eq_expanded, int(self.upratio), 3)
120
+ n1 = plane_eq_expanded[:, 0, :, :]
121
+ n2 = plane_eq_expanded[:, 1, :, :]
122
+ n3 = plane_eq_expanded[:, 2, :, :]
123
+ n4 = plane_eq_expanded[:, 3, :, :]
124
+
125
+ u = self.u.repeat(plane_eq.size(0), plane_eq.size(2) * int(self.upratio), plane_eq.size(3)).cuda()
126
+ u = (u - (self.upratio - 1) * 0.5) / self.upratio
127
+
128
+ v = self.v.repeat(plane_eq.size(0), plane_eq.size(2), plane_eq.size(3) * int(self.upratio)).cuda()
129
+ v = (v - (self.upratio - 1) * 0.5) / self.upratio
130
+
131
+ return n4 / (n1 * u + n2 * v + n3)
132
+
133
+ class bts_gated_fuse(nn.Module):
134
+ def __init__(self, params, feat_out_channels, feat_out_channels_rad, num_features=512):
135
+ super(bts_gated_fuse, self).__init__()
136
+ self.params = params
137
+ self.weight5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[4], feat_out_channels[4], 1, 1, bias=False),
138
+ nn.Sigmoid())
139
+ self.project5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[4], feat_out_channels[4], 1, 1, bias=False),
140
+ nn.ReLU())
141
+ self.upconv5 = upconv(feat_out_channels[4], num_features)
142
+ self.bn5 = nn.BatchNorm2d(num_features, momentum=0.01, affine=True, eps=1.1e-5)
143
+
144
+ self.conv5 = torch.nn.Sequential(nn.Conv2d(num_features + feat_out_channels[3], num_features, 3, 1, 1, bias=False),
145
+ nn.ELU())
146
+
147
+ self.weight4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[3], num_features, 1, 1, bias=False),
148
+ nn.Sigmoid())
149
+ self.project4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[3], num_features, 1, 1, bias=False),
150
+ nn.ReLU())
151
+ self.upconv4 = upconv(num_features, num_features // 2)
152
+ self.bn4 = nn.BatchNorm2d(num_features // 2, momentum=0.01, affine=True, eps=1.1e-5)
153
+ self.conv4 = torch.nn.Sequential(nn.Conv2d(num_features // 2 + feat_out_channels[2], num_features // 2, 3, 1, 1, bias=False),
154
+ nn.ELU())
155
+ self.bn4_2 = nn.BatchNorm2d(num_features // 2, momentum=0.01, affine=True, eps=1.1e-5)
156
+
157
+ self.daspp_3 = atrous_conv(num_features // 2, num_features // 4, 3, apply_bn_first=False)
158
+ self.daspp_6 = atrous_conv(num_features // 2 + num_features // 4 + feat_out_channels[2], num_features // 4, 6)
159
+ self.daspp_12 = atrous_conv(num_features + feat_out_channels[2], num_features // 4, 12)
160
+ self.daspp_18 = atrous_conv(num_features + num_features // 4 + feat_out_channels[2], num_features // 4, 18)
161
+ self.daspp_24 = atrous_conv(num_features + num_features // 2 + feat_out_channels[2], num_features // 4, 24)
162
+ self.daspp_conv = torch.nn.Sequential(nn.Conv2d(num_features + num_features // 2 + num_features // 4, num_features // 4, 3, 1, 1, bias=False),
163
+ nn.ELU())
164
+ self.reduc8x8 = reduction_1x1(num_features // 4, num_features // 4, self.params.max_depth)
165
+ self.lpg8x8 = local_planar_guidance(8)
166
+
167
+ self.weight3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[2], num_features // 4, 1, 1, bias=False),
168
+ nn.Sigmoid())
169
+ self.project3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[2], num_features // 4, 1, 1, bias=False),
170
+ nn.ReLU())
171
+ self.upconv3 = upconv(num_features // 4, num_features // 4)
172
+ self.bn3 = nn.BatchNorm2d(num_features // 4, momentum=0.01, affine=True, eps=1.1e-5)
173
+ self.conv3 = torch.nn.Sequential(nn.Conv2d(num_features // 4 + feat_out_channels[1] + 1, num_features // 4, 3, 1, 1, bias=False),
174
+ nn.ELU())
175
+ self.reduc4x4 = reduction_1x1(num_features // 4, num_features // 8, self.params.max_depth)
176
+ self.lpg4x4 = local_planar_guidance(4)
177
+
178
+ self.weight2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[1], num_features // 4, 1, 1, bias=False),
179
+ nn.Sigmoid())
180
+ self.project2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[1], num_features // 4, 1, 1, bias=False),
181
+ nn.ReLU())
182
+ self.upconv2 = upconv(num_features // 4, num_features // 8)
183
+ self.bn2 = nn.BatchNorm2d(num_features // 8, momentum=0.01, affine=True, eps=1.1e-5)
184
+ self.conv2 = torch.nn.Sequential(nn.Conv2d(num_features // 8 + feat_out_channels[0] + 1, num_features // 8, 3, 1, 1, bias=False),
185
+ nn.ELU())
186
+
187
+ self.reduc2x2 = reduction_1x1(num_features // 8, num_features // 16, self.params.max_depth)
188
+ self.lpg2x2 = local_planar_guidance(2)
189
+
190
+ self.weight1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[0], num_features // 8, 1, 1, bias=False),
191
+ nn.Sigmoid())
192
+ self.project1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_rad[0], num_features // 8, 1, 1, bias=False),
193
+ nn.ReLU())
194
+ self.upconv1 = upconv(num_features // 8, num_features // 16)
195
+ self.reduc1x1 = reduction_1x1(num_features // 16, num_features // 32, self.params.max_depth, is_final=True)
196
+ self.conv1 = torch.nn.Sequential(nn.Conv2d(num_features // 16 + 4, num_features // 16, 3, 1, 1, bias=False),
197
+ nn.ELU())
198
+ self.get_depth = torch.nn.Sequential(nn.Conv2d(num_features // 16, 1, 3, 1, 1, bias=False),
199
+ nn.Sigmoid())
200
+
201
+ self.pool5 = torch.nn.AvgPool2d(32, 32)
202
+ self.pool4 = torch.nn.AvgPool2d(16, 16)
203
+ self.pool3 = torch.nn.AvgPool2d(8, 8)
204
+ self.pool2 = torch.nn.AvgPool2d(4, 4)
205
+ self.pool1 = torch.nn.AvgPool2d(2, 2)
206
+
207
+ def forward(self, img_features, rad_features, focal, radar_confidence):
208
+ skip0, skip1, skip2, skip3 = img_features[0], img_features[1], img_features[2], img_features[3]
209
+ rad_skip0, rad_skip1, rad_skip2, rad_skip3 = rad_features[0], rad_features[1], rad_features[2], rad_features[3]
210
+
211
+ # prepare radar confidence
212
+ radar_confidence5 = self.pool5(radar_confidence)
213
+ radar_confidence4 = self.pool4(radar_confidence)
214
+ radar_confidence3 = self.pool3(radar_confidence)
215
+ radar_confidence2 = self.pool2(radar_confidence)
216
+ radar_confidence1 = self.pool1(radar_confidence)
217
+
218
+
219
+ rad_weight5 = self.weight5(rad_features[4])
220
+ rad_project5 = self.project5(rad_features[4])
221
+
222
+ dense_features = torch.nn.ReLU()(img_features[4])
223
+ dense_features = dense_features + rad_weight5*rad_project5*radar_confidence5
224
+ upconv5 = self.upconv5(dense_features) # H/16
225
+ upconv5 = self.bn5(upconv5)
226
+ concat5 = torch.cat([upconv5, skip3], dim=1)
227
+ iconv5 = self.conv5(concat5)
228
+
229
+ rad_weight4 = self.weight4(rad_skip3)
230
+ rad_project4 = self.project4(rad_skip3)
231
+
232
+ iconv5 = iconv5 + rad_weight4*rad_project4*radar_confidence4
233
+ upconv4 = self.upconv4(iconv5) # H/8
234
+ upconv4 = self.bn4(upconv4)
235
+ concat4 = torch.cat([upconv4, skip2], dim=1)
236
+ iconv4 = self.conv4(concat4)
237
+ iconv4 = self.bn4_2(iconv4)
238
+
239
+ daspp_3 = self.daspp_3(iconv4)
240
+ concat4_2 = torch.cat([concat4, daspp_3], dim=1)
241
+ daspp_6 = self.daspp_6(concat4_2)
242
+ concat4_3 = torch.cat([concat4_2, daspp_6], dim=1)
243
+ daspp_12 = self.daspp_12(concat4_3)
244
+ concat4_4 = torch.cat([concat4_3, daspp_12], dim=1)
245
+ daspp_18 = self.daspp_18(concat4_4)
246
+ concat4_5 = torch.cat([concat4_4, daspp_18], dim=1)
247
+ daspp_24 = self.daspp_24(concat4_5)
248
+ concat4_daspp = torch.cat([iconv4, daspp_3, daspp_6, daspp_12, daspp_18, daspp_24], dim=1)
249
+ daspp_feat = self.daspp_conv(concat4_daspp)
250
+ rad_weight3 = self.weight3(rad_skip2)
251
+ rad_project3 = self.project3(rad_skip2)
252
+ daspp_feat = daspp_feat + rad_weight3*rad_project3*radar_confidence3
253
+
254
+ reduc8x8 = self.reduc8x8(daspp_feat)
255
+ plane_normal_8x8 = reduc8x8[:, :3, :, :]
256
+ plane_normal_8x8 = torch_nn_func.normalize(plane_normal_8x8, 2, 1)
257
+ plane_dist_8x8 = reduc8x8[:, 3, :, :]
258
+ plane_eq_8x8 = torch.cat([plane_normal_8x8, plane_dist_8x8.unsqueeze(1)], 1)
259
+ depth_8x8 = self.lpg8x8(plane_eq_8x8, focal)
260
+ depth_8x8_scaled = depth_8x8.unsqueeze(1) / self.params.max_depth
261
+ depth_8x8_scaled_ds = torch_nn_func.interpolate(depth_8x8_scaled, scale_factor=0.25, mode='nearest')
262
+
263
+ upconv3 = self.upconv3(daspp_feat) # H/4
264
+ upconv3 = self.bn3(upconv3)
265
+ concat3 = torch.cat([upconv3, skip1, depth_8x8_scaled_ds], dim=1)
266
+ iconv3 = self.conv3(concat3)
267
+ rad_weight2 = self.weight2(rad_skip1)
268
+ rad_project2 = self.project2(rad_skip1)
269
+ iconv3 = iconv3 + rad_weight2*rad_project2*radar_confidence2
270
+
271
+ reduc4x4 = self.reduc4x4(iconv3)
272
+ plane_normal_4x4 = reduc4x4[:, :3, :, :]
273
+ plane_normal_4x4 = torch_nn_func.normalize(plane_normal_4x4, 2, 1)
274
+ plane_dist_4x4 = reduc4x4[:, 3, :, :]
275
+ plane_eq_4x4 = torch.cat([plane_normal_4x4, plane_dist_4x4.unsqueeze(1)], 1)
276
+ depth_4x4 = self.lpg4x4(plane_eq_4x4, focal)
277
+ depth_4x4_scaled = depth_4x4.unsqueeze(1) / self.params.max_depth
278
+ depth_4x4_scaled_ds = torch_nn_func.interpolate(depth_4x4_scaled, scale_factor=0.5, mode='nearest')
279
+
280
+ upconv2 = self.upconv2(iconv3) # H/2
281
+ upconv2 = self.bn2(upconv2)
282
+ concat2 = torch.cat([upconv2, skip0, depth_4x4_scaled_ds], dim=1)
283
+ iconv2 = self.conv2(concat2)
284
+ rad_weight1 = self.weight1(rad_skip0)
285
+ rad_project1 = self.project1(rad_skip0)
286
+ iconv2 = iconv2 + rad_weight1*rad_project1*radar_confidence1
287
+
288
+ reduc2x2 = self.reduc2x2(iconv2)
289
+ plane_normal_2x2 = reduc2x2[:, :3, :, :]
290
+ plane_normal_2x2 = torch_nn_func.normalize(plane_normal_2x2, 2, 1)
291
+ plane_dist_2x2 = reduc2x2[:, 3, :, :]
292
+ plane_eq_2x2 = torch.cat([plane_normal_2x2, plane_dist_2x2.unsqueeze(1)], 1)
293
+ depth_2x2 = self.lpg2x2(plane_eq_2x2, focal)
294
+ depth_2x2_scaled = depth_2x2.unsqueeze(1) / self.params.max_depth
295
+
296
+ rad_weight1 = self.weight1(rad_skip0)
297
+ rad_project1 = self.project1(rad_skip0)
298
+
299
+ upconv1 = self.upconv1(iconv2)
300
+ reduc1x1 = self.reduc1x1(upconv1)
301
+ concat1 = torch.cat([upconv1, reduc1x1, depth_2x2_scaled, depth_4x4_scaled, depth_8x8_scaled], dim=1)
302
+ iconv1 = self.conv1(concat1)
303
+ final_depth = self.params.max_depth * self.get_depth(iconv1)
304
+
305
+ return depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth
306
+
307
+ class encoder_image(nn.Module):
308
+ def __init__(self, params):
309
+ super(encoder_image, self).__init__()
310
+ self.params = params
311
+ import torchvision.models as models
312
+ if params.encoder == 'densenet121_bts':
313
+ self.base_model = models.densenet121(pretrained=False).features
314
+ self.feat_names = ['relu0', 'pool0', 'transition1', 'transition2', 'norm5']
315
+ self.feat_out_channels = [64, 64, 128, 256, 1024]
316
+ elif params.encoder == 'densenet161_bts':
317
+ self.base_model = models.densenet161(pretrained=False).features
318
+ self.feat_names = ['relu0', 'pool0', 'transition1', 'transition2', 'norm5']
319
+ self.feat_out_channels = [96, 96, 192, 384, 2208]
320
+ elif params.encoder == 'resnet50_bts':
321
+ self.base_model = models.resnet50(pretrained=False)
322
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
323
+ self.feat_out_channels = [64, 256, 512, 1024, 2048]
324
+ elif params.encoder == 'resnet34_bts':
325
+ self.base_model = models.resnet34(pretrained=False)
326
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
327
+ self.feat_out_channels = [64, 64, 128, 256, 512]
328
+ elif params.encoder == 'resnet18_bts':
329
+ self.base_model = models.resnet18(pretrained=False)
330
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
331
+ self.feat_out_channels = [64, 64, 128, 256, 512]
332
+ elif params.encoder == 'resnet101_bts':
333
+ self.base_model = models.resnet101(pretrained=False)
334
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
335
+ self.feat_out_channels = [64, 256, 512, 1024, 2048]
336
+ elif params.encoder == 'resnext50_bts':
337
+ self.base_model = models.resnext50_32x4d(pretrained=False)
338
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
339
+ self.feat_out_channels = [64, 256, 512, 1024, 2048]
340
+ elif params.encoder == 'resnext101_bts':
341
+ self.base_model = models.resnext101_32x8d(pretrained=False)
342
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
343
+ self.feat_out_channels = [64, 256, 512, 1024, 2048]
344
+ elif params.encoder == 'mobilenetv2_bts':
345
+ self.base_model = models.mobilenet_v2(pretrained=False).features
346
+ self.feat_inds = [2, 4, 7, 11, 19]
347
+ self.feat_out_channels = [16, 24, 32, 64, 1280]
348
+ self.feat_names = []
349
+ else:
350
+ print('Not supported encoder: {}'.format(params.encoder))
351
+
352
+ def forward(self, x):
353
+ feature = x
354
+ skip_feat = []
355
+ i = 1
356
+ for k, v in self.base_model._modules.items():
357
+ if 'fc' in k or 'avgpool' in k:
358
+ continue
359
+ feature = v(feature)
360
+ if self.params.encoder == 'mobilenetv2_bts':
361
+ if i == 2 or i == 4 or i == 7 or i == 11 or i == 19:
362
+ skip_feat.append(feature)
363
+ else:
364
+ if any(x in k for x in self.feat_names):
365
+ skip_feat.append(feature)
366
+ i = i + 1
367
+ return skip_feat
src/Baselines/cafnet_no_smoke/models/model.py ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+ from models.bts import encoder_image, bts_gated_fuse
4
+ from models.radar import encoder_radar_sparse_conv, encoder_radar_sub, decoder_radar
5
+
6
+ class CaFNet(nn.Module):
7
+ def __init__(self, params, threshold=0.4):
8
+ super(CaFNet, self).__init__()
9
+ self.threshold = threshold
10
+ self.encoder = encoder_image(params)
11
+ self.encoder_radar1 = encoder_radar_sparse_conv(params)
12
+ self.decoder_radar = decoder_radar(params, self.encoder.feat_out_channels, self.encoder_radar1.feat_out_channels)
13
+ self.encoder_radar2 = encoder_radar_sub(params)
14
+ self.decoder = bts_gated_fuse(params, self.encoder.feat_out_channels, self.encoder_radar2.feat_out_channels, params.bts_size)
15
+
16
+
17
+ def forward(self, x, radar, focal):
18
+
19
+ skip_feat = self.encoder(x)
20
+ skip_feat_radar = self.encoder_radar1(radar)
21
+ rad_confidence, rad_depth = self.decoder_radar(skip_feat, skip_feat_radar)
22
+ mask = (rad_confidence > self.threshold).float()
23
+ radar_new_input = torch.cat([mask*rad_depth, radar], axis=1)
24
+ skip_feat_radar_new = self.encoder_radar2(radar_new_input)
25
+
26
+ depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth = self.decoder(skip_feat, skip_feat_radar_new, focal, rad_confidence)
27
+
28
+ return depth_8x8_scaled, depth_4x4_scaled, depth_2x2_scaled, reduc1x1, final_depth, rad_confidence, rad_depth
src/Baselines/cafnet_no_smoke/models/radar.py ADDED
@@ -0,0 +1,212 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from models.bts import upconv
2
+ import torch
3
+ import torch.nn as nn
4
+ import torchvision.models as models
5
+
6
+ class encoder_radar_sparse_conv(nn.Module):
7
+ def __init__(self, params):
8
+ # radar encoder for the first stage
9
+ super(encoder_radar_sparse_conv, self).__init__()
10
+
11
+ self.params = params
12
+ self.sparse_conv1 = SparseConv(params.radar_input_channels, 16, 7, activation='elu')
13
+ self.sparse_conv2 = SparseConv(16, 16, 5, activation='elu')
14
+ self.sparse_conv3 = SparseConv(16, 16, 3, activation='elu')
15
+ self.sparse_conv4 = SparseConv(16, 3, 3, activation='elu')
16
+
17
+ if params.encoder_radar == 'resnet34':
18
+ self.base_model_radar = models.resnet34(pretrained=False)
19
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
20
+ self.feat_out_channels = [64, 64, 128, 256, 512]
21
+ elif params.encoder_radar == 'resnet18':
22
+ self.base_model_radar = models.resnet18(pretrained=False)
23
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
24
+ self.feat_out_channels = [64, 64, 128, 256, 512]
25
+ else:
26
+ print('Not supported encoder: {}'.format(params.encoder))
27
+
28
+ def forward(self, x):
29
+ mask = (x[:, 0] > 0).float().unsqueeze(1)
30
+ feature = x
31
+ feature, mask = self.sparse_conv1(feature, mask)
32
+ feature, mask = self.sparse_conv2(feature, mask)
33
+ feature, mask = self.sparse_conv3(feature, mask)
34
+ feature, mask = self.sparse_conv4(feature, mask)
35
+
36
+ skip_feat = []
37
+ i = 1
38
+ for k, v in self.base_model_radar._modules.items():
39
+ if 'fc' in k or 'avgpool' in k:
40
+ continue
41
+ feature = v(feature)
42
+ if any(x in k for x in self.feat_names):
43
+ skip_feat.append(feature)
44
+ i = i + 1
45
+ return skip_feat
46
+
47
+ class encoder_radar_sub(nn.Module):
48
+ def __init__(self, params):
49
+ # radar encoder for the second stage
50
+ super(encoder_radar_sub, self).__init__()
51
+
52
+ self.params = params
53
+ import torchvision.models as models
54
+ self.conv = torch.nn.Sequential(nn.Conv2d(params.radar_input_channels+1, 3, 3, 1, 1, bias=False),
55
+ nn.ELU())
56
+
57
+ if params.encoder_radar == 'resnet34':
58
+ self.base_model_radar = models.resnet34(pretrained=False)
59
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
60
+ self.feat_out_channels = [64, 64, 128, 256, 512]
61
+ elif params.encoder_radar == 'resnet18':
62
+ self.base_model_radar = models.resnet18(pretrained=False)
63
+ self.feat_names = ['relu', 'layer1', 'layer2', 'layer3', 'layer4']
64
+ self.feat_out_channels = [64, 64, 128, 256, 512]
65
+ else:
66
+ print('Not supported encoder: {}'.format(params.encoder))
67
+ def forward(self, x):
68
+ feature = x
69
+ feature = self.conv(feature)
70
+ skip_feat = []
71
+ i = 1
72
+ for k, v in self.base_model_radar._modules.items():
73
+ if 'fc' in k or 'avgpool' in k:
74
+ continue
75
+ feature = v(feature)
76
+ if any(x in k for x in self.feat_names):
77
+ skip_feat.append(feature)
78
+ i = i + 1
79
+ return skip_feat
80
+
81
+
82
+ class decoder_radar(nn.Module):
83
+ def __init__(self, params, feat_out_channels_img, feat_out_channels_radar):
84
+ super(decoder_radar, self).__init__()
85
+ self.params = params
86
+ self.upconv5 = upconv(feat_out_channels_img[4]+feat_out_channels_radar[4], feat_out_channels_radar[4]//2)
87
+ self.bn5 = nn.BatchNorm2d(feat_out_channels_radar[4]//2, momentum=0.01, affine=True, eps=1.1e-5)
88
+ self.conv5 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[4]//2, feat_out_channels_radar[4]//2, 3, 1, 1, bias=False),
89
+ nn.ELU())
90
+
91
+ 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)
92
+ self.bn4 = nn.BatchNorm2d(feat_out_channels_radar[3]//2, momentum=0.01, affine=True, eps=1.1e-5)
93
+ self.conv4 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[3]//2, feat_out_channels_radar[3]//2, 3, 1, 1, bias=False),
94
+ nn.ELU())
95
+
96
+ 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)
97
+ self.bn3 = nn.BatchNorm2d(feat_out_channels_radar[2]//2, momentum=0.01, affine=True, eps=1.1e-5)
98
+ self.conv3 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[2]//2, feat_out_channels_radar[2]//2, 3, 1, 1, bias=False),
99
+ nn.ELU())
100
+
101
+ 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)
102
+ self.bn2 = nn.BatchNorm2d(feat_out_channels_radar[1]//2, momentum=0.01, affine=True, eps=1.1e-5)
103
+ self.conv2 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[1]//2, feat_out_channels_radar[1]//2, 3, 1, 1, bias=False),
104
+ nn.ELU())
105
+
106
+ 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)
107
+ self.bn1 = nn.BatchNorm2d(feat_out_channels_radar[0]//2, momentum=0.01, affine=True, eps=1.1e-5)
108
+ self.conv1 = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, feat_out_channels_radar[0]//2, 3, 1, 1, bias=False),
109
+ nn.ELU())
110
+
111
+ # self.get_depth = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, 1, 3, 1, 1, bias=False),
112
+ # nn.Sigmoid())
113
+
114
+ self.get_depth = torch.nn.Sequential(nn.Conv2d(feat_out_channels_radar[0]//2, 2, 3, 1, 1, bias=False),
115
+ nn.Sigmoid())
116
+
117
+ def forward(self, image_features, radar_features):
118
+ 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]
119
+ 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]
120
+ final = torch.cat([img_final, rad_final], axis=1)
121
+ upconv5 = self.upconv5(final)
122
+ upconv5 = self.bn5(upconv5)
123
+ upconv5 = self.conv5(upconv5)
124
+ upconv5 = torch.cat([img_skip3, rad_skip3, upconv5], axis=1)
125
+
126
+ upconv4 = self.upconv4(upconv5)
127
+ upconv4 = self.bn4(upconv4)
128
+ upconv4 = self.conv4(upconv4)
129
+ upconv4 = torch.cat([img_skip2, rad_skip2, upconv4], axis=1)
130
+
131
+ upconv3 = self.upconv3(upconv4)
132
+ upconv3 = self.bn3(upconv3)
133
+ upconv3 = self.conv3(upconv3)
134
+ upconv3 = torch.cat([img_skip1, rad_skip1, upconv3], axis=1)
135
+
136
+ upconv2 = self.upconv2(upconv3)
137
+ upconv2 = self.bn2(upconv2)
138
+ upconv2 = self.conv2(upconv2)
139
+ upconv2 = torch.cat([img_skip0, rad_skip0, upconv2], axis=1)
140
+
141
+ upconv1 = self.upconv1(upconv2)
142
+ upconv1 = self.bn1(upconv1)
143
+ upconv1 = self.conv1(upconv1)
144
+
145
+ # confidence = self.get_depth(upconv1)
146
+ # depth = self.params.max_depth * confidence
147
+ depth_conf = self.get_depth(upconv1)
148
+ depth = self.params.max_depth * depth_conf[:, 0:1]
149
+ confidence = depth_conf[:, 1:2]
150
+
151
+ return confidence, depth
152
+
153
+
154
+ class SparseConv(nn.Module):
155
+
156
+ def __init__(self,
157
+ in_channels,
158
+ out_channels,
159
+ kernel_size,
160
+ activation='relu'):
161
+ super().__init__()
162
+
163
+ padding = kernel_size//2
164
+
165
+ self.conv = nn.Conv2d(
166
+ in_channels,
167
+ out_channels,
168
+ kernel_size=kernel_size,
169
+ padding=padding,
170
+ bias=False)
171
+
172
+ self.bias = nn.Parameter(
173
+ torch.zeros(out_channels),
174
+ requires_grad=True)
175
+
176
+ self.sparsity = nn.Conv2d(
177
+ in_channels,
178
+ out_channels,
179
+ kernel_size=kernel_size,
180
+ padding=padding,
181
+ bias=False)
182
+
183
+ kernel = torch.FloatTensor(torch.ones([kernel_size, kernel_size])).unsqueeze(0).unsqueeze(0)
184
+
185
+ self.sparsity.weight = nn.Parameter(
186
+ data=kernel,
187
+ requires_grad=False)
188
+
189
+ if activation == 'relu':
190
+ self.act = nn.ReLU(inplace=False)
191
+ elif activation == 'sigmoid':
192
+ self.act = nn.Sigmoid()
193
+ elif activation == 'elu':
194
+ self.act = nn.ELU()
195
+
196
+ self.max_pool = nn.MaxPool2d(
197
+ kernel_size,
198
+ stride=1,
199
+ padding=padding)
200
+
201
+
202
+
203
+ def forward(self, x, mask):
204
+ x = x*mask
205
+ x = self.conv(x)
206
+ normalizer = 1/(self.sparsity(mask)+1e-8)
207
+ x = x * normalizer + self.bias.unsqueeze(0).unsqueeze(2).unsqueeze(3)
208
+ x = self.act(x)
209
+
210
+ mask = self.max_pool(mask)
211
+
212
+ return x, mask
src/Baselines/cafnet_no_smoke/rice_dataset.py ADDED
@@ -0,0 +1,123 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import os
3
+ from typing import Dict, List, Optional, Tuple
4
+
5
+ import numpy as np
6
+ from torch.utils.data import Dataset
7
+
8
+
9
+ class RiceDataset(Dataset):
10
+ """Raw Rice dataset reader for DJI RGB, ZED depth and radar point clouds.
11
+
12
+ This dataset returns raw per-frame arrays and leaves geometric processing to
13
+ `collate_fn_helpers.make_rice_collate_fn`.
14
+ """
15
+
16
+ def __init__(
17
+ self,
18
+ base_dir: str,
19
+ split_json_path: Optional[str] = None,
20
+ split: str = "train",
21
+ input_height: int = 288,
22
+ input_width: int = 512,
23
+ patch_size: Optional[Tuple[int, int]] = None,
24
+ ):
25
+ self.base_dir = base_dir
26
+ self.split = split
27
+ self.input_height = int(input_height)
28
+ self.input_width = int(input_width)
29
+ self.patch_size = self._resolve_patch_size(patch_size)
30
+
31
+ test_sequences = self._load_test_split(split_json_path)
32
+
33
+ all_sequences = sorted(
34
+ d
35
+ for d in os.listdir(base_dir)
36
+ if os.path.isdir(os.path.join(base_dir, d)) and not d.startswith(".")
37
+ )
38
+
39
+ self.sequences: List[str] = []
40
+ for seq in all_sequences:
41
+ if split == "train" and seq in test_sequences:
42
+ continue
43
+ # if split == "train" and seq.lower().startswith("smoke"):
44
+ # continue
45
+ if split == "test" and seq not in test_sequences:
46
+ continue
47
+ if self._is_valid_sequence(os.path.join(base_dir, seq)):
48
+ self.sequences.append(seq)
49
+
50
+ self.dji_rgb_mmaps: Dict[str, np.memmap] = {}
51
+ self.zed_depth_mmaps: Dict[str, np.memmap] = {}
52
+ self.samples: List[Tuple[str, int]] = []
53
+
54
+ for seq in self.sequences:
55
+ seq_dir = os.path.join(self.base_dir, seq)
56
+ dji_rgb_path = os.path.join(seq_dir, "dji_rgb.npy")
57
+ zed_depth_path = os.path.join(seq_dir, "zed_depth.npy")
58
+
59
+ self.dji_rgb_mmaps[seq] = np.load(dji_rgb_path, mmap_mode="r")
60
+ self.zed_depth_mmaps[seq] = np.load(zed_depth_path, mmap_mode="r")
61
+
62
+ n_frames = min(
63
+ len(self.dji_rgb_mmaps[seq]),
64
+ len(self.zed_depth_mmaps[seq]),
65
+ )
66
+ for frame_idx in range(n_frames):
67
+ self.samples.append((seq, frame_idx))
68
+
69
+ def _resolve_patch_size(
70
+ self, patch_size: Optional[Tuple[int, int]]
71
+ ) -> Tuple[int, int]:
72
+ if patch_size is not None:
73
+ return int(patch_size[0]), int(patch_size[1])
74
+
75
+ # Scale default CaFNet patch size (50, 150) from 352x704.
76
+ base_h, base_w = 352, 704
77
+ scale_h = self.input_height / float(base_h)
78
+ scale_w = self.input_width / float(base_w)
79
+ ext_h = max(1, int(round(50 * scale_h)))
80
+ ext_w = max(1, int(round(150 * scale_w)))
81
+ return ext_h, ext_w
82
+
83
+ def _load_test_split(self, split_json_path: Optional[str]) -> set:
84
+ if not split_json_path or not os.path.exists(split_json_path):
85
+ return set()
86
+ with open(split_json_path, "r") as f:
87
+ payload = json.load(f)
88
+ return set(payload.get("test", []))
89
+
90
+ def _is_valid_sequence(self, seq_dir: str) -> bool:
91
+ dji_rgb_path = os.path.join(seq_dir, "dji_rgb.npy")
92
+ zed_depth_path = os.path.join(seq_dir, "zed_depth.npy")
93
+ pcd_dir = os.path.join(seq_dir, "pcd")
94
+ return (
95
+ os.path.exists(dji_rgb_path)
96
+ and os.path.exists(zed_depth_path)
97
+ and os.path.isdir(pcd_dir)
98
+ )
99
+
100
+ def __len__(self) -> int:
101
+ return len(self.samples)
102
+
103
+ def __getitem__(self, idx: int) -> Dict[str, object]:
104
+ seq, frame_idx = self.samples[idx]
105
+ seq_dir = os.path.join(self.base_dir, seq)
106
+
107
+ dji_rgb = np.asarray(self.dji_rgb_mmaps[seq][frame_idx]).copy()
108
+ zed_depth_mm = np.asarray(self.zed_depth_mmaps[seq][frame_idx]).copy()
109
+
110
+ pcd_path = os.path.join(seq_dir, "pcd", f"pcd_{frame_idx}.npy")
111
+ if os.path.exists(pcd_path):
112
+ radar_pcd_xyz = np.asarray(np.load(pcd_path), dtype=np.float32)
113
+ else:
114
+ radar_pcd_xyz = np.zeros((0, 3), dtype=np.float32)
115
+
116
+ return {
117
+ "sample_idx": idx,
118
+ "sequence": seq,
119
+ "frame_idx": frame_idx,
120
+ "dji_rgb": dji_rgb,
121
+ "zed_depth_mm": zed_depth_mm,
122
+ "radar_pcd_xyz": radar_pcd_xyz,
123
+ }
src/Baselines/cafnet_no_smoke/split.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "test": [
3
+ "Dell-1",
4
+ "Dell-2",
5
+ "Smoke-Dell-1",
6
+ "Smoke-Dell-2",
7
+ "Keck-1",
8
+ "Keck-2",
9
+ "Keck-3",
10
+ "Smoke-keck-1",
11
+ "Smoke-keck-2",
12
+ "Smoke-keck-3"
13
+ ]
14
+ }
src/Baselines/da3/inference.py ADDED
@@ -0,0 +1,179 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Depth Anything 3 metric-depth inference for Smoke-Eval sequences.
2
+
3
+ This inference-only adapter follows the official ByteDance-Seed
4
+ Depth-Anything-3 Python API. The upstream package supplies the model
5
+ architecture; this file supplies the artifact's local weights, camera
6
+ calibration, sequence sharding, and output contract.
7
+ """
8
+
9
+ import argparse
10
+ from pathlib import Path
11
+
12
+ import cv2
13
+ import numpy as np
14
+ import torch
15
+ from accelerate import Accelerator
16
+ from safetensors.torch import load_file
17
+ from tqdm.auto import tqdm
18
+
19
+
20
+ INTRINSICS = np.array(
21
+ [[365.13, 0.0, 445.43], [0.0, 365.13, 261.18], [0.0, 0.0, 1.0]],
22
+ dtype=np.float32,
23
+ )
24
+ SCALE_FACTOR = 1.15 * 365.13 / 300.0
25
+ TARGET_SIZE = (896, 504)
26
+
27
+
28
+ class Calibrator:
29
+ """Defish DJI frames and map them to the ZED-aligned view."""
30
+
31
+ def __init__(self):
32
+ k_dji = np.array(
33
+ [
34
+ [718.48555551, 0.0, 963.36465011],
35
+ [0.0, 720.25844189, 537.87569913],
36
+ [0.0, 0.0, 1.0],
37
+ ],
38
+ dtype=np.float64,
39
+ )
40
+ d_dji = np.array(
41
+ [0.19022699, 0.03466753, 0.05858962, -0.07070669],
42
+ dtype=np.float64,
43
+ )
44
+ new_k = cv2.fisheye.estimateNewCameraMatrixForUndistortRectify(
45
+ k_dji,
46
+ d_dji,
47
+ (1920, 1080),
48
+ np.eye(3),
49
+ balance=0.2,
50
+ fov_scale=1.0,
51
+ )
52
+ self.map1, self.map2 = cv2.fisheye.initUndistortRectifyMap(
53
+ k_dji,
54
+ d_dji,
55
+ np.eye(3),
56
+ new_k,
57
+ (1920, 1080),
58
+ cv2.CV_16SC2,
59
+ )
60
+ self.homography = np.array(
61
+ [
62
+ [
63
+ 0.8274446551892256,
64
+ -0.0742944198979625,
65
+ 80.23797348979947,
66
+ ],
67
+ [
68
+ -0.014725864916652691,
69
+ 0.8471179917075127,
70
+ 28.27366063997317,
71
+ ],
72
+ [
73
+ -5.083573451500717e-05,
74
+ -6.846079418201229e-05,
75
+ 1.0,
76
+ ],
77
+ ],
78
+ dtype=np.float64,
79
+ )
80
+
81
+ def __call__(self, rgb: np.ndarray) -> np.ndarray:
82
+ bgr = cv2.cvtColor(rgb, cv2.COLOR_RGB2BGR)
83
+ if bgr.shape[:2] != (1080, 1920):
84
+ bgr = cv2.resize(bgr, (1920, 1080), interpolation=cv2.INTER_LINEAR)
85
+ bgr = cv2.remap(bgr, self.map1, self.map2, cv2.INTER_LINEAR)
86
+ bgr = cv2.warpPerspective(bgr, self.homography, (1918, 1105))
87
+ bgr = bgr[115:760, 255:1400]
88
+ bgr = cv2.resize(bgr, TARGET_SIZE, interpolation=cv2.INTER_AREA)
89
+ return cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)
90
+
91
+
92
+ def parse_args() -> argparse.Namespace:
93
+ parser = argparse.ArgumentParser(description=__doc__)
94
+ parser.add_argument("--data_root", required=True)
95
+ parser.add_argument("--checkpoint", required=True)
96
+ parser.add_argument("--output_dir", required=True)
97
+ parser.add_argument("--model_name", default="da3metric-large")
98
+ parser.add_argument("--batch_size", type=int, default=16)
99
+ parser.add_argument("--sequences", nargs="*", default=None)
100
+ return parser.parse_args()
101
+
102
+
103
+ @torch.no_grad()
104
+ def main() -> None:
105
+ args = parse_args()
106
+ from depth_anything_3.api import DepthAnything3
107
+
108
+ class AccelerateFP16DepthAnything3(DepthAnything3):
109
+ """Use the official API while leaving autocast to Accelerate."""
110
+
111
+ @torch.inference_mode()
112
+ def forward(
113
+ self,
114
+ image,
115
+ extrinsics=None,
116
+ intrinsics=None,
117
+ export_feat_layers=None,
118
+ infer_gs=False,
119
+ use_ray_pose=False,
120
+ ref_view_strategy="saddle_balanced",
121
+ ):
122
+ return self.model(
123
+ image,
124
+ extrinsics,
125
+ intrinsics,
126
+ export_feat_layers,
127
+ infer_gs,
128
+ use_ray_pose,
129
+ ref_view_strategy,
130
+ )
131
+
132
+ accelerator = Accelerator(mixed_precision="fp16")
133
+ data_root = Path(args.data_root)
134
+ output_dir = Path(args.output_dir)
135
+ sequences = sorted(path for path in data_root.iterdir() if path.is_dir())
136
+ if args.sequences:
137
+ requested = set(args.sequences)
138
+ sequences = [path for path in sequences if path.name in requested]
139
+ local_sequences = sequences[
140
+ accelerator.process_index :: accelerator.num_processes
141
+ ]
142
+
143
+ model = AccelerateFP16DepthAnything3(model_name=args.model_name)
144
+ model.load_state_dict(load_file(args.checkpoint, device="cpu"), strict=True)
145
+ model = model.to(accelerator.device).eval()
146
+ calibrate = Calibrator()
147
+ if accelerator.is_main_process:
148
+ output_dir.mkdir(parents=True, exist_ok=True)
149
+ accelerator.wait_for_everyone()
150
+
151
+ for sequence in local_sequences:
152
+ rgb = np.load(sequence / "dji_rgb.npy", mmap_mode="r")
153
+ depth_chunks = []
154
+ for start in tqdm(
155
+ range(0, len(rgb), args.batch_size),
156
+ desc=sequence.name,
157
+ disable=not accelerator.is_local_main_process,
158
+ ):
159
+ end = min(start + args.batch_size, len(rgb))
160
+ images = [
161
+ calibrate(np.asarray(rgb[index])) for index in range(start, end)
162
+ ]
163
+ intrinsics = np.repeat(INTRINSICS[None], len(images), axis=0)
164
+ with accelerator.autocast():
165
+ prediction = model.inference(
166
+ images,
167
+ intrinsics=intrinsics,
168
+ process_res=896,
169
+ process_res_method="upper_bound_resize",
170
+ )
171
+ depth_chunks.append(prediction.depth * SCALE_FACTOR)
172
+ depth = np.concatenate(depth_chunks).astype(np.float32, copy=False)
173
+ np.save(output_dir / f"{sequence.name.lower()}_pred.npy", depth)
174
+
175
+ accelerator.wait_for_everyone()
176
+
177
+
178
+ if __name__ == "__main__":
179
+ main()
src/Baselines/grt/augmentations.py ADDED
@@ -0,0 +1,193 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torchvision.transforms.functional as TF
3
+ from torchvision.transforms import Resize, InterpolationMode
4
+ from typing import Union
5
+ import numpy as np
6
+
7
+ AZIMUTH_RESOLUTION = 128
8
+ ELEVATION_RESOLUTION = 64
9
+
10
+ # Depth output resolution: height=64, width=128
11
+ DEPTH_TARGET_HEIGHT = 64
12
+ DEPTH_TARGET_WIDTH = 128
13
+
14
+ resize_transform = Resize(
15
+ size=[ELEVATION_RESOLUTION, AZIMUTH_RESOLUTION],
16
+ interpolation=InterpolationMode.BILINEAR,
17
+ antialias=True,
18
+ )
19
+
20
+ depth_resize_transform = Resize(
21
+ size=(DEPTH_TARGET_HEIGHT, DEPTH_TARGET_WIDTH),
22
+ interpolation=InterpolationMode.BILINEAR,
23
+ antialias=True,
24
+ )
25
+
26
+
27
+ def translate_radar(radar_data):
28
+ """
29
+ Applies normalization to radar data after batching from dataloader.
30
+ Called before passing data into the model.
31
+
32
+ Args:
33
+ radar_data: Batched radar tensor from dataloader
34
+ Shape: [B, 64, 8, 2, 256, 2] (batch, doppler, azimuth, elevation, range, channels)
35
+ - Channel 0: raw amplitude values
36
+ - Channel 1: phase normalized to [-1, 1] (divided by π)
37
+
38
+ Returns:
39
+ Processed radar tensor with same shape [B, 64, 8, 2, 256, 2]
40
+ - Channel 0: sqrt(amplitude * 1e-3) for magnitude normalization
41
+ - Channel 1: phase * π (converted back to radians [-π, π])
42
+ """
43
+ radar_mag = radar_data[..., 0] # [B, 64, 8, 2, 256] - Extract raw amplitude
44
+ radar_phase = radar_data[..., 1] # [B, 64, 8, 2, 256] - Extract normalized phase
45
+
46
+ # Normalize amplitude: scale then sqrt
47
+ radar_mag_processed = torch.sqrt(radar_mag * 1e-6)
48
+
49
+ # Convert phase back to radians: [-1, 1] -> [-π, π]
50
+ radar_phase_processed = radar_phase * torch.pi
51
+
52
+ # Stack channels back together: [B, 64, 8, 2, 256, 2]
53
+ radar_data_translated = torch.stack(
54
+ [radar_mag_processed, radar_phase_processed], dim=-1
55
+ )
56
+ return radar_data_translated
57
+
58
+
59
+ def resize_depth(
60
+ depth_map: Union[torch.Tensor, np.ndarray],
61
+ ) -> Union[torch.Tensor, np.ndarray]:
62
+ """
63
+ Process depth map from dataloader (same pipeline as denoiser/control crop_depth):
64
+ mm -> meters, clamp [0, 11.2] m, normalize to [0, 1], resize to (64, 128) (h, w).
65
+
66
+ Args:
67
+ depth_map: Depth in millimeters. Torch or numpy.
68
+ Shapes: (H, W), (B, H, W), or (B, 1, H, W).
69
+
70
+ Returns:
71
+ Depth in [0, 1], spatial size (64, 128). Shape [B, 64, 128] for batched input.
72
+ """
73
+ is_numpy = isinstance(depth_map, np.ndarray)
74
+ if is_numpy:
75
+ depth_map = torch.from_numpy(depth_map)
76
+
77
+ depth_map = depth_map.float()
78
+ original_shape = depth_map.shape
79
+
80
+ if depth_map.dim() == 2:
81
+ depth_map = depth_map.unsqueeze(0) # (H, W) -> (1, H, W)
82
+ elif depth_map.dim() == 3:
83
+ depth_map = depth_map.unsqueeze(1) # (B, H, W) -> (B, 1, H, W)
84
+ elif depth_map.dim() != 4:
85
+ raise ValueError(f"Unexpected depth shape: {original_shape}")
86
+
87
+ invalid_mask = ~(torch.isfinite(depth_map) & (depth_map >= 0))
88
+ depth_map[invalid_mask] = 0.0
89
+
90
+ depth_map = depth_map / 1000.0 # mm -> meters
91
+ max_depth_m = 11.2
92
+ depth_map = torch.clamp(depth_map, min=0.0, max=max_depth_m)
93
+ depth_map = depth_map / max_depth_m # [0, 1]
94
+
95
+ invalid_mask = ~torch.isfinite(depth_map)
96
+ depth_map[invalid_mask] = 0.0
97
+
98
+ depth_map = depth_resize_transform(depth_map) # (..., 64, 128)
99
+ depth_values = depth_map.squeeze(1) # [B, 64, 128] or [1, 64, 128]
100
+
101
+ if len(original_shape) == 2:
102
+ depth_values = depth_values.squeeze(0) # (64, 128)
103
+
104
+ if is_numpy:
105
+ depth_values = depth_values.numpy()
106
+ return depth_values
107
+
108
+
109
+ def quantize_depth_to_occupancy(depth_values, num_range_bins=64):
110
+ """
111
+ Quantizes 2D depth values into 3D binary occupancy grid.
112
+
113
+ Args:
114
+ depth_values: Resized depth tensor
115
+ Shape: [B, elevation, azimuth]
116
+ Values: normalized to [0, 1] range
117
+ num_range_bins: Number of range bins for quantization (default: 64)
118
+
119
+ Returns:
120
+ Binary 3D occupancy grid
121
+ Shape: [B, elevation, azimuth, num_range_bins]
122
+ Values: binary (0 or 1) indicating occupied bins
123
+ """
124
+ B, elevation, azimuth = depth_values.shape
125
+
126
+ # Quantize normalized depth [0, 1] directly to range bins [0, num_range_bins-1]
127
+ # Each bin represents 1/num_range_bins of the normalized depth range
128
+ bin_indices = torch.floor(
129
+ depth_values / (1.0 / num_range_bins)
130
+ ).long() # [B, elevation, azimuth]
131
+ bin_indices = torch.clamp(
132
+ bin_indices, 0, num_range_bins - 1
133
+ ) # Handle edge case where depth_values = 1.0
134
+
135
+ # Create binary 3D occupancy grid
136
+ occupancy_grid = torch.zeros(
137
+ B,
138
+ elevation,
139
+ azimuth,
140
+ num_range_bins,
141
+ dtype=torch.float32,
142
+ device=depth_values.device,
143
+ ) # [B, elevation, azimuth, num_range_bins]
144
+
145
+ # Set occupied bins to 1
146
+ # Use advanced indexing to mark the appropriate range bin for each (elevation, azimuth) cell
147
+ batch_idx = torch.arange(B, device=depth_values.device)[:, None, None].expand(
148
+ B, elevation, azimuth
149
+ )
150
+ elevation_idx = torch.arange(elevation, device=depth_values.device)[
151
+ None, :, None
152
+ ].expand(B, elevation, azimuth)
153
+ azimuth_idx = torch.arange(azimuth, device=depth_values.device)[
154
+ None, None, :
155
+ ].expand(B, elevation, azimuth)
156
+
157
+ occupancy_grid[batch_idx, elevation_idx, azimuth_idx, bin_indices] = 1.0
158
+
159
+ return occupancy_grid # [B, elevation, azimuth, num_range_bins]
160
+
161
+
162
+ def dequantize_depth(occupancy_grid):
163
+ """
164
+ Converts 3D binary occupancy grid back to 2D depth map.
165
+ This is the inverse operation of quantize_depth_to_occupancy.
166
+
167
+ Args:
168
+ occupancy_grid: Binary 3D occupancy grid
169
+ Shape: [B, 64, 128, 64] (batch, elevation, azimuth, range)
170
+ Values: binary (0 or 1) or continuous (predicted probabilities)
171
+
172
+ Returns:
173
+ Reconstructed depth map
174
+ Shape: [B, 1, 64, 128] (batch, channel, elevation, azimuth)
175
+ Values: normalized to [0, 1] range
176
+ """
177
+ num_range_bins = occupancy_grid.shape[3]
178
+
179
+ # Find the range bin with maximum value for each (elevation, azimuth) cell
180
+ # For binary: finds the occupied bin
181
+ # For continuous: finds the most likely bin
182
+ bin_indices = torch.argmax(occupancy_grid, dim=3) # [B, 64, 128]
183
+
184
+ # Convert bin indices back to normalized depth values [0, 1]
185
+ # Use bin center: (bin_idx + 0.5) / num_bins
186
+ depth_values = (bin_indices.float() + 1) / num_range_bins # [B, 64, 128]
187
+
188
+ # Add channel dimension: [B, 64, 128] -> [B, 1, 64, 128]
189
+ depth_map = depth_values.unsqueeze(1) # [B, 1, 64, 128]
190
+
191
+ return depth_map
192
+
193
+
src/Baselines/grt/dataloader.py ADDED
@@ -0,0 +1,330 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Dataloader for MobiCom processed dataset (output of processor.py).
3
+
4
+ Uses the optimized format produced by processor.py:
5
+ - radar.npy: (N, doppler, elevation, azimuth, range) complex64
6
+ - dji_rgb.avi: DJI RGB video (FFV1), (N, H, W, 3) uint8
7
+ - zed_depth.npy: (N, H, W) uint16, depth in millimeters
8
+
9
+ This module provides:
10
+ - `RiceDataset`: frame-level dataset returning radar amplitude/phase, DJI RGB,
11
+ and ZED depth (ground truth).
12
+ - `create_rice_dataloader`: generic dataloader for an arbitrary set of sequences.
13
+ - `create_split_dataloaders`: reads train/val/test split from split.json
14
+ (default: radar_model/split.json) and returns train/val/test dataloaders.
15
+ """
16
+
17
+ import json
18
+ from pathlib import Path
19
+ from typing import Dict, List, Optional, Tuple
20
+
21
+ import cv2
22
+ import numpy as np
23
+ import torch
24
+ from torch.utils.data import Dataset, DataLoader, random_split
25
+
26
+
27
+ class RiceDataset(Dataset):
28
+ """
29
+ Dataset for processor.py output: radar, DJI RGB, and ZED depth per frame.
30
+
31
+ Args:
32
+ root_dir: Root directory containing sequence subdirs (e.g. processed/),
33
+ each with radar.npy, dji_rgb.avi, zed_depth.npy.
34
+ sequences: Optional list of sequence names to load. If None, loads all
35
+ subdirs that contain the three required files.
36
+ frame_skip: Sample every frame_skip frames (1 = all frames).
37
+ return_radar_complex: If True, return radar as complex tensor; if False,
38
+ return radar_amplitude and radar_phase as separate float tensors.
39
+ depth_in_meters: If True, convert depth from mm to meters.
40
+ rgb_normalize: If True, return RGB in [0, 1] float; else uint8 [0, 255].
41
+ """
42
+
43
+ # GRT inference consumes radar and depth only. Smoke-Eval packages RGB
44
+ # frames as ``dji_rgb.npy`` rather than the original training video, so
45
+ # requiring the unused video would incorrectly discard every sequence.
46
+ REQUIRED_FILES = ("radar.npy", "zed_depth.npy")
47
+
48
+ def __init__(
49
+ self,
50
+ root_dir: str,
51
+ sequences: Optional[List[str]] = None,
52
+ frame_skip: int = 1,
53
+ return_radar_complex: bool = False,
54
+ depth_in_meters: bool = True,
55
+ rgb_normalize: bool = True,
56
+ ):
57
+ self.root_dir = Path(root_dir)
58
+ self.frame_skip = max(1, frame_skip)
59
+ self.return_radar_complex = return_radar_complex
60
+ self.depth_in_meters = depth_in_meters
61
+ self.rgb_normalize = rgb_normalize
62
+
63
+ self.sequences = self._discover_sequences(sequences)
64
+ self.index_map: List[Tuple[str, int]] = [] # (seq_name, frame_idx)
65
+ self._seq_arrays: Dict[str, Dict] = {} # seq -> {radar, depth, dji_rgb}
66
+
67
+ self._build_index()
68
+
69
+ def _discover_sequences(self, sequences: Optional[List[str]] = None) -> List[str]:
70
+ """Return list of sequence names that have all required files."""
71
+ if not self.root_dir.is_dir():
72
+ raise FileNotFoundError(f"Root directory not found: {self.root_dir}")
73
+
74
+ all_seqs = sorted(
75
+ d.name
76
+ for d in self.root_dir.iterdir()
77
+ if d.is_dir() and not d.name.startswith(".")
78
+ )
79
+ valid = []
80
+ for name in all_seqs:
81
+ seq_dir = self.root_dir / name
82
+ if all((seq_dir / f).exists() for f in self.REQUIRED_FILES):
83
+ valid.append(name)
84
+ if sequences is not None:
85
+ valid = [s for s in valid if s in sequences]
86
+ return valid
87
+
88
+ def _build_index(self) -> None:
89
+ """Build (seq_name, frame_idx) index, using radar.npy for frame count."""
90
+ self.index_map.clear()
91
+ for seq_name in self.sequences:
92
+ seq_dir = self.root_dir / seq_name
93
+ radar_path = seq_dir / "radar.npy"
94
+ radar = np.load(radar_path, mmap_mode="r")
95
+ n_frames = radar.shape[0]
96
+ for i in range(0, n_frames, self.frame_skip):
97
+ self.index_map.append((seq_name, i))
98
+
99
+ # def _load_video_rgb(self, path: Path) -> np.ndarray:
100
+ # """Load RGB AVI (e.g. FFV1) as (N, H, W, 3) uint8 RGB."""
101
+ # cap = cv2.VideoCapture(str(path))
102
+ # if not cap.isOpened():
103
+ # raise RuntimeError(f"Failed to open video: {path}")
104
+ # frames = []
105
+ # while True:
106
+ # ret, frame = cap.read()
107
+ # if not ret:
108
+ # break
109
+ # rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
110
+ # frames.append(rgb)
111
+ # cap.release()
112
+ # if not frames:
113
+ # return np.empty((0, 0, 0, 3), dtype=np.uint8)
114
+ # return np.stack(frames, axis=0)
115
+
116
+ def _load_sequence_arrays(self, seq_name: str) -> Dict:
117
+ """Lazy-load or return cached arrays for a sequence."""
118
+ if seq_name not in self._seq_arrays:
119
+ seq_dir = self.root_dir / seq_name
120
+ # dji_rgb = self._load_video_rgb(seq_dir / "dji_rgb.avi")
121
+ self._seq_arrays[seq_name] = {
122
+ "radar": np.load(seq_dir / "radar.npy", mmap_mode="r"),
123
+ "depth": np.load(seq_dir / "zed_depth.npy", mmap_mode="r"),
124
+ # "dji_rgb": dji_rgb,
125
+ }
126
+ return self._seq_arrays[seq_name]
127
+
128
+ def __len__(self) -> int:
129
+ return len(self.index_map)
130
+
131
+ def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
132
+ seq_name, frame_idx = self.index_map[idx]
133
+ arrs = self._load_sequence_arrays(seq_name)
134
+
135
+ # (H, W, 3) uint8
136
+ # rgb = np.asarray(arrs["dji_rgb"][frame_idx])
137
+ # (H, W) uint16 mm (processor saves as uint16)
138
+ depth = np.asarray(arrs["depth"][frame_idx]).astype(np.float32)
139
+ # (doppler, elevation, azimuth, range) complex64
140
+ radar = np.asarray(arrs["radar"][frame_idx]).copy()
141
+
142
+ # Depth: uint16 mm -> float; optional mm -> m; handle invalid
143
+ if self.depth_in_meters:
144
+ depth = depth / 1000.0
145
+ invalid = ~(np.isfinite(depth) & (depth > 0))
146
+ depth[invalid] = 0.0
147
+ depth = depth[np.newaxis, ...] # (1, H, W)
148
+
149
+ # RGB: (H, W, 3) -> (3, H, W)
150
+ # rgb = np.transpose(rgb, (2, 0, 1))
151
+ # if self.rgb_normalize:
152
+ # rgb = rgb.astype(np.float32) / 255.0
153
+
154
+ # Radar: amplitude and phase
155
+ radar_amplitude = np.abs(radar).astype(np.float32)
156
+ radar_phase = np.angle(radar).astype(np.float32) / np.pi
157
+ out = {
158
+ "radar_amplitude": torch.from_numpy(radar_amplitude),
159
+ "radar_phase": torch.from_numpy(radar_phase),
160
+ # "rgb": torch.from_numpy(rgb),
161
+ "depth": torch.from_numpy(depth),
162
+ "sequence": seq_name,
163
+ "frame_idx": frame_idx,
164
+ }
165
+ if self.return_radar_complex:
166
+ out["radar_cube"] = torch.from_numpy(radar.copy())
167
+ # Depth in mm for optional use (1, H, W) float32
168
+ depth_mm = np.asarray(arrs["depth"][frame_idx]).astype(np.float32)
169
+ out["depth_mm"] = torch.from_numpy(depth_mm[np.newaxis, ...])
170
+ return out
171
+
172
+
173
+ def create_rice_dataloader(
174
+ root_dir: str,
175
+ batch_size: int = 8,
176
+ num_workers: int = 0,
177
+ frame_skip: int = 1,
178
+ sequences: Optional[List[str]] = None,
179
+ return_radar_complex: bool = False,
180
+ depth_in_meters: bool = True,
181
+ rgb_normalize: bool = True,
182
+ shuffle: bool = True,
183
+ ) -> DataLoader:
184
+ """Create a DataLoader for the Rice (processor output) dataset."""
185
+ dataset = RiceDataset(
186
+ root_dir=root_dir,
187
+ sequences=sequences,
188
+ frame_skip=frame_skip,
189
+ return_radar_complex=return_radar_complex,
190
+ depth_in_meters=depth_in_meters,
191
+ rgb_normalize=rgb_normalize,
192
+ )
193
+ return DataLoader(
194
+ dataset,
195
+ batch_size=batch_size,
196
+ shuffle=shuffle,
197
+ num_workers=num_workers,
198
+ pin_memory=True,
199
+ )
200
+
201
+
202
+ def create_split_dataloaders(
203
+ root_dir: str,
204
+ split_json_path: Optional[str] = None,
205
+ batch_size: int = 8,
206
+ num_workers: int = 0,
207
+ frame_skip: int = 1,
208
+ return_radar_complex: bool = False,
209
+ depth_in_meters: bool = True,
210
+ rgb_normalize: bool = True,
211
+ val_ratio: float = 0.2,
212
+ seed: Optional[int] = 42,
213
+ ) -> Tuple[DataLoader, DataLoader, DataLoader]:
214
+ """
215
+ Create train/val/test dataloaders using split.json.
216
+
217
+ Reads the dataset split from split.json. If split_json_path is None,
218
+ uses radar_model/split.json (same directory as this module).
219
+
220
+ Split JSON format:
221
+ { "test": ["seq_x", ...], "train": ["seq_a", ...] } // "train" optional
222
+ If "train" is present and non-empty, only those sequences are used for train/val.
223
+ Otherwise, all sequences under root_dir with required files that are not in "test" are used for training.
224
+ Validation is a random fraction (val_ratio) of the training samples.
225
+
226
+ Returns:
227
+ train_loader, val_loader, test_loader
228
+ """
229
+ if split_json_path is None:
230
+ split_path = Path(__file__).resolve().parent / "split.json"
231
+ else:
232
+ split_path = Path(split_json_path)
233
+ if not split_path.exists() and not split_path.is_absolute():
234
+ # Resolve relative path from this module's directory (e.g. radar_model/)
235
+ fallback = Path(__file__).resolve().parent / split_path.name
236
+ if fallback.exists():
237
+ split_path = fallback
238
+
239
+ with split_path.open("r") as f:
240
+ split = json.load(f)
241
+
242
+ test_sequences = split.get("test", [])
243
+ train_sequences_json = split.get("train", None)
244
+
245
+ # Discover all valid sequences in root_dir
246
+ _discover = RiceDataset(
247
+ root_dir=root_dir,
248
+ sequences=None,
249
+ frame_skip=frame_skip,
250
+ return_radar_complex=return_radar_complex,
251
+ depth_in_meters=depth_in_meters,
252
+ rgb_normalize=rgb_normalize,
253
+ )
254
+ test_set = set(test_sequences)
255
+ if train_sequences_json is not None and len(train_sequences_json) > 0:
256
+ # Use explicit train list (intersect with discovered so only valid seqs are used)
257
+ train_sequences = [s for s in train_sequences_json if s in _discover.sequences]
258
+ else:
259
+ # No "train" key: use all discovered sequences not in test
260
+ train_sequences = [s for s in _discover.sequences if s not in test_set]
261
+
262
+ full_train_dataset = RiceDataset(
263
+ root_dir=root_dir,
264
+ sequences=train_sequences,
265
+ frame_skip=frame_skip,
266
+ return_radar_complex=return_radar_complex,
267
+ depth_in_meters=depth_in_meters,
268
+ rgb_normalize=rgb_normalize,
269
+ )
270
+
271
+ # Random split of training data for validation
272
+ n_total = len(full_train_dataset)
273
+ n_val = int(n_total * val_ratio)
274
+ if n_val == 0 and n_total > 0:
275
+ n_val = 1
276
+ n_train = n_total - n_val
277
+
278
+ if n_total == 0:
279
+ train_dataset = full_train_dataset
280
+ val_dataset = RiceDataset(
281
+ root_dir=root_dir,
282
+ sequences=[],
283
+ frame_skip=frame_skip,
284
+ return_radar_complex=return_radar_complex,
285
+ depth_in_meters=depth_in_meters,
286
+ rgb_normalize=rgb_normalize,
287
+ )
288
+ elif seed is None:
289
+ train_dataset, val_dataset = random_split(
290
+ full_train_dataset, [n_train, n_val]
291
+ )
292
+ else:
293
+ generator = torch.Generator()
294
+ generator.manual_seed(seed)
295
+ train_dataset, val_dataset = random_split(
296
+ full_train_dataset, [n_train, n_val], generator=generator
297
+ )
298
+
299
+ test_dataset = RiceDataset(
300
+ root_dir=root_dir,
301
+ sequences=test_sequences,
302
+ frame_skip=frame_skip,
303
+ return_radar_complex=return_radar_complex,
304
+ depth_in_meters=depth_in_meters,
305
+ rgb_normalize=rgb_normalize,
306
+ )
307
+
308
+ train_loader = DataLoader(
309
+ train_dataset,
310
+ batch_size=batch_size,
311
+ shuffle=True,
312
+ num_workers=num_workers,
313
+ pin_memory=True,
314
+ )
315
+ val_loader = DataLoader(
316
+ val_dataset,
317
+ batch_size=batch_size,
318
+ shuffle=False,
319
+ num_workers=num_workers,
320
+ pin_memory=True,
321
+ )
322
+ test_loader = DataLoader(
323
+ test_dataset,
324
+ batch_size=batch_size,
325
+ shuffle=False,
326
+ num_workers=num_workers,
327
+ pin_memory=True,
328
+ )
329
+
330
+ return train_loader, val_loader, test_loader
src/Baselines/grt/grt_model.py ADDED
@@ -0,0 +1,585 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """GRT-Small Model - from official codebase.
2
+
3
+ This implementation directly copies necessary modules from the official GRT codebase
4
+ (grt/deepradar/modules).
5
+ """
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ from typing import Literal, Optional, Sequence
10
+ import numpy as np
11
+ from einops import rearrange
12
+
13
+ # ============================================================================
14
+ # Official GRT Modules (copied from grt/deepradar/modules/*.py)
15
+ # ============================================================================
16
+
17
+
18
+ class PatchMerge(nn.Module):
19
+ """Merge patches with normalization and nominally reduced projection.
20
+
21
+ From: grt/deepradar/modules/patch.py
22
+ """
23
+
24
+ def __init__(
25
+ self, d_in: int, d_out: int, scale: Sequence[int] = [], norm: bool = True
26
+ ) -> None:
27
+ super().__init__()
28
+
29
+ self.scale = scale
30
+ d_merge = d_in * int(np.prod(scale))
31
+ self.linear = nn.Linear(d_merge, d_out, bias=False)
32
+ self.norm = nn.LayerNorm(d_merge) if norm else None
33
+
34
+ def _merge(self, x: torch.Tensor) -> torch.Tensor:
35
+ """Perform patch merging."""
36
+ n, *t, c = x.shape
37
+ dims = sum(([d // s, s] for d, s in zip(t, self.scale)), start=[n])
38
+ order = (
39
+ [0]
40
+ + [2 * i + 1 for i in range(len(self.scale))]
41
+ + [2 * i + 2 for i in range(len(self.scale))]
42
+ + [-1]
43
+ )
44
+ t2 = [d // s for d, s in zip(t, self.scale)]
45
+ return x.reshape(dims + [c]).permute(order).reshape(n, *t2, -1)
46
+
47
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
48
+ """Merge and project."""
49
+ merged = self._merge(x)
50
+ if self.norm is not None:
51
+ merged = self.norm(merged)
52
+ return self.linear(merged)
53
+
54
+
55
+ class Sinusoid(nn.Module):
56
+ """Centered N-dimensional sinusoidal positional embedding.
57
+
58
+ From: grt/deepradar/modules/position.py
59
+ """
60
+
61
+ def __init__(
62
+ self,
63
+ scale: Optional[Sequence[float]] = None,
64
+ global_scale: float = 1.0,
65
+ coef: float = 10000.0,
66
+ ) -> None:
67
+ super().__init__()
68
+ if scale is None:
69
+ self.scale = [global_scale]
70
+ else:
71
+ self.scale = [s * global_scale for s in scale]
72
+ self.coef = coef
73
+
74
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
75
+ """Apply sinusoidal embedding."""
76
+ # w = coef ** (-i / c)
77
+ nd = len(x.shape) - 2
78
+ c = x.shape[-1] // 2 // nd
79
+ i = torch.arange(c, device=x.device)
80
+ w = self.coef ** (-i / c)
81
+
82
+ start_dim = 0
83
+ for axis, (d, scale) in enumerate(zip(x.shape[1:-1], self.scale * nd)):
84
+ # t = scale * (j - d/2) / (d/2) = scale * (2j / d - 1)
85
+ t = scale * (2 * (torch.arange(d, device=x.device) + 0.5) / d - 1)
86
+ wt = t[:, None] * w[None, :]
87
+
88
+ p_slice = [None] * (len(x.shape) - 1) + [slice(None)]
89
+ p_slice[axis + 1] = slice(None)
90
+
91
+ # pos[2 * i] = sin(w * t)
92
+ x_sin_slice = [slice(None)] * len(x.shape)
93
+ x_sin_slice[-1] = slice(start_dim, start_dim + c * 2, 2)
94
+ x_sin_slice = tuple(x_sin_slice)
95
+ p_slice_tuple = tuple(p_slice)
96
+ x[x_sin_slice] = x[x_sin_slice] + torch.sin(wt)[p_slice_tuple]
97
+
98
+ # pos[2 * i + 1] = cos(w * t)
99
+ x_cos_slice = [slice(None)] * len(x.shape)
100
+ x_cos_slice[-1] = slice(start_dim + 1, start_dim + c * 2 + 1, 2)
101
+ x_cos_slice = tuple(x_cos_slice)
102
+ x[x_cos_slice] = x[x_cos_slice] + torch.cos(wt)[p_slice_tuple]
103
+
104
+ start_dim += c * 2
105
+
106
+ return x
107
+
108
+
109
+ class Readout(nn.Module):
110
+ """Add readout token (concatenating along the spatial axis).
111
+
112
+ From: grt/deepradar/modules/position.py
113
+ """
114
+
115
+ def __init__(self, d_model: int = 512) -> None:
116
+ super().__init__()
117
+ self.readout = nn.Parameter(data=torch.normal(0, 0.02, (d_model,)))
118
+
119
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
120
+ """Concatenate readout token."""
121
+ readout = torch.tile(self.readout[None, None, :], (x.shape[0], 1, 1))
122
+ return torch.concatenate((x, readout), dim=1)
123
+
124
+
125
+ def transformer_mlp(
126
+ d_model: int = 512,
127
+ d_feedforward: int = 2048,
128
+ activation: str = "GELU",
129
+ dropout: float = 0.0,
130
+ eps: float = 1e-5,
131
+ ) -> nn.Module:
132
+ """Create transformer MLP.
133
+
134
+ From: grt/deepradar/modules/transformer.py
135
+ """
136
+ return nn.Sequential(
137
+ nn.LayerNorm(d_model, eps=eps, bias=True),
138
+ nn.Linear(d_model, d_feedforward, bias=True),
139
+ getattr(nn, activation)(),
140
+ nn.Dropout(dropout),
141
+ nn.Linear(d_feedforward, d_model, bias=True),
142
+ nn.Dropout(dropout),
143
+ )
144
+
145
+
146
+ class TransformerLayer(nn.Module):
147
+ """Single transformer (encoder) layer.
148
+
149
+ Uses PyTorch's naming convention to match checkpoint:
150
+ - self_attn (not attn)
151
+ - linear1, linear2 (not feedforward.0, feedforward.4)
152
+ - norm1, norm2 (for attention and feedforward)
153
+ """
154
+
155
+ def __init__(
156
+ self,
157
+ d_model: int = 512,
158
+ n_head: int = 8,
159
+ d_feedforward: int = 2048,
160
+ dropout: float = 0.0,
161
+ activation: str = "GELU",
162
+ ) -> None:
163
+ super().__init__()
164
+
165
+ # Attention with PyTorch naming
166
+ self.self_attn = nn.MultiheadAttention(
167
+ d_model, n_head, dropout=dropout, bias=True, batch_first=True
168
+ )
169
+ self.dropout1 = nn.Dropout(dropout)
170
+
171
+ # Feedforward with PyTorch naming
172
+ self.linear1 = nn.Linear(d_model, d_feedforward, bias=True)
173
+ self.dropout = nn.Dropout(dropout)
174
+ self.linear2 = nn.Linear(d_feedforward, d_model, bias=True)
175
+ self.dropout2 = nn.Dropout(dropout)
176
+
177
+ # Norms
178
+ self.norm1 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
179
+ self.norm2 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
180
+
181
+ # Activation
182
+ self.activation = getattr(nn, activation)()
183
+
184
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
185
+ """Apply transformer with pre-norm (norm_first=True style)."""
186
+ # Self attention block
187
+ x2 = self.norm1(x)
188
+ x2 = self.self_attn(x2, x2, x2, need_weights=False)[0]
189
+ x = x + self.dropout1(x2)
190
+
191
+ # Feedforward block
192
+ x2 = self.norm2(x)
193
+ x2 = self.linear1(x2)
194
+ x2 = self.activation(x2)
195
+ x2 = self.dropout(x2)
196
+ x2 = self.linear2(x2)
197
+ x = x + self.dropout2(x2)
198
+
199
+ return x
200
+
201
+
202
+ class TransformerDecoder(nn.Module):
203
+ """Single transformer (decoder) layer.
204
+
205
+ Uses PyTorch's naming convention to match checkpoint:
206
+ - self_attn, multihead_attn (not attn, attn2)
207
+ - linear1, linear2 (not feedforward.0, feedforward.4)
208
+ - norm1, norm2, norm3 (for self-attn, cross-attn, and feedforward)
209
+ """
210
+
211
+ def __init__(
212
+ self,
213
+ d_model: int = 512,
214
+ n_head: int = 8,
215
+ d_feedforward: int = 2048,
216
+ dropout: float = 0.0,
217
+ activation: str = "GELU",
218
+ ) -> None:
219
+ super().__init__()
220
+
221
+ # Self attention with PyTorch naming
222
+ self.self_attn = nn.MultiheadAttention(
223
+ d_model, n_head, dropout=dropout, bias=True, batch_first=True
224
+ )
225
+ self.dropout1 = nn.Dropout(dropout)
226
+
227
+ # Cross attention with PyTorch naming (multihead_attn, not attn2)
228
+ self.multihead_attn = nn.MultiheadAttention(
229
+ d_model, n_head, dropout=dropout, bias=True, batch_first=True
230
+ )
231
+ self.dropout2 = nn.Dropout(dropout)
232
+
233
+ # Feedforward with PyTorch naming
234
+ self.linear1 = nn.Linear(d_model, d_feedforward, bias=True)
235
+ self.dropout = nn.Dropout(dropout)
236
+ self.linear2 = nn.Linear(d_feedforward, d_model, bias=True)
237
+ self.dropout3 = nn.Dropout(dropout)
238
+
239
+ # Norms (note: norm2 is for cross-attention)
240
+ self.norm1 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
241
+ self.norm2 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
242
+ self.norm3 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
243
+
244
+ # Activation
245
+ self.activation = getattr(nn, activation)()
246
+
247
+ def forward(self, x: torch.Tensor, x_enc: torch.Tensor) -> torch.Tensor:
248
+ """Apply transformer decoder with pre-norm."""
249
+ # Self attention block
250
+ x2 = self.norm1(x)
251
+ x2 = self.self_attn(x2, x2, x2, need_weights=False)[0]
252
+ x = x + self.dropout1(x2)
253
+
254
+ # Cross attention block
255
+ x2 = self.norm2(x)
256
+ x2 = self.multihead_attn(x2, x_enc, x_enc, need_weights=False)[0]
257
+ x = x + self.dropout2(x2)
258
+
259
+ # Feedforward block
260
+ x2 = self.norm3(x)
261
+ x2 = self.linear1(x2)
262
+ x2 = self.activation(x2)
263
+ x2 = self.dropout(x2)
264
+ x2 = self.linear2(x2)
265
+ x = x + self.dropout3(x2)
266
+
267
+ return x
268
+
269
+
270
+ class BasisChange(nn.Module):
271
+ """Create "change-of-basis" query.
272
+
273
+ From: grt/deepradar/modules/transformer.py
274
+ """
275
+
276
+ def __init__(
277
+ self,
278
+ shape: Sequence[int] = [],
279
+ flatten: bool = True,
280
+ scale: Optional[Sequence[float]] = None,
281
+ global_scale: float = 1.0,
282
+ ) -> None:
283
+ super().__init__()
284
+
285
+ self.pos = Sinusoid(scale=scale, global_scale=global_scale)
286
+ self.shape = shape
287
+ self.flatten = flatten
288
+
289
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
290
+ """Apply change of basis."""
291
+ idxs = tuple([slice(None)] + [None] * len(self.shape) + [slice(None)])
292
+ query = self.pos(torch.tile(x[idxs], (1, *self.shape, 1)))
293
+
294
+ if self.flatten:
295
+ query = query.reshape(x.shape[0], -1, x.shape[-1])
296
+ return query
297
+
298
+
299
+ class Unpatch(nn.Module):
300
+ """Unpatch data.
301
+
302
+ Args:
303
+ output_size: output 2D shape.
304
+ features: number of input features; should be `>= size * size`.
305
+ size: patch size as (width, height, channels).
306
+ """
307
+
308
+ def __init__(
309
+ self,
310
+ output_size: Sequence[int],
311
+ features: int = 512,
312
+ size: Sequence[int] = (16, 16),
313
+ ) -> None:
314
+ super().__init__()
315
+
316
+ self.linear = nn.Linear(features, output_size[-1] * int(np.prod(size)))
317
+ self.size = size
318
+ self.output_size = output_size
319
+
320
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
321
+ """Perform 2D unpatching.
322
+
323
+ Operates in batch-spatial-feature order; spatial axes are flattened on
324
+ the input, and unflattened in the output.
325
+ """
326
+ embedding = self.linear(x)
327
+
328
+ if len(self.size) == 2:
329
+ return rearrange(
330
+ embedding,
331
+ "n (x1 x2) (s1 s2 c) -> n (x1 s1) (x2 s2) c",
332
+ x1=self.output_size[0] // self.size[0],
333
+ x2=self.output_size[1] // self.size[1],
334
+ s1=self.size[0],
335
+ s2=self.size[1],
336
+ c=self.output_size[-1],
337
+ )
338
+ elif len(self.size) == 3:
339
+ return rearrange(
340
+ embedding,
341
+ "n (x1 x2 x3) (s1 s2 s3 c) -> n (x1 s1) (x2 s2) (x3 s3) c",
342
+ x1=self.output_size[0] // self.size[0],
343
+ x2=self.output_size[1] // self.size[1],
344
+ x3=self.output_size[2] // self.size[2],
345
+ s1=self.size[0],
346
+ s2=self.size[1],
347
+ s3=self.size[2],
348
+ c=self.output_size[-1],
349
+ )
350
+ else:
351
+ raise ValueError("Unpatch is only implemented for 2D and 3D tensors.")
352
+
353
+
354
+ # ============================================================================
355
+ # GRT Model Components
356
+ # ============================================================================
357
+
358
+
359
+ class GRTEncoder(nn.Module):
360
+ """GRT Transformer Encoder matching official implementation."""
361
+
362
+ def __init__(
363
+ self,
364
+ layers: int = 4,
365
+ dim: int = 512,
366
+ ff_ratio: float = 4.0,
367
+ head_dim: int = 64,
368
+ dropout: float = 0.1,
369
+ activation: str = "GELU",
370
+ patch: list[int] = [2, 8, 2, 4],
371
+ pos_scale: list[float] = [1.0, 1.0, 1.0, 1.0],
372
+ global_scale: float = 16.0,
373
+ input_channels: int = 2,
374
+ positions: Literal["flat", "nd"] = "nd",
375
+ ):
376
+ super().__init__()
377
+
378
+ # Patch embedding
379
+ self.patch = PatchMerge(d_in=input_channels, d_out=dim, scale=patch, norm=False)
380
+
381
+ # Position embedding
382
+ self.positions = positions
383
+ self.pos = Sinusoid(scale=pos_scale, global_scale=global_scale)
384
+
385
+ # Readout token
386
+ self.readout = Readout(d_model=dim)
387
+
388
+ # Encoder layers
389
+ self.layers = nn.ModuleList(
390
+ [
391
+ TransformerLayer(
392
+ d_feedforward=int(ff_ratio * dim),
393
+ d_model=dim,
394
+ n_head=dim // head_dim,
395
+ dropout=dropout,
396
+ activation=activation,
397
+ )
398
+ for _ in range(layers)
399
+ ]
400
+ )
401
+
402
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
403
+ """Forward pass."""
404
+ # Patch embedding
405
+ embedded = self.patch(x)
406
+
407
+ # Apply positional encoding
408
+ if self.positions == "nd":
409
+ embedded = self.pos(embedded)
410
+
411
+ # Flatten spatial dimensions
412
+ flat = embedded.reshape(embedded.shape[0], -1, embedded.shape[-1])
413
+
414
+ # Apply flat positional encoding if needed
415
+ if self.positions == "flat":
416
+ flat = self.pos(flat)
417
+
418
+ # Add readout token
419
+ x = self.readout(flat)
420
+
421
+ # Apply encoder layers
422
+ for layer in self.layers:
423
+ x = layer(x)
424
+
425
+ return x
426
+
427
+
428
+ class GRTDecoder3D(nn.Module):
429
+ """GRT 3D Transformer Decoder matching official implementation."""
430
+
431
+ def __init__(
432
+ self,
433
+ key: str = "map",
434
+ layers: int = 4,
435
+ dim: int = 512,
436
+ ff_ratio: float = 4.0,
437
+ head_dim: int = 64,
438
+ dropout: float = 0.1,
439
+ activation: str = "GELU",
440
+ shape: list[int] = [64, 128, 64],
441
+ pos_scale: list[float] = [1.0, 1.0, 1.0],
442
+ global_scale: float = 16.0,
443
+ patch: list[int] = [8, 8, 8],
444
+ out_dim: int = 0,
445
+ positions: Literal["flat", "nd"] = "nd",
446
+ mode: Literal["last", "pool"] = "last",
447
+ ):
448
+ super().__init__()
449
+
450
+ self.key = key
451
+ self.out_dim = out_dim
452
+ self.mode = mode
453
+
454
+ # Decoder layers
455
+ self.layers = nn.ModuleList(
456
+ [
457
+ TransformerDecoder(
458
+ d_feedforward=int(ff_ratio * dim),
459
+ d_model=dim,
460
+ n_head=dim // head_dim,
461
+ dropout=dropout,
462
+ activation=activation,
463
+ )
464
+ for _ in range(layers)
465
+ ]
466
+ )
467
+
468
+ # Query generation with position encoding
469
+ query_shape = [s // p for s, p in zip(shape, patch)]
470
+ if positions == "flat":
471
+ query_shape = [int(np.prod(query_shape))]
472
+
473
+ self.query = BasisChange(
474
+ shape=query_shape, scale=pos_scale, global_scale=global_scale, flatten=True
475
+ )
476
+
477
+ # Unpatch to reconstruct output
478
+ self.unpatch = Unpatch(
479
+ output_size=(*shape, max(1, self.out_dim)), features=dim, size=patch
480
+ )
481
+
482
+ def forward(self, encoded: torch.Tensor) -> dict[str, torch.Tensor]:
483
+ """Forward pass."""
484
+ # Extract readout token or pool
485
+ if self.mode == "last":
486
+ x = encoded[:, -1, :]
487
+ else:
488
+ x = torch.mean(encoded, dim=1)
489
+
490
+ # Generate query with positional encoding
491
+ x = self.query(x)
492
+
493
+ # Encoded features without readout token
494
+ enc = encoded[:, :-1, :]
495
+
496
+ # Apply decoder layers
497
+ for layer in self.layers:
498
+ x = layer(x, enc)
499
+
500
+ # Unpatch to 3D output
501
+ out = self.unpatch(x)
502
+
503
+ # Squeeze channel dimension if binary output
504
+ if self.out_dim == 0:
505
+ out = out[..., 0]
506
+
507
+ return {self.key: out}
508
+
509
+
510
+ # ============================================================================
511
+ # Complete GRT-Small Model
512
+ # ============================================================================
513
+
514
+
515
+ class GRTSmall(nn.Module):
516
+ """GRT-Small model for 3D occupancy mapping.
517
+
518
+ Input: (batch, doppler, azimuth, elevation, range, 2)
519
+ - doppler: 64
520
+ - azimuth: 8
521
+ - elevation: 2
522
+ - range: 256
523
+ - channels: 2 (I/Q)
524
+
525
+ Output: (batch, elevation, azimuth, range)
526
+ - elevation: 64
527
+ - azimuth: 128
528
+ - range: 64
529
+
530
+ ~29M parameters for GRT-small variant.
531
+ """
532
+
533
+ def __init__(self):
534
+ super().__init__()
535
+
536
+ dim = 512
537
+ layers = 4
538
+
539
+ # Create encoder - stored as "tokenizer" + "encoder" in checkpoint
540
+ # But we organize logically here and handle mapping in load_checkpoint
541
+ self.tokenizer = GRTEncoder(
542
+ layers=layers,
543
+ dim=dim,
544
+ ff_ratio=4.0,
545
+ head_dim=64,
546
+ dropout=0.1,
547
+ activation="GELU",
548
+ patch=[2, 8, 2, 4],
549
+ pos_scale=[1.0, 1.0, 1.0, 1.0],
550
+ global_scale=16.0,
551
+ input_channels=2,
552
+ positions="nd",
553
+ )
554
+
555
+ # Create decoder wrapper
556
+ self.decoder = nn.Module()
557
+ self.decoder.occ3d = GRTDecoder3D(
558
+ key="map",
559
+ layers=layers,
560
+ dim=dim,
561
+ ff_ratio=4.0,
562
+ head_dim=64,
563
+ dropout=0.1,
564
+ activation="GELU",
565
+ shape=[64, 128, 64],
566
+ pos_scale=[1.0, 1.0, 1.0],
567
+ global_scale=16.0,
568
+ patch=[8, 8, 8],
569
+ out_dim=0,
570
+ positions="nd",
571
+ mode="last",
572
+ )
573
+
574
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
575
+ """Forward pass."""
576
+ # Encode
577
+ encoded = self.tokenizer(x)
578
+
579
+ # Decode
580
+ output = self.decoder.occ3d(encoded)
581
+
582
+ # Return just the occupancy map tensor
583
+ return output["map"]
584
+
585
+
src/Baselines/grt/inference.py ADDED
@@ -0,0 +1,220 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Inference Script for GRT-Small (finetuned weights)
4
+
5
+ Runs inference on all valid sequences in the configured Smoke-Eval root by
6
+ default, or on an explicit list supplied with ``--sequences``.
7
+ using weights trained by grt_finetune/train.py.
8
+ For each sequence, saves one .npy file: pred_depth.npy (dequantized predicted depth [T, 64, 128], values in [0, 1]).
9
+
10
+ Single GPU: Each frame is seen exactly once; no duplication or incompleteness.
11
+ Multi-GPU (DDP): Dataloader is sharded; each rank writes its results to a file, then
12
+ main process merges with deduplication by frame_idx (keeps first occurrence) and saves.
13
+ """
14
+
15
+ import os
16
+ import torch
17
+ import numpy as np
18
+ import argparse
19
+ import yaml
20
+ import pickle
21
+ from tqdm import tqdm
22
+ from accelerate import Accelerator
23
+ from accelerate.utils import set_seed
24
+ from collections import defaultdict
25
+ from safetensors.torch import load_file
26
+
27
+ from grt_model import GRTSmall
28
+ from dataloader import create_rice_dataloader
29
+ from augmentations import (
30
+ translate_radar,
31
+ dequantize_depth,
32
+ )
33
+
34
+ def batch_radar_to_spectrum(
35
+ radar_amplitude: torch.Tensor, radar_phase: torch.Tensor
36
+ ) -> torch.Tensor:
37
+ """Restore the GRT spectrum layout from the packaged Smoke-Eval tensors."""
38
+
39
+ amplitude = radar_amplitude.permute(0, 1, 3, 2, 4)
40
+ phase = radar_phase.permute(0, 1, 3, 2, 4)
41
+ return torch.stack((amplitude, phase), dim=-1)
42
+
43
+
44
+ def main():
45
+ parser = argparse.ArgumentParser(
46
+ description="Run GRT inference on Smoke-Eval."
47
+ )
48
+ parser.add_argument(
49
+ "--config", type=str, default="config.yaml", help="Path to config file"
50
+ )
51
+ parser.add_argument(
52
+ "--checkpoint",
53
+ type=str,
54
+ required=True,
55
+ help="Path to weights-only GRT .safetensors file",
56
+ )
57
+ parser.add_argument(
58
+ "--output_dir",
59
+ type=str,
60
+ default="inference_results",
61
+ help="Directory to save results",
62
+ )
63
+ parser.add_argument(
64
+ "--sequences",
65
+ type=str,
66
+ nargs="+",
67
+ default=None,
68
+ help="Optional sequence names; default discovers all valid sequences.",
69
+ )
70
+ parser.add_argument(
71
+ "--debug", action="store_true", help="Run in debug mode (process only 1 batch)"
72
+ )
73
+ args = parser.parse_args()
74
+
75
+ # Load config
76
+ with open(args.config, "r") as f:
77
+ config = yaml.safe_load(f)
78
+
79
+ # Initialize accelerator
80
+ accelerator = Accelerator(mixed_precision="fp16")
81
+ set_seed(config["training"].get("seed", 42))
82
+
83
+ # Create output directory (all ranks so DDP gather_dir can be created)
84
+ os.makedirs(args.output_dir, exist_ok=True)
85
+
86
+ # Create model
87
+ accelerator.print("Creating GRT-Small model...")
88
+ model = GRTSmall()
89
+
90
+ # Safetensors files contain only the model state dictionary.
91
+ accelerator.print(f"Loading checkpoint from {args.checkpoint}")
92
+ model.load_state_dict(load_file(args.checkpoint, device="cpu"), strict=True)
93
+
94
+ # With ``sequences=None`` the public dataset loader discovers every valid
95
+ # sequence under the configured Smoke-Eval root.
96
+ accelerator.print(f"Inference sequences: {args.sequences}")
97
+ inference_loader = create_rice_dataloader(
98
+ root_dir=config["paths"]["data_root"],
99
+ batch_size=config["training"]["batch_size"],
100
+ num_workers=0,
101
+ frame_skip=1,
102
+ sequences=args.sequences,
103
+ shuffle=False,
104
+ )
105
+
106
+ # Prepare model and dataloader
107
+ model, inference_loader = accelerator.prepare(model, inference_loader)
108
+ model.eval()
109
+
110
+ # Dictionary to aggregate results by sequence: sequence_id -> list of (frame_idx, pred_depth)
111
+ results_by_sequence = defaultdict(list)
112
+
113
+ accelerator.print("Starting inference...")
114
+
115
+ with torch.no_grad():
116
+ for batch in tqdm(
117
+ inference_loader, disable=not accelerator.is_local_main_process
118
+ ):
119
+ # Extract data
120
+ rsp_data = batch_radar_to_spectrum(
121
+ batch["radar_amplitude"], batch["radar_phase"]
122
+ )
123
+ sequences = batch["sequence"]
124
+ frame_indices = batch["frame_idx"]
125
+
126
+ # Apply radar augmentation
127
+ rsp_data = translate_radar(rsp_data)
128
+
129
+ # Forward pass
130
+ occupancy_pred_logits = model(rsp_data) # [B, 64, 128, 64]
131
+
132
+ # Dequantize predicted occupancy to depth [B, 1, 64, 128], values in [0, 1]
133
+ pred_depth = dequantize_depth(occupancy_pred_logits)
134
+ pred_depth_np = (
135
+ pred_depth.cpu().numpy().astype(np.float32)
136
+ ) # [B, 1, 64, 128]
137
+
138
+ # Collect results (frame_idx, pred_depth per sample)
139
+ for i in range(len(sequences)):
140
+ seq_id = sequences[i]
141
+ f_idx = frame_indices[i].item()
142
+ # Store [1, 64, 128] per frame; will stack to [T, 64, 128] when saving
143
+ results_by_sequence[seq_id].append(
144
+ {
145
+ "frame_idx": f_idx,
146
+ "pred_depth": pred_depth_np[i],
147
+ }
148
+ )
149
+
150
+ if args.debug:
151
+ break
152
+
153
+ # Single GPU: save directly (each frame seen once, no duplication)
154
+ # Multi-GPU: gather via files, merge with dedupe by frame_idx, then save
155
+ if accelerator.num_processes == 1:
156
+ if accelerator.is_main_process:
157
+ accelerator.print("Saving results (single process)...")
158
+ for seq_id, frames in tqdm(
159
+ results_by_sequence.items(), desc="Saving sequences"
160
+ ):
161
+ frames.sort(key=lambda x: x["frame_idx"])
162
+ pred_depth_stack = np.stack([f["pred_depth"] for f in frames], axis=0)
163
+ pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 64, 128]
164
+ np.save(
165
+ os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"),
166
+ pred_depth_stack,
167
+ )
168
+ accelerator.print(
169
+ f" {seq_id}: saved {pred_depth_stack.shape[0]} frames"
170
+ )
171
+ accelerator.print(f"Processed {len(results_by_sequence)} sequences.")
172
+ accelerator.print(f"Results saved to {args.output_dir}")
173
+ else:
174
+ # DDP: gather results from all ranks via files, dedupe by frame_idx, save on main
175
+ accelerator.wait_for_everyone()
176
+ gather_dir = os.path.join(args.output_dir, "_gather")
177
+ os.makedirs(gather_dir, exist_ok=True)
178
+ rank = accelerator.process_index
179
+ rank_file = os.path.join(gather_dir, f"rank_{rank}_results.pkl")
180
+ with open(rank_file, "wb") as f:
181
+ pickle.dump(dict(results_by_sequence), f, protocol=pickle.HIGHEST_PROTOCOL)
182
+ accelerator.wait_for_everyone()
183
+
184
+ if accelerator.is_main_process:
185
+ accelerator.print("Merging and deduplicating results from all ranks...")
186
+ merged_results = defaultdict(dict) # seq_id -> {frame_idx: pred_depth}
187
+ for r in range(accelerator.num_processes):
188
+ pkl_path = os.path.join(gather_dir, f"rank_{r}_results.pkl")
189
+ with open(pkl_path, "rb") as f:
190
+ rank_results = pickle.load(f)
191
+ for seq_id, frames in rank_results.items():
192
+ for frame_data in frames:
193
+ f_idx = frame_data["frame_idx"]
194
+ if f_idx not in merged_results[seq_id]:
195
+ merged_results[seq_id][f_idx] = frame_data["pred_depth"]
196
+ os.remove(pkl_path)
197
+
198
+ for seq_id, frame_dict in tqdm(
199
+ merged_results.items(), desc="Saving sequences"
200
+ ):
201
+ sorted_items = sorted(frame_dict.items(), key=lambda x: x[0])
202
+ pred_depth_stack = np.stack([item[1] for item in sorted_items], axis=0)
203
+ pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 64, 128]
204
+ np.save(
205
+ os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"),
206
+ pred_depth_stack,
207
+ )
208
+ accelerator.print(
209
+ f" {seq_id}: saved {pred_depth_stack.shape[0]} frames"
210
+ )
211
+ if os.path.isdir(gather_dir) and not os.listdir(gather_dir):
212
+ os.rmdir(gather_dir)
213
+ accelerator.print(f"Processed {len(merged_results)} sequences.")
214
+ accelerator.print(f"Results saved to {args.output_dir}")
215
+
216
+ accelerator.wait_for_everyone()
217
+
218
+
219
+ if __name__ == "__main__":
220
+ main()
src/Baselines/grt/split.json ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "test": [
3
+ "Dell-1",
4
+ "Dell-2",
5
+ "Smoke-Dell-1",
6
+ "Smoke-Dell-2",
7
+ "brk-2",
8
+ "brk-3",
9
+ "Brk-b",
10
+ "brk-basement",
11
+ "Brk-stair",
12
+ "Smoke-brk-2",
13
+ "Smoke-brk-3",
14
+ "Smoke-brk-b"
15
+ ]
16
+ }
src/Baselines/grt_image/augmentations.py ADDED
@@ -0,0 +1,193 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torchvision.transforms.functional as TF
3
+ from torchvision.transforms import Resize, InterpolationMode
4
+ from typing import Union
5
+ import numpy as np
6
+
7
+ AZIMUTH_RESOLUTION = 256
8
+ ELEVATION_RESOLUTION = 128
9
+
10
+ # Depth output resolution: height=128, width=256
11
+ DEPTH_TARGET_HEIGHT = 128
12
+ DEPTH_TARGET_WIDTH = 256
13
+
14
+ resize_transform = Resize(
15
+ size=[ELEVATION_RESOLUTION, AZIMUTH_RESOLUTION],
16
+ interpolation=InterpolationMode.BILINEAR,
17
+ antialias=True,
18
+ )
19
+
20
+ depth_resize_transform = Resize(
21
+ size=(DEPTH_TARGET_HEIGHT, DEPTH_TARGET_WIDTH),
22
+ interpolation=InterpolationMode.BILINEAR,
23
+ antialias=True,
24
+ )
25
+
26
+
27
+ def translate_radar(radar_data):
28
+ """
29
+ Applies normalization to radar data after batching from dataloader.
30
+ Called before passing data into the model.
31
+
32
+ Args:
33
+ radar_data: Batched radar tensor from dataloader
34
+ Shape: [B, 64, 8, 2, 256, 2] (batch, doppler, azimuth, elevation, range, channels)
35
+ - Channel 0: raw amplitude values
36
+ - Channel 1: phase normalized to [-1, 1] (divided by π)
37
+
38
+ Returns:
39
+ Processed radar tensor with same shape [B, 64, 8, 2, 256, 2]
40
+ - Channel 0: sqrt(amplitude * 1e-3) for magnitude normalization
41
+ - Channel 1: phase * π (converted back to radians [-π, π])
42
+ """
43
+ radar_mag = radar_data[..., 0] # [B, 64, 8, 2, 256] - Extract raw amplitude
44
+ radar_phase = radar_data[..., 1] # [B, 64, 8, 2, 256] - Extract normalized phase
45
+
46
+ # Normalize amplitude: scale then sqrt
47
+ radar_mag_processed = torch.sqrt(radar_mag * 1e-6)
48
+
49
+ # Convert phase back to radians: [-1, 1] -> [-π, π]
50
+ radar_phase_processed = radar_phase * torch.pi
51
+
52
+ # Stack channels back together: [B, 64, 8, 2, 256, 2]
53
+ radar_data_translated = torch.stack(
54
+ [radar_mag_processed, radar_phase_processed], dim=-1
55
+ )
56
+ return radar_data_translated
57
+
58
+
59
+ def resize_depth(
60
+ depth_map: Union[torch.Tensor, np.ndarray],
61
+ ) -> Union[torch.Tensor, np.ndarray]:
62
+ """
63
+ Process depth map from dataloader (same pipeline as denoiser/control crop_depth):
64
+ mm -> meters, clamp [0, 11.2] m, normalize to [0, 1], resize to (128, 256) (h, w).
65
+
66
+ Args:
67
+ depth_map: Depth in millimeters. Torch or numpy.
68
+ Shapes: (H, W), (B, H, W), or (B, 1, H, W).
69
+
70
+ Returns:
71
+ Depth in [0, 1], spatial size (128, 256). Shape [B, 128, 256] for batched input.
72
+ """
73
+ is_numpy = isinstance(depth_map, np.ndarray)
74
+ if is_numpy:
75
+ depth_map = torch.from_numpy(depth_map)
76
+
77
+ depth_map = depth_map.float()
78
+ original_shape = depth_map.shape
79
+
80
+ if depth_map.dim() == 2:
81
+ depth_map = depth_map.unsqueeze(0) # (H, W) -> (1, H, W)
82
+ elif depth_map.dim() == 3:
83
+ depth_map = depth_map.unsqueeze(1) # (B, H, W) -> (B, 1, H, W)
84
+ elif depth_map.dim() != 4:
85
+ raise ValueError(f"Unexpected depth shape: {original_shape}")
86
+
87
+ invalid_mask = ~(torch.isfinite(depth_map) & (depth_map >= 0))
88
+ depth_map[invalid_mask] = 0.0
89
+
90
+ depth_map = depth_map / 1000.0 # mm -> meters
91
+ max_depth_m = 11.2
92
+ depth_map = torch.clamp(depth_map, min=0.0, max=max_depth_m)
93
+ depth_map = depth_map / max_depth_m # [0, 1]
94
+
95
+ invalid_mask = ~torch.isfinite(depth_map)
96
+ depth_map[invalid_mask] = 0.0
97
+
98
+ depth_map = depth_resize_transform(depth_map) # (..., 128, 256)
99
+ depth_values = depth_map.squeeze(1) # [B, 128, 256] or [1, 128, 256]
100
+
101
+ if len(original_shape) == 2:
102
+ depth_values = depth_values.squeeze(0) # (128, 256)
103
+
104
+ if is_numpy:
105
+ depth_values = depth_values.numpy()
106
+ return depth_values
107
+
108
+
109
+ def quantize_depth_to_occupancy(depth_values, num_range_bins=64):
110
+ """
111
+ Quantizes 2D depth values into 3D binary occupancy grid.
112
+
113
+ Args:
114
+ depth_values: Resized depth tensor
115
+ Shape: [B, elevation, azimuth]
116
+ Values: normalized to [0, 1] range
117
+ num_range_bins: Number of range bins for quantization (default: 64)
118
+
119
+ Returns:
120
+ Binary 3D occupancy grid
121
+ Shape: [B, elevation, azimuth, num_range_bins]
122
+ Values: binary (0 or 1) indicating occupied bins
123
+ """
124
+ B, elevation, azimuth = depth_values.shape
125
+
126
+ # Quantize normalized depth [0, 1] directly to range bins [0, num_range_bins-1]
127
+ # Each bin represents 1/num_range_bins of the normalized depth range
128
+ bin_indices = torch.floor(
129
+ depth_values / (1.0 / num_range_bins)
130
+ ).long() # [B, elevation, azimuth]
131
+ bin_indices = torch.clamp(
132
+ bin_indices, 0, num_range_bins - 1
133
+ ) # Handle edge case where depth_values = 1.0
134
+
135
+ # Create binary 3D occupancy grid
136
+ occupancy_grid = torch.zeros(
137
+ B,
138
+ elevation,
139
+ azimuth,
140
+ num_range_bins,
141
+ dtype=torch.float32,
142
+ device=depth_values.device,
143
+ ) # [B, elevation, azimuth, num_range_bins]
144
+
145
+ # Set occupied bins to 1
146
+ # Use advanced indexing to mark the appropriate range bin for each (elevation, azimuth) cell
147
+ batch_idx = torch.arange(B, device=depth_values.device)[:, None, None].expand(
148
+ B, elevation, azimuth
149
+ )
150
+ elevation_idx = torch.arange(elevation, device=depth_values.device)[
151
+ None, :, None
152
+ ].expand(B, elevation, azimuth)
153
+ azimuth_idx = torch.arange(azimuth, device=depth_values.device)[
154
+ None, None, :
155
+ ].expand(B, elevation, azimuth)
156
+
157
+ occupancy_grid[batch_idx, elevation_idx, azimuth_idx, bin_indices] = 1.0
158
+
159
+ return occupancy_grid # [B, elevation, azimuth, num_range_bins]
160
+
161
+
162
+ def dequantize_depth(occupancy_grid):
163
+ """
164
+ Converts 3D binary occupancy grid back to 2D depth map.
165
+ This is the inverse operation of quantize_depth_to_occupancy.
166
+
167
+ Args:
168
+ occupancy_grid: Binary 3D occupancy grid
169
+ Shape: [B, 128, 256, 64] (batch, elevation, azimuth, range)
170
+ Values: binary (0 or 1) or continuous (predicted probabilities)
171
+
172
+ Returns:
173
+ Reconstructed depth map
174
+ Shape: [B, 1, 128, 256] (batch, channel, elevation, azimuth)
175
+ Values: normalized to [0, 1] range
176
+ """
177
+ num_range_bins = occupancy_grid.shape[3]
178
+
179
+ # Find the range bin with maximum value for each (elevation, azimuth) cell
180
+ # For binary: finds the occupied bin
181
+ # For continuous: finds the most likely bin
182
+ bin_indices = torch.argmax(occupancy_grid, dim=3) # [B, 128, 256]
183
+
184
+ # Convert bin indices back to normalized depth values [0, 1]
185
+ # Use bin center: (bin_idx + 0.5) / num_bins
186
+ depth_values = (bin_indices.float() + 1) / num_range_bins # [B, 128, 256]
187
+
188
+ # Add channel dimension: [B, 128, 256] -> [B, 1, 128, 256]
189
+ depth_map = depth_values.unsqueeze(1) # [B, 1, 128, 256]
190
+
191
+ return depth_map
192
+
193
+
src/Baselines/grt_image/dataloader.py ADDED
@@ -0,0 +1,344 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ Dataloader for MobiCom processed dataset (output of processor.py).
3
+
4
+ Uses the optimized format produced by processor.py:
5
+ - radar.npy: (N, doppler, elevation, azimuth, range) complex64
6
+ - dji_rgb.npy: (N, H, W, 3) uint8
7
+ - zed_depth.npy: (N, H, W) uint16, depth in millimeters
8
+
9
+ This module provides:
10
+ - `RiceDataset`: frame-level dataset returning radar amplitude/phase, DJI RGB,
11
+ and ZED depth (ground truth).
12
+ - `create_rice_dataloader`: generic dataloader for an arbitrary set of sequences.
13
+ - `create_train_val_test_loaders`: uses the configured split file for fixed
14
+ validation sequences and a separate Smoke-Eval root for testing.
15
+ """
16
+
17
+ import json
18
+ from pathlib import Path
19
+ from typing import Dict, List, Optional, Tuple
20
+
21
+ import numpy as np
22
+ import torch
23
+ import torch.nn.functional as F
24
+ from torch.utils.data import Dataset, DataLoader
25
+
26
+
27
+ class RiceDataset(Dataset):
28
+ """
29
+ Dataset for processor.py output: radar, DJI RGB, and ZED depth per frame.
30
+
31
+ Args:
32
+ root_dir: Root directory containing sequence subdirs (e.g. processed/),
33
+ each with radar.npy, dji_rgb.npy, zed_depth.npy.
34
+ sequences: Optional list of sequence names to load. If None, loads all
35
+ subdirs that contain the three required files.
36
+ frame_skip: Sample every frame_skip frames (1 = all frames).
37
+ return_radar_complex: If True, return radar as complex tensor; if False,
38
+ return radar_amplitude and radar_phase as separate float tensors.
39
+ depth_in_meters: If True, convert depth from mm to meters.
40
+ rgb_normalize: If True, return RGB in [0, 1] float; else uint8 [0, 255].
41
+ """
42
+
43
+ REQUIRED_FILES = ("radar.npy", "dji_rgb.npy", "zed_depth.npy")
44
+
45
+ def __init__(
46
+ self,
47
+ root_dir: str,
48
+ sequences: Optional[List[str]] = None,
49
+ frame_skip: int = 1,
50
+ return_radar_complex: bool = False,
51
+ depth_in_meters: bool = True,
52
+ rgb_normalize: bool = True,
53
+ image_height: int = 288,
54
+ image_width: int = 512,
55
+ ):
56
+ self.root_dir = Path(root_dir)
57
+ self.frame_skip = max(1, frame_skip)
58
+ self.return_radar_complex = return_radar_complex
59
+ self.depth_in_meters = depth_in_meters
60
+ self.rgb_normalize = rgb_normalize
61
+ self.image_height = int(image_height)
62
+ self.image_width = int(image_width)
63
+ if self.image_height <= 0 or self.image_width <= 0:
64
+ raise ValueError("image_height and image_width must be positive")
65
+
66
+ self.sequences = self._discover_sequences(sequences)
67
+ self.index_map: List[Tuple[str, int]] = [] # (seq_name, frame_idx)
68
+ self._seq_arrays: Dict[str, Dict] = {} # seq -> {radar, depth, dji_rgb}
69
+
70
+ self._build_index()
71
+
72
+ def _discover_sequences(self, sequences: Optional[List[str]] = None) -> List[str]:
73
+ """Return list of sequence names that have all required files."""
74
+ if not self.root_dir.is_dir():
75
+ raise FileNotFoundError(f"Root directory not found: {self.root_dir}")
76
+
77
+ all_seqs = sorted(
78
+ d.name
79
+ for d in self.root_dir.iterdir()
80
+ if d.is_dir() and not d.name.startswith(".")
81
+ )
82
+ valid = []
83
+ for name in all_seqs:
84
+ seq_dir = self.root_dir / name
85
+ if all((seq_dir / f).exists() for f in self.REQUIRED_FILES):
86
+ valid.append(name)
87
+ if sequences is not None:
88
+ valid = [s for s in valid if s in sequences]
89
+ return valid
90
+
91
+ def _build_index(self) -> None:
92
+ """Build (seq_name, frame_idx) index, using radar.npy for frame count."""
93
+ self.index_map.clear()
94
+ for seq_name in self.sequences:
95
+ seq_dir = self.root_dir / seq_name
96
+ radar_path = seq_dir / "radar.npy"
97
+ arrays = self._load_sequence_arrays(seq_name)
98
+ n_frames = min(array.shape[0] for array in arrays.values())
99
+ for i in range(0, n_frames, self.frame_skip):
100
+ self.index_map.append((seq_name, i))
101
+
102
+ def _load_sequence_arrays(self, seq_name: str) -> Dict:
103
+ """Lazy-load or return cached arrays for a sequence."""
104
+ if seq_name not in self._seq_arrays:
105
+ seq_dir = self.root_dir / seq_name
106
+ self._seq_arrays[seq_name] = {
107
+ "radar": np.load(seq_dir / "radar.npy", mmap_mode="r"),
108
+ "rgb": np.load(seq_dir / "dji_rgb.npy", mmap_mode="r"),
109
+ "depth": np.load(seq_dir / "zed_depth.npy", mmap_mode="r"),
110
+ }
111
+ return self._seq_arrays[seq_name]
112
+
113
+ def __len__(self) -> int:
114
+ return len(self.index_map)
115
+
116
+ def __getitem__(self, idx: int) -> Dict[str, torch.Tensor]:
117
+ seq_name, frame_idx = self.index_map[idx]
118
+ arrs = self._load_sequence_arrays(seq_name)
119
+
120
+ rgb = np.asarray(arrs["rgb"][frame_idx]).copy()
121
+ if rgb.ndim != 3 or rgb.shape[-1] != 3:
122
+ raise ValueError(f"Expected RGB frame shaped [H, W, 3], got {rgb.shape}")
123
+ # (H, W) uint16 mm (processor saves as uint16)
124
+ depth = np.asarray(arrs["depth"][frame_idx]).astype(np.float32)
125
+ # (doppler, elevation, azimuth, range) complex64
126
+ radar = np.asarray(arrs["radar"][frame_idx]).copy()
127
+
128
+ # Depth: uint16 mm -> float; optional mm -> m; handle invalid
129
+ if self.depth_in_meters:
130
+ depth = depth / 1000.0
131
+ invalid = ~(np.isfinite(depth) & (depth > 0))
132
+ depth[invalid] = 0.0
133
+ depth = depth[np.newaxis, ...] # (1, H, W)
134
+
135
+ # RGB: [H, W, 3] uint8 -> resized [3, image_height, image_width] float.
136
+ image = torch.from_numpy(np.transpose(rgb, (2, 0, 1)).copy()).float()
137
+ if self.rgb_normalize:
138
+ image = image / 255.0
139
+ image = F.interpolate(
140
+ image.unsqueeze(0),
141
+ size=(self.image_height, self.image_width),
142
+ mode="bilinear",
143
+ align_corners=False,
144
+ ).squeeze(0)
145
+
146
+ # Radar: amplitude and phase
147
+ radar_amplitude = np.abs(radar).astype(np.float32)
148
+ radar_phase = np.angle(radar).astype(np.float32) / np.pi
149
+ out = {
150
+ "radar_amplitude": torch.from_numpy(radar_amplitude),
151
+ "radar_phase": torch.from_numpy(radar_phase),
152
+ "image": image,
153
+ "depth": torch.from_numpy(depth),
154
+ "sequence": seq_name,
155
+ "frame_idx": frame_idx,
156
+ }
157
+ if self.return_radar_complex:
158
+ out["radar_cube"] = torch.from_numpy(radar.copy())
159
+ # Depth in mm for optional use (1, H, W) float32
160
+ depth_mm = np.asarray(arrs["depth"][frame_idx]).astype(np.float32)
161
+ out["depth_mm"] = torch.from_numpy(depth_mm[np.newaxis, ...])
162
+ return out
163
+
164
+
165
+ def create_rice_dataloader(
166
+ root_dir: str,
167
+ batch_size: int = 8,
168
+ num_workers: int = 0,
169
+ frame_skip: int = 1,
170
+ sequences: Optional[List[str]] = None,
171
+ return_radar_complex: bool = False,
172
+ depth_in_meters: bool = True,
173
+ rgb_normalize: bool = True,
174
+ image_height: int = 288,
175
+ image_width: int = 512,
176
+ shuffle: bool = True,
177
+ ) -> DataLoader:
178
+ """Create a DataLoader for the Rice (processor output) dataset."""
179
+ dataset = RiceDataset(
180
+ root_dir=root_dir,
181
+ sequences=sequences,
182
+ frame_skip=frame_skip,
183
+ return_radar_complex=return_radar_complex,
184
+ depth_in_meters=depth_in_meters,
185
+ rgb_normalize=rgb_normalize,
186
+ image_height=image_height,
187
+ image_width=image_width,
188
+ )
189
+ return DataLoader(
190
+ dataset,
191
+ batch_size=batch_size,
192
+ shuffle=shuffle,
193
+ num_workers=num_workers,
194
+ pin_memory=True,
195
+ )
196
+
197
+
198
+ def create_train_val_test_loaders(
199
+ train_root: str,
200
+ split_json_path: Optional[str],
201
+ test_root: str,
202
+ batch_size: int = 8,
203
+ num_workers: int = 0,
204
+ frame_skip: int = 1,
205
+ return_radar_complex: bool = False,
206
+ depth_in_meters: bool = True,
207
+ rgb_normalize: bool = True,
208
+ image_height: int = 288,
209
+ image_width: int = 512,
210
+ ) -> Tuple[DataLoader, DataLoader, DataLoader]:
211
+ """Create fixed training/validation and Smoke-Eval test loaders.
212
+
213
+ The ``test`` list in the configured split file is treated as a fixed
214
+ validation sequence list. All other valid training sequences are used
215
+ for training. ``test_root`` is a separately structured Smoke-Eval tree;
216
+ every valid sequence it contains is evaluated only as the test set.
217
+ """
218
+ if split_json_path is None:
219
+ split_path = Path(__file__).resolve().parent / "split.json"
220
+ else:
221
+ split_path = Path(split_json_path)
222
+ if not split_path.exists() and not split_path.is_absolute():
223
+ fallback = Path(__file__).resolve().parent / split_path.name
224
+ if fallback.exists():
225
+ split_path = fallback
226
+
227
+ with split_path.open("r") as f:
228
+ split = json.load(f)
229
+ validation_sequences = split.get("test", [])
230
+
231
+ discovered_train = RiceDataset(
232
+ root_dir=train_root,
233
+ frame_skip=frame_skip,
234
+ return_radar_complex=return_radar_complex,
235
+ depth_in_meters=depth_in_meters,
236
+ rgb_normalize=rgb_normalize,
237
+ image_height=image_height,
238
+ image_width=image_width,
239
+ )
240
+ validation_set = set(validation_sequences)
241
+ train_sequences = [
242
+ sequence
243
+ for sequence in discovered_train.sequences
244
+ if sequence not in validation_set
245
+ ]
246
+ resolved_validation_sequences = [
247
+ sequence
248
+ for sequence in validation_sequences
249
+ if sequence in discovered_train.sequences
250
+ ]
251
+
252
+ dataset_kwargs = {
253
+ "frame_skip": frame_skip,
254
+ "return_radar_complex": return_radar_complex,
255
+ "depth_in_meters": depth_in_meters,
256
+ "rgb_normalize": rgb_normalize,
257
+ "image_height": image_height,
258
+ "image_width": image_width,
259
+ }
260
+ train_dataset = RiceDataset(
261
+ root_dir=train_root, sequences=train_sequences, **dataset_kwargs
262
+ )
263
+ val_dataset = RiceDataset(
264
+ root_dir=train_root,
265
+ sequences=resolved_validation_sequences,
266
+ **dataset_kwargs,
267
+ )
268
+ test_dataset = RiceDataset(root_dir=test_root, sequences=None, **dataset_kwargs)
269
+
270
+ loader_kwargs = {"batch_size": batch_size, "num_workers": num_workers, "pin_memory": True}
271
+ train_loader = DataLoader(train_dataset, shuffle=True, **loader_kwargs)
272
+ val_loader = DataLoader(val_dataset, shuffle=False, **loader_kwargs)
273
+ test_loader = DataLoader(test_dataset, shuffle=False, **loader_kwargs)
274
+ return train_loader, val_loader, test_loader
275
+
276
+
277
+ def create_train_val_loaders(
278
+ train_root: str,
279
+ split_json_path: Optional[str],
280
+ batch_size: int = 8,
281
+ num_workers: int = 0,
282
+ frame_skip: int = 1,
283
+ return_radar_complex: bool = False,
284
+ depth_in_meters: bool = True,
285
+ rgb_normalize: bool = True,
286
+ image_height: int = 288,
287
+ image_width: int = 512,
288
+ ) -> Tuple[DataLoader, DataLoader]:
289
+ """Create training and fixed validation loaders only."""
290
+ if split_json_path is None:
291
+ split_path = Path(__file__).resolve().parent / "split.json"
292
+ else:
293
+ split_path = Path(split_json_path)
294
+ if not split_path.exists() and not split_path.is_absolute():
295
+ fallback = Path(__file__).resolve().parent / split_path.name
296
+ if fallback.exists():
297
+ split_path = fallback
298
+
299
+ with split_path.open("r") as f:
300
+ split = json.load(f)
301
+ validation_sequences = split.get("test", [])
302
+
303
+ discovered = RiceDataset(
304
+ root_dir=train_root,
305
+ frame_skip=frame_skip,
306
+ return_radar_complex=return_radar_complex,
307
+ depth_in_meters=depth_in_meters,
308
+ rgb_normalize=rgb_normalize,
309
+ image_height=image_height,
310
+ image_width=image_width,
311
+ )
312
+ validation_set = set(validation_sequences)
313
+ train_sequences = [
314
+ sequence for sequence in discovered.sequences if sequence not in validation_set
315
+ ]
316
+ resolved_validation_sequences = [
317
+ sequence for sequence in validation_sequences if sequence in discovered.sequences
318
+ ]
319
+
320
+ dataset_kwargs = {
321
+ "frame_skip": frame_skip,
322
+ "return_radar_complex": return_radar_complex,
323
+ "depth_in_meters": depth_in_meters,
324
+ "rgb_normalize": rgb_normalize,
325
+ "image_height": image_height,
326
+ "image_width": image_width,
327
+ }
328
+ train_dataset = RiceDataset(
329
+ root_dir=train_root, sequences=train_sequences, **dataset_kwargs
330
+ )
331
+ val_dataset = RiceDataset(
332
+ root_dir=train_root,
333
+ sequences=resolved_validation_sequences,
334
+ **dataset_kwargs,
335
+ )
336
+ loader_kwargs = {
337
+ "batch_size": batch_size,
338
+ "num_workers": num_workers,
339
+ "pin_memory": True,
340
+ }
341
+ return (
342
+ DataLoader(train_dataset, shuffle=True, **loader_kwargs),
343
+ DataLoader(val_dataset, shuffle=False, **loader_kwargs),
344
+ )
src/Baselines/grt_image/grt_image_resnet_inference.example.yaml ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Anonymous, release-relative configuration for the paper's GRT+Image baseline.
2
+ # This is the ResNet-18 implementation in Baselines/grt_image.
3
+ paths:
4
+ smoke_eval_root: ../../../evaluation_dataset/Smoke-Eval
5
+
6
+ training:
7
+ batch_size: 1
8
+ mixed_precision: fp16
9
+ seed: 42
10
+
11
+ data:
12
+ image_height: 288
13
+ image_width: 512
14
+
15
+ model:
16
+ resnet18_pretrained: true
src/Baselines/grt_image/grt_model.py ADDED
@@ -0,0 +1,799 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """GRT-Small Model - from official codebase.
2
+
3
+ This implementation directly copies necessary modules from the official GRT codebase
4
+ (grt/deepradar/modules).
5
+ """
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ from torchvision.models import ResNet18_Weights, resnet18
10
+ from typing import Literal, Optional, Sequence
11
+ import numpy as np
12
+ from einops import rearrange
13
+ from safetensors.torch import load_file
14
+
15
+ # ============================================================================
16
+ # Official GRT Modules (copied from grt/deepradar/modules/*.py)
17
+ # ============================================================================
18
+
19
+
20
+ class PatchMerge(nn.Module):
21
+ """Merge patches with normalization and nominally reduced projection.
22
+
23
+ From: grt/deepradar/modules/patch.py
24
+ """
25
+
26
+ def __init__(
27
+ self, d_in: int, d_out: int, scale: Sequence[int] = [], norm: bool = True
28
+ ) -> None:
29
+ super().__init__()
30
+
31
+ self.scale = scale
32
+ d_merge = d_in * int(np.prod(scale))
33
+ self.linear = nn.Linear(d_merge, d_out, bias=False)
34
+ self.norm = nn.LayerNorm(d_merge) if norm else None
35
+
36
+ def _merge(self, x: torch.Tensor) -> torch.Tensor:
37
+ """Perform patch merging."""
38
+ n, *t, c = x.shape
39
+ dims = sum(([d // s, s] for d, s in zip(t, self.scale)), start=[n])
40
+ order = (
41
+ [0]
42
+ + [2 * i + 1 for i in range(len(self.scale))]
43
+ + [2 * i + 2 for i in range(len(self.scale))]
44
+ + [-1]
45
+ )
46
+ t2 = [d // s for d, s in zip(t, self.scale)]
47
+ return x.reshape(dims + [c]).permute(order).reshape(n, *t2, -1)
48
+
49
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
50
+ """Merge and project."""
51
+ merged = self._merge(x)
52
+ if self.norm is not None:
53
+ merged = self.norm(merged)
54
+ return self.linear(merged)
55
+
56
+
57
+ class Sinusoid(nn.Module):
58
+ """Centered N-dimensional sinusoidal positional embedding.
59
+
60
+ From: grt/deepradar/modules/position.py
61
+ """
62
+
63
+ def __init__(
64
+ self,
65
+ scale: Optional[Sequence[float]] = None,
66
+ global_scale: float = 1.0,
67
+ coef: float = 10000.0,
68
+ ) -> None:
69
+ super().__init__()
70
+ if scale is None:
71
+ self.scale = [global_scale]
72
+ else:
73
+ self.scale = [s * global_scale for s in scale]
74
+ self.coef = coef
75
+
76
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
77
+ """Apply sinusoidal embedding."""
78
+ # w = coef ** (-i / c)
79
+ nd = len(x.shape) - 2
80
+ c = x.shape[-1] // 2 // nd
81
+ i = torch.arange(c, device=x.device)
82
+ w = self.coef ** (-i / c)
83
+
84
+ start_dim = 0
85
+ for axis, (d, scale) in enumerate(zip(x.shape[1:-1], self.scale * nd)):
86
+ # t = scale * (j - d/2) / (d/2) = scale * (2j / d - 1)
87
+ t = scale * (2 * (torch.arange(d, device=x.device) + 0.5) / d - 1)
88
+ wt = t[:, None] * w[None, :]
89
+
90
+ p_slice = [None] * (len(x.shape) - 1) + [slice(None)]
91
+ p_slice[axis + 1] = slice(None)
92
+
93
+ # pos[2 * i] = sin(w * t)
94
+ x_sin_slice = [slice(None)] * len(x.shape)
95
+ x_sin_slice[-1] = slice(start_dim, start_dim + c * 2, 2)
96
+ x_sin_slice = tuple(x_sin_slice)
97
+ p_slice_tuple = tuple(p_slice)
98
+ x[x_sin_slice] = x[x_sin_slice] + torch.sin(wt)[p_slice_tuple]
99
+
100
+ # pos[2 * i + 1] = cos(w * t)
101
+ x_cos_slice = [slice(None)] * len(x.shape)
102
+ x_cos_slice[-1] = slice(start_dim + 1, start_dim + c * 2 + 1, 2)
103
+ x_cos_slice = tuple(x_cos_slice)
104
+ x[x_cos_slice] = x[x_cos_slice] + torch.cos(wt)[p_slice_tuple]
105
+
106
+ start_dim += c * 2
107
+
108
+ return x
109
+
110
+
111
+ class Readout(nn.Module):
112
+ """Add readout token (concatenating along the spatial axis).
113
+
114
+ From: grt/deepradar/modules/position.py
115
+ """
116
+
117
+ def __init__(self, d_model: int = 512) -> None:
118
+ super().__init__()
119
+ self.readout = nn.Parameter(data=torch.normal(0, 0.02, (d_model,)))
120
+
121
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
122
+ """Concatenate readout token."""
123
+ readout = torch.tile(self.readout[None, None, :], (x.shape[0], 1, 1))
124
+ return torch.concatenate((x, readout), dim=1)
125
+
126
+
127
+ def transformer_mlp(
128
+ d_model: int = 512,
129
+ d_feedforward: int = 2048,
130
+ activation: str = "GELU",
131
+ dropout: float = 0.0,
132
+ eps: float = 1e-5,
133
+ ) -> nn.Module:
134
+ """Create transformer MLP.
135
+
136
+ From: grt/deepradar/modules/transformer.py
137
+ """
138
+ return nn.Sequential(
139
+ nn.LayerNorm(d_model, eps=eps, bias=True),
140
+ nn.Linear(d_model, d_feedforward, bias=True),
141
+ getattr(nn, activation)(),
142
+ nn.Dropout(dropout),
143
+ nn.Linear(d_feedforward, d_model, bias=True),
144
+ nn.Dropout(dropout),
145
+ )
146
+
147
+
148
+ class TransformerLayer(nn.Module):
149
+ """Single transformer (encoder) layer.
150
+
151
+ Uses PyTorch's naming convention to match checkpoint:
152
+ - self_attn (not attn)
153
+ - linear1, linear2 (not feedforward.0, feedforward.4)
154
+ - norm1, norm2 (for attention and feedforward)
155
+ """
156
+
157
+ def __init__(
158
+ self,
159
+ d_model: int = 512,
160
+ n_head: int = 8,
161
+ d_feedforward: int = 2048,
162
+ dropout: float = 0.0,
163
+ activation: str = "GELU",
164
+ ) -> None:
165
+ super().__init__()
166
+
167
+ # Attention with PyTorch naming
168
+ self.self_attn = nn.MultiheadAttention(
169
+ d_model, n_head, dropout=dropout, bias=True, batch_first=True
170
+ )
171
+ self.dropout1 = nn.Dropout(dropout)
172
+
173
+ # Feedforward with PyTorch naming
174
+ self.linear1 = nn.Linear(d_model, d_feedforward, bias=True)
175
+ self.dropout = nn.Dropout(dropout)
176
+ self.linear2 = nn.Linear(d_feedforward, d_model, bias=True)
177
+ self.dropout2 = nn.Dropout(dropout)
178
+
179
+ # Norms
180
+ self.norm1 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
181
+ self.norm2 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
182
+
183
+ # Activation
184
+ self.activation = getattr(nn, activation)()
185
+
186
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
187
+ """Apply transformer with pre-norm (norm_first=True style)."""
188
+ # Self attention block
189
+ x2 = self.norm1(x)
190
+ x2 = self.self_attn(x2, x2, x2, need_weights=False)[0]
191
+ x = x + self.dropout1(x2)
192
+
193
+ # Feedforward block
194
+ x2 = self.norm2(x)
195
+ x2 = self.linear1(x2)
196
+ x2 = self.activation(x2)
197
+ x2 = self.dropout(x2)
198
+ x2 = self.linear2(x2)
199
+ x = x + self.dropout2(x2)
200
+
201
+ return x
202
+
203
+
204
+ class TransformerDecoder(nn.Module):
205
+ """Single transformer (decoder) layer.
206
+
207
+ Uses PyTorch's naming convention to match checkpoint:
208
+ - self_attn, multihead_attn (not attn, attn2)
209
+ - linear1, linear2 (not feedforward.0, feedforward.4)
210
+ - norm1, norm2, norm3 (for self-attn, cross-attn, and feedforward)
211
+ """
212
+
213
+ def __init__(
214
+ self,
215
+ d_model: int = 512,
216
+ n_head: int = 8,
217
+ d_feedforward: int = 2048,
218
+ dropout: float = 0.0,
219
+ activation: str = "GELU",
220
+ ) -> None:
221
+ super().__init__()
222
+
223
+ # Self attention with PyTorch naming
224
+ self.self_attn = nn.MultiheadAttention(
225
+ d_model, n_head, dropout=dropout, bias=True, batch_first=True
226
+ )
227
+ self.dropout1 = nn.Dropout(dropout)
228
+
229
+ # Cross attention with PyTorch naming (multihead_attn, not attn2)
230
+ self.multihead_attn = nn.MultiheadAttention(
231
+ d_model, n_head, dropout=dropout, bias=True, batch_first=True
232
+ )
233
+ self.dropout2 = nn.Dropout(dropout)
234
+
235
+ # Feedforward with PyTorch naming
236
+ self.linear1 = nn.Linear(d_model, d_feedforward, bias=True)
237
+ self.dropout = nn.Dropout(dropout)
238
+ self.linear2 = nn.Linear(d_feedforward, d_model, bias=True)
239
+ self.dropout3 = nn.Dropout(dropout)
240
+
241
+ # Norms (note: norm2 is for cross-attention)
242
+ self.norm1 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
243
+ self.norm2 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
244
+ self.norm3 = nn.LayerNorm(d_model, eps=1e-5, bias=True)
245
+
246
+ # Activation
247
+ self.activation = getattr(nn, activation)()
248
+
249
+ def forward(self, x: torch.Tensor, x_enc: torch.Tensor) -> torch.Tensor:
250
+ """Apply transformer decoder with pre-norm."""
251
+ # Self attention block
252
+ x2 = self.norm1(x)
253
+ x2 = self.self_attn(x2, x2, x2, need_weights=False)[0]
254
+ x = x + self.dropout1(x2)
255
+
256
+ # Cross attention block
257
+ x2 = self.norm2(x)
258
+ x2 = self.multihead_attn(x2, x_enc, x_enc, need_weights=False)[0]
259
+ x = x + self.dropout2(x2)
260
+
261
+ # Feedforward block
262
+ x2 = self.norm3(x)
263
+ x2 = self.linear1(x2)
264
+ x2 = self.activation(x2)
265
+ x2 = self.dropout(x2)
266
+ x2 = self.linear2(x2)
267
+ x = x + self.dropout3(x2)
268
+
269
+ return x
270
+
271
+
272
+ class BasisChange(nn.Module):
273
+ """Create "change-of-basis" query.
274
+
275
+ From: grt/deepradar/modules/transformer.py
276
+ """
277
+
278
+ def __init__(
279
+ self,
280
+ shape: Sequence[int] = [],
281
+ flatten: bool = True,
282
+ scale: Optional[Sequence[float]] = None,
283
+ global_scale: float = 1.0,
284
+ ) -> None:
285
+ super().__init__()
286
+
287
+ self.pos = Sinusoid(scale=scale, global_scale=global_scale)
288
+ self.shape = shape
289
+ self.flatten = flatten
290
+
291
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
292
+ """Apply change of basis."""
293
+ idxs = tuple([slice(None)] + [None] * len(self.shape) + [slice(None)])
294
+ query = self.pos(torch.tile(x[idxs], (1, *self.shape, 1)))
295
+
296
+ if self.flatten:
297
+ query = query.reshape(x.shape[0], -1, x.shape[-1])
298
+ return query
299
+
300
+
301
+ class Unpatch(nn.Module):
302
+ """Unpatch data.
303
+
304
+ Args:
305
+ output_size: output 2D shape.
306
+ features: number of input features; should be `>= size * size`.
307
+ size: patch size as (width, height, channels).
308
+ """
309
+
310
+ def __init__(
311
+ self,
312
+ output_size: Sequence[int],
313
+ features: int = 512,
314
+ size: Sequence[int] = (16, 16),
315
+ ) -> None:
316
+ super().__init__()
317
+
318
+ self.linear = nn.Linear(features, output_size[-1] * int(np.prod(size)))
319
+ self.size = size
320
+ self.output_size = output_size
321
+
322
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
323
+ """Perform 2D unpatching.
324
+
325
+ Operates in batch-spatial-feature order; spatial axes are flattened on
326
+ the input, and unflattened in the output.
327
+ """
328
+ embedding = self.linear(x)
329
+
330
+ if len(self.size) == 2:
331
+ return rearrange(
332
+ embedding,
333
+ "n (x1 x2) (s1 s2 c) -> n (x1 s1) (x2 s2) c",
334
+ x1=self.output_size[0] // self.size[0],
335
+ x2=self.output_size[1] // self.size[1],
336
+ s1=self.size[0],
337
+ s2=self.size[1],
338
+ c=self.output_size[-1],
339
+ )
340
+ elif len(self.size) == 3:
341
+ return rearrange(
342
+ embedding,
343
+ "n (x1 x2 x3) (s1 s2 s3 c) -> n (x1 s1) (x2 s2) (x3 s3) c",
344
+ x1=self.output_size[0] // self.size[0],
345
+ x2=self.output_size[1] // self.size[1],
346
+ x3=self.output_size[2] // self.size[2],
347
+ s1=self.size[0],
348
+ s2=self.size[1],
349
+ s3=self.size[2],
350
+ c=self.output_size[-1],
351
+ )
352
+ else:
353
+ raise ValueError("Unpatch is only implemented for 2D and 3D tensors.")
354
+
355
+
356
+ # ============================================================================
357
+ # GRT Model Components
358
+ # ============================================================================
359
+
360
+
361
+ class GRTEncoder(nn.Module):
362
+ """GRT Transformer Encoder matching official implementation."""
363
+
364
+ def __init__(
365
+ self,
366
+ layers: int = 4,
367
+ dim: int = 512,
368
+ ff_ratio: float = 4.0,
369
+ head_dim: int = 64,
370
+ dropout: float = 0.1,
371
+ activation: str = "GELU",
372
+ patch: list[int] = [2, 8, 2, 4],
373
+ pos_scale: list[float] = [1.0, 1.0, 1.0, 1.0],
374
+ global_scale: float = 16.0,
375
+ input_channels: int = 2,
376
+ positions: Literal["flat", "nd"] = "nd",
377
+ ):
378
+ super().__init__()
379
+
380
+ # Patch embedding
381
+ self.patch = PatchMerge(d_in=input_channels, d_out=dim, scale=patch, norm=False)
382
+
383
+ # Position embedding
384
+ self.positions = positions
385
+ self.pos = Sinusoid(scale=pos_scale, global_scale=global_scale)
386
+
387
+ # Readout token
388
+ self.readout = Readout(d_model=dim)
389
+
390
+ # Encoder layers
391
+ self.layers = nn.ModuleList(
392
+ [
393
+ TransformerLayer(
394
+ d_feedforward=int(ff_ratio * dim),
395
+ d_model=dim,
396
+ n_head=dim // head_dim,
397
+ dropout=dropout,
398
+ activation=activation,
399
+ )
400
+ for _ in range(layers)
401
+ ]
402
+ )
403
+
404
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
405
+ """Forward pass."""
406
+ # Patch embedding
407
+ embedded = self.patch(x)
408
+
409
+ # Apply positional encoding
410
+ if self.positions == "nd":
411
+ embedded = self.pos(embedded)
412
+
413
+ # Flatten spatial dimensions
414
+ flat = embedded.reshape(embedded.shape[0], -1, embedded.shape[-1])
415
+
416
+ # Apply flat positional encoding if needed
417
+ if self.positions == "flat":
418
+ flat = self.pos(flat)
419
+
420
+ # Add readout token
421
+ x = self.readout(flat)
422
+
423
+ # Apply encoder layers
424
+ for layer in self.layers:
425
+ x = layer(x)
426
+
427
+ return x
428
+
429
+
430
+ class GRTDecoder3D(nn.Module):
431
+ """GRT 3D Transformer Decoder matching official implementation."""
432
+
433
+ def __init__(
434
+ self,
435
+ key: str = "map",
436
+ layers: int = 4,
437
+ dim: int = 512,
438
+ ff_ratio: float = 4.0,
439
+ head_dim: int = 64,
440
+ dropout: float = 0.1,
441
+ activation: str = "GELU",
442
+ shape: list[int] = [64, 128, 64],
443
+ pos_scale: list[float] = [1.0, 1.0, 1.0],
444
+ global_scale: float = 16.0,
445
+ patch: list[int] = [8, 8, 8],
446
+ out_dim: int = 0,
447
+ positions: Literal["flat", "nd"] = "nd",
448
+ mode: Literal["last", "pool"] = "last",
449
+ ):
450
+ super().__init__()
451
+
452
+ self.key = key
453
+ self.out_dim = out_dim
454
+ self.mode = mode
455
+
456
+ # Decoder layers
457
+ self.layers = nn.ModuleList(
458
+ [
459
+ TransformerDecoder(
460
+ d_feedforward=int(ff_ratio * dim),
461
+ d_model=dim,
462
+ n_head=dim // head_dim,
463
+ dropout=dropout,
464
+ activation=activation,
465
+ )
466
+ for _ in range(layers)
467
+ ]
468
+ )
469
+
470
+ # Query generation with position encoding
471
+ query_shape = [s // p for s, p in zip(shape, patch)]
472
+ if positions == "flat":
473
+ query_shape = [int(np.prod(query_shape))]
474
+
475
+ self.query = BasisChange(
476
+ shape=query_shape, scale=pos_scale, global_scale=global_scale, flatten=True
477
+ )
478
+
479
+ # Unpatch to reconstruct output
480
+ self.unpatch = Unpatch(
481
+ output_size=(*shape, max(1, self.out_dim)), features=dim, size=patch
482
+ )
483
+
484
+ def forward(self, encoded: torch.Tensor) -> dict[str, torch.Tensor]:
485
+ """Forward pass."""
486
+ # Extract readout token or pool
487
+ if self.mode == "last":
488
+ x = encoded[:, -1, :]
489
+ else:
490
+ x = torch.mean(encoded, dim=1)
491
+
492
+ # Generate query with positional encoding
493
+ x = self.query(x)
494
+
495
+ # Encoded features without readout token
496
+ enc = encoded[:, :-1, :]
497
+
498
+ # Apply decoder layers
499
+ for layer in self.layers:
500
+ x = layer(x, enc)
501
+
502
+ # Unpatch to 3D output
503
+ out = self.unpatch(x)
504
+
505
+ # Squeeze channel dimension if binary output
506
+ if self.out_dim == 0:
507
+ out = out[..., 0]
508
+
509
+ return {self.key: out}
510
+
511
+
512
+ # ============================================================================
513
+ # Complete GRT-Small Model
514
+ # ============================================================================
515
+
516
+
517
+ class GRTSmall(nn.Module):
518
+ """GRT-Small model for 3D occupancy mapping.
519
+
520
+ Input: (batch, doppler, azimuth, elevation, range, 2)
521
+ - doppler: 64
522
+ - azimuth: 8
523
+ - elevation: 2
524
+ - range: 256
525
+ - channels: 2 (I/Q)
526
+
527
+ Output: (batch, elevation, azimuth, range)
528
+ - elevation: 64
529
+ - azimuth: 128
530
+ - range: 64
531
+
532
+ ~29M parameters for GRT-small variant.
533
+ """
534
+
535
+ def __init__(self):
536
+ super().__init__()
537
+
538
+ dim = 512
539
+ layers = 4
540
+
541
+ # Create encoder - stored as "tokenizer" + "encoder" in checkpoint
542
+ # But we organize logically here and handle mapping in load_checkpoint
543
+ self.tokenizer = GRTEncoder(
544
+ layers=layers,
545
+ dim=dim,
546
+ ff_ratio=4.0,
547
+ head_dim=64,
548
+ dropout=0.1,
549
+ activation="GELU",
550
+ patch=[2, 8, 2, 4],
551
+ pos_scale=[1.0, 1.0, 1.0, 1.0],
552
+ global_scale=16.0,
553
+ input_channels=2,
554
+ positions="nd",
555
+ )
556
+
557
+ # Create decoder wrapper
558
+ self.decoder = nn.Module()
559
+ self.decoder.occ3d = GRTDecoder3D(
560
+ key="map",
561
+ layers=layers,
562
+ dim=dim,
563
+ ff_ratio=4.0,
564
+ head_dim=64,
565
+ dropout=0.1,
566
+ activation="GELU",
567
+ shape=[64, 128, 64],
568
+ pos_scale=[1.0, 1.0, 1.0],
569
+ global_scale=16.0,
570
+ patch=[8, 8, 8],
571
+ out_dim=0,
572
+ positions="nd",
573
+ mode="last",
574
+ )
575
+
576
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
577
+ """Forward pass."""
578
+ # Encode
579
+ encoded = self.tokenizer(x)
580
+
581
+ # Decode
582
+ output = self.decoder.occ3d(encoded)
583
+
584
+ # Return just the occupancy map tensor
585
+ return output["map"]
586
+
587
+
588
+ class ResNet18ImageTokenizer(nn.Module):
589
+ """Coarse ResNet-18 spatial tokens projected into GRT's 512-D memory."""
590
+
591
+ def __init__(
592
+ self,
593
+ pretrained: bool,
594
+ image_height: int,
595
+ image_width: int,
596
+ output_dim: int = 512,
597
+ ):
598
+ super().__init__()
599
+ self.pretrained = bool(pretrained)
600
+ self.image_height = int(image_height)
601
+ self.image_width = int(image_width)
602
+ self.output_stride = 32
603
+ if (
604
+ self.image_height % self.output_stride
605
+ or self.image_width % self.output_stride
606
+ ):
607
+ raise ValueError(
608
+ "ResNet-18 tokenization requires image dimensions divisible by 32, "
609
+ f"got {(self.image_height, self.image_width)}"
610
+ )
611
+
612
+ weights = ResNet18_Weights.DEFAULT if self.pretrained else None
613
+ resnet = resnet18(weights=weights)
614
+ self.backbone = nn.Sequential(
615
+ resnet.conv1,
616
+ resnet.bn1,
617
+ resnet.relu,
618
+ resnet.maxpool,
619
+ resnet.layer1,
620
+ resnet.layer2,
621
+ resnet.layer3,
622
+ resnet.layer4,
623
+ )
624
+ self.register_buffer(
625
+ "image_mean",
626
+ torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1),
627
+ )
628
+ self.register_buffer(
629
+ "image_std",
630
+ torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1),
631
+ )
632
+ self.projection = nn.Sequential(
633
+ nn.LayerNorm(512),
634
+ nn.Linear(512, output_dim),
635
+ )
636
+ self.modality = nn.Parameter(torch.empty(1, 1, output_dim))
637
+ nn.init.normal_(self.modality, mean=0.0, std=0.02)
638
+
639
+ if self.pretrained:
640
+ for parameter in self.backbone.parameters():
641
+ parameter.requires_grad_(False)
642
+ self.backbone.eval()
643
+
644
+ def train(self, mode: bool = True):
645
+ super().train(mode)
646
+ if self.pretrained:
647
+ self.backbone.eval()
648
+ return self
649
+
650
+ def forward(self, image: torch.Tensor) -> torch.Tensor:
651
+ """Return layer-4 spatial features as [B, H/32 * W/32, 512]."""
652
+ if image.ndim != 4 or image.shape[1] != 3:
653
+ raise ValueError(
654
+ "ResNet18ImageTokenizer expects RGB images shaped [B, 3, H, W], "
655
+ f"got {tuple(image.shape)}"
656
+ )
657
+ height, width = image.shape[-2:]
658
+ if height != self.image_height or width != self.image_width:
659
+ raise ValueError(
660
+ "Image size must match the configured ResNet-18 size "
661
+ f"{(self.image_height, self.image_width)}, got {(height, width)}"
662
+ )
663
+
664
+ image = (image - self.image_mean) / self.image_std
665
+ if self.pretrained:
666
+ with torch.no_grad():
667
+ features = self.backbone(image)
668
+ else:
669
+ features = self.backbone(image)
670
+
671
+ spatial_tokens = features.flatten(2).transpose(1, 2)
672
+ return self.projection(spatial_tokens) + self.modality
673
+
674
+
675
+ def fuse_decoder_memory(
676
+ radar_encoded: torch.Tensor, image_tokens: torch.Tensor
677
+ ) -> torch.Tensor:
678
+ """Insert image memory before GRT's final readout token.
679
+
680
+ GRTDecoder3D uses the final token as its query seed and every preceding
681
+ token as cross-attention memory. Keeping the readout last is therefore a
682
+ required part of the fusion contract.
683
+ """
684
+ if radar_encoded.ndim != 3 or image_tokens.ndim != 3:
685
+ raise ValueError("radar_encoded and image_tokens must both be [B, N, C]")
686
+ if radar_encoded.shape[1] < 1:
687
+ raise ValueError("radar_encoded must contain the GRT readout token")
688
+ if (
689
+ radar_encoded.shape[0] != image_tokens.shape[0]
690
+ or radar_encoded.shape[2] != image_tokens.shape[2]
691
+ ):
692
+ raise ValueError(
693
+ "radar and image token batches must have matching batch and channel dimensions"
694
+ )
695
+ return torch.cat(
696
+ [radar_encoded[:, :-1, :], image_tokens, radar_encoded[:, -1:, :]],
697
+ dim=1,
698
+ )
699
+
700
+
701
+ class GRTImageNaiveSmall(nn.Module):
702
+ """Naive GRT+Image model with a fresh joint occupancy decoder."""
703
+
704
+ def __init__(
705
+ self,
706
+ resnet18_pretrained: bool = False,
707
+ image_height: int = 288,
708
+ image_width: int = 512,
709
+ ):
710
+ super().__init__()
711
+
712
+ dim = 512
713
+ layers = 4
714
+ self.tokenizer = GRTEncoder(
715
+ layers=layers,
716
+ dim=dim,
717
+ ff_ratio=4.0,
718
+ head_dim=64,
719
+ dropout=0.1,
720
+ activation="GELU",
721
+ patch=[2, 8, 2, 4],
722
+ pos_scale=[1.0, 1.0, 1.0, 1.0],
723
+ global_scale=16.0,
724
+ input_channels=2,
725
+ positions="nd",
726
+ )
727
+ self.image_tokenizer = ResNet18ImageTokenizer(
728
+ pretrained=resnet18_pretrained,
729
+ image_height=image_height,
730
+ image_width=image_width,
731
+ output_dim=dim,
732
+ )
733
+
734
+ self.decoder = nn.Module()
735
+ self.decoder.occ3d = GRTDecoder3D(
736
+ key="map",
737
+ layers=layers,
738
+ dim=dim,
739
+ ff_ratio=4.0,
740
+ head_dim=64,
741
+ dropout=0.1,
742
+ activation="GELU",
743
+ shape=[128, 256, 64],
744
+ pos_scale=[1.0, 1.0, 1.0],
745
+ global_scale=16.0,
746
+ patch=[8, 8, 8],
747
+ out_dim=0,
748
+ positions="nd",
749
+ mode="last",
750
+ )
751
+ self._radar_encoder_frozen = False
752
+
753
+ def freeze_radar_encoder(self) -> None:
754
+ """Freeze GRT feature extraction and keep its dropout disabled."""
755
+ self._radar_encoder_frozen = True
756
+ for parameter in self.tokenizer.parameters():
757
+ parameter.requires_grad_(False)
758
+ self.tokenizer.eval()
759
+
760
+ def train(self, mode: bool = True):
761
+ super().train(mode)
762
+ if self._radar_encoder_frozen:
763
+ self.tokenizer.eval()
764
+ return self
765
+
766
+ def forward(self, radar: torch.Tensor, image: torch.Tensor) -> torch.Tensor:
767
+ radar_encoded = self.tokenizer(radar)
768
+ image_tokens = self.image_tokenizer(image)
769
+ fused_encoded = fuse_decoder_memory(radar_encoded, image_tokens)
770
+ return self.decoder.occ3d(fused_encoded)["map"]
771
+
772
+
773
+ def load_radar_encoder_checkpoint(
774
+ model: GRTImageNaiveSmall, checkpoint_path, map_location="cpu"
775
+ ) -> dict:
776
+ """Load only the pretrained GRT tokenizer/encoder and leave fusion fresh."""
777
+ state_dict = load_file(checkpoint_path, device="cpu")
778
+
779
+ encoder_state = {
780
+ key: value for key, value in state_dict.items() if key.startswith("tokenizer.")
781
+ }
782
+ if not encoder_state:
783
+ raise RuntimeError(
784
+ "Radar checkpoint does not contain any tokenizer.* encoder parameters"
785
+ )
786
+
787
+ missing_keys, unexpected_keys = model.load_state_dict(encoder_state, strict=False)
788
+ missing_encoder_keys = [
789
+ key for key in missing_keys if key.startswith("tokenizer.")
790
+ ]
791
+ if missing_encoder_keys or unexpected_keys:
792
+ raise RuntimeError(
793
+ "Radar checkpoint is not compatible with the GRT encoder: "
794
+ f"missing encoder keys {missing_encoder_keys}; "
795
+ f"unexpected keys {list(unexpected_keys)}"
796
+ )
797
+ return checkpoint
798
+
799
+
src/Baselines/grt_image/inference.py ADDED
@@ -0,0 +1,224 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Inference for the naive GRT+Image baseline.
4
+
5
+ Runs inference on specified sequences (default: brk_3rd, brk_3rd_fog, brk_3rd_fog2)
6
+ using weights trained by train.py.
7
+ For each sequence, saves pred_depth.npy with shape [T, 128, 256] in [0, 1].
8
+
9
+ Single GPU: Each frame is seen exactly once; no duplication or incompleteness.
10
+ Multi-GPU (DDP): Dataloader is sharded; each rank writes its results to a file, then
11
+ main process merges with deduplication by frame_idx (keeps first occurrence) and saves.
12
+ """
13
+
14
+ import os
15
+ import torch
16
+ import numpy as np
17
+ import argparse
18
+ import yaml
19
+ import pickle
20
+ from tqdm import tqdm
21
+ from accelerate import Accelerator
22
+ from accelerate.utils import set_seed
23
+ from collections import defaultdict
24
+ from safetensors.torch import load_file
25
+
26
+ from grt_model import GRTImageNaiveSmall
27
+ from dataloader import create_rice_dataloader
28
+ from augmentations import (
29
+ translate_radar,
30
+ dequantize_depth,
31
+ )
32
+
33
+ def batch_radar_to_spectrum(
34
+ radar_amplitude: torch.Tensor, radar_phase: torch.Tensor
35
+ ) -> torch.Tensor:
36
+ """Build GRT's [B, D, A, E, R, 2] spectrum from loader tensors."""
37
+ amplitude = radar_amplitude.permute(0, 1, 3, 2, 4)
38
+ phase = radar_phase.permute(0, 1, 3, 2, 4)
39
+ return torch.stack([amplitude, phase], dim=-1)
40
+
41
+
42
+ def main():
43
+ parser = argparse.ArgumentParser(
44
+ description="Run naive GRT+Image inference on Smoke-Eval sequences"
45
+ )
46
+ parser.add_argument(
47
+ "--config", type=str, default="config.yaml", help="Path to config file"
48
+ )
49
+ parser.add_argument(
50
+ "--checkpoint",
51
+ type=str,
52
+ required=True,
53
+ help="Path to validation-selected GRT+Image .safetensors file",
54
+ )
55
+ parser.add_argument(
56
+ "--output_dir",
57
+ type=str,
58
+ default="inference_results",
59
+ help="Directory to save results",
60
+ )
61
+ parser.add_argument(
62
+ "--sequences",
63
+ type=str,
64
+ nargs="+",
65
+ default=None,
66
+ help="Optional Smoke-Eval sequence subset (default: every valid sequence)",
67
+ )
68
+ parser.add_argument(
69
+ "--debug", action="store_true", help="Run in debug mode (process only 1 batch)"
70
+ )
71
+ args = parser.parse_args()
72
+
73
+ # Load config
74
+ with open(args.config, "r") as f:
75
+ config = yaml.safe_load(f)
76
+
77
+ # Initialize accelerator
78
+ accelerator = Accelerator(mixed_precision="fp16")
79
+ set_seed(config["training"].get("seed", 42))
80
+
81
+ # Create output directory (all ranks so DDP gather_dir can be created)
82
+ os.makedirs(args.output_dir, exist_ok=True)
83
+
84
+ # Create model
85
+ accelerator.print("Creating naive GRT+Image model...")
86
+ model = GRTImageNaiveSmall(
87
+ resnet18_pretrained=config["model"].get("resnet18_pretrained", True),
88
+ image_height=config["data"].get("image_height", 288),
89
+ image_width=config["data"].get("image_width", 512),
90
+ )
91
+
92
+ # Safetensors files contain only the model state dictionary.
93
+ accelerator.print(f"Loading checkpoint from {args.checkpoint}")
94
+ model.load_state_dict(load_file(args.checkpoint, device="cpu"), strict=True)
95
+
96
+ sequence_description = args.sequences if args.sequences else "all valid Smoke-Eval sequences"
97
+ accelerator.print(f"Inference sequences: {sequence_description}")
98
+ inference_loader = create_rice_dataloader(
99
+ root_dir=config["paths"]["smoke_eval_root"],
100
+ batch_size=config["training"]["batch_size"],
101
+ num_workers=0,
102
+ frame_skip=1,
103
+ sequences=args.sequences,
104
+ image_height=config["data"].get("image_height", 288),
105
+ image_width=config["data"].get("image_width", 512),
106
+ shuffle=False,
107
+ )
108
+
109
+ # Prepare model and dataloader
110
+ model, inference_loader = accelerator.prepare(model, inference_loader)
111
+ model.eval()
112
+
113
+ # Dictionary to aggregate results by sequence: sequence_id -> list of (frame_idx, pred_depth)
114
+ results_by_sequence = defaultdict(list)
115
+
116
+ accelerator.print("Starting inference...")
117
+
118
+ with torch.no_grad():
119
+ for batch in tqdm(
120
+ inference_loader, disable=not accelerator.is_local_main_process
121
+ ):
122
+ # Extract data
123
+ rsp_data = batch_radar_to_spectrum(
124
+ batch["radar_amplitude"], batch["radar_phase"]
125
+ )
126
+ image = batch["image"]
127
+ sequences = batch["sequence"]
128
+ frame_indices = batch["frame_idx"]
129
+
130
+ # Apply radar augmentation
131
+ rsp_data = translate_radar(rsp_data)
132
+
133
+ # Forward pass
134
+ occupancy_pred_logits = model(rsp_data, image) # [B, 128, 256, 64]
135
+
136
+ # Dequantize to depth [B, 1, 128, 256], values in [0, 1].
137
+ pred_depth = dequantize_depth(occupancy_pred_logits)
138
+ pred_depth_np = (
139
+ pred_depth.cpu().numpy().astype(np.float32)
140
+ ) # [B, 1, 128, 256]
141
+
142
+ # Collect results (frame_idx, pred_depth per sample)
143
+ for i in range(len(sequences)):
144
+ seq_id = sequences[i]
145
+ f_idx = frame_indices[i].item()
146
+ # Store [1, 128, 256] per frame.
147
+ results_by_sequence[seq_id].append(
148
+ {
149
+ "frame_idx": f_idx,
150
+ "pred_depth": pred_depth_np[i],
151
+ }
152
+ )
153
+
154
+ if args.debug:
155
+ break
156
+
157
+ # Single GPU: save directly (each frame seen once, no duplication)
158
+ # Multi-GPU: gather via files, merge with dedupe by frame_idx, then save
159
+ if accelerator.num_processes == 1:
160
+ if accelerator.is_main_process:
161
+ accelerator.print("Saving results (single process)...")
162
+ for seq_id, frames in tqdm(
163
+ results_by_sequence.items(), desc="Saving sequences"
164
+ ):
165
+ frames.sort(key=lambda x: x["frame_idx"])
166
+ pred_depth_stack = np.stack([f["pred_depth"] for f in frames], axis=0)
167
+ pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 128, 256]
168
+ np.save(
169
+ os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"),
170
+ pred_depth_stack,
171
+ )
172
+ accelerator.print(
173
+ f" {seq_id}: saved {pred_depth_stack.shape[0]} frames"
174
+ )
175
+ accelerator.print(f"Processed {len(results_by_sequence)} sequences.")
176
+ accelerator.print(f"Results saved to {args.output_dir}")
177
+ else:
178
+ # DDP: gather results from all ranks via files, dedupe by frame_idx, save on main
179
+ accelerator.wait_for_everyone()
180
+ gather_dir = os.path.join(args.output_dir, "_gather")
181
+ os.makedirs(gather_dir, exist_ok=True)
182
+ rank = accelerator.process_index
183
+ rank_file = os.path.join(gather_dir, f"rank_{rank}_results.pkl")
184
+ with open(rank_file, "wb") as f:
185
+ pickle.dump(dict(results_by_sequence), f, protocol=pickle.HIGHEST_PROTOCOL)
186
+ accelerator.wait_for_everyone()
187
+
188
+ if accelerator.is_main_process:
189
+ accelerator.print("Merging and deduplicating results from all ranks...")
190
+ merged_results = defaultdict(dict) # seq_id -> {frame_idx: pred_depth}
191
+ for r in range(accelerator.num_processes):
192
+ pkl_path = os.path.join(gather_dir, f"rank_{r}_results.pkl")
193
+ with open(pkl_path, "rb") as f:
194
+ rank_results = pickle.load(f)
195
+ for seq_id, frames in rank_results.items():
196
+ for frame_data in frames:
197
+ f_idx = frame_data["frame_idx"]
198
+ if f_idx not in merged_results[seq_id]:
199
+ merged_results[seq_id][f_idx] = frame_data["pred_depth"]
200
+ os.remove(pkl_path)
201
+
202
+ for seq_id, frame_dict in tqdm(
203
+ merged_results.items(), desc="Saving sequences"
204
+ ):
205
+ sorted_items = sorted(frame_dict.items(), key=lambda x: x[0])
206
+ pred_depth_stack = np.stack([item[1] for item in sorted_items], axis=0)
207
+ pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 128, 256]
208
+ np.save(
209
+ os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"),
210
+ pred_depth_stack,
211
+ )
212
+ accelerator.print(
213
+ f" {seq_id}: saved {pred_depth_stack.shape[0]} frames"
214
+ )
215
+ if os.path.isdir(gather_dir) and not os.listdir(gather_dir):
216
+ os.rmdir(gather_dir)
217
+ accelerator.print(f"Processed {len(merged_results)} sequences.")
218
+ accelerator.print(f"Results saved to {args.output_dir}")
219
+
220
+ accelerator.wait_for_everyone()
221
+
222
+
223
+ if __name__ == "__main__":
224
+ main()
src/Baselines/grt_image/split.json ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "test": [
3
+ "Dell-1",
4
+ "Dell-2",
5
+ "Smoke-Dell-1",
6
+ "Smoke-Dell-2",
7
+ "brk-2",
8
+ "brk-3",
9
+ "Brk-b",
10
+ "brk-basement",
11
+ "Brk-stair",
12
+ "Smoke-brk-2",
13
+ "Smoke-brk-3",
14
+ "Smoke-brk-b"
15
+ ]
16
+ }
src/Baselines/radarcam-depth/data/SML_dataset.py ADDED
@@ -0,0 +1,83 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch.utils.data
2
+ import numpy as np
3
+ import modules.midas.utils as utils
4
+ from PIL import Image
5
+
6
+ def load_input_image(input_image_fp):
7
+ return utils.read_image(input_image_fp)
8
+
9
+
10
+ def load_sparse_depth(input_sparse_depth_fp):
11
+ input_sparse_depth = np.array(Image.open(input_sparse_depth_fp), dtype=np.float32) / 256.0
12
+ input_sparse_depth[input_sparse_depth <= 0] = 0.0
13
+ return input_sparse_depth
14
+
15
+
16
+ class SML_dataset(torch.utils.data.Dataset):
17
+ def __init__(self,
18
+ image_paths,
19
+ radar_paths,
20
+ gt_paths,
21
+ sparse_gt_paths,
22
+ rcnet_paths,
23
+ mono_pred_paths = None,
24
+ mono_ga_paths = None,
25
+ ):
26
+
27
+ self.n_sample = len(image_paths)
28
+
29
+ for paths in [image_paths, radar_paths, gt_paths, sparse_gt_paths,
30
+ rcnet_paths, mono_pred_paths, mono_ga_paths]:
31
+ if paths is not None:
32
+ assert len(paths) == self.n_sample
33
+
34
+ self.image_paths = image_paths
35
+ self.radar_paths = radar_paths
36
+ self.gt_paths = gt_paths
37
+ self.sparse_gt_paths = sparse_gt_paths
38
+ self.rcnet_paths = rcnet_paths
39
+ self.mono_pred_paths = mono_pred_paths
40
+ self.mono_ga_paths = mono_ga_paths
41
+
42
+
43
+ def __getitem__(self, index):
44
+ image = load_input_image(self.image_paths[index])
45
+ radar = load_sparse_depth(self.radar_paths[index])
46
+ gt = load_sparse_depth(self.gt_paths[index])
47
+ sparse_gt = load_sparse_depth(self.sparse_gt_paths[index])
48
+ rcnet = load_sparse_depth(self.rcnet_paths[index])
49
+
50
+ image, radar, gt, sparse_gt, rcnet = [
51
+ T.astype(np.float32)
52
+ for T in [image, radar, gt, sparse_gt, rcnet]
53
+ ]
54
+
55
+ # Crop the image for ZJU dataset
56
+ if image.shape[0] == 720:
57
+ image = image[720 // 3: 720 // 4 * 3, :, :]
58
+ radar = radar[720 // 3: 720 // 4 * 3, :]
59
+ gt = gt[720 // 3: 720 // 4 * 3, :]
60
+ sparse_gt = sparse_gt[720 // 3: 720 // 4 * 3, :]
61
+
62
+
63
+ if self.mono_ga_paths is not None:
64
+ mono_pred = load_sparse_depth(self.mono_ga_paths[index])
65
+ mono_pred = mono_pred.astype(np.float32)
66
+ if mono_pred.shape[0] == 720:
67
+ mono_pred = mono_pred[720 // 3: 720 // 4 * 3, :]
68
+ else:
69
+ mono_pred = None
70
+
71
+ if self.mono_ga_paths is not None:
72
+ mono_ga = load_sparse_depth(self.mono_ga_paths[index])
73
+ mono_ga = mono_ga.astype(np.float32)
74
+ if mono_ga.shape[0] == 720:
75
+ mono_ga = mono_ga[720 // 3: 720 // 4 * 3, :]
76
+ else:
77
+ mono_ga = None
78
+
79
+ return image, mono_pred, radar, gt, sparse_gt, rcnet, mono_ga
80
+
81
+
82
+ def __len__(self):
83
+ return self.n_sample
src/Baselines/radarcam-depth/data/data_utils.py ADDED
@@ -0,0 +1,326 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ from scipy.interpolate import LinearNDInterpolator
3
+ from PIL import Image
4
+ import matplotlib.pyplot as plt
5
+
6
+
7
+
8
+ def load_data_path(root, file_name_txt, data_type):
9
+ with open(file_name_txt, 'r') as f:
10
+ data_path = f.readlines()
11
+ data_path = [root + x.strip() + data_type for x in data_path]
12
+ return data_path
13
+
14
+
15
+ def load_data_path_nu(root, name_list, data_type):
16
+ data_path = [root + x.strip() + data_type for x in name_list]
17
+ return data_path
18
+
19
+
20
+ def read_paths(filepath):
21
+ '''
22
+ Reads a newline delimited file containing paths
23
+
24
+ Arg(s):
25
+ filepath : str
26
+ path to file to be read
27
+ Return:
28
+ list[str] : list of paths
29
+ '''
30
+
31
+ path_list = []
32
+ with open(filepath) as f:
33
+ while True:
34
+ path = f.readline().rstrip('\n')
35
+
36
+ # If there was nothing to read
37
+ if path == '':
38
+ break
39
+
40
+ path_list.append(path)
41
+
42
+ return path_list
43
+
44
+
45
+ def write_paths(filepath, paths):
46
+ '''
47
+ Stores line delimited paths into file
48
+
49
+ Arg(s):
50
+ filepath : str
51
+ path to file to save paths
52
+ paths : list[str]
53
+ paths to write into file
54
+ '''
55
+
56
+ with open(filepath, 'w') as o:
57
+ for idx in range(len(paths)):
58
+ o.write(paths[idx] + '\n')
59
+
60
+
61
+ def load_image(path, normalize=False, data_format='HWC'):
62
+ '''
63
+ Loads an RGB image
64
+
65
+ Arg(s):
66
+ path : str
67
+ path to RGB image
68
+ normalize : bool
69
+ if set, then normalize image between [0, 1]
70
+ data_format : str
71
+ 'CHW', or 'HWC'
72
+ Returns:
73
+ numpy[float32] : H x W x C or C x H x W image
74
+ '''
75
+
76
+ # Load image
77
+ image = Image.open(path).convert('RGB')
78
+
79
+ # Convert to numpy
80
+ image = np.asarray(image, np.float32)
81
+
82
+ if data_format == 'HWC':
83
+ pass
84
+ elif data_format == 'CHW':
85
+ image = np.transpose(image, (2, 0, 1))
86
+ else:
87
+ raise ValueError('Unsupported data format: {}'.format(data_format))
88
+
89
+ # Normalize
90
+ image = image / 255.0 if normalize else image #255.0
91
+
92
+ return image
93
+
94
+
95
+
96
+ def load_depth(path, multiplier=256.0, data_format='HW'):
97
+ '''
98
+ Loads a depth map from a 16-bit PNG file
99
+
100
+ Arg(s):
101
+ path : str
102
+ path to 16-bit PNG file
103
+ multiplier : float
104
+ multiplier for encoding float as 16/32 bit unsigned integer
105
+ data_format : str
106
+ HW, CHW, HWC
107
+ Returns:
108
+ numpy[float32] : depth map
109
+ '''
110
+
111
+ # Loads depth map from 16-bit PNG file
112
+ z = np.array(Image.open(path), dtype=np.float32)
113
+
114
+ # Assert 16-bit (not 8-bit) depth map
115
+ z = z / multiplier
116
+ z[z <= 0] = 0.0
117
+
118
+ if data_format == 'HW':
119
+ pass
120
+ elif data_format == 'CHW':
121
+ z = np.expand_dims(z, axis=0)
122
+ elif data_format == 'HWC':
123
+ z = np.expand_dims(z, axis=-1)
124
+ else:
125
+ raise ValueError('Unsupported data format: {}'.format(data_format))
126
+
127
+ return z
128
+
129
+
130
+ def save_depth(z, path, multiplier=256.0):
131
+ '''
132
+ Saves a depth map to a 16-bit PNG file
133
+
134
+ Arg(s):
135
+ z : numpy[float32]
136
+ depth map
137
+ path : str
138
+ path to store depth map
139
+ multiplier : float
140
+ multiplier for encoding float as 16/32 bit unsigned integer
141
+ '''
142
+
143
+ z = np.uint32(z * multiplier)
144
+ z = Image.fromarray(z, mode='I')
145
+ z.save(path)
146
+
147
+
148
+ def save_color_depth(z, path):
149
+ '''
150
+ Saves a color depth map to a 16-bit PNG file
151
+
152
+ Arg(s):
153
+ z : numpy[float32]
154
+ depth map
155
+ path : str
156
+ path to store depth map
157
+ multiplier : float
158
+ multiplier for encoding float as 16/32 bit unsigned integer
159
+ '''
160
+
161
+ # Normalize depth map to the range [0, 1]
162
+ z_normalized = (z - np.min(z)) / (np.max(z) - np.min(z))
163
+
164
+ # Convert depth map to color
165
+ # colormap = plt.cm.jet # Choose a colormap (e.g., jet)
166
+ colormap = plt.cm.viridis
167
+ z_color = colormap(z_normalized)
168
+
169
+ # Scale color values to the range [0, 255] and convert to uint8
170
+ z_color = np.uint8(z_color * 255)
171
+
172
+ # Save color depth map as an image
173
+ image = Image.fromarray(z_color)
174
+ image.save(path)
175
+
176
+
177
+ def load_response(path, multiplier=2**14, data_format='HW'):
178
+ '''
179
+ Loads a response map from a 16-bit PNG file
180
+
181
+ Arg(s):
182
+ path : str
183
+ path to 16-bit PNG file
184
+ multiplier : float
185
+ multiplier for encoding float as 16/32 bit unsigned integer
186
+ data_format : str
187
+ HW, CHW, HWC
188
+ Returns:
189
+ numpy[float32] : response map
190
+ '''
191
+
192
+ # Loads response map from 16-bit PNG file
193
+ response = np.array(Image.open(path), dtype=np.float32)
194
+
195
+ # Convert using encodering multiplier
196
+ response = response / multiplier
197
+
198
+ if data_format == 'HW':
199
+ pass
200
+ elif data_format == 'CHW':
201
+ response = np.expand_dims(response, axis=0)
202
+ elif data_format == 'HWC':
203
+ response = np.expand_dims(response, axis=-1)
204
+ else:
205
+ raise ValueError('Unsupported data format: {}'.format(data_format))
206
+
207
+ return response
208
+
209
+
210
+ def save_response(response, path, multiplier=2**14):
211
+ '''
212
+ Saves a response map to a 16-bit PNG file
213
+
214
+ Arg(s):
215
+ response : numpy[float32]
216
+ depth map
217
+ path : str
218
+ path to store depth map
219
+ multiplier : float
220
+ multiplier for encoding float as 16/32 bit unsigned integer
221
+ '''
222
+
223
+ response = np.uint32(response * multiplier)
224
+ response = Image.fromarray(response, mode='I')
225
+ response.save(path)
226
+
227
+
228
+ def interpolate_depth(depth_map, validity_map, log_space=False):
229
+ '''
230
+ Interpolate sparse depth with barycentric coordinates
231
+
232
+ Arg(s):
233
+ depth_map : np.float32
234
+ H x W depth map
235
+ validity_map : np.float32
236
+ H x W depth map
237
+ log_space : bool
238
+ if set then produce in log space
239
+ Returns:
240
+ np.float32 : H x W interpolated depth map
241
+ '''
242
+
243
+ assert depth_map.ndim == 2 and validity_map.ndim == 2
244
+
245
+ rows, cols = depth_map.shape
246
+ data_row_idx, data_col_idx = np.where(validity_map)
247
+ depth_values = depth_map[data_row_idx, data_col_idx]
248
+
249
+ # Perform linear interpolation in log space
250
+ if log_space:
251
+ depth_values = np.log(depth_values)
252
+
253
+ interpolator = LinearNDInterpolator(
254
+ # points=Delaunay(np.stack([data_row_idx, data_col_idx], axis=1).astype(np.float32)),
255
+ points=np.stack([data_row_idx, data_col_idx], axis=1),
256
+ values=depth_values,
257
+ fill_value=0 if not log_space else np.log(1e-3))
258
+
259
+ query_row_idx, query_col_idx = np.meshgrid(
260
+ np.arange(rows), np.arange(cols), indexing='ij')
261
+
262
+ query_coord = np.stack(
263
+ [query_row_idx.ravel(), query_col_idx.ravel()], axis=1)
264
+
265
+ Z = interpolator(query_coord).reshape([rows, cols])
266
+
267
+ if log_space:
268
+ Z = np.exp(Z)
269
+ Z[Z < 1e-1] = 0.0
270
+
271
+ return Z
272
+
273
+
274
+ def interpolate_depth_ZJU(depth_map, validity_map=None, log_space=False, window_size=12):
275
+ '''
276
+ Interpolate sparse depth with barycentric coordinates
277
+ Args:
278
+ depth_map : np.float32
279
+ H x W depth map
280
+ validity_map : np.float32
281
+ H x W depth map
282
+ log_space : bool
283
+ if set then produce in log space
284
+ window_size : int
285
+ size of the window for checking validity
286
+ Returns:
287
+ np.float32 : H x W interpolated depth map
288
+ '''
289
+ assert depth_map.ndim == 2
290
+ if validity_map is None:
291
+ validity_map = depth_map > 0.0
292
+ rows, cols = depth_map.shape
293
+ data_row_idx, data_col_idx = np.where(validity_map)
294
+ depth_values = depth_map[data_row_idx, data_col_idx]
295
+ # Perform linear interpolation in log space
296
+ if log_space:
297
+ depth_values = np.log(depth_values)
298
+ interpolator = LinearNDInterpolator(
299
+ points=np.stack([data_row_idx, data_col_idx], axis=1),
300
+ values=depth_values,
301
+ fill_value=0 if not log_space else np.log(1e-3))
302
+ query_row_idx, query_col_idx = np.meshgrid(np.arange(rows), np.arange(cols), indexing='ij')
303
+ Z = np.zeros_like(depth_map)
304
+
305
+ # Create window indices for each query point
306
+ query_indices = np.stack([query_row_idx.ravel(), query_col_idx.ravel()], axis=1)
307
+ window_indices = np.indices((window_size, window_size)).reshape(2, -1) - window_size // 2
308
+
309
+ # Calculate window indices for each query point
310
+ window_row_indices = np.clip(query_indices[:, 0, None] + window_indices[0], 0, rows - 1)
311
+ window_col_indices = np.clip(query_indices[:, 1, None] + window_indices[1], 0, cols - 1)
312
+
313
+ # Get window values and check validity
314
+ window_values = depth_map[window_row_indices, window_col_indices]
315
+ valid_indices = np.any(window_values > 0, axis=1)
316
+
317
+ # Interpolate for valid query points
318
+ valid_query_indices = np.where(valid_indices)[0]
319
+ valid_query_coords = query_indices[valid_query_indices]
320
+ Z.ravel()[valid_query_indices] = interpolator(valid_query_coords)
321
+
322
+ if log_space:
323
+ Z = np.exp(Z)
324
+ Z[Z < 1e-1] = 0.0
325
+
326
+ return Z
src/Baselines/radarcam-depth/data/datasets.py ADDED
@@ -0,0 +1,392 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.utils.data
3
+ from torch.utils.data import Dataset
4
+ import numpy as np
5
+ import data.data_utils as data_utils
6
+ import random
7
+ import os
8
+ from PIL import Image
9
+ from data.data_utils import load_depth
10
+
11
+
12
+ def random_sample(T):
13
+ '''
14
+ Arg(s):
15
+ T : numpy[float32]
16
+ C x N array
17
+ Returns:
18
+ numpy[float32] : random sample from T
19
+ '''
20
+
21
+ index = np.random.randint(0, T.shape[0])
22
+ return T[index, :]
23
+
24
+
25
+ def random_crop(inputs, shape, crop_type=['none']):
26
+ '''
27
+ Apply crop to inputs e.g. images, depth
28
+
29
+ Arg(s):
30
+ inputs : list[numpy[float32]]
31
+ list of numpy arrays e.g. images, depth, and validity maps
32
+ shape : list[int]
33
+ shape (height, width) to crop inputs
34
+ crop_type : str
35
+ none, horizontal, vertical, anchored, top, bottom, left, right, center
36
+ Return:
37
+ list[numpy[float32]] : list of cropped inputs
38
+ '''
39
+
40
+ n_height, n_width = shape
41
+ _, o_height, o_width = inputs[0].shape
42
+
43
+ # Get delta of crop and original height and width
44
+
45
+ d_height = o_height - n_height
46
+ d_width = o_width - n_width
47
+
48
+ # By default, perform center crop
49
+ y_start = d_height // 2
50
+ x_start = d_width // 2
51
+
52
+ # If left alignment, then set starting height to 0
53
+ if 'left' in crop_type:
54
+ x_start = 0
55
+
56
+ # If right alignment, then set starting height to right most position
57
+ elif 'right' in crop_type:
58
+ x_start = d_width
59
+
60
+ elif 'horizontal' in crop_type:
61
+
62
+ # Select from one of the pre-defined anchored locations
63
+ if 'anchored' in crop_type:
64
+ # Create anchor positions
65
+ crop_anchors = [
66
+ 0.0, 0.50, 1.0
67
+ ]
68
+
69
+ widths = [
70
+ anchor * d_width for anchor in crop_anchors
71
+ ]
72
+ x_start = int(widths[np.random.randint(low=0, high=len(widths))])
73
+
74
+ # Randomly select a crop location
75
+ else:
76
+ x_start = np.random.randint(low=0, high=d_width)
77
+
78
+ # If top alignment, then set starting height to 0
79
+ if 'top' in crop_type:
80
+ y_start = 0
81
+
82
+ # If bottom alignment, then set starting height to lowest position
83
+ elif 'bottom' in crop_type:
84
+ y_start = d_height
85
+
86
+ elif 'vertical' in crop_type and np.random.rand() <= 0.30:
87
+
88
+ # Select from one of the pre-defined anchored locations
89
+ if 'anchored' in crop_type:
90
+ # Create anchor positions
91
+ crop_anchors = [
92
+ 0.0, 0.50, 1.0
93
+ ]
94
+
95
+ heights = [
96
+ anchor * d_height for anchor in crop_anchors
97
+ ]
98
+ y_start = int(heights[np.random.randint(low=0, high=len(heights))])
99
+
100
+ # Randomly select a crop location
101
+ else:
102
+ y_start = np.random.randint(low=0, high=d_height)
103
+
104
+ elif 'center' in crop_type:
105
+ pass
106
+
107
+ # Crop each input into (n_height, n_width)
108
+ y_end = y_start + n_height
109
+ x_end = x_start + n_width
110
+
111
+ outputs = [
112
+ T[:, y_start:y_end, x_start:x_end] for T in inputs
113
+ ]
114
+
115
+ return outputs
116
+
117
+
118
+ class RCNetTrainingDataset(torch.utils.data.Dataset):
119
+ '''
120
+ Dataset for fetching:
121
+ (1) image
122
+ (2) radar point
123
+ (3) ground truth
124
+ (4) bounding boxes for the points
125
+ (5) image crops for summary part of the code
126
+
127
+ Arg(s):
128
+ image_paths : list[str]
129
+ paths to images
130
+ radar_paths : list[str]
131
+ paths to radar points
132
+ ground_truth_paths : list[str]
133
+ paths to ground truth depth maps
134
+ crop_width : int
135
+ width of crop centered at the radar point
136
+ total_points_sampled: int
137
+ 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
138
+ sample_probability_of_lidar: int
139
+ randomly sample lidar with this probability and add noise to it instead of using radar points
140
+ min_radar_depth_m: float
141
+ minimum depth accepted for synthetic radar sampling
142
+ max_radar_depth_m: float
143
+ maximum depth accepted for synthetic radar sampling
144
+ '''
145
+
146
+ def __init__(self,
147
+ image_paths,
148
+ radar_paths,
149
+ ground_truth_paths,
150
+ patch_size,
151
+ total_points_sampled,
152
+ sample_probability_of_lidar,
153
+ min_radar_depth_m=0.05,
154
+ max_radar_depth_m=11.2):
155
+
156
+ self.n_sample = len(image_paths)
157
+
158
+ assert self.n_sample == len(ground_truth_paths)
159
+ assert self.n_sample == len(radar_paths)
160
+
161
+ self.image_paths = image_paths
162
+ self.radar_paths = radar_paths
163
+ self.ground_truth_paths = ground_truth_paths
164
+
165
+ self.patch_size = patch_size
166
+ self.pad_size_x = patch_size[1] // 2
167
+ self.padding = ((0, 0), (0, 0), (self.pad_size_x, self.pad_size_x))
168
+
169
+ self.data_format = 'CHW'
170
+ self.total_points_sampled = total_points_sampled
171
+ self.sample_probability_of_lidar = sample_probability_of_lidar
172
+ self.min_radar_depth_m = min_radar_depth_m
173
+ self.max_radar_depth_m = max_radar_depth_m
174
+
175
+ def __getitem__(self, index):
176
+
177
+ # Load image
178
+ image = data_utils.load_image(
179
+ self.image_paths[index],
180
+ normalize=False,
181
+ data_format=self.data_format)
182
+
183
+ height, width = image.shape[1:]
184
+ if height == 720: # ZJU dataset
185
+ image = image[:, 720 // 3: 720 // 4 * 3, :]
186
+
187
+ image = np.pad(
188
+ image,
189
+ pad_width=self.padding,
190
+ mode='edge')
191
+
192
+ # Load radar points N x 3
193
+ radar_points = np.load(self.radar_paths[index])
194
+
195
+ if height == 720:
196
+ radar_points = radar_points[radar_points[:, 1] < 720 // 4 * 3]
197
+ radar_points[:, 1] = radar_points[:, 1] - 720 // 3
198
+ radar_points = radar_points[radar_points[:, 1] >= 0]
199
+
200
+ if radar_points.ndim == 1:
201
+ # Only one point (,3), expand to 1 x 3
202
+ radar_points = np.expand_dims(radar_points, axis=0)
203
+
204
+ # Store bounding boxes for all radar points
205
+ bounding_boxes_list = []
206
+
207
+ # randomly sample radar points to output
208
+ if radar_points.shape[0] <= self.total_points_sampled:
209
+ radar_points = np.repeat(radar_points, 100, axis=0)
210
+ random_idx = np.random.randint(radar_points.shape[0], size=self.total_points_sampled)
211
+ radar_points = radar_points[random_idx, :]
212
+
213
+ # Load ground truth depth
214
+ ground_truth = data_utils.load_depth(
215
+ self.ground_truth_paths[index],
216
+ data_format=self.data_format)
217
+
218
+ if height == 720:
219
+ ground_truth = ground_truth[:, 720 // 3: 720 // 4 * 3]
220
+
221
+ if random.random() < self.sample_probability_of_lidar:
222
+ ground_truth_for_sampling = np.copy(ground_truth)
223
+ ground_truth_for_sampling = ground_truth_for_sampling.squeeze()
224
+ valid_lidar = np.isfinite(ground_truth_for_sampling)
225
+ valid_lidar &= ground_truth_for_sampling >= self.min_radar_depth_m
226
+ valid_lidar &= ground_truth_for_sampling <= self.max_radar_depth_m
227
+ idx_lidar_samples = np.where(valid_lidar)
228
+ n_lidar_samples = len(idx_lidar_samples[0])
229
+
230
+ if n_lidar_samples > 0:
231
+ # Keep the fixed point count required by RC-Net. Replacement
232
+ # handles frames with fewer valid GT pixels than requested.
233
+ if n_lidar_samples >= self.total_points_sampled:
234
+ random_indices = random.sample(
235
+ range(n_lidar_samples), self.total_points_sampled
236
+ )
237
+ else:
238
+ random_indices = np.random.choice(
239
+ n_lidar_samples,
240
+ size=self.total_points_sampled,
241
+ replace=True,
242
+ )
243
+
244
+ points_x = idx_lidar_samples[1][random_indices]
245
+ points_y = idx_lidar_samples[0][random_indices]
246
+ points_z = ground_truth_for_sampling[points_y, points_x]
247
+
248
+ noise_for_fake_radar_x = np.random.normal(0, 25, radar_points.shape[0])
249
+ noise_for_fake_radar_z = np.random.uniform(low=0.0, high=0.4, size=radar_points.shape[0])
250
+
251
+ fake_radar_points = np.copy(radar_points)
252
+ fake_radar_points[:, 0] = points_x + noise_for_fake_radar_x
253
+ fake_radar_points[:, 0] = np.clip(fake_radar_points[:, 0], 0, ground_truth_for_sampling.shape[1])
254
+ fake_radar_points[:, 2] = points_z + noise_for_fake_radar_z
255
+ # we keep the y as the same it is since it is erroneous
256
+
257
+ # convert x and y indices back to int after adding noise
258
+ fake_radar_points[:, 0] = fake_radar_points[:, 0].astype(int)
259
+ fake_radar_points[:, 1] = fake_radar_points[:, 1].astype(int)
260
+
261
+ radar_points = np.copy(fake_radar_points)
262
+
263
+ # get the shifted radar points after padding
264
+ for radar_point_idx in range(0, radar_points.shape[0]):
265
+ # Set radar point to the center of the patch
266
+ radar_points[radar_point_idx, 0] = radar_points[radar_point_idx, 0] + self.pad_size_x
267
+
268
+ bounding_box = [0, 0, 0, 0]
269
+ bounding_box[0] = radar_points[radar_point_idx, 0] - self.pad_size_x
270
+ bounding_box[1] = 0
271
+ bounding_box[2] = radar_points[radar_point_idx, 0] + self.pad_size_x
272
+ bounding_box[3] = self.patch_size[0]
273
+ bounding_boxes_list.append(np.asarray(bounding_box))
274
+
275
+ ground_truth = np.pad(
276
+ ground_truth,
277
+ pad_width=self.padding,
278
+ mode='constant',
279
+ constant_values=0)
280
+
281
+ ground_truth_crops = []
282
+
283
+ # Crop image and ground truth
284
+ for radar_point_idx in range(0, radar_points.shape[0]):
285
+ start_x = int(radar_points[radar_point_idx, 0] - self.pad_size_x)
286
+ end_x = int(radar_points[radar_point_idx, 0] + self.pad_size_x)
287
+ start_y = image.shape[-2] - self.patch_size[0]
288
+
289
+ ground_truth_cropped = ground_truth[:, start_y:, start_x:end_x]
290
+ ground_truth_crops.append(ground_truth_cropped)
291
+
292
+ image = image[:, start_y:, ...]
293
+
294
+ ground_truth = np.asarray(ground_truth_crops)
295
+
296
+ # Convert to float32
297
+ image, radar_points, ground_truth = [
298
+ T.astype(np.float32)
299
+ for T in [image, radar_points, ground_truth]
300
+ ]
301
+
302
+ bounding_boxes_list = [T.astype(np.float32) for T in bounding_boxes_list]
303
+
304
+ bounding_boxes_list = np.stack(bounding_boxes_list, axis=0)
305
+
306
+ return image, radar_points, bounding_boxes_list, ground_truth
307
+
308
+ def __len__(self):
309
+ return self.n_sample
310
+
311
+
312
+ class RCNetInferenceDataset(torch.utils.data.Dataset):
313
+ '''
314
+ Dataset for fetching:
315
+ (1) image
316
+ (2) radar points
317
+ (3) ground truth (if available)
318
+
319
+ Arg(s):
320
+ image_paths : list[str]
321
+ paths to images
322
+ radar_paths : list[str]
323
+ paths to radar points
324
+ ground_truth_paths : list[str]
325
+ paths to ground truth paths
326
+ '''
327
+
328
+ def __init__(self, image_paths, radar_paths, ground_truth_paths=None):
329
+
330
+ self.n_sample = len(image_paths)
331
+
332
+ assert self.n_sample == len(radar_paths)
333
+
334
+ self.image_paths = image_paths
335
+ self.radar_paths = radar_paths
336
+
337
+ if ground_truth_paths is not None and None not in ground_truth_paths:
338
+ assert self.n_sample == len(ground_truth_paths)
339
+ self.ground_truth_available = True
340
+ else:
341
+ self.ground_truth_available = False
342
+
343
+ self.ground_truth_paths = ground_truth_paths
344
+
345
+ self.data_format = 'CHW'
346
+
347
+ def __getitem__(self, index):
348
+
349
+ # Load image
350
+ image = data_utils.load_image(
351
+ self.image_paths[index],
352
+ normalize=False,
353
+ data_format=self.data_format)
354
+
355
+ height, width = image.shape[1:]
356
+ if height == 720: # ZJU dataset
357
+ image = image[:, 720 // 3: 720 // 4 * 3, :]
358
+
359
+ # Load radar points N x 3
360
+ radar_points = np.load(self.radar_paths[index])
361
+
362
+ if height == 720:
363
+ radar_points = radar_points[radar_points[:, 1] < 720 // 4 * 3]
364
+ radar_points[:, 1] = radar_points[:, 1] - 720 // 3
365
+ radar_points = radar_points[radar_points[:, 1] >= 0]
366
+
367
+ if radar_points.ndim == 1:
368
+ # Expand to 1 x 3
369
+ radar_points = np.expand_dims(radar_points, axis=0)
370
+
371
+ inputs = [image, radar_points]
372
+
373
+ if self.ground_truth_available:
374
+ # Load ground truth depth
375
+ ground_truth = data_utils.load_depth(
376
+ self.ground_truth_paths[index],
377
+ data_format=self.data_format)
378
+ if height == 720:
379
+ ground_truth = ground_truth[:, 720 // 3: 720 // 4 * 3]
380
+
381
+ inputs.append(ground_truth)
382
+
383
+ # Convert to float32
384
+ inputs = [
385
+ T.astype(np.float32)
386
+ for T in inputs
387
+ ]
388
+
389
+ return inputs
390
+
391
+ def __len__(self):
392
+ return self.n_sample
src/Baselines/radarcam-depth/linear_attention.py ADDED
@@ -0,0 +1,184 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from torch.nn import Module, Dropout
3
+ import torch.nn as nn
4
+ import copy
5
+
6
+
7
+ def elu_feature_map(x):
8
+ return torch.nn.functional.elu(x) + 1
9
+
10
+
11
+
12
+ class LinearAttention(Module):
13
+ def __init__(self, eps=1e-6):
14
+ super().__init__()
15
+ self.feature_map = elu_feature_map
16
+ self.eps = eps
17
+
18
+ def forward(self, queries, keys, values, q_mask=None, kv_mask=None):
19
+ """ Multi-Head linear attention proposed in "Transformers are RNNs"
20
+ Args:
21
+ queries: [N, L, H, D]
22
+ keys: [N, S, H, D]
23
+ values: [N, S, H, D]
24
+ q_mask: [N, L]
25
+ kv_mask: [N, S]
26
+ Returns:
27
+ queried_values: (N, L, H, D)
28
+ """
29
+ Q = self.feature_map(queries)
30
+ K = self.feature_map(keys)
31
+
32
+ # set padded position to zero
33
+ if q_mask is not None:
34
+ Q = Q * q_mask[:, :, None, None]
35
+ if kv_mask is not None:
36
+ K = K * kv_mask[:, :, None, None]
37
+ values = values * kv_mask[:, :, None, None]
38
+
39
+ v_length = values.size(1)
40
+ values = values / v_length # prevent fp16 overflow
41
+ KV = torch.einsum("nshd,nshv->nhdv", K, values) # (S,D)' @ S,V
42
+ Z = 1 / (torch.einsum("nlhd,nhd->nlh", Q, K.sum(dim=1)) + self.eps)
43
+ queried_values = torch.einsum("nlhd,nhdv,nlh->nlhv", Q, KV, Z) * v_length
44
+
45
+ return queried_values.contiguous()
46
+
47
+
48
+
49
+ class FullAttention(Module):
50
+ def __init__(self, use_dropout=False, attention_dropout=0.1):
51
+ super().__init__()
52
+ self.use_dropout = use_dropout
53
+ self.dropout = Dropout(attention_dropout)
54
+
55
+ def forward(self, queries, keys, values, q_mask=None, kv_mask=None):
56
+ """ Multi-head scaled dot-product attention, a.k.a full attention.
57
+ Args:
58
+ queries: [N, L, H, D]
59
+ keys: [N, S, H, D]
60
+ values: [N, S, H, D]
61
+ q_mask: [N, L]
62
+ kv_mask: [N, S]
63
+ Returns:
64
+ queried_values: (N, L, H, D)
65
+ """
66
+
67
+ # Compute the unnormalized attention and apply the masks
68
+ QK = torch.einsum("nlhd,nshd->nlsh", queries, keys)
69
+ if kv_mask is not None:
70
+ QK.masked_fill_(~(q_mask[:, :, None, None] * kv_mask[:, None, :, None]), float('-inf'))
71
+
72
+ # Compute the attention and the weighted average
73
+ softmax_temp = 1. / queries.size(3)**.5 # sqrt(D)
74
+ A = torch.softmax(softmax_temp * QK, dim=2)
75
+ if self.use_dropout:
76
+ A = self.dropout(A)
77
+
78
+ queried_values = torch.einsum("nlsh,nshd->nlhd", A, values)
79
+
80
+ return queried_values.contiguous()
81
+
82
+
83
+
84
+ class LoFTREncoderLayer(nn.Module):
85
+ def __init__(self,
86
+ d_model,
87
+ nhead,
88
+ attention='linear'):
89
+ super(LoFTREncoderLayer, self).__init__()
90
+
91
+ self.dim = d_model // nhead
92
+ self.nhead = nhead
93
+
94
+ # multi-head attention
95
+ self.q_proj = nn.Linear(d_model, d_model, bias=False)
96
+ self.k_proj = nn.Linear(d_model, d_model, bias=False)
97
+ self.v_proj = nn.Linear(d_model, d_model, bias=False)
98
+ self.attention = LinearAttention() if attention == 'linear' else FullAttention()
99
+ self.merge = nn.Linear(d_model, d_model, bias=False)
100
+
101
+ # feed-forward network
102
+ self.mlp = nn.Sequential(
103
+ nn.Linear(d_model*2, d_model*2, bias=False),
104
+ nn.ReLU(True),
105
+ nn.Linear(d_model*2, d_model, bias=False),
106
+ )
107
+
108
+ # norm and dropout
109
+ self.norm1 = nn.LayerNorm(d_model)
110
+ self.norm2 = nn.LayerNorm(d_model)
111
+
112
+ def forward(self, x, source, x_mask=None, source_mask=None):
113
+ """
114
+ Args:
115
+ x (torch.Tensor): [N, L, C]
116
+ source (torch.Tensor): [N, S, C]
117
+ x_mask (torch.Tensor): [N, L] (optional)
118
+ source_mask (torch.Tensor): [N, S] (optional)
119
+ """
120
+ bs = x.size(0)
121
+ query, key, value = x, source, source
122
+
123
+ # multi-head attention
124
+ query = self.q_proj(query).view(bs, -1, self.nhead, self.dim) # [N, L, (H, D)]
125
+ key = self.k_proj(key).view(bs, -1, self.nhead, self.dim) # [N, S, (H, D)]
126
+ value = self.v_proj(value).view(bs, -1, self.nhead, self.dim)
127
+ message = self.attention(query, key, value, q_mask=x_mask, kv_mask=source_mask) # [N, L, (H, D)]
128
+ message = self.merge(message.view(bs, -1, self.nhead*self.dim)) # [N, L, C]
129
+ message = self.norm1(message)
130
+
131
+ # feed-forward network
132
+ message = self.mlp(torch.cat([x, message], dim=2))
133
+ message = self.norm2(message)
134
+
135
+ return x + message
136
+
137
+
138
+
139
+ class LocalFeatureTransformer(nn.Module):
140
+ """A Local Feature Transformer (LoFTR) module."""
141
+
142
+ def __init__(self, type, n_layers=1, d_model=256, nhead=8, attention='linear'):
143
+ super(LocalFeatureTransformer, self).__init__()
144
+
145
+ self.d_model = d_model
146
+ self.nhead = nhead
147
+ self.layer_names = type * n_layers
148
+ self.attention = attention
149
+ encoder_layer = LoFTREncoderLayer(self.d_model, self.nhead, self.attention)
150
+
151
+ self.layers = nn.ModuleList([copy.deepcopy(encoder_layer) for _ in range(len(self.layer_names))])
152
+ self._reset_parameters()
153
+
154
+ def _reset_parameters(self):
155
+ for p in self.parameters():
156
+ if p.dim() > 1:
157
+ nn.init.xavier_uniform_(p)
158
+
159
+ def forward(self, feat0, feat1, mask0=None, mask1=None):
160
+ """
161
+ Args:
162
+ feat0 (torch.Tensor): [N, L, C]
163
+ feat1 (torch.Tensor): [N, S, C]
164
+ mask0 (torch.Tensor): [N, L] (optional)
165
+ mask1 (torch.Tensor): [N, S] (optional)
166
+ """
167
+
168
+ assert self.d_model == feat0.size(2), "the feature number of src and transformer must be equal"
169
+
170
+ for layer, name in zip(self.layers, self.layer_names):
171
+ # if name == 'self0':
172
+ # feat0 = layer(feat0, feat0, mask0, mask0)
173
+ # elif name == 'self1':
174
+ # feat1 = layer(feat1, feat1, mask1, mask1)
175
+ if name == 'self':
176
+ feat0 = layer(feat0, feat0, mask0, mask0)
177
+ feat1 = layer(feat1, feat1, mask1, mask1)
178
+ elif name == 'cross':
179
+ feat0 = layer(feat0, feat1, mask0, mask1)
180
+ feat1 = layer(feat1, feat0, mask1, mask0)
181
+ else:
182
+ raise KeyError
183
+
184
+ return feat0, feat1
src/Baselines/radarcam-depth/modules/estimator.py ADDED
@@ -0,0 +1,188 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import time
3
+ from scipy.optimize import minimize_scalar
4
+
5
+ def compute_scale_and_shift_ls(prediction, target, mask):
6
+ # tuple specifying with axes to sum
7
+ sum_axes = (0, 1)
8
+
9
+ # system matrix: A = [[a_00, a_01], [a_10, a_11]]
10
+ a_00 = np.sum(mask * prediction * prediction, sum_axes)
11
+ a_01 = np.sum(mask * prediction, sum_axes)
12
+ a_11 = np.sum(mask, sum_axes)
13
+
14
+ # right hand side: b = [b_0, b_1]
15
+ b_0 = np.sum(mask * prediction * target, sum_axes)
16
+ b_1 = np.sum(mask * target, sum_axes)
17
+
18
+ # solution: x = A^-1 . b = [[a_11, -a_01], [-a_10, a_00]] / (a_00 * a_11 - a_01 * a_10) . b
19
+ x_0 = np.zeros_like(b_0)
20
+ x_1 = np.zeros_like(b_1)
21
+
22
+ det = a_00 * a_11 - a_01 * a_01
23
+ # A needs to be a positive definite matrix.
24
+ valid = det > 0
25
+
26
+ x_0[valid] = (a_11[valid] * b_0[valid] - a_01[valid] * b_1[valid]) / det[valid]
27
+ x_1[valid] = (-a_01[valid] * b_0[valid] + a_00[valid] * b_1[valid]) / det[valid]
28
+
29
+ return x_0, x_1
30
+
31
+
32
+
33
+ def compute_scale_and_shift_ransac(prediction, target, mask,
34
+ num_iterations, sample_size,
35
+ inlier_threshold, inlier_ratio_threshold):
36
+ # start = time.time()
37
+ best_scale = 0.0
38
+ best_shift = 0.0
39
+ best_inlier_count = 0
40
+
41
+ valid_indices = np.where(mask)
42
+ valid_count = len(valid_indices[0])
43
+ # print('valid_count: ', valid_count)
44
+
45
+ for _ in range(num_iterations):
46
+ if valid_count < sample_size:
47
+ break
48
+
49
+ # Randomly sample from valid indices
50
+ indices = np.random.choice(valid_count, size=sample_size, replace=False)
51
+ mask_sample = np.zeros_like(mask)
52
+ mask_sample[valid_indices[0][indices], valid_indices[1][indices]] = 1
53
+
54
+ # Calculate x_0 and x_1 for the sampled data
55
+ sum_axes = (0, 1)
56
+ a_00 = np.sum(mask_sample * prediction * prediction, sum_axes)
57
+ a_01 = np.sum(mask_sample * prediction, sum_axes)
58
+ a_11 = np.sum(mask_sample, sum_axes)
59
+ b_0 = np.sum(mask_sample * prediction * target, sum_axes)
60
+ b_1 = np.sum(mask_sample * target, sum_axes)
61
+ det = a_00 * a_11 - a_01 * a_01
62
+ valid = det > 0
63
+ x_0 = np.zeros_like(b_0)
64
+ x_1 = np.zeros_like(b_1)
65
+ x_0[valid] = (a_11[valid] * b_0[valid] - a_01[valid] * b_1[valid]) / det[valid]
66
+ x_1[valid] = (-a_01[valid] * b_0[valid] + a_00[valid] * b_1[valid]) / det[valid]
67
+
68
+ # Calculate residuals and count inliers
69
+ residuals = np.abs(mask * prediction * x_0 + x_1 - mask * target)
70
+ residuals = residuals[mask]
71
+
72
+ inlier_count = np.sum(residuals < inlier_threshold)
73
+
74
+ # Update best model if current model has more inliers
75
+ if inlier_count > best_inlier_count:
76
+ best_scale = x_0
77
+ best_shift = x_1
78
+ best_inlier_count = inlier_count
79
+ inlier_ratio = inlier_count / valid_count
80
+ if inlier_ratio > inlier_ratio_threshold:
81
+ break
82
+
83
+ print('best_inlier_count: ', best_inlier_count)
84
+ print('inlier_ratio: ', best_inlier_count / valid_count)
85
+ # print('best_scale: ', best_scale)
86
+ # print('best_shift: ', best_shift)
87
+ # print('time', time.time() - start)
88
+ return best_scale, best_shift
89
+
90
+
91
+
92
+ class LeastSquaresEstimator(object):
93
+ def __init__(self, estimate, target, valid):
94
+ self.estimate = estimate
95
+ self.target = target
96
+ self.valid = valid
97
+
98
+ # to be computed
99
+ self.scale = 1.0
100
+ self.shift = 0.0
101
+ self.output = None
102
+
103
+ def compute_scale_and_shift_ran(self,
104
+ num_iterations=60, sample_size=5,
105
+ inlier_threshold=0.02, inlier_ratio_threshold=0.8):
106
+ self.scale, self.shift = compute_scale_and_shift_ransac(self.estimate, self.target, self.valid,
107
+ num_iterations, sample_size,
108
+ inlier_threshold, inlier_ratio_threshold)
109
+
110
+ def compute_scale_and_shift(self):
111
+ self.scale, self.shift = compute_scale_and_shift_ls(self.estimate, self.target, self.valid)
112
+
113
+
114
+ def apply_scale_and_shift(self):
115
+ self.output = self.estimate * self.scale + self.shift
116
+
117
+ def clamp_min_max(self, clamp_min=None, clamp_max=None):
118
+ if clamp_min is not None:
119
+ if clamp_min > 0:
120
+ clamp_min_inv = 1.0/clamp_min
121
+ self.output[self.output > clamp_min_inv] = clamp_min_inv
122
+ assert np.max(self.output) <= clamp_min_inv
123
+ else: # divide by zero, so skip
124
+ pass
125
+ if clamp_max is not None:
126
+ clamp_max_inv = 1.0/clamp_max
127
+ self.output[self.output < clamp_max_inv] = clamp_max_inv
128
+
129
+
130
+
131
+ def objective_function(x_0, prediction, target, mask):
132
+ # Calculate x_0 * prediction
133
+ x_0_prediction = x_0 * prediction
134
+ # Calculate the error between x_0 * prediction and target, using the mask
135
+ error = np.sum(mask * abs(x_0_prediction - target))
136
+ return error
137
+
138
+
139
+
140
+ class Optimizer(object):
141
+ def __init__(self, estimate, target, valid, depth_type):
142
+ self.estimate = estimate
143
+ self.target = target
144
+ self.valid = valid
145
+ self.depth_type = depth_type
146
+ # to be computed
147
+ self.scale = 1.0
148
+ self.output = None
149
+
150
+ def optimize_scale(self):
151
+ if self.depth_type == 'inv':
152
+ bounds = (0.0003, 0.01)
153
+ else:
154
+ bounds = (0.5, 1.6) # pos
155
+
156
+ # Minimize the objective function using scipy.optimize.minimize_scalar
157
+ result = minimize_scalar(
158
+ objective_function, args=(self.estimate, self.target, self.valid),
159
+ bounds=bounds
160
+ )
161
+
162
+ # Extract the optimized x_0 value from the result
163
+ optimized_x_0 = result.x
164
+ self.scale = optimized_x_0
165
+
166
+ def apply_scale(self):
167
+ self.output = self.estimate * self.scale
168
+
169
+ def clamp_min_max(self, clamp_min=None, clamp_max=None):
170
+ if clamp_min is not None:
171
+ if clamp_min > 0:
172
+ clamp_min_inv = 1.0/clamp_min
173
+ self.output[self.output > clamp_min_inv] = clamp_min_inv
174
+ assert np.max(self.output) <= clamp_min_inv
175
+ else: # divide by zero, so skip
176
+ pass
177
+ if clamp_max is not None:
178
+ clamp_max_inv = 1.0/clamp_max
179
+ self.output[self.output < clamp_max_inv] = clamp_max_inv
180
+
181
+ def clamp_min_max_pos(self, clamp_min=None, clamp_max=None):
182
+ if clamp_min is not None:
183
+ if clamp_min >= 0:
184
+ self.output[self.output < clamp_min] = clamp_min
185
+ else:
186
+ pass
187
+ if clamp_max is not None:
188
+ self.output[self.output > clamp_max] = clamp_max
src/Baselines/radarcam-depth/modules/midas/base_model.py ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from safetensors.torch import load_file
3
+
4
+
5
+ class BaseModel(torch.nn.Module):
6
+ def load(self, path):
7
+ """Load model from file.
8
+
9
+ Args:
10
+ path (str): file path
11
+ """
12
+ self.load_state_dict(load_file(path, device="cpu"), strict=True)
src/Baselines/radarcam-depth/modules/midas/blocks.py ADDED
@@ -0,0 +1,197 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+
4
+ def _make_encoder(backbone, features, use_pretrained, groups=1, expand=False, exportable=True):
5
+ if backbone == "efficientnet_lite3":
6
+ pretrained = _make_pretrained_efficientnet_lite3(use_pretrained, exportable=exportable)
7
+ scratch = _make_scratch([32, 48, 136, 384], features, groups=groups, expand=expand) # efficientnet_lite3
8
+ else:
9
+ print(f"Backbone '{backbone}' not implemented")
10
+ assert False
11
+
12
+ return pretrained, scratch
13
+
14
+
15
+ def _make_scratch(in_shape, out_shape, groups=1, expand=False):
16
+ scratch = nn.Module()
17
+
18
+ out_shape1 = out_shape
19
+ out_shape2 = out_shape
20
+ out_shape3 = out_shape
21
+ out_shape4 = out_shape
22
+ if expand==True:
23
+ out_shape1 = out_shape
24
+ out_shape2 = out_shape*2
25
+ out_shape3 = out_shape*4
26
+ out_shape4 = out_shape*8
27
+
28
+ scratch.layer1_rn = nn.Conv2d(
29
+ in_shape[0], out_shape1, kernel_size=3, stride=1, padding=1, bias=False, groups=groups
30
+ )
31
+ scratch.layer2_rn = nn.Conv2d(
32
+ in_shape[1], out_shape2, kernel_size=3, stride=1, padding=1, bias=False, groups=groups
33
+ )
34
+ scratch.layer3_rn = nn.Conv2d(
35
+ in_shape[2], out_shape3, kernel_size=3, stride=1, padding=1, bias=False, groups=groups
36
+ )
37
+ scratch.layer4_rn = nn.Conv2d(
38
+ in_shape[3], out_shape4, kernel_size=3, stride=1, padding=1, bias=False, groups=groups
39
+ )
40
+
41
+ return scratch
42
+
43
+
44
+ def _make_pretrained_efficientnet_lite3(use_pretrained, exportable=False):
45
+ efficientnet = torch.hub.load(
46
+ "rwightman/gen-efficientnet-pytorch",
47
+ "tf_efficientnet_lite3",
48
+ pretrained=use_pretrained,
49
+ exportable=exportable,
50
+ trust_repo=True,
51
+ )
52
+ return _make_efficientnet_backbone(efficientnet)
53
+
54
+
55
+ def _make_efficientnet_backbone(effnet):
56
+ pretrained = nn.Module()
57
+
58
+ pretrained.layer1 = nn.Sequential(
59
+ effnet.conv_stem, effnet.bn1, effnet.act1, *effnet.blocks[0:2]
60
+ )
61
+ pretrained.layer2 = nn.Sequential(*effnet.blocks[2:3])
62
+ pretrained.layer3 = nn.Sequential(*effnet.blocks[3:5])
63
+ pretrained.layer4 = nn.Sequential(*effnet.blocks[5:9])
64
+
65
+ return pretrained
66
+
67
+
68
+ class ResidualConvUnit_custom(nn.Module):
69
+ """Residual convolution module.
70
+ """
71
+
72
+ def __init__(self, features, activation, bn):
73
+ """Init.
74
+
75
+ Args:
76
+ features (int): number of features
77
+ """
78
+ super().__init__()
79
+
80
+ self.bn = bn
81
+
82
+ self.groups=1
83
+
84
+ self.conv1 = nn.Conv2d(
85
+ features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups
86
+ )
87
+
88
+ self.conv2 = nn.Conv2d(
89
+ features, features, kernel_size=3, stride=1, padding=1, bias=True, groups=self.groups
90
+ )
91
+
92
+ if self.bn==True:
93
+ self.bn1 = nn.BatchNorm2d(features)
94
+ self.bn2 = nn.BatchNorm2d(features)
95
+
96
+ self.activation = activation
97
+
98
+ self.skip_add = nn.quantized.FloatFunctional()
99
+
100
+ def forward(self, x):
101
+ """Forward pass.
102
+
103
+ Args:
104
+ x (tensor): input
105
+
106
+ Returns:
107
+ tensor: output
108
+ """
109
+
110
+ out = self.activation(x)
111
+ out = self.conv1(out)
112
+ if self.bn==True:
113
+ out = self.bn1(out)
114
+
115
+ out = self.activation(out)
116
+ out = self.conv2(out)
117
+ if self.bn==True:
118
+ out = self.bn2(out)
119
+
120
+ if self.groups > 1:
121
+ out = self.conv_merge(out)
122
+
123
+ return self.skip_add.add(out, x)
124
+
125
+
126
+ class FeatureFusionBlock_custom(nn.Module):
127
+ """Feature fusion block.
128
+ """
129
+
130
+ def __init__(self, features, activation, deconv=False, bn=False, expand=False, align_corners=True):
131
+ """Init.
132
+
133
+ Args:
134
+ features (int): number of features
135
+ """
136
+ super(FeatureFusionBlock_custom, self).__init__()
137
+
138
+ self.deconv = deconv
139
+ self.align_corners = align_corners
140
+
141
+ self.groups=1
142
+
143
+ self.expand = expand
144
+ out_features = features
145
+ if self.expand==True:
146
+ out_features = features//2
147
+
148
+ self.out_conv = nn.Conv2d(features, out_features, kernel_size=1, stride=1, padding=0, bias=True, groups=1)
149
+
150
+ self.resConfUnit1 = ResidualConvUnit_custom(features, activation, bn)
151
+ self.resConfUnit2 = ResidualConvUnit_custom(features, activation, bn)
152
+
153
+ self.skip_add = nn.quantized.FloatFunctional()
154
+
155
+ def forward(self, *xs):
156
+ """Forward pass.
157
+
158
+ Returns:
159
+ tensor: output
160
+ """
161
+ output = xs[0]
162
+
163
+ if len(xs) == 2:
164
+ res = self.resConfUnit1(xs[1])
165
+ output = self.skip_add.add(output, res)
166
+
167
+ output = self.resConfUnit2(output)
168
+
169
+ output = nn.functional.interpolate(
170
+ output, scale_factor=2, mode="bilinear", align_corners=self.align_corners
171
+ )
172
+
173
+ output = self.out_conv(output)
174
+
175
+ return output
176
+
177
+
178
+ class OutputConv(nn.Module):
179
+ """Output conv block.
180
+ """
181
+
182
+ def __init__(self, features, groups, activation, non_negative):
183
+
184
+ super(OutputConv, self).__init__()
185
+
186
+ self.output_conv = nn.Sequential(
187
+ nn.Conv2d(features, features//2, kernel_size=3, stride=1, padding=1, groups=groups),
188
+ nn.Upsample(scale_factor=2, mode="bilinear"),
189
+ nn.Conv2d(features//2, 32, kernel_size=3, stride=1, padding=1),
190
+ activation,
191
+ nn.Conv2d(32, 1, kernel_size=1, stride=1, padding=0),
192
+ nn.ReLU(True) if non_negative else nn.Identity(),
193
+ nn.Identity(),
194
+ )
195
+
196
+ def forward(self, x):
197
+ return self.output_conv(x)
src/Baselines/radarcam-depth/modules/midas/midas_net_custom.py ADDED
@@ -0,0 +1,138 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.nn as nn
3
+
4
+ from torch.nn import functional as F
5
+
6
+ from .base_model import BaseModel
7
+ from .blocks import FeatureFusionBlock_custom, _make_encoder, OutputConv
8
+
9
+ def weights_init(m):
10
+ import math
11
+ # initialize from normal (Gaussian) distribution
12
+ if isinstance(m, nn.Conv2d):
13
+ n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
14
+ m.weight.data.normal_(0, math.sqrt(2.0 / n))
15
+ if m.bias is not None:
16
+ m.bias.data.zero_()
17
+ elif isinstance(m, nn.BatchNorm2d):
18
+ m.weight.data.fill_(1)
19
+ m.bias.data.zero_()
20
+
21
+
22
+ class MidasNet_small_videpth(BaseModel):
23
+ """Network for monocular depth estimation.
24
+ """
25
+
26
+ def __init__(self, device = 'cpu', path=None, features=64, backbone="efficientnet_lite3", non_negative=False, exportable=True, channels_last=False, align_corners=True,
27
+ blocks={'expand': True}, in_channels=2, regress='r', min_pred=None, max_pred=None):
28
+ """Init.
29
+
30
+ Args:
31
+ path (str, optional): Path to saved model. Defaults to None.
32
+ features (int, optional): Number of features. Defaults to 64.
33
+ backbone (str, optional): Backbone network for encoder. Defaults to efficientnet_lite3.
34
+ """
35
+ print("Loading weights: ", path)
36
+
37
+ super(MidasNet_small_videpth, self).__init__()
38
+
39
+ use_pretrained = False
40
+
41
+ self.channels_last = channels_last
42
+ self.blocks = blocks
43
+ self.backbone = backbone
44
+
45
+ self.groups = 1
46
+
47
+ # for model output
48
+ self.regress = regress
49
+ self.min_pred = min_pred
50
+ self.max_pred = max_pred
51
+
52
+ features1=features
53
+ features2=features
54
+ features3=features
55
+ features4=features
56
+ self.expand = False
57
+ if "expand" in self.blocks and self.blocks['expand'] == True:
58
+ self.expand = True
59
+ features1=features
60
+ features2=features*2
61
+ features3=features*4
62
+ features4=features*8
63
+
64
+ self.first = nn.Sequential(
65
+ nn.Conv2d(in_channels, 3, kernel_size=3, stride=1, padding=1),
66
+ nn.BatchNorm2d(3),
67
+ nn.ReLU(inplace=True)
68
+ )
69
+ self.first.apply(weights_init)
70
+
71
+ self.pretrained, self.scratch = _make_encoder(self.backbone, features, use_pretrained, groups=self.groups, expand=self.expand, exportable=exportable)
72
+
73
+ self.scratch.activation = nn.ReLU(False)
74
+
75
+ self.scratch.refinenet4 = FeatureFusionBlock_custom(features4, self.scratch.activation, deconv=False, bn=False, expand=self.expand, align_corners=align_corners)
76
+ self.scratch.refinenet3 = FeatureFusionBlock_custom(features3, self.scratch.activation, deconv=False, bn=False, expand=self.expand, align_corners=align_corners)
77
+ self.scratch.refinenet2 = FeatureFusionBlock_custom(features2, self.scratch.activation, deconv=False, bn=False, expand=self.expand, align_corners=align_corners)
78
+ self.scratch.refinenet1 = FeatureFusionBlock_custom(features1, self.scratch.activation, deconv=False, bn=False, align_corners=align_corners)
79
+
80
+ self.scratch.output_conv = OutputConv(features, self.groups, self.scratch.activation, non_negative)
81
+
82
+ if path:
83
+ self.load(path)
84
+
85
+ self.to(device)
86
+
87
+
88
+ def forward(self, x, d):
89
+ """Forward pass.
90
+
91
+ Args:
92
+ x (tensor): input data (image)
93
+ d (tensor): unalterated input depth
94
+
95
+ Returns:
96
+ tensor: depth
97
+ """
98
+ if self.channels_last==True:
99
+ print("self.channels_last = ", self.channels_last)
100
+ x.contiguous(memory_format=torch.channels_last)
101
+
102
+ layer_0 = self.first(x)
103
+
104
+ layer_1 = self.pretrained.layer1(layer_0)
105
+ layer_2 = self.pretrained.layer2(layer_1)
106
+ layer_3 = self.pretrained.layer3(layer_2)
107
+ layer_4 = self.pretrained.layer4(layer_3)
108
+
109
+ layer_1_rn = self.scratch.layer1_rn(layer_1)
110
+ layer_2_rn = self.scratch.layer2_rn(layer_2)
111
+ layer_3_rn = self.scratch.layer3_rn(layer_3)
112
+ layer_4_rn = self.scratch.layer4_rn(layer_4)
113
+
114
+ path_4 = self.scratch.refinenet4(layer_4_rn)
115
+ path_3 = self.scratch.refinenet3(path_4, layer_3_rn)
116
+ path_2 = self.scratch.refinenet2(path_3, layer_2_rn)
117
+ path_1 = self.scratch.refinenet1(path_2, layer_1_rn)
118
+
119
+ out = self.scratch.output_conv(path_1)
120
+
121
+ scales = F.relu(1.0 + out)
122
+ pred = d * scales
123
+
124
+ # clamp pred to min and max
125
+ if self.min_pred is not None:
126
+ min_pred_inv = 1.0/self.min_pred
127
+ pred[pred > min_pred_inv] = min_pred_inv
128
+ if self.max_pred is not None:
129
+ max_pred_inv = 1.0/self.max_pred
130
+ pred[pred < max_pred_inv] = max_pred_inv
131
+
132
+ # also return scales
133
+ return (pred, scales)
134
+
135
+
136
+
137
+
138
+
src/Baselines/radarcam-depth/modules/midas/normalization.py ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ VOID_INTERMEDIATE = {
2
+
3
+ "dpt_beit_large_512" : {
4
+ "void_150" : {
5
+ "mean" : {"int_depth" : 0.730, "int_scales" : 0.380},
6
+ "std" : {"int_depth" : 0.226, "int_scales" : 0.102},
7
+ },
8
+ "void_500" : {
9
+ "mean" : {"int_depth" : 0.736, "int_scales" : 0.366},
10
+ "std" : {"int_depth" : 0.232, "int_scales" : 0.099},
11
+ },
12
+ "void_1500" : {
13
+ "mean" : {"int_depth" : 0.730, "int_scales" : 0.355},
14
+ "std" : {"int_depth" : 0.232, "int_scales" : 0.096},
15
+ },
16
+ },
17
+
18
+ "dpt_swin2_large_384" : {
19
+ "void_150" : {
20
+ "mean" : {"int_depth" : 0.730, "int_scales" : 0.402},
21
+ "std" : {"int_depth" : 0.219, "int_scales" : 0.107},
22
+ },
23
+ "void_500" : {
24
+ "mean" : {"int_depth" : 0.736, "int_scales" : 0.389},
25
+ "std" : {"int_depth" : 0.224, "int_scales" : 0.106},
26
+ },
27
+ "void_1500" : {
28
+ "mean" : {"int_depth" : 0.730, "int_scales" : 0.377},
29
+ "std" : {"int_depth" : 0.226, "int_scales" : 0.103},
30
+ },
31
+ },
32
+
33
+ "dpt_large" : {
34
+ "void_150" : {
35
+ "mean" : {"int_depth" : 0.729, "int_scales" : 0.403},
36
+ "std" : {"int_depth" : 0.213, "int_scales" : 0.116},
37
+ },
38
+ "void_500" : {
39
+ "mean" : {"int_depth" : 0.735, "int_scales" : 0.390},
40
+ "std" : {"int_depth" : 0.219, "int_scales" : 0.116},
41
+ },
42
+ "void_1500" : {
43
+ "mean" : {"int_depth" : 0.730, "int_scales" : 0.380},
44
+ "std" : {"int_depth" : 0.221, "int_scales" : 0.116},
45
+ },
46
+ },
47
+
48
+ "dpt_hybrid": {
49
+ "void_150" : {
50
+ "mean" : {"int_depth" : 0.729, "int_scales" : 0.404},
51
+ "std" : {"int_depth" : 0.210, "int_scales" : 0.117},
52
+ },
53
+ "void_500" : {
54
+ "mean" : {"int_depth" : 0.735, "int_scales" : 0.392},
55
+ "std" : {"int_depth" : 0.215, "int_scales" : 0.118},
56
+ },
57
+ "void_1500" : {
58
+ "mean" : {"int_depth" : 0.730, "int_scales" : 0.381},
59
+ "std" : {"int_depth" : 0.218, "int_scales" : 0.117},
60
+ },
61
+ },
62
+
63
+ "dpt_swin2_tiny_256" : {
64
+ "void_150" : {
65
+ "mean" : {"int_depth" : 0.735, "int_scales" : 0.419},
66
+ "std" : {"int_depth" : 0.207, "int_scales" : 0.122},
67
+ },
68
+ "void_500" : {
69
+ "mean" : {"int_depth" : 0.741, "int_scales" : 0.406},
70
+ "std" : {"int_depth" : 0.212, "int_scales" : 0.124},
71
+ },
72
+ "void_1500" : {
73
+ "mean" : {"int_depth" : 0.733, "int_scales" : 0.396},
74
+ "std" : {"int_depth" : 0.213, "int_scales" : 0.125},
75
+ },
76
+ },
77
+
78
+ "dpt_levit_224" : {
79
+ "void_150" : {
80
+ "mean" : {"int_depth" : 0.734, "int_scales" : 0.421},
81
+ "std" : {"int_depth" : 0.198, "int_scales" : 0.129},
82
+ },
83
+ "void_500" : {
84
+ "mean" : {"int_depth" : 0.740, "int_scales" : 0.410},
85
+ "std" : {"int_depth" : 0.202, "int_scales" : 0.134},
86
+ },
87
+ "void_1500" : {
88
+ "mean" : {"int_depth" : 0.734, "int_scales" : 0.400},
89
+ "std" : {"int_depth" : 0.204, "int_scales" : 0.137},
90
+ },
91
+ },
92
+
93
+ "midas_small" : {
94
+ "void_150" : {
95
+ "mean" : {"int_depth" : 0.723, "int_scales" : 0.402},
96
+ "std" : {"int_depth" : 0.190, "int_scales" : 0.132},
97
+ },
98
+ "void_500" : {
99
+ "mean" : {"int_depth" : 0.731, "int_scales" : 0.393},
100
+ "std" : {"int_depth" : 0.196, "int_scales" : 0.136},
101
+ },
102
+ "void_1500" : {
103
+ "mean" : {"int_depth" : 0.728, "int_scales" : 0.385},
104
+ "std" : {"int_depth" : 0.199, "int_scales" : 0.140},
105
+ },
106
+ },
107
+
108
+ }
109
+
src/Baselines/radarcam-depth/modules/midas/transforms.py ADDED
@@ -0,0 +1,263 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import cv2
3
+ import math
4
+ import torch
5
+ import torchvision.transforms as transforms
6
+
7
+ from modules.midas.utils import normalize_unit_range
8
+ import modules.midas.normalization as normalization
9
+
10
+ class Resize(object):
11
+ """Resize sample to given size (width, height).
12
+ """
13
+
14
+ def __init__(
15
+ self,
16
+ width,
17
+ height,
18
+ resize_target=True,
19
+ keep_aspect_ratio=False,
20
+ ensure_multiple_of=1,
21
+ resize_method="lower_bound",
22
+ image_interpolation_method=cv2.INTER_AREA,
23
+ ):
24
+ """Init.
25
+
26
+ Args:
27
+ width (int): desired output width
28
+ height (int): desired output height
29
+ resize_target (bool, optional):
30
+ True: Resize the full sample (image, mask, target).
31
+ False: Resize image only.
32
+ Defaults to True.
33
+ keep_aspect_ratio (bool, optional):
34
+ True: Keep the aspect ratio of the input sample.
35
+ Output sample might not have the given width and height, and
36
+ resize behaviour depends on the parameter 'resize_method'.
37
+ Defaults to False.
38
+ ensure_multiple_of (int, optional):
39
+ Output width and height is constrained to be multiple of this parameter.
40
+ Defaults to 1.
41
+ resize_method (str, optional):
42
+ "lower_bound": Output will be at least as large as the given size.
43
+ "upper_bound": Output will be at max as large as the given size. (Output size might be smaller than given size.)
44
+ "minimal": Scale as least as possible. (Output size might be smaller than given size.)
45
+ Defaults to "lower_bound".
46
+ """
47
+ self.__width = width
48
+ self.__height = height
49
+
50
+ self.__resize_target = resize_target
51
+ self.__keep_aspect_ratio = keep_aspect_ratio
52
+ self.__multiple_of = ensure_multiple_of
53
+ self.__resize_method = resize_method
54
+ self.__image_interpolation_method = image_interpolation_method
55
+
56
+ def constrain_to_multiple_of(self, x, min_val=0, max_val=None):
57
+ y = (np.round(x / self.__multiple_of) * self.__multiple_of).astype(int)
58
+
59
+ if max_val is not None and y > max_val:
60
+ y = (np.floor(x / self.__multiple_of) * self.__multiple_of).astype(int)
61
+
62
+ if y < min_val:
63
+ y = (np.ceil(x / self.__multiple_of) * self.__multiple_of).astype(int)
64
+
65
+ return y
66
+
67
+ def get_size(self, width, height):
68
+ # determine new height and width
69
+ scale_height = self.__height / height
70
+ scale_width = self.__width / width
71
+
72
+ if self.__keep_aspect_ratio:
73
+ if self.__resize_method == "lower_bound":
74
+ # scale such that output size is lower bound
75
+ if scale_width > scale_height:
76
+ # fit width
77
+ scale_height = scale_width
78
+ else:
79
+ # fit height
80
+ scale_width = scale_height
81
+ elif self.__resize_method == "upper_bound":
82
+ # scale such that output size is upper bound
83
+ if scale_width < scale_height:
84
+ # fit width
85
+ scale_height = scale_width
86
+ else:
87
+ # fit height
88
+ scale_width = scale_height
89
+ elif self.__resize_method == "minimal":
90
+ # scale as least as possbile
91
+ if abs(1 - scale_width) < abs(1 - scale_height):
92
+ # fit width
93
+ scale_height = scale_width
94
+ else:
95
+ # fit height
96
+ scale_width = scale_height
97
+ else:
98
+ raise ValueError(
99
+ f"resize_method {self.__resize_method} not implemented"
100
+ )
101
+
102
+ if self.__resize_method == "lower_bound":
103
+ new_height = self.constrain_to_multiple_of(
104
+ scale_height * height, min_val=self.__height
105
+ )
106
+ new_width = self.constrain_to_multiple_of(
107
+ scale_width * width, min_val=self.__width
108
+ )
109
+ elif self.__resize_method == "upper_bound":
110
+ new_height = self.constrain_to_multiple_of(
111
+ scale_height * height, max_val=self.__height
112
+ )
113
+ new_width = self.constrain_to_multiple_of(
114
+ scale_width * width, max_val=self.__width
115
+ )
116
+ elif self.__resize_method == "minimal":
117
+ new_height = self.constrain_to_multiple_of(scale_height * height)
118
+ new_width = self.constrain_to_multiple_of(scale_width * width)
119
+ else:
120
+ raise ValueError(f"resize_method {self.__resize_method} not implemented")
121
+
122
+ return (new_width, new_height)
123
+
124
+ def __call__(self, sample):
125
+ width, height = self.get_size(
126
+ sample["image"].shape[1], sample["image"].shape[0]
127
+ )
128
+
129
+ # resize sample
130
+ for item in sample.keys():
131
+ interpolation_method = self.__image_interpolation_method
132
+ sample[item] = cv2.resize(
133
+ sample[item],
134
+ (width, height),
135
+ interpolation=interpolation_method,
136
+ )
137
+
138
+ if self.__resize_target:
139
+
140
+ if "gt" in sample:
141
+ sample["gt"] = cv2.resize(
142
+ sample["gt"],
143
+ (width, height),
144
+ interpolation=cv2.INTER_NEAREST
145
+ )
146
+
147
+ if "sparse_gt" in sample:
148
+ sample["sparse_gt"] = cv2.resize(
149
+ sample["sparse_gt"],
150
+ (width, height),
151
+ interpolation=cv2.INTER_NEAREST
152
+ )
153
+ if "gt_sky" in sample:
154
+ sample["gt_sky"] = cv2.resize(
155
+ sample["gt_sky"],
156
+ (width, height),
157
+ interpolation=cv2.INTER_NEAREST
158
+ )
159
+
160
+ return sample
161
+
162
+
163
+
164
+ class NormalizeIntermediate(object):
165
+ """Normalize intermediate data by given mean and std.
166
+ """
167
+
168
+ def __init__(self, mean, std):
169
+
170
+ self.__int_depth_mean = mean["int_depth"]
171
+ self.__int_depth_std = std["int_depth"]
172
+
173
+ self.__int_scales_mean = mean["int_scales"]
174
+ self.__int_scales_std = std["int_scales"]
175
+
176
+ def __call__(self, sample):
177
+
178
+ if "int_depth" in sample and sample["int_depth"] is not None:
179
+ sample["int_depth"] = (sample["int_depth"] - self.__int_depth_mean) / self.__int_depth_std
180
+
181
+ if "int_scales" in sample and sample["int_scales"] is not None:
182
+ sample["int_scales"] = (sample["int_scales"] - self.__int_scales_mean) / self.__int_scales_std
183
+
184
+ return sample
185
+
186
+
187
+ class PrepareForNet(object):
188
+ """Prepare sample for usage as network input.
189
+ """
190
+
191
+ def __init__(self):
192
+ pass
193
+
194
+ def __call__(self, sample):
195
+
196
+ for item in sample.keys():
197
+
198
+ if sample[item] is None:
199
+ pass
200
+ elif item == "image":
201
+ image = np.transpose(sample["image"], (2, 0, 1))
202
+ sample["image"] = np.ascontiguousarray(image).astype(np.float32)
203
+ else:
204
+ array = sample[item].astype(np.float32)
205
+ array = np.expand_dims(array, axis=0) # add channel dim
206
+ sample[item] = np.ascontiguousarray(array)
207
+
208
+ return sample
209
+
210
+
211
+ class Tensorize(object):
212
+ """Convert sample to tensor.
213
+ """
214
+
215
+ def __init__(self):
216
+ pass
217
+
218
+ def __call__(self, sample):
219
+
220
+ for item in sample.keys():
221
+
222
+ if sample[item] is None:
223
+ pass
224
+ else:
225
+ # before tensorizing, verify that data is clean
226
+ assert not np.any(np.isnan(sample[item]))
227
+ sample[item] = torch.Tensor(sample[item])
228
+
229
+ return sample
230
+
231
+
232
+ def get_transforms(depth_predictor, sparsifier, nsamples):
233
+
234
+ resize_method_dict = {
235
+ "dpt_beit_large_512" : "minimal",
236
+ "dpt_swin2_large_384" : "minimal",
237
+ "dpt_large" : "minimal",
238
+ "dpt_hybrid" : "minimal",
239
+ "dpt_swin2_tiny_256" : "minimal",
240
+ "dpt_levit_224" : "minimal",
241
+ "midas_small" : "upper_bound",
242
+ }
243
+
244
+ sml_model_transform_steps = [
245
+ Resize(
246
+ width=288,
247
+ height=288,
248
+ resize_target=False,
249
+ keep_aspect_ratio=True,
250
+ ensure_multiple_of=32,
251
+ resize_method=resize_method_dict["dpt_hybrid"],
252
+ image_interpolation_method=cv2.INTER_NEAREST,
253
+ ),
254
+ NormalizeIntermediate(
255
+ mean=normalization.VOID_INTERMEDIATE[depth_predictor][f"{sparsifier}_{nsamples}"]["mean"],
256
+ std=normalization.VOID_INTERMEDIATE[depth_predictor][f"{sparsifier}_{nsamples}"]["std"],
257
+ ),
258
+ PrepareForNet(),
259
+ Tensorize(),
260
+ ]
261
+
262
+ return transforms.Compose(sml_model_transform_steps)
263
+
src/Baselines/radarcam-depth/modules/midas/utils.py ADDED
@@ -0,0 +1,237 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Utils for monoDepth.
2
+ """
3
+ import sys
4
+ import re
5
+ import numpy as np
6
+ import cv2
7
+ import torch
8
+
9
+
10
+ def read_pfm(path):
11
+ """Read pfm file.
12
+
13
+ Args:
14
+ path (str): path to file
15
+
16
+ Returns:
17
+ tuple: (data, scale)
18
+ """
19
+ with open(path, "rb") as file:
20
+
21
+ color = None
22
+ width = None
23
+ height = None
24
+ scale = None
25
+ endian = None
26
+
27
+ header = file.readline().rstrip()
28
+ if header.decode("ascii") == "PF":
29
+ color = True
30
+ elif header.decode("ascii") == "Pf":
31
+ color = False
32
+ else:
33
+ raise Exception("Not a PFM file: " + path)
34
+
35
+ dim_match = re.match(r"^(\d+)\s(\d+)\s$", file.readline().decode("ascii"))
36
+ if dim_match:
37
+ width, height = list(map(int, dim_match.groups()))
38
+ else:
39
+ raise Exception("Malformed PFM header.")
40
+
41
+ scale = float(file.readline().decode("ascii").rstrip())
42
+ if scale < 0:
43
+ # little-endian
44
+ endian = "<"
45
+ scale = -scale
46
+ else:
47
+ # big-endian
48
+ endian = ">"
49
+
50
+ data = np.fromfile(file, endian + "f")
51
+ shape = (height, width, 3) if color else (height, width)
52
+
53
+ data = np.reshape(data, shape)
54
+ data = np.flipud(data)
55
+
56
+ return data, scale
57
+
58
+
59
+ def write_pfm(path, image, scale=1):
60
+ """Write pfm file.
61
+
62
+ Args:
63
+ path (str): pathto file
64
+ image (array): data
65
+ scale (int, optional): Scale. Defaults to 1.
66
+ """
67
+
68
+ with open(path, "wb") as file:
69
+ color = None
70
+
71
+ if image.dtype.name != "float32":
72
+ raise Exception("Image dtype must be float32.")
73
+
74
+ image = np.flipud(image)
75
+
76
+ if len(image.shape) == 3 and image.shape[2] == 3: # color image
77
+ color = True
78
+ elif (
79
+ len(image.shape) == 2 or len(image.shape) == 3 and image.shape[2] == 1
80
+ ): # greyscale
81
+ color = False
82
+ else:
83
+ raise Exception("Image must have H x W x 3, H x W x 1 or H x W dimensions.")
84
+
85
+ file.write("PF\n" if color else "Pf\n".encode())
86
+ file.write("%d %d\n".encode() % (image.shape[1], image.shape[0]))
87
+
88
+ endian = image.dtype.byteorder
89
+
90
+ if endian == "<" or endian == "=" and sys.byteorder == "little":
91
+ scale = -scale
92
+
93
+ file.write("%f\n".encode() % scale)
94
+
95
+ image.tofile(file)
96
+
97
+
98
+ def read_image(path):
99
+ """Read image and output RGB image (0-1).
100
+
101
+ Args:
102
+ path (str): path to file
103
+
104
+ Returns:
105
+ array: RGB image (0-1)
106
+ """
107
+ img = cv2.imread(path)
108
+
109
+ if img.ndim == 2:
110
+ img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)
111
+
112
+ img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) / 255.0
113
+
114
+ return img
115
+
116
+
117
+ def resize_image(img):
118
+ """Resize image and make it fit for network.
119
+
120
+ Args:
121
+ img (array): image
122
+
123
+ Returns:
124
+ tensor: data ready for network
125
+ """
126
+ height_orig = img.shape[0]
127
+ width_orig = img.shape[1]
128
+
129
+ if width_orig > height_orig:
130
+ scale = width_orig / 384
131
+ else:
132
+ scale = height_orig / 384
133
+
134
+ height = (np.ceil(height_orig / scale / 32) * 32).astype(int)
135
+ width = (np.ceil(width_orig / scale / 32) * 32).astype(int)
136
+
137
+ img_resized = cv2.resize(img, (width, height), interpolation=cv2.INTER_AREA)
138
+
139
+ img_resized = (
140
+ torch.from_numpy(np.transpose(img_resized, (2, 0, 1))).contiguous().float()
141
+ )
142
+ img_resized = img_resized.unsqueeze(0)
143
+
144
+ return img_resized
145
+
146
+
147
+ def resize_depth(depth, width, height):
148
+ """Resize depth map and bring to CPU (numpy).
149
+
150
+ Args:
151
+ depth (tensor): depth
152
+ width (int): image width
153
+ height (int): image height
154
+
155
+ Returns:
156
+ array: processed depth
157
+ """
158
+ depth = torch.squeeze(depth[0, :, :, :]).to("cpu")
159
+
160
+ depth_resized = cv2.resize(
161
+ depth.numpy(), (width, height), interpolation=cv2.INTER_CUBIC
162
+ )
163
+
164
+ return depth_resized
165
+
166
+
167
+ def write_depth(path, depth, bits=1):
168
+ """Write depth map to pfm and png file.
169
+
170
+ Args:
171
+ path (str): filepath without extension
172
+ depth (array): depth
173
+ """
174
+ write_pfm(path + ".pfm", depth.astype(np.float32))
175
+
176
+ depth_min = depth.min()
177
+ depth_max = depth.max()
178
+
179
+ max_val = (2**(8*bits))-1
180
+
181
+ if depth_max - depth_min > np.finfo("float").eps:
182
+ out = max_val * (depth - depth_min) / (depth_max - depth_min)
183
+ else:
184
+ out = np.zeros(depth.shape, dtype=depth.type)
185
+
186
+ if bits == 1:
187
+ cv2.imwrite(path + ".png", out.astype("uint8"))
188
+ elif bits == 2:
189
+ cv2.imwrite(path + ".png", out.astype("uint16"))
190
+
191
+ return
192
+
193
+
194
+ def write_png(path, array, bits=2, absolute=True):
195
+ """Write array to png file.
196
+
197
+ Args:
198
+ path (str): filepath without extension
199
+ array (array): array to be saved
200
+ """
201
+ if absolute:
202
+ out = array
203
+ else:
204
+ array_min = np.min(array)
205
+ array_max = np.max(array)
206
+
207
+ max_val = (2**(8*bits))-1
208
+
209
+ if array_max - array_min > np.finfo("float").eps:
210
+ out = max_val * (array - array_min) / (array_max - array_min)
211
+ else:
212
+ print(f"zero array not being saved at {path}")
213
+ return
214
+
215
+ if bits == 1:
216
+ cv2.imwrite(path + ".png", out.astype("uint8"), [cv2.IMWRITE_PNG_COMPRESSION, 0])
217
+ elif bits == 2:
218
+ cv2.imwrite(path + ".png", out.astype("uint16"), [cv2.IMWRITE_PNG_COMPRESSION, 0])
219
+
220
+ return
221
+
222
+
223
+ def normalize_unit_range(data):
224
+ """Normalize data array to [0, 1] range.
225
+
226
+ Args:
227
+ data (array): input array
228
+
229
+ Returns:
230
+ array: normalized array
231
+ """
232
+ if np.max(data) - np.min(data) > np.finfo("float").eps:
233
+ normalized = (data - np.min(data)) / (np.max(data) - np.min(data))
234
+ else:
235
+ raise ValueError("cannot normalize array, max-min range is 0")
236
+
237
+ return normalized
src/Baselines/radarcam-depth/networks.py ADDED
@@ -0,0 +1,1516 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from utils import net_utils
3
+ import torchvision
4
+ from linear_attention import LocalFeatureTransformer
5
+
6
+ '''
7
+ Encoders
8
+ '''
9
+
10
+
11
+ class ResNetEncoder(torch.nn.Module):
12
+ '''
13
+ ResNet encoder with skip connections
14
+ Arg(s):
15
+ n_layer : int
16
+ architecture type based on layers: 18, 34, 50
17
+ input_channels : int
18
+ number of channels in input data
19
+ n_filters : list
20
+ number of filters to use for each block
21
+ weight_initializer : str
22
+ kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
23
+ activation_func : func
24
+ activation function after convolution
25
+ use_batch_norm : bool
26
+ if set, then applied batch normalization
27
+ '''
28
+
29
+ def __init__(self,
30
+ n_layer,
31
+ input_channels=3,
32
+ n_filters=[32, 64, 128, 256, 256],
33
+ weight_initializer='kaiming_uniform',
34
+ activation_func='leaky_relu',
35
+ use_batch_norm=False):
36
+ super(ResNetEncoder, self).__init__()
37
+
38
+ if n_layer == 18:
39
+ n_blocks = [2, 2, 2, 2]
40
+ resnet_block = net_utils.ResNetBlock
41
+ elif n_layer == 34:
42
+ n_blocks = [3, 4, 6, 3]
43
+ resnet_block = net_utils.ResNetBlock
44
+ else:
45
+ raise ValueError('Only supports 18, 34 layer architecture')
46
+
47
+ for n in range(len(n_filters) - len(n_blocks) - 1):
48
+ n_blocks = n_blocks + [n_blocks[-1]]
49
+
50
+ network_depth = len(n_filters)
51
+
52
+ assert network_depth < 8, 'Does not support network depth of 8 or more'
53
+ assert network_depth == len(n_blocks) + 1
54
+
55
+ # Keep track on current block
56
+ block_idx = 0
57
+ filter_idx = 0
58
+
59
+ activation_func = net_utils.activation_func(activation_func)
60
+
61
+ in_channels, out_channels = [input_channels, n_filters[filter_idx]]
62
+
63
+ # Resolution 1/1 -> 1/2
64
+ self.conv1 = net_utils.Conv2d(
65
+ in_channels,
66
+ out_channels,
67
+ kernel_size=7,
68
+ stride=2,
69
+ weight_initializer=weight_initializer,
70
+ activation_func=activation_func,
71
+ use_batch_norm=use_batch_norm)
72
+
73
+ # Resolution 1/2 -> 1/4
74
+ self.max_pool = torch.nn.MaxPool2d(
75
+ kernel_size=3,
76
+ stride=2,
77
+ padding=1)
78
+
79
+ filter_idx = filter_idx + 1
80
+
81
+ in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]]
82
+
83
+ self.blocks2 = self._make_layer(
84
+ network_block=resnet_block,
85
+ n_block=n_blocks[block_idx],
86
+ in_channels=in_channels,
87
+ out_channels=out_channels,
88
+ stride=1,
89
+ weight_initializer=weight_initializer,
90
+ activation_func=activation_func,
91
+ use_batch_norm=use_batch_norm)
92
+
93
+ # Resolution 1/4 -> 1/8
94
+ block_idx = block_idx + 1
95
+ filter_idx = filter_idx + 1
96
+
97
+ in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]]
98
+
99
+ self.blocks3 = self._make_layer(
100
+ network_block=resnet_block,
101
+ n_block=n_blocks[block_idx],
102
+ in_channels=in_channels,
103
+ out_channels=out_channels,
104
+ stride=2,
105
+ weight_initializer=weight_initializer,
106
+ activation_func=activation_func,
107
+ use_batch_norm=use_batch_norm)
108
+
109
+ # Resolution 1/8 -> 1/16
110
+ block_idx = block_idx + 1
111
+ filter_idx = filter_idx + 1
112
+
113
+ in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]]
114
+
115
+ self.blocks4 = self._make_layer(
116
+ network_block=resnet_block,
117
+ n_block=n_blocks[block_idx],
118
+ in_channels=in_channels,
119
+ out_channels=out_channels,
120
+ stride=2,
121
+ weight_initializer=weight_initializer,
122
+ activation_func=activation_func,
123
+ use_batch_norm=use_batch_norm)
124
+
125
+ # Resolution 1/16 -> 1/32
126
+ block_idx = block_idx + 1
127
+ filter_idx = filter_idx + 1
128
+
129
+ in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]]
130
+
131
+ self.blocks5 = self._make_layer(
132
+ network_block=resnet_block,
133
+ n_block=n_blocks[block_idx],
134
+ in_channels=in_channels,
135
+ out_channels=out_channels,
136
+ stride=2,
137
+ weight_initializer=weight_initializer,
138
+ activation_func=activation_func,
139
+ use_batch_norm=use_batch_norm)
140
+
141
+ # Resolution 1/32 -> 1/64
142
+ block_idx = block_idx + 1
143
+ filter_idx = filter_idx + 1
144
+
145
+ if filter_idx < len(n_filters):
146
+
147
+ in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]]
148
+
149
+ self.blocks6 = self._make_layer(
150
+ network_block=resnet_block,
151
+ n_block=n_blocks[block_idx],
152
+ in_channels=in_channels,
153
+ out_channels=out_channels,
154
+ stride=2,
155
+ weight_initializer=weight_initializer,
156
+ activation_func=activation_func,
157
+ use_batch_norm=use_batch_norm)
158
+ else:
159
+ self.blocks6 = None
160
+
161
+ # Resolution 1/64 -> 1/128
162
+ block_idx = block_idx + 1
163
+ filter_idx = filter_idx + 1
164
+
165
+ if filter_idx < len(n_filters):
166
+
167
+ in_channels, out_channels = [n_filters[filter_idx - 1], n_filters[filter_idx]]
168
+
169
+ self.blocks7 = self._make_layer(
170
+ network_block=resnet_block,
171
+ n_block=n_blocks[block_idx],
172
+ in_channels=in_channels,
173
+ out_channels=out_channels,
174
+ stride=2,
175
+ weight_initializer=weight_initializer,
176
+ activation_func=activation_func,
177
+ use_batch_norm=use_batch_norm)
178
+ else:
179
+ self.blocks7 = None
180
+
181
+ def _make_layer(self,
182
+ network_block,
183
+ n_block,
184
+ in_channels,
185
+ out_channels,
186
+ stride,
187
+ weight_initializer,
188
+ activation_func,
189
+ use_batch_norm):
190
+ '''
191
+ Creates a layer
192
+ Arg(s):
193
+ network_block : Object
194
+ block type
195
+ n_block : int
196
+ number of blocks to use in layer
197
+ in_channels : int
198
+ number of channels
199
+ out_channels : int
200
+ number of output channels
201
+ stride : int
202
+ stride of convolution
203
+ weight_initializer : str
204
+ kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
205
+ activation_func : func
206
+ activation function after convolution
207
+ use_batch_norm : bool
208
+ if set, then applied batch normalization
209
+ '''
210
+
211
+ blocks = []
212
+
213
+ for n in range(n_block):
214
+
215
+ if n == 0:
216
+ stride = stride
217
+ else:
218
+ in_channels = out_channels
219
+ stride = 1
220
+
221
+ block = network_block(
222
+ in_channels=in_channels,
223
+ out_channels=out_channels,
224
+ stride=stride,
225
+ weight_initializer=weight_initializer,
226
+ activation_func=activation_func,
227
+ use_batch_norm=use_batch_norm)
228
+
229
+ blocks.append(block)
230
+
231
+ blocks = torch.nn.Sequential(*blocks)
232
+
233
+ return blocks
234
+
235
+ def forward(self, x):
236
+ '''
237
+ Forward input x through the ResNet model
238
+ Arg(s):
239
+ x : torch.Tensor
240
+ Returns:
241
+ torch.Tensor[float32] : latent vector
242
+ list[torch.Tensor[float32]] : skip connections
243
+ '''
244
+
245
+ layers = [x]
246
+
247
+ # Resolution 1/1 -> 1/2
248
+ layers.append(self.conv1(layers[-1]))
249
+
250
+ # Resolution 1/2 -> 1/4
251
+ max_pool = self.max_pool(layers[-1])
252
+ layers.append(self.blocks2(max_pool))
253
+
254
+ # Resolution 1/4 -> 1/8
255
+ layers.append(self.blocks3(layers[-1]))
256
+
257
+ # Resolution 1/8 -> 1/16
258
+ layers.append(self.blocks4(layers[-1]))
259
+
260
+ # Resolution 1/16 -> 1/32
261
+ layers.append(self.blocks5(layers[-1]))
262
+
263
+ # Resolution 1/32 -> 1/64
264
+ if self.blocks6 is not None:
265
+ layers.append(self.blocks6(layers[-1]))
266
+
267
+ # Resolution 1/64 -> 1/128
268
+ if self.blocks7 is not None:
269
+ layers.append(self.blocks7(layers[-1]))
270
+
271
+ return layers[-1], layers[1:-1]
272
+
273
+
274
+ class FullyConnectedEncoder(torch.nn.Module):
275
+ '''
276
+ Fully connected encoder
277
+ Arg(s):
278
+ input_channels : int
279
+ number of input channels
280
+ n_neurons : list[int]
281
+ number of filters to use per layer
282
+ latent_size : int
283
+ number of output neuron
284
+ weight_initializer : str
285
+ kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
286
+ activation_func : str
287
+ activation function after convolution
288
+ '''
289
+
290
+ def __init__(self,
291
+ input_channels=3,
292
+ n_neurons=[32, 64, 96, 128, 256],
293
+ latent_size=29 * 10,
294
+ weight_initializer='kaiming_uniform',
295
+ activation_func='leaky_relu'):
296
+ super(FullyConnectedEncoder, self).__init__()
297
+
298
+ activation_func = net_utils.activation_func(activation_func)
299
+
300
+ self.mlp = torch.nn.Sequential(
301
+ net_utils.FullyConnected(
302
+ in_features=input_channels,
303
+ out_features=n_neurons[0],
304
+ weight_initializer=weight_initializer,
305
+ activation_func=activation_func),
306
+ net_utils.FullyConnected(
307
+ in_features=n_neurons[0],
308
+ out_features=n_neurons[1],
309
+ weight_initializer=weight_initializer,
310
+ activation_func=activation_func),
311
+ net_utils.FullyConnected(
312
+ in_features=n_neurons[1],
313
+ out_features=n_neurons[2],
314
+ weight_initializer=weight_initializer,
315
+ activation_func=activation_func),
316
+ net_utils.FullyConnected(
317
+ in_features=n_neurons[2],
318
+ out_features=n_neurons[3],
319
+ weight_initializer=weight_initializer,
320
+ activation_func=activation_func),
321
+ net_utils.FullyConnected(
322
+ in_features=n_neurons[3],
323
+ out_features=n_neurons[4],
324
+ weight_initializer=weight_initializer,
325
+ activation_func=activation_func),
326
+ net_utils.FullyConnected(
327
+ in_features=n_neurons[4],
328
+ out_features=latent_size,
329
+ weight_initializer=weight_initializer,
330
+ activation_func=activation_func))
331
+
332
+ def forward(self, x):
333
+ return self.mlp(x)
334
+
335
+
336
+ class FusionNetEncoder(torch.nn.Module):
337
+ '''
338
+ FusionNet encoder with skip connections
339
+ Arg(s):
340
+ n_layer : int
341
+ number of layer for encoder
342
+ input_channels_image : int
343
+ number of channels in input data
344
+ input_channels_depth : int
345
+ number of channels in input data
346
+ n_filters_per_block : list[int]
347
+ number of filters to use for each block
348
+ weight_initializer : str
349
+ kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
350
+ activation_func : func
351
+ activation function after convolution
352
+ use_batch_norm : bool
353
+ if set, then applied batch normalization
354
+ fusion_type : str
355
+ add, weight
356
+ '''
357
+
358
+ def __init__(self,
359
+ n_layer=18,
360
+ input_channels_image=3,
361
+ input_channels_depth=3,
362
+ n_filters_encoder_image=[32, 64, 128, 256, 256],
363
+ n_filters_encoder_depth=[32, 64, 128, 256, 256],
364
+ weight_initializer='kaiming_uniform',
365
+ activation_func='leaky_relu',
366
+ use_batch_norm=False,
367
+ fusion_type='add'):
368
+ super(FusionNetEncoder, self).__init__()
369
+
370
+ self.fusion_type = fusion_type
371
+
372
+ if n_layer == 18:
373
+ n_blocks = [2, 2, 2, 2]
374
+ elif n_layer == 34:
375
+ n_blocks = [3, 4, 6, 3]
376
+ else:
377
+ raise ValueError('Only supports 18, 34 layer architecture')
378
+
379
+ resnet_block = net_utils.ResNetBlock
380
+
381
+ assert len(n_filters_encoder_image) == len(n_filters_encoder_depth)
382
+
383
+ for n in range(len(n_filters_encoder_image) - len(n_blocks) - 1):
384
+ n_blocks = n_blocks + [n_blocks[-1]]
385
+
386
+ network_depth = len(n_filters_encoder_image)
387
+
388
+ assert network_depth < 8, 'Does not support network depth of 8 or more'
389
+ assert network_depth == len(n_blocks) + 1
390
+
391
+ # Keep track on current block
392
+ block_idx = 0
393
+ filter_idx = 0
394
+
395
+ activation_func = net_utils.activation_func(activation_func)
396
+
397
+ # Resolution 1/1 -> 1/2
398
+ self.conv1_image = net_utils.Conv2d(
399
+ input_channels_image,
400
+ n_filters_encoder_image[filter_idx],
401
+ kernel_size=7,
402
+ stride=2,
403
+ weight_initializer=weight_initializer,
404
+ activation_func=activation_func,
405
+ use_batch_norm=use_batch_norm)
406
+
407
+ self.conv1_depth = net_utils.Conv2d(
408
+ input_channels_depth,
409
+ n_filters_encoder_depth[filter_idx],
410
+ kernel_size=7,
411
+ stride=2,
412
+ weight_initializer=weight_initializer,
413
+ activation_func=activation_func,
414
+ use_batch_norm=use_batch_norm)
415
+
416
+ if fusion_type == 'add':
417
+ self.conv1_project = net_utils.Conv2d(
418
+ n_filters_encoder_depth[filter_idx],
419
+ n_filters_encoder_image[filter_idx],
420
+ kernel_size=1,
421
+ stride=1,
422
+ weight_initializer=weight_initializer,
423
+ activation_func=net_utils.activation_func('linear'),
424
+ use_batch_norm=use_batch_norm)
425
+
426
+ elif fusion_type == 'weight':
427
+
428
+ self.conv1_weight = net_utils.Conv2d(
429
+ n_filters_encoder_depth[filter_idx],
430
+ n_filters_encoder_depth[filter_idx],
431
+ kernel_size=3,
432
+ stride=1,
433
+ weight_initializer=weight_initializer,
434
+ activation_func=net_utils.activation_func('sigmoid'),
435
+ use_batch_norm=use_batch_norm)
436
+
437
+ elif fusion_type == 'weight_and_project':
438
+
439
+ self.conv1_weight = net_utils.Conv2d(
440
+ n_filters_encoder_depth[filter_idx],
441
+ n_filters_encoder_image[filter_idx],
442
+ kernel_size=1,
443
+ stride=1,
444
+ weight_initializer=weight_initializer,
445
+ activation_func=net_utils.activation_func('sigmoid'),
446
+ use_batch_norm=use_batch_norm)
447
+
448
+ self.conv1_project = net_utils.Conv2d(
449
+ n_filters_encoder_depth[filter_idx],
450
+ n_filters_encoder_image[filter_idx],
451
+ kernel_size=1,
452
+ stride=1,
453
+ weight_initializer=weight_initializer,
454
+ activation_func=net_utils.activation_func('linear'),
455
+ use_batch_norm=use_batch_norm)
456
+
457
+ # Resolution 1/2 -> 1/4
458
+ self.max_pool = torch.nn.MaxPool2d(
459
+ kernel_size=3,
460
+ stride=2,
461
+ padding=1)
462
+
463
+ filter_idx = filter_idx + 1
464
+
465
+ in_channels_image, out_channels_image = [
466
+ n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx]
467
+ ]
468
+
469
+ in_channels_depth, out_channels_depth = [
470
+ n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx]
471
+ ]
472
+
473
+ self.blocks2_image, self.blocks2_depth = self._make_layer(
474
+ network_block=resnet_block,
475
+ n_block=n_blocks[block_idx],
476
+ in_channels_image=in_channels_image,
477
+ in_channels_depth=in_channels_depth,
478
+ out_channels_image=out_channels_image,
479
+ out_channels_depth=out_channels_depth,
480
+ stride=1,
481
+ weight_initializer=weight_initializer,
482
+ activation_func=activation_func,
483
+ use_batch_norm=use_batch_norm)
484
+
485
+ if fusion_type == 'add':
486
+ self.conv2_project = net_utils.Conv2d(
487
+ out_channels_depth,
488
+ out_channels_image,
489
+ kernel_size=1,
490
+ stride=1,
491
+ weight_initializer=weight_initializer,
492
+ activation_func=net_utils.activation_func('linear'),
493
+ use_batch_norm=use_batch_norm)
494
+
495
+ elif fusion_type == 'weight':
496
+
497
+ self.conv2_weight = net_utils.Conv2d(
498
+ out_channels_depth,
499
+ out_channels_depth,
500
+ kernel_size=3,
501
+ stride=1,
502
+ weight_initializer=weight_initializer,
503
+ activation_func=net_utils.activation_func('sigmoid'),
504
+ use_batch_norm=use_batch_norm)
505
+
506
+ elif fusion_type == 'weight_and_project':
507
+
508
+ self.conv2_weight = net_utils.Conv2d(
509
+ out_channels_depth,
510
+ out_channels_image,
511
+ kernel_size=1,
512
+ stride=1,
513
+ weight_initializer=weight_initializer,
514
+ activation_func=net_utils.activation_func('sigmoid'),
515
+ use_batch_norm=use_batch_norm)
516
+
517
+ self.conv2_project = net_utils.Conv2d(
518
+ out_channels_depth,
519
+ out_channels_image,
520
+ kernel_size=1,
521
+ stride=1,
522
+ weight_initializer=weight_initializer,
523
+ activation_func=net_utils.activation_func('linear'),
524
+ use_batch_norm=use_batch_norm)
525
+
526
+ # Resolution 1/4 -> 1/8
527
+ block_idx = block_idx + 1
528
+ filter_idx = filter_idx + 1
529
+
530
+ in_channels_image, out_channels_image = [
531
+ n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx]
532
+ ]
533
+
534
+ in_channels_depth, out_channels_depth = [
535
+ n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx]
536
+ ]
537
+
538
+ self.blocks3_image, self.blocks3_depth = self._make_layer(
539
+ network_block=resnet_block,
540
+ n_block=n_blocks[block_idx],
541
+ in_channels_image=in_channels_image,
542
+ in_channels_depth=in_channels_depth,
543
+ out_channels_image=out_channels_image,
544
+ out_channels_depth=out_channels_depth,
545
+ stride=2,
546
+ weight_initializer=weight_initializer,
547
+ activation_func=activation_func,
548
+ use_batch_norm=use_batch_norm)
549
+
550
+ if fusion_type == 'add':
551
+ self.conv3_project = net_utils.Conv2d(
552
+ out_channels_depth,
553
+ out_channels_image,
554
+ kernel_size=1,
555
+ stride=1,
556
+ weight_initializer=weight_initializer,
557
+ activation_func=net_utils.activation_func('linear'),
558
+ use_batch_norm=use_batch_norm)
559
+
560
+ elif fusion_type == 'weight':
561
+
562
+ self.conv3_weight = net_utils.Conv2d(
563
+ out_channels_depth,
564
+ out_channels_depth,
565
+ kernel_size=3,
566
+ stride=1,
567
+ weight_initializer=weight_initializer,
568
+ activation_func=net_utils.activation_func('sigmoid'),
569
+ use_batch_norm=use_batch_norm)
570
+
571
+ elif fusion_type == 'weight_and_project':
572
+
573
+ self.conv3_weight = net_utils.Conv2d(
574
+ out_channels_depth,
575
+ out_channels_image,
576
+ kernel_size=1,
577
+ stride=1,
578
+ weight_initializer=weight_initializer,
579
+ activation_func=net_utils.activation_func('sigmoid'),
580
+ use_batch_norm=use_batch_norm)
581
+
582
+ self.conv3_project = net_utils.Conv2d(
583
+ out_channels_depth,
584
+ out_channels_image,
585
+ kernel_size=1,
586
+ stride=1,
587
+ weight_initializer=weight_initializer,
588
+ activation_func=net_utils.activation_func('linear'),
589
+ use_batch_norm=use_batch_norm)
590
+
591
+ # Resolution 1/8 -> 1/16
592
+ block_idx = block_idx + 1
593
+ filter_idx = filter_idx + 1
594
+
595
+ in_channels_image, out_channels_image = [
596
+ n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx]
597
+ ]
598
+
599
+ in_channels_depth, out_channels_depth = [
600
+ n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx]
601
+ ]
602
+
603
+ self.blocks4_image, self.blocks4_depth = self._make_layer(
604
+ network_block=resnet_block,
605
+ n_block=n_blocks[block_idx],
606
+ in_channels_image=in_channels_image,
607
+ in_channels_depth=in_channels_depth,
608
+ out_channels_image=out_channels_image,
609
+ out_channels_depth=out_channels_depth,
610
+ stride=2,
611
+ weight_initializer=weight_initializer,
612
+ activation_func=activation_func,
613
+ use_batch_norm=use_batch_norm)
614
+
615
+ if fusion_type == 'add':
616
+ self.conv4_project = net_utils.Conv2d(
617
+ out_channels_depth,
618
+ out_channels_image,
619
+ kernel_size=1,
620
+ stride=1,
621
+ weight_initializer=weight_initializer,
622
+ activation_func=net_utils.activation_func('linear'),
623
+ use_batch_norm=use_batch_norm)
624
+
625
+ elif fusion_type == 'weight':
626
+
627
+ self.conv4_weight = net_utils.Conv2d(
628
+ out_channels_depth,
629
+ out_channels_depth,
630
+ kernel_size=3,
631
+ stride=1,
632
+ weight_initializer=weight_initializer,
633
+ activation_func=net_utils.activation_func('sigmoid'),
634
+ use_batch_norm=use_batch_norm)
635
+
636
+ elif fusion_type == 'weight_and_project':
637
+
638
+ self.conv4_weight = net_utils.Conv2d(
639
+ out_channels_depth,
640
+ out_channels_image,
641
+ kernel_size=1,
642
+ stride=1,
643
+ weight_initializer=weight_initializer,
644
+ activation_func=net_utils.activation_func('sigmoid'),
645
+ use_batch_norm=use_batch_norm)
646
+
647
+ self.conv4_project = net_utils.Conv2d(
648
+ out_channels_depth,
649
+ out_channels_image,
650
+ kernel_size=1,
651
+ stride=1,
652
+ weight_initializer=weight_initializer,
653
+ activation_func=net_utils.activation_func('linear'),
654
+ use_batch_norm=use_batch_norm)
655
+
656
+ # Resolution 1/16 -> 1/32
657
+ block_idx = block_idx + 1
658
+ filter_idx = filter_idx + 1
659
+
660
+ in_channels_image, out_channels_image = [
661
+ n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx]
662
+ ]
663
+
664
+ in_channels_depth, out_channels_depth = [
665
+ n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx]
666
+ ]
667
+
668
+ self.blocks5_image, self.blocks5_depth = self._make_layer(
669
+ network_block=resnet_block,
670
+ n_block=n_blocks[block_idx],
671
+ in_channels_image=in_channels_image,
672
+ in_channels_depth=in_channels_depth,
673
+ out_channels_image=out_channels_image,
674
+ out_channels_depth=out_channels_depth,
675
+ stride=2,
676
+ weight_initializer=weight_initializer,
677
+ activation_func=activation_func,
678
+ use_batch_norm=use_batch_norm)
679
+
680
+ if fusion_type == 'add':
681
+ self.conv5_project = net_utils.Conv2d(
682
+ out_channels_depth,
683
+ out_channels_image,
684
+ kernel_size=1,
685
+ stride=1,
686
+ weight_initializer=weight_initializer,
687
+ activation_func=net_utils.activation_func('linear'),
688
+ use_batch_norm=use_batch_norm)
689
+
690
+ elif fusion_type == 'weight':
691
+
692
+ self.conv5_weight = net_utils.Conv2d(
693
+ out_channels_depth,
694
+ out_channels_depth,
695
+ kernel_size=3,
696
+ stride=1,
697
+ weight_initializer=weight_initializer,
698
+ activation_func=net_utils.activation_func('sigmoid'),
699
+ use_batch_norm=use_batch_norm)
700
+
701
+ if fusion_type == 'weight_and_project':
702
+ self.conv5_weight = net_utils.Conv2d(
703
+ out_channels_depth,
704
+ out_channels_image,
705
+ kernel_size=1,
706
+ stride=1,
707
+ weight_initializer=weight_initializer,
708
+ activation_func=net_utils.activation_func('sigmoid'),
709
+ use_batch_norm=use_batch_norm)
710
+
711
+ self.conv5_project = net_utils.Conv2d(
712
+ out_channels_depth,
713
+ out_channels_image,
714
+ kernel_size=1,
715
+ stride=1,
716
+ weight_initializer=weight_initializer,
717
+ activation_func=net_utils.activation_func('linear'),
718
+ use_batch_norm=use_batch_norm)
719
+
720
+ # Resolution 1/32 -> 1/64
721
+ block_idx = block_idx + 1
722
+ filter_idx = filter_idx + 1
723
+
724
+ if filter_idx < len(n_filters_encoder_image):
725
+
726
+ in_channels_image, out_channels_image = [
727
+ n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx]
728
+ ]
729
+
730
+ in_channels_depth, out_channels_depth = [
731
+ n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx]
732
+ ]
733
+
734
+ self.blocks6_image, self.blocks6_depth = self._make_layer(
735
+ network_block=resnet_block,
736
+ n_block=n_blocks[block_idx],
737
+ in_channels_image=in_channels_image,
738
+ in_channels_depth=in_channels_depth,
739
+ out_channels_image=out_channels_image,
740
+ out_channels_depth=out_channels_depth,
741
+ stride=2,
742
+ weight_initializer=weight_initializer,
743
+ activation_func=activation_func,
744
+ use_batch_norm=use_batch_norm)
745
+
746
+ if fusion_type == 'add':
747
+ self.conv6_project = net_utils.Conv2d(
748
+ out_channels_depth,
749
+ out_channels_image,
750
+ kernel_size=1,
751
+ stride=1,
752
+ weight_initializer=weight_initializer,
753
+ activation_func=net_utils.activation_func('linear'),
754
+ use_batch_norm=use_batch_norm)
755
+
756
+ if fusion_type == 'weight_and_project':
757
+ self.conv6_weight = net_utils.Conv2d(
758
+ out_channels_depth,
759
+ out_channels_image,
760
+ kernel_size=1,
761
+ stride=1,
762
+ weight_initializer=weight_initializer,
763
+ activation_func=net_utils.activation_func('sigmoid'),
764
+ use_batch_norm=use_batch_norm)
765
+
766
+ self.conv6_project = net_utils.Conv2d(
767
+ out_channels_depth,
768
+ out_channels_image,
769
+ kernel_size=1,
770
+ stride=1,
771
+ weight_initializer=weight_initializer,
772
+ activation_func=net_utils.activation_func('linear'),
773
+ use_batch_norm=use_batch_norm)
774
+ else:
775
+ self.blocks6_image = None
776
+ self.blocks6_depth = None
777
+ self.conv6_weight = None
778
+ self.conv6_project = None
779
+
780
+ # Resolution 1/64 -> 1/128
781
+ block_idx = block_idx + 1
782
+ filter_idx = filter_idx + 1
783
+
784
+ if filter_idx < len(n_filters_encoder_image):
785
+
786
+ in_channels_image, out_channels_image = [
787
+ n_filters_encoder_image[filter_idx - 1], n_filters_encoder_image[filter_idx]
788
+ ]
789
+
790
+ in_channels_depth, out_channels_depth = [
791
+ n_filters_encoder_depth[filter_idx - 1], n_filters_encoder_depth[filter_idx]
792
+ ]
793
+
794
+ self.blocks7_image, self.blocks7_depth = self._make_layer(
795
+ network_block=resnet_block,
796
+ n_block=n_blocks[block_idx],
797
+ in_channels_image=in_channels_image,
798
+ in_channels_depth=in_channels_depth,
799
+ out_channels_image=out_channels_image,
800
+ out_channels_depth=out_channels_depth,
801
+ stride=2,
802
+ weight_initializer=weight_initializer,
803
+ activation_func=activation_func,
804
+ use_batch_norm=use_batch_norm)
805
+
806
+ if fusion_type == 'weight_and_project':
807
+ self.conv7_weight = net_utils.Conv2d(
808
+ out_channels_depth,
809
+ out_channels_image,
810
+ kernel_size=1,
811
+ stride=1,
812
+ weight_initializer=weight_initializer,
813
+ activation_func=net_utils.activation_func('sigmoid'),
814
+ use_batch_norm=use_batch_norm)
815
+
816
+ self.conv7_project = net_utils.Conv2d(
817
+ out_channels_depth,
818
+ out_channels_image,
819
+ kernel_size=1,
820
+ stride=1,
821
+ weight_initializer=weight_initializer,
822
+ activation_func=net_utils.activation_func('linear'),
823
+ use_batch_norm=use_batch_norm)
824
+ else:
825
+ self.blocks7_image = None
826
+ self.blocks7_depth = None
827
+ self.conv7_weight = None
828
+ self.conv7_project = None
829
+
830
+ def _make_layer(self,
831
+ network_block,
832
+ n_block,
833
+ in_channels_image,
834
+ in_channels_depth,
835
+ out_channels_image,
836
+ out_channels_depth,
837
+ stride,
838
+ weight_initializer,
839
+ activation_func,
840
+ use_batch_norm):
841
+ '''
842
+ Creates a layer
843
+ Arg(s):
844
+ network_block : Object
845
+ block type
846
+ n_block : int
847
+ number of blocks to use in layer
848
+ in_channels_image : int
849
+ number of channels in image branch
850
+ in_channels_depth : int
851
+ number of channels in depth branch
852
+ out_channels_image : int
853
+ number of output channels in image branch
854
+ out_channels_depth : int
855
+ number of output channels in depth branch
856
+ stride : int
857
+ stride of convolution
858
+ weight_initializer : str
859
+ kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
860
+ activation_func : func
861
+ activation function after convolution
862
+ use_batch_norm : bool
863
+ if set, then applied batch normalization
864
+ '''
865
+
866
+ blocks_image = []
867
+ blocks_depth = []
868
+
869
+ for n in range(n_block):
870
+
871
+ if n == 0:
872
+ stride = stride
873
+ else:
874
+ in_channels_image = out_channels_image
875
+ in_channels_depth = out_channels_depth
876
+ stride = 1
877
+
878
+ block_image = network_block(
879
+ in_channels=in_channels_image,
880
+ out_channels=out_channels_image,
881
+ stride=stride,
882
+ weight_initializer=weight_initializer,
883
+ activation_func=activation_func,
884
+ use_batch_norm=use_batch_norm)
885
+
886
+ blocks_image.append(block_image)
887
+
888
+ block_depth = network_block(
889
+ in_channels=in_channels_depth,
890
+ out_channels=out_channels_depth,
891
+ stride=stride,
892
+ weight_initializer=weight_initializer,
893
+ activation_func=activation_func,
894
+ use_batch_norm=use_batch_norm)
895
+
896
+ blocks_depth.append(block_depth)
897
+
898
+ blocks_image = torch.nn.Sequential(*blocks_image)
899
+ blocks_depth = torch.nn.Sequential(*blocks_depth)
900
+
901
+ return blocks_image, blocks_depth
902
+
903
+ def forward(self, image, depth):
904
+ '''
905
+ Forward input x through the ResNet model
906
+ Arg(s):
907
+ image : torch.Tensor
908
+ depth : torch.Tensor
909
+ Returns:
910
+ torch.Tensor[float32] : latent vector
911
+ list[torch.Tensor[float32]] : skip connections
912
+ '''
913
+
914
+ layers = []
915
+
916
+ # Resolution 1/1 -> 1/2
917
+ conv1_image = self.conv1_image(image)
918
+ conv1_depth = self.conv1_depth(depth)
919
+
920
+ if self.fusion_type == 'add':
921
+ conv1_project = self.conv1_project(conv1_depth)
922
+ conv1 = conv1_project + conv1_image
923
+ elif self.fusion_type == 'weight':
924
+ conv1_weight = self.conv1_weight(conv1_depth)
925
+ conv1 = conv1_weight * conv1_depth + conv1_image
926
+ elif self.fusion_type == 'weight_and_project':
927
+ conv1_weight = self.conv1_weight(conv1_depth)
928
+ conv1_project = self.conv1_project(conv1_depth)
929
+ conv1 = conv1_weight * conv1_project + conv1_image
930
+ elif self.fusion_type == 'concat':
931
+ conv1 = torch.cat([conv1_depth, conv1_image], dim=1)
932
+ else:
933
+ raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type))
934
+
935
+ layers.append(conv1)
936
+
937
+ # Resolution 1/2 -> 1/4
938
+ max_pool_image = self.max_pool(conv1_image)
939
+ max_pool_depth = self.max_pool(conv1_depth)
940
+
941
+ blocks2_image = self.blocks2_image(max_pool_image)
942
+ blocks2_depth = self.blocks2_depth(max_pool_depth)
943
+
944
+ if self.fusion_type == 'add':
945
+ conv2_project = self.conv2_project(blocks2_depth)
946
+ blocks2 = conv2_project + blocks2_image
947
+ elif self.fusion_type == 'weight':
948
+ conv2_weight = self.conv2_weight(blocks2_depth)
949
+ blocks2 = conv2_weight * blocks2_depth + blocks2_image
950
+ elif self.fusion_type == 'weight_and_project':
951
+ conv2_weight = self.conv2_weight(blocks2_depth)
952
+ conv2_project = self.conv2_project(blocks2_depth)
953
+ blocks2 = conv2_weight * conv2_project + blocks2_image
954
+ elif self.fusion_type == 'concat':
955
+ blocks2 = torch.cat([blocks2_image, blocks2_depth], dim=1)
956
+ else:
957
+ raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type))
958
+
959
+ layers.append(blocks2)
960
+
961
+ # Resolution 1/4 -> 1/8
962
+ blocks3_image = self.blocks3_image(blocks2_image)
963
+ blocks3_depth = self.blocks3_depth(blocks2_depth)
964
+
965
+ if self.fusion_type == 'add':
966
+ conv3_project = self.conv3_project(blocks3_depth)
967
+ blocks3 = conv3_project + blocks3_image
968
+ elif self.fusion_type == 'weight':
969
+ conv3_weight = self.conv3_weight(blocks3_depth)
970
+ blocks3 = conv3_weight * blocks3_depth + blocks3_image
971
+ elif self.fusion_type == 'weight_and_project':
972
+ conv3_weight = self.conv3_weight(blocks3_depth)
973
+ conv3_project = self.conv3_project(blocks3_depth)
974
+ blocks3 = conv3_weight * conv3_project + blocks3_image
975
+ elif self.fusion_type == 'concat':
976
+ blocks3 = torch.cat([blocks3_image, blocks3_depth], dim=1)
977
+ else:
978
+ raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type))
979
+
980
+ layers.append(blocks3)
981
+
982
+ # Resolution 1/8 -> 1/16
983
+ blocks4_image = self.blocks4_image(blocks3_image)
984
+ blocks4_depth = self.blocks4_depth(blocks3_depth)
985
+
986
+ if self.fusion_type == 'add':
987
+ conv4_project = self.conv4_project(blocks4_depth)
988
+ blocks4 = conv4_project + blocks4_image
989
+ elif self.fusion_type == 'weight':
990
+ conv4_weight = self.conv4_weight(blocks4_depth)
991
+ blocks4 = conv4_weight * blocks4_depth + blocks4_image
992
+ elif self.fusion_type == 'weight_and_project':
993
+ conv4_weight = self.conv4_weight(blocks4_depth)
994
+ conv4_project = self.conv4_project(blocks4_depth)
995
+ blocks4 = conv4_weight * conv4_project + blocks4_image
996
+ elif self.fusion_type == 'concat':
997
+ blocks4 = torch.cat([blocks4_image, blocks4_depth], dim=1)
998
+ else:
999
+ raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type))
1000
+
1001
+ layers.append(blocks4)
1002
+
1003
+ # Resolution 1/16 -> 1/32
1004
+ blocks5_image = self.blocks5_image(blocks4_image)
1005
+ blocks5_depth = self.blocks5_depth(blocks4_depth)
1006
+
1007
+ if self.fusion_type == 'add':
1008
+ conv5_project = self.conv5_project(blocks5_depth)
1009
+ blocks5 = conv5_project + blocks5_image
1010
+ elif self.fusion_type == 'weight':
1011
+ conv5_weight = self.conv5_weight(blocks5_depth)
1012
+ blocks5 = conv5_weight * blocks5_depth + blocks5_image
1013
+ elif self.fusion_type == 'weight_and_project':
1014
+ conv5_weight = self.conv5_weight(blocks5_depth)
1015
+ conv5_project = self.conv5_project(blocks5_depth)
1016
+ blocks5 = conv5_weight * conv5_project + blocks5_image
1017
+ elif self.fusion_type == 'concat':
1018
+ blocks5 = torch.cat([blocks5_image, blocks5_depth], dim=1)
1019
+ else:
1020
+ raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type))
1021
+
1022
+ layers.append(blocks5)
1023
+
1024
+ # Resolution 1/32 -> 1/64
1025
+ if self.blocks6_image is not None and self.blocks6_depth is not None:
1026
+ blocks6_image = self.blocks6_image(blocks5_image)
1027
+ blocks6_depth = self.blocks6_depth(blocks5_depth)
1028
+
1029
+ if self.fusion_type == 'add':
1030
+ conv6_project = self.conv6_project(blocks6_depth)
1031
+ blocks6 = conv6_project + blocks6_image
1032
+ elif self.fusion_type == 'weight':
1033
+ conv6_weight = self.conv6_weight(blocks6_depth)
1034
+ blocks6 = conv6_weight * blocks6_depth + blocks6_image
1035
+ elif self.fusion_type == 'weight_and_project':
1036
+ conv6_weight = self.conv6_weight(blocks6_depth)
1037
+ conv6_project = self.conv6_project(blocks6_depth)
1038
+ blocks6 = conv6_weight * conv6_project + blocks6_image
1039
+ elif self.fusion_type == 'concat':
1040
+ blocks6 = torch.cat([blocks6_image, blocks6_depth], dim=1)
1041
+ else:
1042
+ raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type))
1043
+
1044
+ layers.append(blocks6)
1045
+
1046
+ # Resolution 1/64 -> 1/128
1047
+ if self.blocks7_image is not None and self.blocks7_depth is not None:
1048
+ blocks7_image = self.blocks7_image(blocks6_image)
1049
+ blocks7_depth = self.blocks7_depth(blocks6_depth)
1050
+
1051
+ if self.fusion_type == 'add':
1052
+ conv7_project = self.conv7_project(blocks7_depth)
1053
+ blocks7 = conv7_project + blocks7_image
1054
+ elif self.fusion_type == 'weight':
1055
+ conv7_weight = self.conv7_weight(blocks7_depth)
1056
+ blocks7 = conv7_weight * blocks7_depth + blocks7_image
1057
+ elif self.fusion_type == 'weight_and_project':
1058
+ conv7_weight = self.conv7_weight(blocks7_depth)
1059
+ conv7_project = self.conv7_project(blocks7_depth)
1060
+ blocks7 = conv7_weight * conv7_project + blocks7_image
1061
+ elif self.fusion_type == 'concat':
1062
+ blocks7 = torch.cat([blocks7_image, blocks7_depth], dim=1)
1063
+ else:
1064
+ raise ValueError('Unsupported fusion type: {}'.format(self.fusion_type))
1065
+
1066
+ layers.append(blocks7)
1067
+
1068
+ return layers[-1], layers[:-1]
1069
+
1070
+
1071
+ class RCNetEncoder(torch.nn.Module):
1072
+ '''
1073
+ Radar association network
1074
+ Arg(s):
1075
+ in_channels_image : int
1076
+ number of input channels for image (RGB) branch
1077
+ in_channels_depth : int
1078
+ number of input channels for depth branch
1079
+ n_filters_encoder_image : int
1080
+ number of filters for image (RGB) branch
1081
+ n_neurons_encoder_depth : int
1082
+ number of neurons for depth (radar) branch
1083
+ latent_size_depth : int
1084
+ size of latent vector
1085
+ weight_initializer : str
1086
+ kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
1087
+ activation_func : func
1088
+ activation function after convolution
1089
+ use_batch_norm : bool
1090
+ if set, then applied batch normalization
1091
+ '''
1092
+ def __init__(self,
1093
+ input_channels_image=3,
1094
+ input_channels_depth=3,
1095
+ input_patch_size_image=(900, 288),
1096
+ n_filters_encoder_image=[32, 64, 128, 128, 128],
1097
+ n_neurons_encoder_depth=[32, 64, 128, 128, 128],
1098
+ latent_size_depth=128 * 29 * 10,
1099
+ weight_initializer='kaiming_uniform',
1100
+ activation_func='leaky_relu',
1101
+ use_batch_norm=False):
1102
+ super(RCNetEncoder, self).__init__()
1103
+
1104
+ self.n_neuron_latent_depth = n_neurons_encoder_depth[-1]
1105
+
1106
+ self.encoder_image = ResNetEncoder(
1107
+ n_layer=18,
1108
+ input_channels=input_channels_image,
1109
+ n_filters=n_filters_encoder_image,
1110
+ weight_initializer=weight_initializer,
1111
+ activation_func=activation_func,
1112
+ use_batch_norm=use_batch_norm)
1113
+
1114
+ self.attention = LocalFeatureTransformer(['self','cross'], n_layers=4, d_model=self.n_neuron_latent_depth)
1115
+
1116
+ self.encoder_depth = FullyConnectedEncoder(
1117
+ input_channels=input_channels_depth,
1118
+ n_neurons=n_neurons_encoder_depth,
1119
+ latent_size=latent_size_depth,
1120
+ weight_initializer=weight_initializer,
1121
+ activation_func=activation_func)
1122
+
1123
+ self.input_patch_size_image =input_patch_size_image
1124
+
1125
+ def forward(self, image, points, b_boxes):
1126
+ # Image shape: (B, C, H, W) # Should be (B, 3, 768, 288)
1127
+ # points shape: (B*K, X)
1128
+ # b_boxes: [(K, 4) * B], this should be a list with B elements, and each element is (K, 4) size
1129
+ # K is the number of radar points per image
1130
+ # X is the radar dimension
1131
+
1132
+
1133
+ # Define dimensions
1134
+ shape = self.input_patch_size_image
1135
+ latent_height = int(shape[-2] // 32.0)
1136
+ latent_width = int(shape[-1] // 32.0)
1137
+ batch_size = image.shape[0]
1138
+
1139
+ # Define scales and feature sizes
1140
+ skip_scales = [ 1 /2.0, 1/ 4.0, 1 / 8.0, 1 / 16.0, 1 / 32.0, 1 / 64.0, 1 / 128.0]
1141
+ skip_feature_sizes = [
1142
+ (int(shape[-2] * skip_scale),
1143
+ int(shape[-1] * skip_scale))
1144
+ for skip_scale in skip_scales
1145
+ ] # Should be [(384, 144), (192, 72), (96, 36), (48, 18)]
1146
+
1147
+ latent_scale = 1 / 32.0
1148
+ latent_feature_size = (latent_height, latent_width) # Should be (24, 9)
1149
+
1150
+ # Forward the entire image
1151
+ latent_image, skips_image = self.encoder_image(image)
1152
+
1153
+ # ROI pooling on latent images
1154
+ latent_image_pooled = torchvision.ops.roi_pool(
1155
+ latent_image, b_boxes,
1156
+ spatial_scale=latent_scale,
1157
+ output_size=latent_feature_size
1158
+ ) # (N*K, C, H, W)
1159
+
1160
+ # ROI pooling on the skips
1161
+ skips_image_pooled = []
1162
+ for skip_image_idx in range(len(skips_image)):
1163
+ skips_image_pooled.append(
1164
+ torchvision.ops.roi_pool(
1165
+ skips_image[skip_image_idx], b_boxes,
1166
+ spatial_scale=skip_scales[skip_image_idx],
1167
+ output_size=skip_feature_sizes[skip_image_idx]
1168
+ ) # (N*K, C, H, W)
1169
+ )
1170
+
1171
+ # Radar points size: (bath_size * total_points_sampled, 3)
1172
+ # latent_depth size: (batch_size * total_points_sampled, n_neuron_latent_depth, patch_w//32, patch_h//32)
1173
+ # latent_image_pooled size = latent_depth size
1174
+ latent_depth = self.encoder_depth(points)
1175
+ latent_depth = latent_depth.view(points.shape[0], self.n_neuron_latent_depth, -1, latent_width)
1176
+
1177
+ latent_depth_reshape = latent_depth.view(latent_depth.shape[0], latent_depth.shape[1], -1).permute(0, 2, 1)
1178
+ latent_image_pooled_reshape = latent_image_pooled.view(latent_image_pooled.shape[0],
1179
+ latent_image_pooled.shape[1], -1).permute(0, 2, 1)
1180
+ latent_depth_tf, latent_image_pooled_tf = self.attention(latent_depth_reshape, latent_image_pooled_reshape)
1181
+ latent_depth_tf = latent_depth_tf.permute(0, 2, 1).view(latent_depth.shape)
1182
+ latent_image_pooled_tf = latent_image_pooled_tf.permute(0, 2, 1).view(latent_image_pooled.shape)
1183
+
1184
+ # Concatenate the features
1185
+ # latent = torch.cat([latent_image_pooled, latent_depth], dim=1)
1186
+ latent = torch.cat([latent_image_pooled_tf, latent_depth_tf], dim=1)
1187
+ return latent, skips_image_pooled
1188
+
1189
+
1190
+ '''
1191
+ Decoder
1192
+ '''
1193
+
1194
+
1195
+ class MultiScaleDecoder(torch.nn.Module):
1196
+ '''
1197
+ Multi-scale decoder with skip connections
1198
+ Arg(s):
1199
+ input_channels : int
1200
+ number of channels in input latent vector
1201
+ output_channels : int
1202
+ number of channels or classes in output
1203
+ n_resolution : int
1204
+ number of output resolutions (scales) for multi-scale prediction
1205
+ n_filters : int list
1206
+ number of filters to use at each decoder block
1207
+ n_skips : int list
1208
+ number of filters from skip connections
1209
+ weight_initializer : str
1210
+ kaiming_normal, kaiming_uniform, xavier_normal, xavier_uniform
1211
+ activation_func : func
1212
+ activation function after convolution
1213
+ output_func : func
1214
+ activation function for output
1215
+ use_batch_norm : bool
1216
+ if set, then applied batch normalization
1217
+ deconv_type : str
1218
+ deconvolution types available: transpose, up
1219
+ '''
1220
+
1221
+ def __init__(self,
1222
+ input_channels=256,
1223
+ output_channels=1,
1224
+ n_resolution=1,
1225
+ n_filters=[256, 128, 64, 32, 16],
1226
+ n_skips=[256, 128, 64, 32, 0],
1227
+ weight_initializer='kaiming_uniform',
1228
+ activation_func='leaky_relu',
1229
+ output_func='linear',
1230
+ use_batch_norm=False,
1231
+ deconv_type='up'):
1232
+ super(MultiScaleDecoder, self).__init__()
1233
+
1234
+ network_depth = len(n_filters)
1235
+
1236
+ assert network_depth < 8, 'Does not support network depth of 8 or more'
1237
+ assert n_resolution > 0 and n_resolution < network_depth
1238
+
1239
+ self.n_resolution = n_resolution
1240
+ self.output_func = output_func
1241
+
1242
+ activation_func = net_utils.activation_func(activation_func)
1243
+ output_func = net_utils.activation_func(output_func)
1244
+
1245
+ # Upsampling from lower to full resolution requires multi-scale
1246
+ if 'upsample' in self.output_func and self.n_resolution < 2:
1247
+ self.n_resolution = 2
1248
+
1249
+ filter_idx = 0
1250
+
1251
+ in_channels, skip_channels, out_channels = [
1252
+ input_channels, n_skips[filter_idx], n_filters[filter_idx]
1253
+ ]
1254
+
1255
+ # Resolution 1/128 -> 1/64
1256
+ if network_depth > 6:
1257
+ self.deconv6 = net_utils.DecoderBlock(
1258
+ in_channels,
1259
+ skip_channels,
1260
+ out_channels,
1261
+ weight_initializer=weight_initializer,
1262
+ activation_func=activation_func,
1263
+ use_batch_norm=use_batch_norm,
1264
+ deconv_type=deconv_type)
1265
+
1266
+ filter_idx = filter_idx + 1
1267
+
1268
+ in_channels, skip_channels, out_channels = [
1269
+ n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx]
1270
+ ]
1271
+ else:
1272
+ self.deconv6 = None
1273
+
1274
+ # Resolution 1/64 -> 1/32
1275
+ if network_depth > 5:
1276
+ self.deconv5 = net_utils.DecoderBlock(
1277
+ in_channels,
1278
+ skip_channels,
1279
+ out_channels,
1280
+ weight_initializer=weight_initializer,
1281
+ activation_func=activation_func,
1282
+ use_batch_norm=use_batch_norm,
1283
+ deconv_type=deconv_type)
1284
+
1285
+ filter_idx = filter_idx + 1
1286
+
1287
+ in_channels, skip_channels, out_channels = [
1288
+ n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx]
1289
+ ]
1290
+ else:
1291
+ self.deconv5 = None
1292
+
1293
+ # Resolution 1/32 -> 1/16
1294
+ self.deconv4 = net_utils.DecoderBlock(
1295
+ in_channels,
1296
+ skip_channels,
1297
+ out_channels,
1298
+ weight_initializer=weight_initializer,
1299
+ activation_func=activation_func,
1300
+ use_batch_norm=use_batch_norm,
1301
+ deconv_type=deconv_type)
1302
+
1303
+ # Resolution 1/16 -> 1/8
1304
+ filter_idx = filter_idx + 1
1305
+
1306
+ in_channels, skip_channels, out_channels = [
1307
+ n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx]
1308
+ ]
1309
+
1310
+ self.deconv3 = net_utils.DecoderBlock(
1311
+ in_channels,
1312
+ skip_channels,
1313
+ out_channels,
1314
+ weight_initializer=weight_initializer,
1315
+ activation_func=activation_func,
1316
+ use_batch_norm=use_batch_norm,
1317
+ deconv_type=deconv_type)
1318
+
1319
+ if self.n_resolution > 3:
1320
+ self.output3 = net_utils.Conv2d(
1321
+ out_channels,
1322
+ output_channels,
1323
+ kernel_size=3,
1324
+ stride=1,
1325
+ weight_initializer=weight_initializer,
1326
+ activation_func=output_func,
1327
+ use_batch_norm=False)
1328
+
1329
+ # Resolution 1/8 -> 1/4
1330
+ filter_idx = filter_idx + 1
1331
+
1332
+ in_channels, skip_channels, out_channels = [
1333
+ n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx]
1334
+ ]
1335
+
1336
+ if self.n_resolution > 3:
1337
+ skip_channels = skip_channels + output_channels
1338
+
1339
+ self.deconv2 = net_utils.DecoderBlock(
1340
+ in_channels,
1341
+ skip_channels,
1342
+ out_channels,
1343
+ weight_initializer=weight_initializer,
1344
+ activation_func=activation_func,
1345
+ use_batch_norm=use_batch_norm,
1346
+ deconv_type=deconv_type)
1347
+
1348
+ if self.n_resolution > 2:
1349
+ self.output2 = net_utils.Conv2d(
1350
+ out_channels,
1351
+ output_channels,
1352
+ kernel_size=3,
1353
+ stride=1,
1354
+ weight_initializer=weight_initializer,
1355
+ activation_func=output_func,
1356
+ use_batch_norm=False)
1357
+
1358
+ # Resolution 1/4 -> 1/2
1359
+ filter_idx = filter_idx + 1
1360
+
1361
+ in_channels, skip_channels, out_channels = [
1362
+ n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx]
1363
+ ]
1364
+
1365
+ if self.n_resolution > 2:
1366
+ skip_channels = skip_channels + output_channels
1367
+
1368
+ self.deconv1 = net_utils.DecoderBlock(
1369
+ in_channels,
1370
+ skip_channels,
1371
+ out_channels,
1372
+ weight_initializer=weight_initializer,
1373
+ activation_func=activation_func,
1374
+ use_batch_norm=use_batch_norm,
1375
+ deconv_type=deconv_type)
1376
+
1377
+ if self.n_resolution > 1:
1378
+ self.output1 = net_utils.Conv2d(
1379
+ out_channels,
1380
+ output_channels,
1381
+ kernel_size=3,
1382
+ stride=1,
1383
+ weight_initializer=weight_initializer,
1384
+ activation_func=output_func,
1385
+ use_batch_norm=False)
1386
+
1387
+ # Resolution 1/2 -> 1/1
1388
+ filter_idx = filter_idx + 1
1389
+
1390
+ in_channels, skip_channels, out_channels = [
1391
+ n_filters[filter_idx - 1], n_skips[filter_idx], n_filters[filter_idx]
1392
+ ]
1393
+
1394
+ if self.n_resolution > 1:
1395
+ skip_channels = skip_channels + output_channels
1396
+
1397
+ self.deconv0 = net_utils.DecoderBlock(
1398
+ in_channels,
1399
+ skip_channels,
1400
+ out_channels,
1401
+ weight_initializer=weight_initializer,
1402
+ activation_func=activation_func,
1403
+ use_batch_norm=use_batch_norm,
1404
+ deconv_type=deconv_type)
1405
+
1406
+ self.output0 = net_utils.Conv2d(
1407
+ out_channels,
1408
+ output_channels,
1409
+ kernel_size=3,
1410
+ stride=1,
1411
+ weight_initializer=weight_initializer,
1412
+ activation_func=output_func,
1413
+ use_batch_norm=False)
1414
+
1415
+ def forward(self, x, skips, shape=None):
1416
+ '''
1417
+ Forward latent vector x through decoder network
1418
+ Arg(s):
1419
+ x : torch.Tensor[float32]
1420
+ latent vector
1421
+ skips : list[torch.Tensor[float32]]
1422
+ list of skip connection tensors (earlier are larger resolution)
1423
+ shape : tuple[int]
1424
+ (height, width) tuple denoting output size
1425
+ Returns:
1426
+ list[torch.Tensor[float32]] : list of outputs at multiple scales
1427
+ '''
1428
+
1429
+ layers = [x]
1430
+ outputs = []
1431
+
1432
+ # Start at the end and walk backwards through skip connections
1433
+ n = len(skips) - 1
1434
+
1435
+ # Resolution 1/128 -> 1/64
1436
+ if self.deconv6 is not None:
1437
+ layers.append(self.deconv6(layers[-1], skips[n]))
1438
+ n = n - 1
1439
+
1440
+ # Resolution 1/64 -> 1/32
1441
+ if self.deconv5 is not None:
1442
+ layers.append(self.deconv5(layers[-1], skips[n]))
1443
+ n = n - 1
1444
+
1445
+ # Resolution 1/32 -> 1/16
1446
+ layers.append(self.deconv4(layers[-1], skips[n]))
1447
+
1448
+ # Resolution 1/16 -> 1/8
1449
+ n = n - 1
1450
+
1451
+ layers.append(self.deconv3(layers[-1], skips[n]))
1452
+
1453
+ if self.n_resolution > 3:
1454
+ output3 = self.output3(layers[-1])
1455
+ outputs.append(output3)
1456
+
1457
+ upsample_output3 = torch.nn.functional.interpolate(
1458
+ input=outputs[-1],
1459
+ scale_factor=2,
1460
+ mode='bilinear',
1461
+ align_corners=True)
1462
+
1463
+ # Resolution 1/8 -> 1/4
1464
+ n = n - 1
1465
+
1466
+ skip = torch.cat([skips[n], upsample_output3], dim=1) if self.n_resolution > 3 else skips[n]
1467
+ layers.append(self.deconv2(layers[-1], skip))
1468
+
1469
+ if self.n_resolution > 2:
1470
+ output2 = self.output2(layers[-1])
1471
+ outputs.append(output2)
1472
+
1473
+ upsample_output2 = torch.nn.functional.interpolate(
1474
+ input=outputs[-1],
1475
+ scale_factor=2,
1476
+ mode='bilinear',
1477
+ align_corners=True)
1478
+
1479
+ # Resolution 1/4 -> 1/2
1480
+ n = n - 1
1481
+
1482
+ skip = torch.cat([skips[n], upsample_output2], dim=1) if self.n_resolution > 2 else skips[n]
1483
+ layers.append(self.deconv1(layers[-1], skip))
1484
+
1485
+ if self.n_resolution > 1:
1486
+ output1 = self.output1(layers[-1])
1487
+ outputs.append(output1)
1488
+
1489
+ upsample_output1 = torch.nn.functional.interpolate(
1490
+ input=outputs[-1],
1491
+ scale_factor=2,
1492
+ mode='bilinear',
1493
+ align_corners=True)
1494
+
1495
+ # Resolution 1/2 -> 1/1
1496
+ n = n - 1
1497
+
1498
+ if 'upsample' in self.output_func:
1499
+ output0 = upsample_output1
1500
+ else:
1501
+ if self.n_resolution > 1:
1502
+ # If there is skip connection at layer 0
1503
+ skip = torch.cat([skips[n], upsample_output1], dim=1) if n == 0 else upsample_output1
1504
+ layers.append(self.deconv0(layers[-1], skip))
1505
+ else:
1506
+
1507
+ if n == 0:
1508
+ layers.append(self.deconv0(layers[-1], skips[n]))
1509
+ else:
1510
+ layers.append(self.deconv0(layers[-1], shape=shape[-2:]))
1511
+
1512
+ output0 = self.output0(layers[-1])
1513
+
1514
+ outputs.append(output0)
1515
+
1516
+ return outputs