| """Run SkySense inference and save arrays for evaluation.""" |
|
|
| import importlib.util |
| import argparse |
| 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("skysense_model", ROOT / "model" / "skysense.py") |
| module = importlib.util.module_from_spec(spec) |
| spec.loader.exec_module(module) |
| return module.SkySense |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Run batched SkySense segmentation inference") |
| parser.add_argument("--batch-size", type=int) |
| args = parser.parse_args() |
| with (ROOT / "conf" / "config.yaml").open(encoding="utf-8") as handle: |
| config = yaml.safe_load(handle) |
| checkpoint_path = ROOT / config["paths"]["checkpoint"] |
| if not checkpoint_path.exists(): |
| raise FileNotFoundError( |
| f"Missing checkpoint: {checkpoint_path.relative_to(ROOT)}. " |
| "Run `python scripts/train.py` first." |
| ) |
| use_accelerator = torch.cuda.is_available() and config["runtime"].get("device", "auto") != "cpu" |
| device = torch.device("cuda" if use_accelerator else "cpu") |
| checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) |
| SkySense = load_model_class() |
| model = SkySense( |
| **config["model"], |
| hr_channels=config["data"]["hr_channels"], |
| s2_channels=config["data"]["s2_channels"], |
| s1_channels=config["data"]["s1_channels"], |
| num_classes=config["data"]["num_classes"], |
| ).to(device) |
| model.load_state_dict(checkpoint["model"]) |
| model.eval() |
| test_path = ROOT / config["data"]["root"] / "test.npz" |
| if not test_path.exists(): |
| raise FileNotFoundError( |
| f"Missing inference data: {test_path.relative_to(ROOT)}. " |
| "Run `python scripts/fake_data.py` first." |
| ) |
| archive = np.load(test_path) |
| keys = ["hr", "s2", "s1", "dates_hr", "dates_s2", "dates_s1", "region"] |
| arrays = {key: archive[key] for key in keys} |
| expected = { |
| "hr": (config["data"]["hr_timesteps"], config["data"]["hr_channels"], config["data"]["hr_size"], config["data"]["hr_size"]), |
| "s2": (config["data"]["s2_timesteps"], config["data"]["s2_channels"], config["data"]["s2_size"], config["data"]["s2_size"]), |
| "s1": (config["data"]["s1_timesteps"], config["data"]["s1_channels"], config["data"]["s1_size"], config["data"]["s1_size"]), |
| "dates_hr": (config["data"]["hr_timesteps"],), |
| "dates_s2": (config["data"]["s2_timesteps"],), |
| "dates_s1": (config["data"]["s1_timesteps"],), |
| "region": (), |
| } |
| sample_count = len(arrays["hr"]) |
| for key, shape in expected.items(): |
| if len(arrays[key]) != sample_count or tuple(arrays[key].shape[1:]) != shape: |
| raise ValueError(f"Invalid test {key} shape {arrays[key].shape}; expected [N,{','.join(map(str, shape))}]") |
| for key in ("hr", "s2", "s1"): |
| if not np.issubdtype(arrays[key].dtype, np.floating): |
| raise TypeError(f"{key} must use a floating dtype") |
| for key in ("dates_hr", "dates_s2", "dates_s1", "region"): |
| if arrays[key].dtype != np.int64: |
| raise TypeError(f"{key} must use int64") |
| if any(np.any((arrays[key] < 0) | (arrays[key] > 364)) for key in ("dates_hr", "dates_s2", "dates_s1")): |
| raise ValueError("Test dates must be in [0, 364]") |
| if np.any((arrays["region"] < 0) | (arrays["region"] >= config["model"]["num_regions"])): |
| raise ValueError("Test region IDs are out of range") |
| labels = archive["labels"] |
| if labels.dtype != np.int64 or labels.shape != (sample_count, config["data"]["hr_size"], config["data"]["hr_size"]): |
| raise ValueError("Test labels must be int64 [N,hr_size,hr_size]") |
| batch_size = args.batch_size or config["train"]["batch_size"] |
| predictions = [] |
| all_probabilities = [] |
| with torch.inference_mode(): |
| for start in range(0, len(arrays["hr"]), batch_size): |
| tensors = {key: torch.from_numpy(value[start:start + batch_size]).to(device) |
| for key, value in arrays.items()} |
| output = model(tensors["hr"], tensors["s2"], tensors["s1"], tensors["dates_hr"], tensors["dates_s2"], tensors["dates_s1"], tensors["region"]) |
| probabilities = output["logits"].softmax(dim=1).cpu().numpy() |
| all_probabilities.append(probabilities) |
| predictions.append(probabilities.argmax(axis=1)) |
| output_dir = ROOT / config["paths"]["inference_dir"] |
| output_dir.mkdir(parents=True, exist_ok=True) |
| np.save(output_dir / "predictions.npy", np.concatenate(predictions)) |
| np.save(output_dir / "probabilities.npy", np.concatenate(all_probabilities)) |
| np.save(output_dir / "targets.npy", labels) |
| data_source = str(archive["data_source"]) if "data_source" in archive.files else "unknown" |
| protocol = str(archive["protocol"]) if "protocol" in archive.files else "unknown" |
| np.savez(output_dir / "metadata.npz", data_source=data_source, protocol=protocol) |
| print( |
| f"output={output_dir.relative_to(ROOT)} samples={len(archive['hr'])} " |
| f"data_source={data_source} protocol={protocol}" |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|