W-MAE / scripts /inference.py
yzt15806542928's picture
Upload folder using huggingface_hub
80cf062 verified
Raw
History Blame Contribute Delete
7.47 kB
from __future__ import annotations
import argparse
from pathlib import Path
import sys
from typing import Any
import numpy as np
import torch
from torch.utils.data import DataLoader
PROJECT_ROOT = Path(__file__).resolve().parents[1]
if str(PROJECT_ROOT) not in sys.path:
sys.path.insert(0, str(PROJECT_ROOT))
SCRIPTS_DIR = PROJECT_ROOT / "scripts"
if str(SCRIPTS_DIR) not in sys.path:
sys.path.insert(0, str(SCRIPTS_DIR))
from model.checkpoint import load_checkpoint
from era5_adapter import WMAEERA5Dataset
from train import build_model, device_summary, load_config, resolve_device, resolve_project_path, validate_config
def build_inference_dataset(config: dict[str, Any], config_path: Path) -> WMAEERA5Dataset:
data = config["data"]
if int(data["input_steps"]) != 1 or int(data["output_steps"]) != 1:
raise ValueError("W-MAE inference currently requires input_steps=1 and output_steps=1.")
if not data["variables"]:
raise ValueError(
"data.variables is empty. Supply the verified, explicitly ordered 20-channel list before inference."
)
return WMAEERA5Dataset(
dataset_dir=resolve_project_path(data["dataset_dir"], config_path),
years=data["test_years"],
variables=data["variables"],
task="pretrain",
input_steps=1,
output_steps=1,
normalize=bool(data["normalize"]),
)
def sample_time_index(time_index: Any, batch_index: int, batch_size: int) -> Any:
"""Extract one sample from DataLoader's sequence-major default collation."""
if isinstance(time_index, (list, tuple)):
if time_index and isinstance(time_index[0], (list, tuple)):
return time_index[0][batch_index]
if len(time_index) == batch_size:
return time_index[batch_index]
return time_index
def run_inference(
model: torch.nn.Module,
loader: DataLoader,
output_dir: Path,
mask_ratio: float,
device: torch.device,
max_samples: int | None = None,
log_interval: int = 1,
) -> int:
output_dir.mkdir(parents=True, exist_ok=True)
model.eval()
written = 0
with torch.no_grad():
total_batches = len(loader)
for batch_index, (inputs, _, _, step_idx, time_index) in enumerate(loader, start=1):
if inputs.ndim != 4:
raise ValueError(f"W-MAE inference expects a 4D batch, got {tuple(inputs.shape)}.")
inputs = inputs.to(device)
result = model(inputs, mask_ratio=mask_ratio)
reconstruction = model.unpatchify(result.prediction)
error = reconstruction - inputs
mask = result.mask.reshape(result.mask.shape[0], *model.patch_embed.grid_size)
for sample_index in range(inputs.shape[0]):
if max_samples is not None and written >= max_samples:
print(
f"inference batch={batch_index}/{total_batches} samples_written={written}",
flush=True,
)
return written
sample_time = sample_time_index(time_index, sample_index, inputs.shape[0])
if isinstance(sample_time, (list, tuple)):
sample_time = sample_time[0]
np.savez_compressed(
output_dir / f"sample_{written:06d}.npz",
input=inputs[sample_index].cpu().numpy(),
reconstruction=reconstruction[sample_index].cpu().numpy(),
error=error[sample_index].cpu().numpy(),
mask=mask[sample_index].cpu().numpy(),
step_idx=np.asarray(step_idx[sample_index].item()),
time_index=np.asarray(sample_time),
)
written += 1
if batch_index == 1 or batch_index % log_interval == 0 or batch_index == total_batches:
print(
f"inference batch={batch_index}/{total_batches} samples_written={written}",
flush=True,
)
return written
def main() -> None:
parser = argparse.ArgumentParser(description="Run W-MAE reconstruction inference on ERA5 samples.")
parser.add_argument("--config", type=Path, default=PROJECT_ROOT / "conf" / "config.yaml")
parser.add_argument("--checkpoint", type=Path, default="./data/checkpoint/model_bak.pth")
parser.add_argument("--checkpoint-source-root", type=Path, default=None)
parser.add_argument("--non-strict-checkpoint", action="store_true")
parser.add_argument("--output-dir", type=Path, default=None)
parser.add_argument(
"--device",
default="auto",
help="Inference device: auto (default), cuda (AMD DCU through HIP), or cpu.",
)
parser.add_argument("--mask-ratio", type=float, default=None)
parser.add_argument("--max-samples", type=int, default=None)
parser.add_argument("--log-interval", type=int, default=1)
args = parser.parse_args()
if args.max_samples is not None and args.max_samples <= 0:
raise ValueError("--max-samples must be positive when provided.")
if args.log_interval <= 0:
raise ValueError("--log-interval must be positive.")
config_path = args.config.resolve()
print(
f"inference started: config={config_path} checkpoint={args.checkpoint} device={args.device}",
flush=True,
)
config = load_config(config_path)
validate_config(config)
if args.checkpoint is None:
raise ValueError("Inference requires an explicit --checkpoint path.")
device = resolve_device(args.device)
print(f"runtime: {device_summary(device)}", flush=True)
print("building model", flush=True)
model = build_model(config).to(device)
print(f"model ready: parameters_device={next(model.parameters()).device}", flush=True)
source_root = (
resolve_project_path(args.checkpoint_source_root, config_path)
if args.checkpoint_source_root
else None
)
report = load_checkpoint(
model,
resolve_project_path(args.checkpoint, config_path),
strict=not args.non_strict_checkpoint,
map_location=device,
source_root=source_root,
)
print(f"loaded checkpoint: {report}", flush=True)
data_config = config["data"]
print("building inference dataset", flush=True)
dataset = build_inference_dataset(config, config_path)
loader = DataLoader(
dataset,
batch_size=int(data_config["batch_size"]),
shuffle=False,
num_workers=int(data_config["num_workers"]),
)
print(f"dataset ready: samples={len(dataset)} batches={len(loader)}", flush=True)
output_dir = resolve_project_path(args.output_dir, config_path) if args.output_dir else (
resolve_project_path(config["project"]["output_dir"], config_path) / "inference"
)
mask_ratio = float(config["model"]["mask_ratio"] if args.mask_ratio is None else args.mask_ratio)
if mask_ratio not in {0.0, 0.75}:
raise ValueError("W-MAE supports mask_ratio 0.0 or 0.75 only.")
print(f"inference running: mask_ratio={mask_ratio} output_dir={output_dir}", flush=True)
count = run_inference(model, loader, output_dir, mask_ratio, device, args.max_samples, args.log_interval)
print(f"wrote {count} reconstruction samples to {output_dir}", flush=True)
print("inference completed successfully", flush=True)
if __name__ == "__main__":
main()