| """Reconstruct validated sensor NPZ files in bounded device batches.""" |
|
|
| import importlib.util |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import yaml |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
|
|
|
|
| def load_model_class(): |
| spec = importlib.util.spec_from_file_location("dofa_model", ROOT / "model" / "dofa.py") |
| module = importlib.util.module_from_spec(spec) |
| spec.loader.exec_module(module) |
| return module.DOFA |
|
|
|
|
| def expected_wavelengths(modality): |
| if modality.get("wavelength_mode") == "synthetic_uniform": |
| return modality["wavelength_start"] + np.arange(modality["channels"], dtype=np.float32) * modality["wavelength_step"] |
| return np.asarray(modality["wavelengths"], dtype=np.float32) |
|
|
|
|
| def scalar(archive, key, default=None): |
| if key not in archive: |
| if default is not None: |
| return default |
| raise ValueError(f"NPZ is missing required metadata: {key}") |
| if archive[key].ndim != 0: |
| raise ValueError(f"NPZ metadata {key} must be a scalar") |
| return archive[key].item() |
|
|
|
|
| def validate_archive(archive, path, name, data_config): |
| if "images" not in archive or "wavelengths" not in archive: |
| raise ValueError(f"{path}: images and wavelengths are required") |
| images, wavelengths = archive["images"], archive["wavelengths"] |
| modality_config = data_config["modalities"][name] |
| if images.ndim != 4 or images.dtype != np.float32: |
| raise ValueError(f"{path}: images must be float32 NCHW") |
| expected_shape = (modality_config["channels"], data_config["image_size"], data_config["image_size"]) |
| if images.shape[1:] != expected_shape: |
| raise ValueError(f"{path}: expected [N,{expected_shape[0]},224,224], got {images.shape}") |
| expected = expected_wavelengths(modality_config) |
| if wavelengths.shape != (modality_config["channels"],) or not np.issubdtype(wavelengths.dtype, np.floating) or not np.isfinite(wavelengths).all() or not np.allclose(wavelengths, expected, rtol=1e-5, atol=1e-6): |
| raise ValueError(f"{path}: wavelengths do not match configured sensor wavelengths") |
| modality = str(scalar(archive, "modality")) |
| protocol = str(scalar(archive, "protocol")) |
| source = str(scalar(archive, "data_source", "unknown")) |
| if modality != name or protocol != data_config["protocol"]: |
| raise ValueError(f"{path}: modality/protocol metadata does not match config") |
| data_range = float(scalar(archive, "data_range", data_config.get("data_range"))) |
| if not np.isfinite(data_range) or data_range <= 0: |
| raise ValueError(f"{path}: PSNR requires a positive data_range metadata or config value") |
| return images, wavelengths.astype("float32", copy=False), protocol, source, data_range |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text(encoding="utf-8")) |
| checkpoint_path = ROOT / config["paths"]["checkpoint"] |
| if not checkpoint_path.exists(): |
| raise FileNotFoundError("Missing checkpoint. Run `python scripts/train.py` first.") |
| device = torch.device("cuda" if torch.cuda.is_available() and config["runtime"]["device"] != "cpu" else "cpu") |
| checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False) |
| if checkpoint.get("protocol") != config["data"]["protocol"]: |
| raise ValueError("Checkpoint and configured protocols do not match") |
| model = load_model_class()(**config["model"]).to(device) |
| model.load_state_dict(checkpoint["model"]) |
| model.eval() |
| output_dir = ROOT / config["paths"]["inference_dir"] |
| output_dir.mkdir(parents=True, exist_ok=True) |
| data_root = ROOT / config["data"]["root"] |
| batch_size = config["runtime"]["inference_batch_size"] |
| torch.manual_seed(config["seed"]) |
| for modality in config["data"]["modalities"]: |
| path = data_root / f"test_{modality}.npz" |
| if not path.exists(): |
| raise FileNotFoundError(f"Missing test data: {path.relative_to(ROOT)}") |
| archive = np.load(path) |
| images, wavelengths, protocol, source, data_range = validate_archive( |
| archive, path, modality, config["data"]) |
| reconstructions = np.empty_like(images) |
| masks = np.empty((len(images), model.num_patches), dtype=bool) |
| wavelength_tensor = torch.from_numpy(wavelengths).to(device) |
| with torch.inference_mode(): |
| for start in range(0, len(images), batch_size): |
| stop = min(start + batch_size, len(images)) |
| batch = torch.from_numpy(images[start:stop]).to(device) |
| output = model(batch, wavelength_tensor) |
| reconstructions[start:stop] = output["reconstruction"].cpu().numpy() |
| masks[start:stop] = output["mask"].cpu().numpy() |
| del batch, output |
| target = output_dir / f"{modality}_reconstruction.npz" |
| np.savez_compressed(target, inputs=images, reconstructions=reconstructions, masks=masks, |
| wavelengths=wavelengths, modality=np.asarray(modality), |
| data_source=np.asarray(source), protocol=np.asarray(protocol), |
| data_range=np.asarray(data_range, dtype=np.float32)) |
| print(f"output={target.relative_to(ROOT)} channels={images.shape[1]} batch_size={batch_size}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|