SkySense / scripts /inference.py
zhangrenchao's picture
Update SkySense model package
2d7eead verified
Raw
History Blame Contribute Delete
5.29 kB
"""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()