Instructions to use phi-lab-rice/GRADE with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use phi-lab-rice/GRADE with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("phi-lab-rice/GRADE", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Release all GRADE models, checkpoints, and reviewed evaluation code (part 2)
Browse filesAdd 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
- src/Ablation/ours_radar_no_doppler/inference.py +267 -0
- src/Ablation/ours_radar_no_doppler/iq1m_dataset.py +305 -0
- src/Ablation/ours_radar_no_doppler/radar_depth.py +406 -0
- src/Ablation/ours_radar_no_doppler/rice_dataset.py +159 -0
- src/Ablation/ours_radar_no_doppler/split.json +34 -0
- src/Ablation/ours_radar_no_grad/inference.py +16 -0
- src/Baselines/cafnet/collate_fn_helpers.py +404 -0
- src/Baselines/cafnet/dataloader.py +100 -0
- src/Baselines/cafnet/extract_pcd_from_depth.py +96 -0
- src/Baselines/cafnet/inference.py +224 -0
- src/Baselines/cafnet/inference_config.yaml +29 -0
- src/Baselines/cafnet/models/bts.py +367 -0
- src/Baselines/cafnet/models/model.py +28 -0
- src/Baselines/cafnet/models/radar.py +212 -0
- src/Baselines/cafnet/rice_dataset.py +121 -0
- src/Baselines/cafnet/split.json +14 -0
- src/Baselines/cafnet_no_smoke/collate_fn_helpers.py +404 -0
- src/Baselines/cafnet_no_smoke/dataloader.py +100 -0
- src/Baselines/cafnet_no_smoke/extract_pcd_from_depth.py +96 -0
- src/Baselines/cafnet_no_smoke/inference.py +224 -0
- src/Baselines/cafnet_no_smoke/inference_config.yaml +29 -0
- src/Baselines/cafnet_no_smoke/models/bts.py +367 -0
- src/Baselines/cafnet_no_smoke/models/model.py +28 -0
- src/Baselines/cafnet_no_smoke/models/radar.py +212 -0
- src/Baselines/cafnet_no_smoke/rice_dataset.py +123 -0
- src/Baselines/cafnet_no_smoke/split.json +14 -0
- src/Baselines/da3/inference.py +179 -0
- src/Baselines/grt/augmentations.py +193 -0
- src/Baselines/grt/dataloader.py +330 -0
- src/Baselines/grt/grt_model.py +585 -0
- src/Baselines/grt/inference.py +220 -0
- src/Baselines/grt/split.json +16 -0
- src/Baselines/grt_image/augmentations.py +193 -0
- src/Baselines/grt_image/dataloader.py +344 -0
- src/Baselines/grt_image/grt_image_resnet_inference.example.yaml +16 -0
- src/Baselines/grt_image/grt_model.py +799 -0
- src/Baselines/grt_image/inference.py +224 -0
- src/Baselines/grt_image/split.json +16 -0
- src/Baselines/radarcam-depth/data/SML_dataset.py +83 -0
- src/Baselines/radarcam-depth/data/data_utils.py +326 -0
- src/Baselines/radarcam-depth/data/datasets.py +392 -0
- src/Baselines/radarcam-depth/linear_attention.py +184 -0
- src/Baselines/radarcam-depth/modules/estimator.py +188 -0
- src/Baselines/radarcam-depth/modules/midas/base_model.py +12 -0
- src/Baselines/radarcam-depth/modules/midas/blocks.py +197 -0
- src/Baselines/radarcam-depth/modules/midas/midas_net_custom.py +138 -0
- src/Baselines/radarcam-depth/modules/midas/normalization.py +109 -0
- src/Baselines/radarcam-depth/modules/midas/transforms.py +263 -0
- src/Baselines/radarcam-depth/modules/midas/utils.py +237 -0
- 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
|