DOFA / scripts /inference.py
zhangrenchao's picture
Upload DOFA model package
1d4cac8 verified
Raw
History Blame Contribute Delete
5.33 kB
"""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()