| |
| """Run autoregressive FengWu-W2S inference and save forecast fields.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import sys |
| from datetime import datetime, timedelta |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import yaml |
|
|
| PROJECT_ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(PROJECT_ROOT)) |
|
|
| from model.fengwu_w2s import FengWuW2S |
| from scripts.data_loader import ERA5WindowDataset, read_metadata, resolve_data_dir |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--config", default=str(PROJECT_ROOT / "conf/config.yaml")) |
| parser.add_argument("--data-dir") |
| parser.add_argument("--checkpoint") |
| parser.add_argument("--output-dir") |
| parser.add_argument("--split", choices=("train", "val", "test"), default="test") |
| parser.add_argument("--steps", type=int) |
| parser.add_argument("--limit", type=int, help="Limit windows for a quick smoke test") |
| parser.add_argument("--stochastic", action="store_true") |
| parser.add_argument("--seed", type=int, default=0) |
| return parser.parse_args() |
|
|
|
|
| def _resolve(path: str | Path, root: Path = PROJECT_ROOT) -> Path: |
| candidate = Path(path).expanduser() |
| return candidate if candidate.is_absolute() else (root / candidate).resolve() |
|
|
|
|
| def _build_model(config: dict) -> FengWuW2S: |
| model_cfg = config["model"] |
| data_cfg = config["data"] |
| return FengWuW2S( |
| in_channels=len(data_cfg["channels"]), |
| group_indices=data_cfg["groups"], |
| hidden_channels=int(model_cfg.get("hidden_channels", 32)), |
| latent_channels=int(model_cfg.get("latent_channels", 64)), |
| patch_size=int(model_cfg.get("patch_size", 4)), |
| num_blocks=int(model_cfg.get("num_blocks", 2)), |
| perturbation_scale=float(model_cfg.get("perturbation_scale", 0.01)), |
| ) |
|
|
|
|
| def main() -> None: |
| args = parse_args() |
| with Path(args.config).open(encoding="utf-8") as source: |
| config = yaml.safe_load(source) |
| data_cfg = config["data"] |
| model_cfg = config["model"] |
| data_dir = resolve_data_dir(args.data_dir or data_cfg["data_dir"], PROJECT_ROOT) |
| checkpoint = _resolve(args.checkpoint or model_cfg.get("default_checkpoint", "./data/checkpoints/model_bak.pth")) |
| if not checkpoint.exists(): |
| raise FileNotFoundError(f"Checkpoint not found: {checkpoint}; run scripts/train.py first") |
| split_years = { |
| "train": data_cfg["train_years"], |
| "val": data_cfg["val_years"], |
| "test": data_cfg["test_years"], |
| }[args.split] |
| steps = int(args.steps or model_cfg.get("inference_steps", 1)) |
| dataset = ERA5WindowDataset( |
| data_dir=data_dir, |
| years=split_years, |
| channels=data_cfg["channels"], |
| input_steps=int(model_cfg.get("input_steps", 2)), |
| rollout_steps=max(1, steps), |
| ) |
| metadata = read_metadata(data_dir, data_cfg["channels"]) |
| means = torch.from_numpy(metadata["means"]) |
| stds = torch.from_numpy(metadata["stds"]) |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| model = _build_model(config).to(device) |
| state = torch.load(checkpoint, map_location=device, weights_only=False) |
| model.load_state_dict(state["model_state_dict"]) |
| model.eval() |
| output_dir = _resolve(args.output_dir or "./result/output") |
| output_dir.mkdir(parents=True, exist_ok=True) |
| index = [] |
| limit = len(dataset) if args.limit is None else min(len(dataset), max(0, int(args.limit))) |
| for sample_index in range(limit): |
| inputs, _, timestamp = dataset[sample_index] |
| initial = inputs.unsqueeze(0).to(device) |
| with torch.no_grad(): |
| prediction = model.rollout(initial, steps=steps, stochastic=args.stochastic, seed=args.seed + sample_index) |
| prediction = prediction.squeeze(0).cpu() * stds + means |
| year = timestamp[:4] |
| year_dir = output_dir / year |
| year_dir.mkdir(parents=True, exist_ok=True) |
| start = datetime.strptime(timestamp, "%Y%m%d%H") |
| for lead, field in enumerate(prediction): |
| valid_time = start + timedelta(hours=lead * int(metadata["time_step"])) |
| valid_timestamp = valid_time.strftime("%Y%m%d%H") |
| path = year_dir / f"{timestamp}_lead{lead:03d}.npy" |
| np.save(path, field.numpy().astype(np.float32)) |
| index.append({ |
| "source_timestamp": timestamp, |
| "valid_timestamp": valid_timestamp, |
| "lead": lead, |
| |
| |
| "path": str(path.relative_to(output_dir)), |
| }) |
| (output_dir / "index.json").write_text(json.dumps(index, indent=2), encoding="utf-8") |
| print(f"Saved {len(index)} forecast fields under {output_dir}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|