English
DeepChoice / dataset_generation /gridnet_hd_batch_generation.py
antoine.carreaud67
Restore paper experiment plans and update README
4ccb60a
Raw
History Blame Contribute Delete
84.9 kB
import argparse
import json
import logging
import os
import subprocess
import sys
import time
from pathlib import Path
import cv2 as cv
import h5py
import jax
import jax.numpy as jnp
import laspy
import numpy as np
import torch
import yaml
from tqdm import tqdm
from dataset_generation.fast_proj import (
compute_depth_map,
f_frame_agi,
parse_calibration_xml,
read_camera_file,
)
from dataset_generation.batch_serialization import (
save_compact_payloads_to_pt,
save_dataset_to_pt_parallel,
save_tile_observations_to_pt,
)
from dataset_generation.gridnet_hd_manifest import load_or_build_manifest, resolve_logits_dir
from utils.utilities import load_normals_from_h5, map_main_class_indices
from utils.visibility_criteria import (
compute_criteria_maps,
compute_visibility_criteria_numpy,
radian_to_degree,
round_to_three_digits,
)
logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
_PROJECTION_FN_CACHE = {}
_DEPTH_FN_CACHE = {}
def morton2D(x, y, bits=21):
z = 0
for i in range(bits):
z |= ((x >> i) & 1) << (2 * i) | ((y >> i) & 1) << (2 * i + 1)
return z
def load_calibration(calibration_file):
try:
return parse_calibration_xml(calibration_file)
except Exception as exc:
logging.error("Failed to load calibration %s: %s", calibration_file, exc)
sys.exit(1)
def load_camera_parameters(camera_file, offset):
try:
return read_camera_file(camera_file, offset)
except Exception as exc:
logging.error("Failed to load camera poses %s: %s", camera_file, exc)
sys.exit(1)
def _extract_labels(las_data):
if hasattr(las_data, "ground_truth"):
return np.asarray(las_data.ground_truth)
return None
def load_las_points(las_file_path):
try:
las_data = laspy.read(las_file_path)
offset = np.array([np.min(las_data.x), np.min(las_data.y), np.min(las_data.z)])
labels = _extract_labels(las_data)
points = {
"X": np.array(las_data.x - offset[0], dtype=np.float32),
"Y": np.array(las_data.y - offset[1], dtype=np.float32),
"Z": np.array(las_data.z - offset[2], dtype=np.float32),
"GT": labels,
"GT_AVAILABLE": labels is not None,
}
codes = morton2D(points["X"].astype(int), points["Y"].astype(int))
perm = np.argsort(codes)
for key in ("X", "Y", "Z"):
points[key] = points[key][perm]
if points["GT"] is not None:
points["GT"] = points["GT"][perm]
points_for_normals = np.vstack((points["X"], points["Y"], points["Z"])).T.astype(np.float32)
return points, offset, las_data, points_for_normals
except Exception as exc:
logging.error("Failed to load LAS %s: %s", las_file_path, exc)
sys.exit(1)
def _get_projection_fn(param_keys):
cache_key = tuple(param_keys)
if cache_key not in _PROJECTION_FN_CACHE:
in_axes_dict = {key: None for key in param_keys}
for key in ("X", "Y", "Z"):
in_axes_dict[key] = 0
_PROJECTION_FN_CACHE[cache_key] = jax.jit(jax.vmap(f_frame_agi, in_axes=(in_axes_dict,)))
return _PROJECTION_FN_CACHE[cache_key]
def _get_depth_fn(buffer_size, threshold):
cache_key = (int(buffer_size), float(threshold))
if cache_key not in _DEPTH_FN_CACHE:
_DEPTH_FN_CACHE[cache_key] = jax.jit(
lambda px, py, pz, dmap: compute_depth_map(
px,
py,
pz,
dmap,
buffer_size=buffer_size,
threshold=threshold,
)
)
return _DEPTH_FN_CACHE[cache_key]
def empty_compact_observations(cam_key, logits_dim):
return {
"camera": cam_key,
"point_indices": np.empty((0,), dtype=np.int64),
"pixel_coords": np.empty((0, 2), dtype=np.int32),
"visibility": np.empty((0, 6), dtype=np.float32),
"ground_truth": np.empty((0,), dtype=np.int16),
"logit_vectors": np.empty((0, logits_dim), dtype=np.float32),
}
def filter_compact_observations_to_valid_points(observations, valid_point_mask):
if not observations or "point_indices" not in observations:
return observations
if np.any(valid_point_mask):
return {
"camera": observations["camera"],
"point_indices": observations["point_indices"][valid_point_mask],
"pixel_coords": observations["pixel_coords"][valid_point_mask],
"visibility": observations["visibility"][valid_point_mask],
"ground_truth": observations["ground_truth"][valid_point_mask],
"logit_vectors": observations["logit_vectors"][valid_point_mask],
}
return empty_compact_observations(observations["camera"], observations["logit_vectors"].shape[1])
def project_points_raw(cam_params, global_params, subscale_factor, target_device=None):
params = global_params.copy()
params.update(cam_params)
f_proj = _get_projection_fn(params.keys())
if target_device is None:
x, y, z, in_bounds = f_proj(params)
else:
with jax.default_device(target_device):
x, y, z, in_bounds = f_proj(params)
full_i = x.astype(int)
full_j = y.astype(int)
i = (x / subscale_factor).astype(int)
j = (y / subscale_factor).astype(int)
return x, y, z, in_bounds, full_i, full_j, i, j
def project_points_for_camera(cam_params, global_params, subscale_factor, buffer_size, threshold, target_device=None):
_, _, z, in_bounds, full_i, full_j, i, j = project_points_raw(
cam_params,
global_params,
subscale_factor,
target_device=target_device,
)
if jnp.sum(in_bounds) == 0:
return None, None, None, None, None, None
depth_map = jnp.full((global_params["height"] // subscale_factor, global_params["width"] // subscale_factor), jnp.inf)
compute_depth_map_jit = _get_depth_fn(buffer_size, threshold)
if target_device is None:
_, visible_in_bounds = compute_depth_map_jit(i[in_bounds], j[in_bounds], z[in_bounds], depth_map)
else:
with jax.default_device(target_device):
_, visible_in_bounds = compute_depth_map_jit(i[in_bounds], j[in_bounds], z[in_bounds], depth_map)
return in_bounds, visible_in_bounds, full_i, full_j, i, j
def rot_x_np(omega):
return np.array(
[[1.0, 0.0, 0.0], [0.0, np.cos(omega), np.sin(omega)], [0.0, -np.sin(omega), np.cos(omega)]],
dtype=np.float32,
)
def rot_y_np(phi):
return np.array(
[[np.cos(phi), 0.0, -np.sin(phi)], [0.0, 1.0, 0.0], [np.sin(phi), 0.0, np.cos(phi)]],
dtype=np.float32,
)
def rot_z_np(kappa):
return np.array(
[[np.cos(kappa), np.sin(kappa), 0.0], [-np.sin(kappa), np.cos(kappa), 0.0], [0.0, 0.0, 1.0]],
dtype=np.float32,
)
def rot_zyx_np(omega, phi, kappa):
return rot_z_np(kappa) @ rot_y_np(phi) @ rot_x_np(omega)
def prefilter_points_for_camera(points_xyz, cam_params, calibration_params, cfg):
if not cfg["load"].get("prefilter_enabled", True):
return np.ones(points_xyz.shape[0], dtype=bool)
camera_position = np.array([cam_params["Xs"], cam_params["Ys"], cam_params["Zs"]], dtype=np.float32)
rotation = rot_zyx_np(cam_params["omega"], cam_params["phi"], cam_params["kappa"])
relative = points_xyz - camera_position[None, :]
rms = relative @ rotation.T
z_cam = -rms[:, 2]
mask = z_cam > 0.0
if not np.any(mask):
return mask
max_distance = cfg["load"].get("max_distance")
if max_distance is not None:
mask &= np.linalg.norm(relative, axis=1) <= float(max_distance)
if not np.any(mask):
return mask
margin = float(cfg["load"].get("prefilter_margin", 0.15))
width = float(calibration_params["width"])
height = float(calibration_params["height"])
focal = float(calibration_params["f"])
w_half = (width / focal / 2.0) * (1.0 + margin)
h_half = (height / focal / 2.0) * (1.0 + margin)
x_norm = np.zeros(points_xyz.shape[0], dtype=np.float32)
y_norm = np.zeros(points_xyz.shape[0], dtype=np.float32)
valid = mask
x_norm[valid] = -rms[valid, 0] / rms[valid, 2]
y_norm[valid] = -rms[valid, 1] / rms[valid, 2]
return valid & (x_norm >= -w_half) & (x_norm < w_half) & (y_norm >= -h_half) & (y_norm < h_half)
def load_logits_image(logits_path, expected_shape=None):
logits_img = np.load(logits_path)
if logits_img.ndim != 3:
raise ValueError(f"Expected logits image of shape [H, W, C], got {logits_img.shape}")
if expected_shape is not None and logits_img.shape[:2] != expected_shape:
raise ValueError(f"Logits/image size mismatch for {logits_path}: {logits_img.shape} vs {expected_shape}")
return logits_img.astype(np.float32)
def build_compact_observations(
cam_key,
points_for_normals,
normals,
points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cam_params,
in_bounds,
visible_in_bounds,
full_i,
full_j,
logits_i,
logits_j,
prefilter_mask,
allow_missing_ground_truth=False,
global_indices=None,
remap_ground_truth=True,
):
logits_dim = logits_img.shape[2]
in_bounds_indices = np.flatnonzero(in_bounds)
visible_indices = in_bounds_indices[np.asarray(visible_in_bounds, dtype=bool)]
if visible_indices.size == 0:
return empty_compact_observations(cam_key, logits_dim), 0, 0
filtered_points = points_for_normals[prefilter_mask]
filtered_normals = normals[prefilter_mask]
filtered_local_indices = np.nonzero(prefilter_mask)[0]
crit = compute_visibility_criteria_numpy(filtered_points[visible_indices], filtered_normals[visible_indices], cam_params)
angles = np.degrees(crit[:, 0]).astype(np.float32)
distances = np.round(crit[:, 1], 3).astype(np.float32)
local_chunk_indices = filtered_local_indices[visible_indices]
original_point_indices = (
local_chunk_indices.astype(np.int64)
if global_indices is None
else global_indices[local_chunk_indices].astype(np.int64)
)
u_full = full_i[visible_indices].astype(np.int32)
v_full = full_j[visible_indices].astype(np.int32)
u_logit = logits_i[visible_indices].astype(np.int32)
v_logit = logits_j[visible_indices].astype(np.int32)
contrast_values = contrast_map[v_full, u_full].astype(np.float32)
blur_values = blur_map[v_full, u_full].astype(np.float32)
snr_values = snr_map[v_full, u_full].astype(np.float32)
saturation_values = saturation_map[v_full, u_full].astype(np.float32)
logits_values = logits_img[v_logit, u_logit, :].astype(np.float32)
if points["GT_AVAILABLE"]:
if remap_ground_truth:
class_ids = map_main_class_indices(points["GT"][original_point_indices])
else:
class_ids = np.asarray(points["GT"][original_point_indices], dtype=np.int16)
valid_mask = class_ids != 255
else:
class_ids = np.full(original_point_indices.shape[0], -1, dtype=np.int16)
valid_mask = np.ones(original_point_indices.shape[0], dtype=bool) if allow_missing_ground_truth else np.zeros(
original_point_indices.shape[0], dtype=bool
)
valid_indices = np.nonzero(valid_mask)[0]
if valid_indices.size == 0:
return empty_compact_observations(cam_key, logits_dim), int(visible_indices.size), 0
point_indices = original_point_indices[valid_indices]
order = np.argsort(point_indices, kind="mergesort")
sorted_indices = valid_indices[order]
compact_observations = {
"camera": cam_key,
"point_indices": original_point_indices[sorted_indices].astype(np.int64),
"pixel_coords": np.stack((u_full[sorted_indices], v_full[sorted_indices]), axis=1).astype(np.int32),
"visibility": np.stack(
(
angles[sorted_indices],
distances[sorted_indices],
contrast_values[sorted_indices],
blur_values[sorted_indices],
snr_values[sorted_indices],
saturation_values[sorted_indices],
),
axis=1,
).astype(np.float32),
"ground_truth": class_ids[sorted_indices].astype(np.int16),
"logit_vectors": logits_values[sorted_indices].astype(np.float32),
}
selected_points = int(np.unique(compact_observations["point_indices"]).size)
return compact_observations, int(visible_indices.size), selected_points
def build_camera_observations(
cam_key,
cam_params,
calibration_params,
points_for_normals,
normals,
points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cfg,
allow_missing_ground_truth=False,
global_indices=None,
target_device=None,
remap_ground_truth=True,
):
step_start = time.perf_counter()
prefilter_mask = prefilter_points_for_camera(points_for_normals, cam_params, calibration_params, cfg)
if not np.any(prefilter_mask):
return {}, {
"prefilter_s": time.perf_counter() - step_start,
"projection_s": 0.0,
"criteria_s": 0.0,
"visible_points": 0,
"selected_points": 0,
"prefilter_kept": 0,
}
filtered_params = calibration_params.copy()
filtered_params["X"] = calibration_params["X"][prefilter_mask]
filtered_params["Y"] = calibration_params["Y"][prefilter_mask]
filtered_params["Z"] = calibration_params["Z"][prefilter_mask]
prefilter_s = time.perf_counter() - step_start
projection_start = time.perf_counter()
in_bounds_jax, visible_in_bounds_jax, full_i_jax, full_j_jax, logits_i_jax, logits_j_jax = project_points_for_camera(
cam_params,
filtered_params,
cfg["load"]["subscale"],
cfg["load"]["buffer_size"],
cfg["load"]["threshold"],
target_device=target_device,
)
if in_bounds_jax is None:
return {}, {
"prefilter_s": prefilter_s,
"projection_s": time.perf_counter() - projection_start,
"criteria_s": 0.0,
"visible_points": 0,
"selected_points": 0,
"prefilter_kept": int(prefilter_mask.sum()),
}
in_bounds, visible_in_bounds, full_i, full_j, logits_i, logits_j = jax.device_get(
(in_bounds_jax, visible_in_bounds_jax, full_i_jax, full_j_jax, logits_i_jax, logits_j_jax)
)
if np.count_nonzero(np.asarray(visible_in_bounds, dtype=bool)) == 0:
return {}, {
"prefilter_s": prefilter_s,
"projection_s": time.perf_counter() - projection_start,
"criteria_s": 0.0,
"visible_points": 0,
"selected_points": 0,
"prefilter_kept": int(prefilter_mask.sum()),
}
projection_s = time.perf_counter() - projection_start
criteria_start = time.perf_counter()
compact_observations, visible_points, selected_points = build_compact_observations(
cam_key,
points_for_normals,
normals,
points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cam_params,
in_bounds,
visible_in_bounds,
full_i,
full_j,
logits_i,
logits_j,
prefilter_mask,
allow_missing_ground_truth=allow_missing_ground_truth,
global_indices=global_indices,
remap_ground_truth=remap_ground_truth,
)
return compact_observations, {
"prefilter_s": prefilter_s,
"projection_s": projection_s,
"criteria_s": time.perf_counter() - criteria_start,
"visible_points": visible_points,
"selected_points": selected_points,
"prefilter_kept": int(prefilter_mask.sum()),
}
def process_camera_chunk(
start_idx,
end_idx,
calibration_params,
points,
coords,
normals,
points_gt,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cam_key,
cam_params,
cfg,
allow_missing_ground_truth,
global_indices,
target_device,
):
chunk_params = calibration_params.copy()
chunk_params["X"] = points["X"][start_idx:end_idx]
chunk_params["Y"] = points["Y"][start_idx:end_idx]
chunk_params["Z"] = points["Z"][start_idx:end_idx]
chunk_points = {
"GT": points_gt,
"GT_AVAILABLE": points_gt is not None,
}
observations, stats = build_camera_observations(
cam_key,
cam_params,
chunk_params,
coords[start_idx:end_idx],
normals[start_idx:end_idx],
chunk_points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cfg,
allow_missing_ground_truth=allow_missing_ground_truth,
global_indices=global_indices[start_idx:end_idx],
target_device=target_device,
)
return observations, stats
def compute_camera_global_depth_map(
cam_key,
cam_params,
calibration_params,
points_xyz,
cfg,
chunk_size,
target_device=None,
):
subscale_factor = int(cfg["load"]["subscale"])
buffer_size = int(cfg["load"]["buffer_size"])
threshold = float(cfg["load"]["threshold"])
depth_map = jnp.full(
(
int(calibration_params["height"]) // subscale_factor,
int(calibration_params["width"]) // subscale_factor,
),
jnp.inf,
)
compute_depth_map_jit = _get_depth_fn(buffer_size, threshold)
total_prefilter_s = 0.0
total_projection_s = 0.0
total_prefilter_kept = 0
for start_idx in range(0, points_xyz.shape[0], chunk_size):
end_idx = start_idx + chunk_size
step_start = time.perf_counter()
chunk_points_xyz = points_xyz[start_idx:end_idx]
prefilter_mask = prefilter_points_for_camera(chunk_points_xyz, cam_params, calibration_params, cfg)
total_prefilter_s += time.perf_counter() - step_start
if not np.any(prefilter_mask):
continue
filtered_params = calibration_params.copy()
filtered_params["X"] = calibration_params["X"][start_idx:end_idx][prefilter_mask]
filtered_params["Y"] = calibration_params["Y"][start_idx:end_idx][prefilter_mask]
filtered_params["Z"] = calibration_params["Z"][start_idx:end_idx][prefilter_mask]
total_prefilter_kept += int(prefilter_mask.sum())
projection_start = time.perf_counter()
_, _, z, in_bounds, _, _, i, j = project_points_raw(
cam_params,
filtered_params,
subscale_factor,
target_device=target_device,
)
if jnp.sum(in_bounds) != 0:
if target_device is None:
depth_map, _ = compute_depth_map_jit(i[in_bounds], j[in_bounds], z[in_bounds], depth_map)
else:
with jax.default_device(target_device):
depth_map, _ = compute_depth_map_jit(i[in_bounds], j[in_bounds], z[in_bounds], depth_map)
total_projection_s += time.perf_counter() - projection_start
logging.debug("Built global depth map for camera %s", cam_key)
return depth_map, total_prefilter_s, total_projection_s, total_prefilter_kept
def build_camera_observations_global_depth(
cam_key,
cam_params,
calibration_params,
points_for_normals,
normals,
points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cfg,
global_depth_map,
allow_missing_ground_truth=False,
global_indices=None,
target_device=None,
):
step_start = time.perf_counter()
prefilter_mask = prefilter_points_for_camera(points_for_normals, cam_params, calibration_params, cfg)
if not np.any(prefilter_mask):
return {}, {
"prefilter_s": time.perf_counter() - step_start,
"projection_s": 0.0,
"criteria_s": 0.0,
"visible_points": 0,
"selected_points": 0,
"prefilter_kept": 0,
}
filtered_params = calibration_params.copy()
filtered_params["X"] = calibration_params["X"][prefilter_mask]
filtered_params["Y"] = calibration_params["Y"][prefilter_mask]
filtered_params["Z"] = calibration_params["Z"][prefilter_mask]
prefilter_s = time.perf_counter() - step_start
projection_start = time.perf_counter()
_, _, z_jax, in_bounds_jax, full_i_jax, full_j_jax, logits_i_jax, logits_j_jax = project_points_raw(
cam_params,
filtered_params,
cfg["load"]["subscale"],
target_device=target_device,
)
if jnp.sum(in_bounds_jax) == 0:
return {}, {
"prefilter_s": prefilter_s,
"projection_s": time.perf_counter() - projection_start,
"criteria_s": 0.0,
"visible_points": 0,
"selected_points": 0,
"prefilter_kept": int(prefilter_mask.sum()),
}
compute_depth_map_jit = _get_depth_fn(cfg["load"]["buffer_size"], cfg["load"]["threshold"])
if target_device is None:
_, visible_in_bounds_jax = compute_depth_map_jit(
logits_i_jax[in_bounds_jax],
logits_j_jax[in_bounds_jax],
z_jax[in_bounds_jax],
global_depth_map,
)
else:
with jax.default_device(target_device):
_, visible_in_bounds_jax = compute_depth_map_jit(
logits_i_jax[in_bounds_jax],
logits_j_jax[in_bounds_jax],
z_jax[in_bounds_jax],
global_depth_map,
)
in_bounds, visible_in_bounds, full_i, full_j, logits_i, logits_j = jax.device_get(
(in_bounds_jax, visible_in_bounds_jax, full_i_jax, full_j_jax, logits_i_jax, logits_j_jax)
)
projection_s = time.perf_counter() - projection_start
criteria_start = time.perf_counter()
compact_observations, visible_points, selected_points = build_compact_observations(
cam_key,
points_for_normals,
normals,
points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cam_params,
in_bounds,
visible_in_bounds,
full_i,
full_j,
logits_i,
logits_j,
prefilter_mask,
allow_missing_ground_truth=allow_missing_ground_truth,
global_indices=global_indices,
)
return compact_observations, {
"prefilter_s": prefilter_s,
"projection_s": projection_s,
"criteria_s": time.perf_counter() - criteria_start,
"visible_points": visible_points,
"selected_points": selected_points,
"prefilter_kept": int(prefilter_mask.sum()),
}
def process_camera_items_for_tile_global_depth(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
):
output_mode = cfg.get("runtime", {}).get("tile_output_format", "h5_then_pt")
use_dense = output_mode in {"pt_direct", "pt_sharded_direct"}
all_observations = None if use_dense else {}
extension = cfg["logits"].get("extension", ".npy")
max_views = int(cfg["selection"]["max_views"])
selection_strategy = cfg["selection"]["strategy"]
chunk_size = int(cfg["load"].get("point_chunk_size", len(points["X"])))
total_points = len(points["X"])
global_indices = np.arange(total_points, dtype=np.int64)
pad_count = (-total_points) % chunk_size
padded_points = {key: value.copy() if isinstance(value, np.ndarray) else value for key, value in points.items()}
padded_coords = coords.copy()
padded_normals = normals.copy()
if pad_count:
padded_points["X"] = np.pad(padded_points["X"], (0, pad_count), mode="constant", constant_values=0)
padded_points["Y"] = np.pad(padded_points["Y"], (0, pad_count), mode="constant", constant_values=0)
padded_points["Z"] = np.pad(padded_points["Z"], (0, pad_count), mode="constant", constant_values=0)
padded_coords = np.pad(padded_coords, ((0, pad_count), (0, 0)), mode="constant", constant_values=0)
padded_normals = np.pad(padded_normals, ((0, pad_count), (0, 0)), mode="constant", constant_values=0)
global_indices = np.pad(global_indices, (0, pad_count), mode="constant", constant_values=-1)
padded_total_points = len(padded_points["X"])
logging.info(
"Using JAX backend=%s for tile %s with chunk_size=%s max_distance=%s visibility_mode=two_pass_global_depth",
jax.default_backend(),
tile_info["tile"],
chunk_size,
cfg["load"].get("max_distance"),
)
for cam_key, cam_params in tqdm(camera_items, desc=f"{tile_info['tile']} cameras", unit="camera", leave=False):
camera_start = time.perf_counter()
logits_path = logits_dir / f"{cam_key}{extension}"
image_path = Path(tile_info["image_dir"]) / f"{cam_key}.JPG"
if not logits_path.exists():
logging.warning("Missing logits for %s in tile %s", cam_key, tile_info["tile"])
continue
if not image_path.exists():
logging.warning("Missing image for %s in tile %s", cam_key, tile_info["tile"])
continue
io_start = time.perf_counter()
image = cv.imread(str(image_path), cv.IMREAD_COLOR)
if image is None:
raise FileNotFoundError(f"Unable to read image {image_path}")
expected_logits_shape = (image.shape[0] // cfg["load"]["subscale"], image.shape[1] // cfg["load"]["subscale"])
logits_img = load_logits_image(logits_path, expected_shape=expected_logits_shape)
if use_dense and all_observations is None:
all_observations = create_dense_observations(
total_points,
max_views,
logits_img.shape[2],
cfg["data"]["vmin"],
cfg["data"]["vmax"],
get_storage_config(cfg)["visibility_quant_max"],
get_storage_config(cfg)["logits_quant_max"],
)
io_s = time.perf_counter() - io_start
maps_start = time.perf_counter()
contrast_map, blur_map, snr_map, saturation_map = compute_criteria_maps(image)
maps_s = time.perf_counter() - maps_start
depth_map, pass1_prefilter_s, pass1_projection_s, pass1_prefilter_kept = compute_camera_global_depth_map(
cam_key,
cam_params,
calibration_params,
padded_coords,
cfg,
chunk_size,
)
chunk_prefilter_s = 0.0
chunk_projection_s = 0.0
chunk_criteria_s = 0.0
chunk_merge_s = 0.0
prefilter_kept = 0
visible_points = 0
selected_points = 0
num_chunks = padded_total_points // chunk_size
for start_idx in range(0, padded_total_points, chunk_size):
end_idx = start_idx + chunk_size
chunk_params = calibration_params.copy()
chunk_params["X"] = padded_points["X"][start_idx:end_idx]
chunk_params["Y"] = padded_points["Y"][start_idx:end_idx]
chunk_params["Z"] = padded_points["Z"][start_idx:end_idx]
observations, stats = build_camera_observations_global_depth(
cam_key,
cam_params,
chunk_params,
padded_coords[start_idx:end_idx],
padded_normals[start_idx:end_idx],
padded_points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cfg,
depth_map,
allow_missing_ground_truth=allow_missing_ground_truth,
global_indices=global_indices[start_idx:end_idx],
)
if pad_count and end_idx > total_points and observations and "point_indices" in observations:
observations = filter_compact_observations_to_valid_points(observations, observations["point_indices"] >= 0)
chunk_prefilter_s += stats["prefilter_s"]
chunk_projection_s += stats["projection_s"]
chunk_criteria_s += stats["criteria_s"]
prefilter_kept += stats["prefilter_kept"]
visible_points += stats["visible_points"]
selected_points += stats["selected_points"]
merge_start = time.perf_counter()
if use_dense:
merge_compact_observations_dense(all_observations, observations, max_views, selection_strategy)
else:
merge_compact_observations_bounded(all_observations, observations, max_views, selection_strategy)
chunk_merge_s += time.perf_counter() - merge_start
logging.info(
"Tile %s camera %s done in %.2fs | io=%.2fs maps=%.2fs pass1_prefilter=%.2fs pass1_projection=%.2fs "
"prefilter=%.2fs projection=%.2fs criteria=%.2fs merge=%.2fs prefilter_kept=%s pass1_prefilter_kept=%s visible=%s "
"selected=%s chunks=%s",
tile_info["tile"],
cam_key,
time.perf_counter() - camera_start,
io_s,
maps_s,
pass1_prefilter_s,
pass1_projection_s,
chunk_prefilter_s,
chunk_projection_s,
chunk_criteria_s,
chunk_merge_s,
prefilter_kept,
pass1_prefilter_kept,
visible_points,
selected_points,
num_chunks,
)
if use_dense and all_observations is None:
all_observations = create_dense_observations(
total_points,
max_views,
0,
cfg["data"]["vmin"],
cfg["data"]["vmax"],
get_storage_config(cfg)["visibility_quant_max"],
get_storage_config(cfg)["logits_quant_max"],
)
return all_observations
def select_observations(observations, max_views, strategy):
if strategy == "topk_distance":
ranked = sorted(observations, key=lambda obs: (obs["visibility"][1], obs["visibility"][0]))
else:
ranked = observations
return ranked[:max_views]
def trim_bucket(bucket, max_views, strategy):
if len(bucket["camera"]) <= max_views:
return
if strategy == "topk_distance":
order = sorted(
range(len(bucket["camera"])),
key=lambda idx: (bucket["visibility"][idx][1], bucket["visibility"][idx][0]),
)[:max_views]
else:
order = list(range(max_views))
bucket["camera"] = [bucket["camera"][idx] for idx in order]
bucket["pixel_coords"] = [bucket["pixel_coords"][idx] for idx in order]
bucket["visibility"] = [bucket["visibility"][idx] for idx in order]
bucket["logit_vectors"] = [bucket["logit_vectors"][idx] for idx in order]
def merge_compact_observations_bounded(all_observations, compact_observations, max_views, strategy):
if not compact_observations or "point_indices" not in compact_observations:
return
point_indices = compact_observations["point_indices"]
if point_indices.size == 0:
return
group_starts = np.r_[0, np.flatnonzero(np.diff(point_indices)) + 1]
group_ends = np.r_[group_starts[1:], point_indices.size]
camera_key = compact_observations["camera"]
for start, end in zip(group_starts, group_ends):
point_idx = int(point_indices[start])
bucket = all_observations.setdefault(
point_idx,
{
"ground_truth": None,
"camera": [],
"pixel_coords": [],
"visibility": [],
"logit_vectors": [],
},
)
gt_slice = compact_observations["ground_truth"][start:end]
valid_gt = gt_slice[gt_slice >= 0]
if bucket["ground_truth"] is None and valid_gt.size > 0:
bucket["ground_truth"] = int(valid_gt[0])
count = end - start
bucket["camera"].extend([camera_key] * count)
bucket["pixel_coords"].extend(compact_observations["pixel_coords"][start:end].tolist())
bucket["visibility"].extend(compact_observations["visibility"][start:end].tolist())
bucket["logit_vectors"].extend([row for row in compact_observations["logit_vectors"][start:end]])
trim_bucket(bucket, max_views, strategy)
def merge_worker_observations_bounded(all_observations, worker_observations, max_views, strategy):
for point_idx, worker_bucket in worker_observations.items():
bucket = all_observations.setdefault(
point_idx,
{
"ground_truth": None,
"camera": [],
"pixel_coords": [],
"visibility": [],
"logit_vectors": [],
},
)
if bucket["ground_truth"] is None and worker_bucket.get("ground_truth") is not None:
bucket["ground_truth"] = int(worker_bucket["ground_truth"])
bucket["camera"].extend(worker_bucket["camera"])
bucket["pixel_coords"].extend(worker_bucket["pixel_coords"])
bucket["visibility"].extend(worker_bucket["visibility"])
bucket["logit_vectors"].extend(worker_bucket["logit_vectors"])
trim_bucket(bucket, max_views, strategy)
def create_dense_observations(num_points, max_views, logits_dim, visibility_vmin, visibility_vmax, visibility_quant_max, logits_quant_max):
return {
"mode": "dense",
"counts": np.zeros(num_points, dtype=np.uint8),
"ground_truth": np.full(num_points, -1, dtype=np.int16),
"pixel_coords": np.zeros((num_points, max_views, 2), dtype=np.uint16),
"visibility": np.zeros((num_points, max_views, 6), dtype=np.uint16),
"logit_vectors": np.zeros((num_points, max_views, logits_dim), dtype=np.uint8),
"visibility_vmin": np.asarray(visibility_vmin, dtype=np.float32),
"visibility_vmax": np.asarray(visibility_vmax, dtype=np.float32),
"visibility_quant_max": int(visibility_quant_max),
"logits_quant_max": int(logits_quant_max),
}
def count_nonempty_observations(observations):
if isinstance(observations, dict) and observations.get("mode") == "dense":
return int(np.count_nonzero(observations["counts"]))
return len(observations)
def slice_point_data(points, coords, normals, start_idx, end_idx):
sliced_points = {"GT_AVAILABLE": points["GT_AVAILABLE"]}
for key in ("X", "Y", "Z"):
sliced_points[key] = points[key][start_idx:end_idx]
sliced_points["GT"] = None if points["GT"] is None else points["GT"][start_idx:end_idx]
return sliced_points, coords[start_idx:end_idx], normals[start_idx:end_idx]
def _sort_dense_point_slots(dense_observations, point_idx):
count = int(dense_observations["counts"][point_idx])
if count <= 1:
return
order = np.lexsort(
(
dense_observations["visibility"][point_idx, :count, 0],
dense_observations["visibility"][point_idx, :count, 1],
)
)
dense_observations["pixel_coords"][point_idx, :count] = dense_observations["pixel_coords"][point_idx, :count][order]
dense_observations["visibility"][point_idx, :count] = dense_observations["visibility"][point_idx, :count][order]
dense_observations["logit_vectors"][point_idx, :count] = dense_observations["logit_vectors"][point_idx, :count][order]
def _select_topk_indices(visibility, max_views, strategy):
if visibility.shape[0] <= max_views:
order = np.lexsort((visibility[:, 0], visibility[:, 1]))
return order
if strategy == "topk_distance":
order = np.lexsort((visibility[:, 0], visibility[:, 1]))
return order[:max_views]
return np.arange(max_views, dtype=np.int64)
def _merge_dense_point_arrays(dense_observations, point_idx, new_pixel_coords, new_visibility, new_logit_vectors, max_views, strategy):
current_count = int(dense_observations["counts"][point_idx])
if current_count > 0:
merged_pixel_coords = np.concatenate(
(dense_observations["pixel_coords"][point_idx, :current_count], new_pixel_coords),
axis=0,
)
merged_visibility = np.concatenate(
(dense_observations["visibility"][point_idx, :current_count], new_visibility),
axis=0,
)
merged_logit_vectors = np.concatenate(
(dense_observations["logit_vectors"][point_idx, :current_count], new_logit_vectors),
axis=0,
)
else:
merged_pixel_coords = new_pixel_coords
merged_visibility = new_visibility
merged_logit_vectors = new_logit_vectors
keep = _select_topk_indices(merged_visibility, max_views, strategy)
keep_count = len(keep)
dense_observations["pixel_coords"][point_idx, :keep_count] = merged_pixel_coords[keep]
dense_observations["visibility"][point_idx, :keep_count] = merged_visibility[keep]
dense_observations["logit_vectors"][point_idx, :keep_count] = merged_logit_vectors[keep]
dense_observations["counts"][point_idx] = keep_count
def merge_compact_observations_dense(dense_observations, compact_observations, max_views, strategy):
if not compact_observations or "point_indices" not in compact_observations:
return
point_indices = compact_observations["point_indices"]
if point_indices.size == 0:
return
quantized_visibility = quantize_visibility(
compact_observations["visibility"],
dense_observations["visibility_vmin"],
dense_observations["visibility_vmax"],
dense_observations["visibility_quant_max"],
)
quantized_logits = quantize_logits(
compact_observations["logit_vectors"],
dense_observations["logits_quant_max"],
)
group_starts = np.r_[0, np.flatnonzero(np.diff(point_indices)) + 1]
group_ends = np.r_[group_starts[1:], point_indices.size]
for start, end in zip(group_starts, group_ends):
point_idx = int(point_indices[start])
gt_slice = compact_observations["ground_truth"][start:end]
valid_gt = gt_slice[gt_slice >= 0]
if valid_gt.size > 0 and dense_observations["ground_truth"][point_idx] < 0:
dense_observations["ground_truth"][point_idx] = int(valid_gt[0])
_merge_dense_point_arrays(
dense_observations,
point_idx,
compact_observations["pixel_coords"][start:end].astype(np.uint16, copy=False),
quantized_visibility[start:end],
quantized_logits[start:end],
max_views,
strategy,
)
def merge_dense_observations(target, source, max_views, strategy):
source_counts = source["counts"]
for point_idx in np.flatnonzero(source_counts):
point_idx = int(point_idx)
src_count = int(source_counts[point_idx])
if src_count == 0:
continue
if target["ground_truth"][point_idx] < 0 and source["ground_truth"][point_idx] >= 0:
target["ground_truth"][point_idx] = source["ground_truth"][point_idx]
_merge_dense_point_arrays(
target,
point_idx,
source["pixel_coords"][point_idx, :src_count],
source["visibility"][point_idx, :src_count],
source["logit_vectors"][point_idx, :src_count],
max_views,
strategy,
)
def create_compact_payload_template(cfg, logits_dim):
storage_cfg = get_storage_config(cfg)
return {
"point_indices": [],
"visibility": [],
"logit_vectors": [],
"ground_truth": [],
"visibility_vmin": np.asarray(cfg["data"]["vmin"], dtype=np.float32),
"visibility_vmax": np.asarray(cfg["data"]["vmax"], dtype=np.float32),
"visibility_quant_max": int(storage_cfg["visibility_quant_max"]),
"logits_quant_max": int(storage_cfg["logits_quant_max"]),
"logits_dim": int(logits_dim),
}
def append_compact_payload(payload, compact_observations):
if not compact_observations or "point_indices" not in compact_observations:
return
if compact_observations["point_indices"].size == 0:
return
payload["point_indices"].append(compact_observations["point_indices"].astype(np.uint32, copy=False))
payload["visibility"].append(
quantize_visibility(
compact_observations["visibility"],
payload["visibility_vmin"],
payload["visibility_vmax"],
payload["visibility_quant_max"],
)
)
payload["logit_vectors"].append(
quantize_logits(compact_observations["logit_vectors"], payload["logits_quant_max"])
)
payload["ground_truth"].append(compact_observations["ground_truth"].astype(np.int16, copy=False))
def finalize_compact_payload(payload):
if payload is None:
return None
logits_dim = payload["logits_dim"]
if payload["point_indices"]:
return {
"point_indices": np.concatenate(payload["point_indices"], axis=0),
"visibility": np.concatenate(payload["visibility"], axis=0),
"logit_vectors": np.concatenate(payload["logit_vectors"], axis=0),
"ground_truth": np.concatenate(payload["ground_truth"], axis=0),
"visibility_vmin": payload["visibility_vmin"],
"visibility_vmax": payload["visibility_vmax"],
"visibility_quant_max": payload["visibility_quant_max"],
"logits_quant_max": payload["logits_quant_max"],
}
return {
"point_indices": np.empty((0,), dtype=np.uint32),
"visibility": np.empty((0, 6), dtype=np.uint16),
"logit_vectors": np.empty((0, logits_dim), dtype=np.uint8),
"ground_truth": np.empty((0,), dtype=np.int16),
"visibility_vmin": payload["visibility_vmin"],
"visibility_vmax": payload["visibility_vmax"],
"visibility_quant_max": payload["visibility_quant_max"],
"logits_quant_max": payload["logits_quant_max"],
}
def point_ground_truth(selected):
labels = [obs["ground_truth"] for obs in selected if obs["ground_truth"] is not None]
if not labels:
return None
values, counts = np.unique(np.asarray(labels, dtype=np.int64), return_counts=True)
return int(values[np.argmax(counts)])
def get_storage_config(cfg):
storage_cfg = cfg.get("storage", {})
return {
"coord_scale": float(storage_cfg.get("coord_scale", 0.001)),
"pixel_coords_dtype": storage_cfg.get("pixel_coords_dtype", "uint16"),
"visibility_quant_max": int(storage_cfg.get("visibility_quant_max", 65535)),
"logits_quant_max": int(storage_cfg.get("logits_quant_max", 255)),
"compression": storage_cfg.get("h5_compression", "lzf"),
}
def quantize_coordinates(coords_row, coord_scale):
return np.rint(coords_row / coord_scale).astype(np.int32)
def quantize_visibility(visibility, vmin, vmax, quant_max):
vmin_arr = np.asarray(vmin, dtype=np.float32)
vmax_arr = np.asarray(vmax, dtype=np.float32)
clipped = np.clip(visibility.astype(np.float32), vmin_arr, vmax_arr)
normalized = (clipped - vmin_arr) / np.maximum(vmax_arr - vmin_arr, 1e-8)
return np.rint(normalized * quant_max).astype(np.uint16)
def quantize_logits(logits, quant_max):
logits = np.asarray(logits)
if np.issubdtype(logits.dtype, np.integer) or float(np.nanmax(logits)) > 1.0:
return np.clip(logits, 0, quant_max).astype(np.uint8)
clipped = np.clip(logits.astype(np.float32), 0.0, 1.0)
return np.rint(clipped * quant_max).astype(np.uint8)
def write_tile_h5(output_h5_file, tile_name, split_name, coords, all_observations, cfg):
output_h5_file.parent.mkdir(parents=True, exist_ok=True)
min_views = cfg["selection"]["min_views"]
max_views = cfg["selection"]["max_views"]
strategy = cfg["selection"]["strategy"]
total_points = len(all_observations)
storage_cfg = get_storage_config(cfg)
compression = storage_cfg["compression"]
vmin = cfg["data"]["vmin"]
vmax = cfg["data"]["vmax"]
logging.info(
"Writing tile %s to %s with %s aggregated point buckets",
tile_name,
output_h5_file,
total_points,
)
with h5py.File(output_h5_file, "w") as handle:
handle.attrs["tile"] = tile_name
handle.attrs["split"] = split_name
handle.attrs["feature_names"] = np.array(cfg["data"]["visibility_feature_names"], dtype="S")
handle.attrs["coords_scale"] = storage_cfg["coord_scale"]
handle.attrs["coords_offset"] = np.zeros(3, dtype=np.float32)
handle.attrs["visibility_vmin"] = np.asarray(vmin, dtype=np.float32)
handle.attrs["visibility_vmax"] = np.asarray(vmax, dtype=np.float32)
handle.attrs["visibility_quant_max"] = storage_cfg["visibility_quant_max"]
handle.attrs["logits_quant_max"] = storage_cfg["logits_quant_max"]
points_group = handle.create_group("points")
kept = 0
skipped = 0
for point_idx, observations in all_observations.items():
if len(observations["camera"]) < min_views:
skipped += 1
continue
gt = observations["ground_truth"]
if gt is None:
skipped += 1
continue
group = points_group.create_group(str(point_idx))
group.create_dataset(
"coordinates",
data=quantize_coordinates(coords[point_idx], storage_cfg["coord_scale"]),
dtype="int32",
compression=compression,
)
group.create_dataset(
"image_ids",
data=np.array(observations["camera"], dtype=h5py.string_dtype("utf-8")),
compression=compression,
)
group.create_dataset(
"pixel_coords",
data=np.array(observations["pixel_coords"], dtype=np.uint16),
dtype=storage_cfg["pixel_coords_dtype"],
compression=compression,
)
group.create_dataset(
"visibility",
data=quantize_visibility(
np.array(observations["visibility"], dtype=np.float32),
vmin,
vmax,
storage_cfg["visibility_quant_max"],
),
dtype="uint16",
compression=compression,
)
group.create_dataset(
"logit_vectors",
data=quantize_logits(
np.array(observations["logit_vectors"], dtype=np.float32),
storage_cfg["logits_quant_max"],
),
dtype="uint8",
compression=compression,
)
group.create_dataset("ground_truth", data=np.uint8(gt), dtype="uint8")
kept += 1
handle.attrs["num_points_kept"] = kept
handle.attrs["num_points_skipped"] = skipped
logging.info(
"Finished writing tile %s: kept=%s skipped=%s",
tile_name,
kept,
skipped,
)
def process_camera_items_for_tile(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
):
output_mode = cfg.get("runtime", {}).get("tile_output_format", "h5_then_pt")
use_dense = output_mode in {"pt_direct", "pt_sharded_direct"}
all_observations = None if use_dense else {}
extension = cfg["logits"].get("extension", ".npy")
max_views = int(cfg["selection"]["max_views"])
selection_strategy = cfg["selection"]["strategy"]
chunk_size = int(cfg["load"].get("point_chunk_size", len(points["X"])))
total_points = len(points["X"])
global_indices = np.arange(total_points, dtype=np.int64)
pad_count = (-total_points) % chunk_size
padded_points = {key: value.copy() if isinstance(value, np.ndarray) else value for key, value in points.items()}
padded_coords = coords.copy()
padded_normals = normals.copy()
if pad_count:
padded_points["X"] = np.pad(padded_points["X"], (0, pad_count), mode="constant", constant_values=0)
padded_points["Y"] = np.pad(padded_points["Y"], (0, pad_count), mode="constant", constant_values=0)
padded_points["Z"] = np.pad(padded_points["Z"], (0, pad_count), mode="constant", constant_values=0)
padded_coords = np.pad(padded_coords, ((0, pad_count), (0, 0)), mode="constant", constant_values=0)
padded_normals = np.pad(padded_normals, ((0, pad_count), (0, 0)), mode="constant", constant_values=0)
global_indices = np.pad(global_indices, (0, pad_count), mode="constant", constant_values=-1)
padded_total_points = len(padded_points["X"])
logging.info(
"Using JAX backend=%s for tile %s with chunk_size=%s max_distance=%s visibility_mode=chunk_local_depth",
jax.default_backend(),
tile_info["tile"],
chunk_size,
cfg["load"].get("max_distance"),
)
for cam_key, cam_params in tqdm(camera_items, desc=f"{tile_info['tile']} cameras", unit="camera", leave=False):
camera_start = time.perf_counter()
logits_path = logits_dir / f"{cam_key}{extension}"
image_path = Path(tile_info["image_dir"]) / f"{cam_key}.JPG"
if not logits_path.exists():
logging.warning("Missing logits for %s in tile %s", cam_key, tile_info["tile"])
continue
if not image_path.exists():
logging.warning("Missing image for %s in tile %s", cam_key, tile_info["tile"])
continue
io_start = time.perf_counter()
image = cv.imread(str(image_path), cv.IMREAD_COLOR)
if image is None:
raise FileNotFoundError(f"Unable to read image {image_path}")
expected_logits_shape = (image.shape[0] // cfg["load"]["subscale"], image.shape[1] // cfg["load"]["subscale"])
logits_img = load_logits_image(logits_path, expected_shape=expected_logits_shape)
if use_dense and all_observations is None:
all_observations = create_dense_observations(
total_points,
max_views,
logits_img.shape[2],
cfg["data"]["vmin"],
cfg["data"]["vmax"],
get_storage_config(cfg)["visibility_quant_max"],
get_storage_config(cfg)["logits_quant_max"],
)
io_s = time.perf_counter() - io_start
maps_start = time.perf_counter()
contrast_map, blur_map, snr_map, saturation_map = compute_criteria_maps(image)
maps_s = time.perf_counter() - maps_start
chunk_prefilter_s = 0.0
chunk_projection_s = 0.0
chunk_criteria_s = 0.0
chunk_merge_s = 0.0
prefilter_kept = 0
visible_points = 0
selected_points = 0
num_chunks = padded_total_points // chunk_size
for start_idx in range(0, padded_total_points, chunk_size):
end_idx = start_idx + chunk_size
chunk_params = calibration_params.copy()
chunk_params["X"] = padded_points["X"][start_idx:end_idx]
chunk_params["Y"] = padded_points["Y"][start_idx:end_idx]
chunk_params["Z"] = padded_points["Z"][start_idx:end_idx]
observations, stats = build_camera_observations(
cam_key,
cam_params,
chunk_params,
padded_coords[start_idx:end_idx],
padded_normals[start_idx:end_idx],
padded_points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cfg,
allow_missing_ground_truth=allow_missing_ground_truth,
global_indices=global_indices[start_idx:end_idx],
)
if pad_count and end_idx > total_points and observations and "point_indices" in observations:
observations = filter_compact_observations_to_valid_points(observations, observations["point_indices"] >= 0)
chunk_prefilter_s += stats["prefilter_s"]
chunk_projection_s += stats["projection_s"]
chunk_criteria_s += stats["criteria_s"]
prefilter_kept += stats["prefilter_kept"]
visible_points += stats["visible_points"]
selected_points += stats["selected_points"]
merge_start = time.perf_counter()
if use_dense:
merge_compact_observations_dense(all_observations, observations, max_views, selection_strategy)
else:
merge_compact_observations_bounded(all_observations, observations, max_views, selection_strategy)
chunk_merge_s += time.perf_counter() - merge_start
logging.info(
"Tile %s camera %s done in %.2fs | io=%.2fs maps=%.2fs prefilter=%.2fs projection=%.2fs criteria=%.2fs merge=%.2fs prefilter_kept=%s visible=%s selected=%s chunks=%s",
tile_info["tile"],
cam_key,
time.perf_counter() - camera_start,
io_s,
maps_s,
chunk_prefilter_s,
chunk_projection_s,
chunk_criteria_s,
chunk_merge_s,
prefilter_kept,
visible_points,
selected_points,
num_chunks,
)
if use_dense and all_observations is None:
all_observations = create_dense_observations(
total_points,
max_views,
0,
cfg["data"]["vmin"],
cfg["data"]["vmax"],
get_storage_config(cfg)["visibility_quant_max"],
get_storage_config(cfg)["logits_quant_max"],
)
return all_observations
def process_camera_items_for_tile_compact(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
):
extension = cfg["logits"].get("extension", ".npy")
chunk_size = int(cfg["load"].get("point_chunk_size", len(points["X"])))
total_points = len(points["X"])
global_indices = np.arange(total_points, dtype=np.int64)
pad_count = (-total_points) % chunk_size
padded_points = {key: value.copy() if isinstance(value, np.ndarray) else value for key, value in points.items()}
padded_coords = coords.copy()
padded_normals = normals.copy()
if pad_count:
padded_points["X"] = np.pad(padded_points["X"], (0, pad_count), mode="constant", constant_values=0)
padded_points["Y"] = np.pad(padded_points["Y"], (0, pad_count), mode="constant", constant_values=0)
padded_points["Z"] = np.pad(padded_points["Z"], (0, pad_count), mode="constant", constant_values=0)
padded_coords = np.pad(padded_coords, ((0, pad_count), (0, 0)), mode="constant", constant_values=0)
padded_normals = np.pad(padded_normals, ((0, pad_count), (0, 0)), mode="constant", constant_values=0)
global_indices = np.pad(global_indices, (0, pad_count), mode="constant", constant_values=-1)
padded_total_points = len(padded_points["X"])
compact_payload = None
logging.info(
"Using JAX backend=%s for tile %s with chunk_size=%s max_distance=%s visibility_mode=%s compact_reduce=true",
jax.default_backend(),
tile_info["tile"],
chunk_size,
cfg["load"].get("max_distance"),
get_visibility_mode(cfg),
)
for cam_key, cam_params in tqdm(camera_items, desc=f"{tile_info['tile']} cameras", unit="camera", leave=False):
camera_start = time.perf_counter()
logits_path = logits_dir / f"{cam_key}{extension}"
image_path = Path(tile_info["image_dir"]) / f"{cam_key}.JPG"
if not logits_path.exists():
logging.warning("Missing logits for %s in tile %s", cam_key, tile_info["tile"])
continue
if not image_path.exists():
logging.warning("Missing image for %s in tile %s", cam_key, tile_info["tile"])
continue
io_start = time.perf_counter()
image = cv.imread(str(image_path), cv.IMREAD_COLOR)
if image is None:
raise FileNotFoundError(f"Unable to read image {image_path}")
expected_logits_shape = (image.shape[0] // cfg["load"]["subscale"], image.shape[1] // cfg["load"]["subscale"])
logits_img = load_logits_image(logits_path, expected_shape=expected_logits_shape)
if compact_payload is None:
compact_payload = create_compact_payload_template(cfg, logits_img.shape[2])
io_s = time.perf_counter() - io_start
maps_start = time.perf_counter()
contrast_map, blur_map, snr_map, saturation_map = compute_criteria_maps(image)
maps_s = time.perf_counter() - maps_start
chunk_prefilter_s = 0.0
chunk_projection_s = 0.0
chunk_criteria_s = 0.0
chunk_append_s = 0.0
prefilter_kept = 0
visible_points = 0
selected_points = 0
num_chunks = padded_total_points // chunk_size
for start_idx in range(0, padded_total_points, chunk_size):
end_idx = start_idx + chunk_size
chunk_params = calibration_params.copy()
chunk_params["X"] = padded_points["X"][start_idx:end_idx]
chunk_params["Y"] = padded_points["Y"][start_idx:end_idx]
chunk_params["Z"] = padded_points["Z"][start_idx:end_idx]
observations, stats = build_camera_observations(
cam_key,
cam_params,
chunk_params,
padded_coords[start_idx:end_idx],
padded_normals[start_idx:end_idx],
padded_points,
logits_img,
contrast_map,
blur_map,
snr_map,
saturation_map,
cfg,
allow_missing_ground_truth=allow_missing_ground_truth,
global_indices=global_indices[start_idx:end_idx],
)
if pad_count and end_idx > total_points and observations and "point_indices" in observations:
observations = filter_compact_observations_to_valid_points(observations, observations["point_indices"] >= 0)
chunk_prefilter_s += stats["prefilter_s"]
chunk_projection_s += stats["projection_s"]
chunk_criteria_s += stats["criteria_s"]
prefilter_kept += stats["prefilter_kept"]
visible_points += stats["visible_points"]
selected_points += stats["selected_points"]
append_start = time.perf_counter()
append_compact_payload(compact_payload, observations)
chunk_append_s += time.perf_counter() - append_start
logging.info(
"Tile %s camera %s done in %.2fs | io=%.2fs maps=%.2fs prefilter=%.2fs projection=%.2fs criteria=%.2fs append=%.2fs prefilter_kept=%s visible=%s selected=%s chunks=%s",
tile_info["tile"],
cam_key,
time.perf_counter() - camera_start,
io_s,
maps_s,
chunk_prefilter_s,
chunk_projection_s,
chunk_criteria_s,
chunk_append_s,
prefilter_kept,
visible_points,
selected_points,
num_chunks,
)
if compact_payload is None:
compact_payload = create_compact_payload_template(cfg, 0)
return finalize_compact_payload(compact_payload)
def get_visibility_mode(cfg):
return cfg["load"].get("visibility_mode", "chunk_local_depth")
def get_worker_output_mode(cfg):
return cfg.get("runtime", {}).get("worker_output_mode", "inherit")
def get_point_shard_size(cfg, total_points):
return int(cfg.get("runtime", {}).get("point_shard_size", total_points))
def _launch_camera_workers(
cfg,
tile_info,
camera_items,
temp_dir,
worker_count,
shard_start=None,
shard_end=None,
):
worker_output_mode = get_worker_output_mode(cfg)
camera_groups = [[] for _ in range(worker_count)]
for idx, camera_item in enumerate(camera_items):
camera_groups[idx % worker_count].append(camera_item[0])
processes = []
output_files = []
for worker_idx, camera_group in enumerate(camera_groups):
if not camera_group:
continue
camera_list_file = temp_dir / f"worker_{worker_idx}_cameras.json"
output_file = temp_dir / f"worker_{worker_idx}_observations.pt"
stdout_file = temp_dir / f"worker_{worker_idx}.stdout.log"
stderr_file = temp_dir / f"worker_{worker_idx}.stderr.log"
camera_list_file.write_text(json.dumps(camera_group), encoding="utf-8")
env = os.environ.copy()
env["CUDA_VISIBLE_DEVICES"] = str(worker_idx)
env.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false")
cmd = [
sys.executable,
str(Path(__file__).resolve().parents[1] / "scripts" / "run_batch_generation.py"),
"--config",
cfg["_config_path"],
"--worker-mode",
"camera-batch",
"--tile",
tile_info["tile"],
"--camera-list-file",
str(camera_list_file),
"--worker-output",
str(output_file),
]
if shard_start is not None and shard_end is not None:
cmd.extend(["--shard-start", str(shard_start), "--shard-end", str(shard_end)])
logging.info(
"Launching worker %s for tile %s%s on GPU %s with %s cameras",
worker_idx,
tile_info["tile"],
"" if shard_start is None else f" shard {shard_start}:{shard_end}",
worker_idx,
len(camera_group),
)
stdout_handle = None
stderr_handle = None
if worker_output_mode == "files":
stdout_handle = open(stdout_file, "w", encoding="utf-8")
stderr_handle = open(stderr_file, "w", encoding="utf-8")
process = subprocess.Popen(
cmd,
cwd=str(Path(__file__).resolve().parents[1]),
env=env,
stdout=stdout_handle,
stderr=stderr_handle,
)
elif worker_output_mode == "inherit":
process = subprocess.Popen(
cmd,
cwd=str(Path(__file__).resolve().parents[1]),
env=env,
)
else:
raise ValueError(f"Unsupported worker_output_mode: {worker_output_mode}")
processes.append((worker_idx, process, stdout_handle, stderr_handle, stdout_file, stderr_file))
output_files.append(output_file)
for worker_idx, process, stdout_handle, stderr_handle, stdout_file, stderr_file in processes:
return_code = process.wait()
if stdout_handle is not None:
stdout_handle.close()
if stderr_handle is not None:
stderr_handle.close()
if return_code != 0:
log_hint = (
f" Worker logs: stdout={stdout_file} stderr={stderr_file}"
if worker_output_mode == "files"
else ""
)
if return_code == -9:
raise RuntimeError(
f"Camera worker {worker_idx} was killed with code -9 for tile {tile_info['tile']}. "
f"This usually indicates an out-of-memory kill (GPU or system RAM).{log_hint}"
)
raise RuntimeError(
f"Camera worker {worker_idx} failed for tile {tile_info['tile']} with code {return_code}.{log_hint}"
)
return output_files
def process_tile_sharded_to_pt(
tile_info,
cfg,
calibration_params,
points,
coords,
offset,
normals,
camera_items,
logits_dir,
allow_missing_ground_truth,
):
max_views = int(cfg["selection"]["max_views"])
selection_strategy = cfg["selection"]["strategy"]
shard_size = get_point_shard_size(cfg, len(points["X"]))
total_points = len(points["X"])
parallel_camera_workers = int(cfg.get("runtime", {}).get("parallel_camera_workers", 1))
gpu_count = len(jax.devices("gpu")) if jax.default_backend() == "gpu" else 0
visibility_mode = get_visibility_mode(cfg)
saved_paths = []
temp_root = Path(cfg["data"]["batches_dir"]) / ".camera_workers" / tile_info["tile"]
for shard_idx, shard_start in enumerate(range(0, total_points, shard_size)):
shard_end = min(shard_start + shard_size, total_points)
shard_points, shard_coords, shard_normals = slice_point_data(points, coords, normals, shard_start, shard_end)
shard_label = f"{tile_info['tile']}_shard_{shard_idx:04d}"
logging.info(
"Processing tile %s shard %s (%s:%s, %s points)",
tile_info["tile"],
shard_idx,
shard_start,
shard_end,
shard_end - shard_start,
)
use_parallel_workers = (
parallel_camera_workers > 1
and gpu_count > 0
and len(camera_items) > 1
and cfg.get("_worker_mode") is None
)
if use_parallel_workers:
worker_count = min(parallel_camera_workers, gpu_count, len(camera_items))
temp_dir = temp_root / f"shard_{shard_idx:04d}"
temp_dir.mkdir(parents=True, exist_ok=True)
output_files = _launch_camera_workers(
cfg,
tile_info,
camera_items,
temp_dir,
worker_count,
shard_start=shard_start,
shard_end=shard_end,
)
shard_observations = None
for output_file in output_files:
worker_payload = torch.load(output_file, map_location="cpu", weights_only=False)
if shard_observations is None:
shard_observations = worker_payload["observations"]
else:
merge_dense_observations(shard_observations, worker_payload["observations"], max_views, selection_strategy)
if shard_observations is None:
shard_observations = create_dense_observations(
shard_end - shard_start,
max_views,
0,
cfg["data"]["vmin"],
cfg["data"]["vmax"],
get_storage_config(cfg)["visibility_quant_max"],
get_storage_config(cfg)["logits_quant_max"],
)
else:
if visibility_mode == "two_pass_global_depth":
shard_observations = process_camera_items_for_tile_global_depth(
tile_info,
camera_items,
cfg,
calibration_params,
shard_points,
shard_coords,
shard_normals,
logits_dir,
allow_missing_ground_truth,
)
elif visibility_mode == "chunk_local_depth":
shard_observations = process_camera_items_for_tile(
tile_info,
camera_items,
cfg,
calibration_params,
shard_points,
shard_coords,
shard_normals,
logits_dir,
allow_missing_ground_truth,
)
else:
raise ValueError(f"Unsupported visibility_mode: {visibility_mode}")
shard_saved_paths = save_tile_observations_to_pt(
cfg,
shard_label,
tile_info["split"],
shard_coords,
offset,
shard_observations,
cfg["data"]["batches_dir"],
batch_size=cfg["data"]["dataloader_batch_size"],
)
saved_paths.extend(shard_saved_paths)
logging.info(
"Finished tile %s shard %s: %s non-empty point buckets -> %s pt batches",
tile_info["tile"],
shard_idx,
count_nonempty_observations(shard_observations),
len(shard_saved_paths),
)
return saved_paths
def build_tile_dataset(tile_info, cfg, overwrite=False):
prepared_root = Path(cfg["data"]["prepared_h5_dir"]) / tile_info["split"]
output_h5_file = prepared_root / f"{tile_info['tile']}.h5"
output_mode = cfg.get("runtime", {}).get("tile_output_format", "h5_then_pt")
if output_mode == "pt_direct":
pt_output_dir = Path(cfg["data"]["batches_dir"]) / tile_info["split"]
existing_pt_batches = sorted(pt_output_dir.glob(f"{tile_info['tile']}_batch_*.pt"))
if existing_pt_batches and not overwrite:
logging.info("Skipping prepared tile %s", tile_info["tile"])
return [str(path) for path in existing_pt_batches]
elif output_mode == "pt_compact_reduce":
pt_output_dir = Path(cfg["data"]["batches_dir"]) / tile_info["split"]
existing_pt_batches = sorted(pt_output_dir.glob(f"{tile_info['tile']}_batch_*.pt"))
if existing_pt_batches and not overwrite:
logging.info("Skipping prepared tile %s", tile_info["tile"])
return [str(path) for path in existing_pt_batches]
elif output_mode == "pt_sharded_direct":
pt_output_dir = Path(cfg["data"]["batches_dir"]) / tile_info["split"]
existing_pt_batches = sorted(pt_output_dir.glob(f"{tile_info['tile']}_shard_*_batch_*.pt"))
if existing_pt_batches and not overwrite:
logging.info("Skipping prepared tile %s", tile_info["tile"])
return [str(path) for path in existing_pt_batches]
elif output_h5_file.exists() and not overwrite:
logging.info("Skipping prepared tile %s", tile_info["tile"])
return output_h5_file
normals_path = Path(cfg["data"]["normals_dir"]) / tile_info["split"] / f"{tile_info['tile']}.h5"
if not normals_path.exists():
raise FileNotFoundError(f"Normals not found for tile {tile_info['tile']}: {normals_path}")
logits_dir = resolve_logits_dir(cfg, tile_info)
if not logits_dir.exists():
raise FileNotFoundError(f"Logits directory not found for tile {tile_info['tile']}: {logits_dir}")
calibration_params = load_calibration(tile_info["calibration_file"])
points, offset, _, coords = load_las_points(tile_info["lidar_path"])
camera_dict = load_camera_parameters(tile_info["camera_file"], offset)
normals = load_normals_from_h5(normals_path)
allow_missing_ground_truth = tile_info["split"] == "test"
if not points["GT_AVAILABLE"] and not allow_missing_ground_truth:
raise ValueError(
f"Tile {tile_info['tile']} in split {tile_info['split']} has no ground_truth field in LAS."
)
camera_items = list(camera_dict.items())
max_cameras = cfg.get("debug", {}).get("max_cameras_per_tile")
if max_cameras is not None:
camera_items = camera_items[: int(max_cameras)]
if output_mode == "pt_sharded_direct":
saved_paths = process_tile_sharded_to_pt(
tile_info,
cfg,
calibration_params,
points,
coords,
offset,
normals,
camera_items,
logits_dir,
allow_missing_ground_truth,
)
logging.info("Prepared tile %s directly to %s pt shard batches", tile_info["tile"], len(saved_paths))
return saved_paths
parallel_camera_workers = int(cfg.get("runtime", {}).get("parallel_camera_workers", 1))
gpu_count = len(jax.devices("gpu")) if jax.default_backend() == "gpu" else 0
visibility_mode = get_visibility_mode(cfg)
worker_output_mode = get_worker_output_mode(cfg)
if parallel_camera_workers > 1 and gpu_count > 0 and len(camera_items) > 1 and cfg.get("_worker_mode") is None:
worker_count = min(parallel_camera_workers, gpu_count, len(camera_items))
temp_dir = prepared_root / ".camera_workers" / tile_info["tile"]
temp_dir.mkdir(parents=True, exist_ok=True)
camera_groups = [[] for _ in range(worker_count)]
for idx, camera_item in enumerate(camera_items):
camera_groups[idx % worker_count].append(camera_item[0])
processes = []
output_files = []
for worker_idx, camera_group in enumerate(camera_groups):
if not camera_group:
continue
camera_list_file = temp_dir / f"worker_{worker_idx}_cameras.json"
output_file = temp_dir / f"worker_{worker_idx}_observations.pt"
stdout_file = temp_dir / f"worker_{worker_idx}.stdout.log"
stderr_file = temp_dir / f"worker_{worker_idx}.stderr.log"
camera_list_file.write_text(json.dumps(camera_group), encoding="utf-8")
env = os.environ.copy()
env["CUDA_VISIBLE_DEVICES"] = str(worker_idx)
env.setdefault("XLA_PYTHON_CLIENT_PREALLOCATE", "false")
cmd = [
sys.executable,
str(Path(__file__).resolve().parents[1] / "scripts" / "run_batch_generation.py"),
"--config",
cfg["_config_path"],
"--worker-mode",
"camera-batch",
"--tile",
tile_info["tile"],
"--camera-list-file",
str(camera_list_file),
"--worker-output",
str(output_file),
]
logging.info("Launching worker %s for tile %s on GPU %s with %s cameras", worker_idx, tile_info["tile"], worker_idx, len(camera_group))
stdout_handle = None
stderr_handle = None
if worker_output_mode == "files":
stdout_handle = open(stdout_file, "w", encoding="utf-8")
stderr_handle = open(stderr_file, "w", encoding="utf-8")
process = subprocess.Popen(
cmd,
cwd=str(Path(__file__).resolve().parents[1]),
env=env,
stdout=stdout_handle,
stderr=stderr_handle,
)
elif worker_output_mode == "inherit":
process = subprocess.Popen(
cmd,
cwd=str(Path(__file__).resolve().parents[1]),
env=env,
)
else:
raise ValueError(f"Unsupported worker_output_mode: {worker_output_mode}")
processes.append((worker_idx, process, stdout_handle, stderr_handle, stdout_file, stderr_file))
output_files.append(output_file)
for worker_idx, process, stdout_handle, stderr_handle, stdout_file, stderr_file in processes:
return_code = process.wait()
if stdout_handle is not None:
stdout_handle.close()
if stderr_handle is not None:
stderr_handle.close()
if return_code != 0:
log_hint = (
f" Worker logs: stdout={stdout_file} stderr={stderr_file}"
if worker_output_mode == "files"
else ""
)
if return_code == -9:
raise RuntimeError(
f"Camera worker {worker_idx} was killed with code -9 for tile {tile_info['tile']}. "
f"This usually indicates an out-of-memory kill (GPU or system RAM). "
f"{log_hint}"
)
raise RuntimeError(
f"Camera worker {worker_idx} failed for tile {tile_info['tile']} with code {return_code}. "
f"{log_hint}"
)
if output_mode == "pt_compact_reduce":
compact_payloads = []
for output_file in output_files:
worker_payload = torch.load(output_file, map_location="cpu", weights_only=False)
compact_payloads.append(worker_payload["observations"])
saved_paths = save_compact_payloads_to_pt(
cfg,
tile_info["tile"],
tile_info["split"],
coords,
offset,
compact_payloads,
cfg["data"]["batches_dir"],
batch_size=cfg["data"]["dataloader_batch_size"],
)
logging.info("Prepared tile %s with compact_reduce to %s pt batches", tile_info["tile"], len(saved_paths))
return saved_paths
all_observations = {}
max_views = int(cfg["selection"]["max_views"])
selection_strategy = cfg["selection"]["strategy"]
for output_file in output_files:
worker_payload = torch.load(output_file, map_location="cpu", weights_only=False)
merge_worker_observations_bounded(
all_observations,
worker_payload["observations"],
max_views,
selection_strategy,
)
else:
if output_mode == "pt_compact_reduce":
compact_payload = process_camera_items_for_tile_compact(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
)
saved_paths = save_compact_payloads_to_pt(
cfg,
tile_info["tile"],
tile_info["split"],
coords,
offset,
[compact_payload],
cfg["data"]["batches_dir"],
batch_size=cfg["data"]["dataloader_batch_size"],
)
logging.info("Prepared tile %s with compact_reduce to %s pt batches", tile_info["tile"], len(saved_paths))
return saved_paths
if visibility_mode == "two_pass_global_depth":
all_observations = process_camera_items_for_tile_global_depth(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
)
elif visibility_mode == "chunk_local_depth":
all_observations = process_camera_items_for_tile(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
)
else:
raise ValueError(f"Unsupported visibility_mode: {visibility_mode}")
logging.info(
"Tile %s finished camera processing; aggregated observations for %s point buckets",
tile_info["tile"],
count_nonempty_observations(all_observations),
)
if output_mode == "pt_direct":
saved_paths = save_tile_observations_to_pt(
cfg,
tile_info["tile"],
tile_info["split"],
coords,
offset,
all_observations,
cfg["data"]["batches_dir"],
batch_size=cfg["data"]["dataloader_batch_size"],
)
logging.info("Prepared tile %s directly to %s pt batches", tile_info["tile"], len(saved_paths))
return saved_paths
write_tile_h5(output_h5_file, tile_info["tile"], tile_info["split"], coords, all_observations, cfg)
logging.info("Prepared tile %s -> %s", tile_info["tile"], output_h5_file)
return output_h5_file
def run_camera_batch_worker(cfg, tile_name, camera_list_file, worker_output, shard_start=None, shard_end=None):
manifest = load_or_build_manifest(cfg)
tile_info = next(item for item in manifest if item["tile"] == tile_name)
logits_dir = resolve_logits_dir(cfg, tile_info)
calibration_params = load_calibration(tile_info["calibration_file"])
points, offset, _, coords = load_las_points(tile_info["lidar_path"])
camera_dict = load_camera_parameters(tile_info["camera_file"], offset)
normals_path = Path(cfg["data"]["normals_dir"]) / tile_info["split"] / f"{tile_info['tile']}.h5"
normals = load_normals_from_h5(normals_path)
allow_missing_ground_truth = tile_info["split"] == "test"
if not points["GT_AVAILABLE"] and not allow_missing_ground_truth:
raise ValueError(
f"Tile {tile_info['tile']} in split {tile_info['split']} has no ground_truth field in LAS."
)
selected_camera_names = json.loads(Path(camera_list_file).read_text(encoding="utf-8"))
camera_items = [(name, camera_dict[name]) for name in selected_camera_names if name in camera_dict]
if shard_start is not None and shard_end is not None:
points, coords, normals = slice_point_data(points, coords, normals, int(shard_start), int(shard_end))
output_mode = cfg.get("runtime", {}).get("tile_output_format", "h5_then_pt")
if output_mode == "pt_compact_reduce":
observations = process_camera_items_for_tile_compact(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
)
Path(worker_output).parent.mkdir(parents=True, exist_ok=True)
torch.save({"observations": observations}, worker_output, pickle_protocol=4)
return
visibility_mode = get_visibility_mode(cfg)
if visibility_mode == "two_pass_global_depth":
observations = process_camera_items_for_tile_global_depth(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
)
elif visibility_mode == "chunk_local_depth":
observations = process_camera_items_for_tile(
tile_info,
camera_items,
cfg,
calibration_params,
points,
coords,
normals,
logits_dir,
allow_missing_ground_truth,
)
else:
raise ValueError(f"Unsupported visibility_mode: {visibility_mode}")
Path(worker_output).parent.mkdir(parents=True, exist_ok=True)
torch.save({"observations": observations}, worker_output, pickle_protocol=4)
def batch_generation_main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=Path, default="configs/config_gridnet_hd_dataset_generation.yaml")
parser.add_argument("--worker-mode", choices=["camera-batch"], default=None)
parser.add_argument("--tile", default=None)
parser.add_argument("--camera-list-file", default=None)
parser.add_argument("--worker-output", default=None)
parser.add_argument("--shard-start", type=int, default=None)
parser.add_argument("--shard-end", type=int, default=None)
args = parser.parse_args()
cfg = yaml.safe_load(args.config.read_text())
cfg["_config_path"] = str(args.config.resolve())
cfg["_worker_mode"] = args.worker_mode
if args.worker_mode == "camera-batch":
run_camera_batch_worker(
cfg,
args.tile,
args.camera_list_file,
args.worker_output,
shard_start=args.shard_start,
shard_end=args.shard_end,
)
return
manifest = load_or_build_manifest(cfg)
debug_tiles = cfg.get("debug", {}).get("tiles")
if debug_tiles:
wanted_tiles = set(debug_tiles)
manifest = [tile_info for tile_info in manifest if tile_info["tile"] in wanted_tiles]
overwrite = cfg.get("runtime", {}).get("overwrite_existing", False)
output_mode = cfg.get("runtime", {}).get("tile_output_format", "h5_then_pt")
h5_paths_by_split = {"train": [], "val": [], "test": []}
for tile_info in tqdm(manifest, desc="Tiles", unit="tile"):
output_path = build_tile_dataset(tile_info, cfg, overwrite=overwrite)
if output_mode == "h5_then_pt":
h5_paths_by_split[tile_info["split"]].append(str(output_path))
if output_mode == "h5_then_pt":
save_dataset_to_pt_parallel(
cfg,
h5_paths_by_split["train"],
h5_paths_by_split["val"],
h5_paths_by_split["test"],
cfg["data"]["batches_dir"],
batch_size=cfg["data"]["dataloader_batch_size"],
num_workers=cfg["data"]["max_workers"],
)