| """Autoregressive tile inference with the unchanged scientific global-grid contract.""" |
|
|
| import json |
| from pathlib import Path |
| import sys |
|
|
| import numpy as np |
| import torch |
| import yaml |
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| sys.path.insert(0, str(ROOT / "scripts")) |
| from model.fuxi_ocean import FORMAT_VERSION, FuXiOcean, assemble_tiles |
| from train import TileDataset, validate_data_contract |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| torch.set_num_threads(config["runtime"]["num_threads"]) |
| device = torch.device("cuda" if torch.cuda.is_available() and config["runtime"]["device"] != "cpu" else "cpu") |
| checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=False) |
| required = {"model", "model_config", "format_version"} |
| if not required.issubset(checkpoint): |
| raise ValueError(f"checkpoint missing keys: {sorted(required - checkpoint.keys())}") |
| model_config = checkpoint["model_config"] |
| model_args = model_config["architecture"] |
| model = FuXiOcean(**model_args).to(device) |
| model.load_state_dict(checkpoint["model"]); model.eval() |
| data = np.load(ROOT / config["data"]["path"]) |
| validate_data_contract(data, config) |
| shape_keys = ("input_shape", "atmosphere_shape", "output_shape") |
| if checkpoint["format_version"] != FORMAT_VERSION or model_config["data_format_version"] != str(data["format_version"]) or any(model_config[key] != data[key].tolist() for key in shape_keys): |
| raise ValueError("checkpoint, model configuration, and data contract are incompatible") |
| predictions, truths, initial, latitudes, masks, origins = [], [], [], [], [], [] |
| with torch.no_grad(): |
| for index in range(int(data["train_count"]), len(data["ocean"])): |
| batch = TileDataset(data, [index])[0] |
| ocean, atmosphere, coordinates, bathymetry, mask, time_info, target, latitude = [value[None].to(device) for value in batch] |
| history = ocean; sample_predictions, sample_truths = [], [] |
| for lead in range(config["inference"]["rollout_steps"]): |
| time_info[:, 2] = lead |
| prediction = model(history, atmosphere, coordinates, bathymetry, mask, time_info) |
| truth = target + lead * 0.005 |
| sample_predictions.append(prediction.cpu().numpy()[0]); sample_truths.append(truth.cpu().numpy()[0]) |
| history = torch.cat((history[:, 1:], prediction[:, None]), dim=1) |
| predictions.append(sample_predictions); truths.append(sample_truths); initial.append(ocean[0, -1].cpu().numpy()) |
| latitudes.append(latitude[0].cpu().numpy()); masks.append(mask[0].cpu().numpy()); origins.append(data["tile_origins"][index]) |
| output = ROOT / config["paths"]["inference"]; output.parent.mkdir(parents=True, exist_ok=True) |
| records = data["selected_tile_records"][int(data["train_count"]):] |
| |
| test_tile = np.ones((1, 1, config["data"]["tile_height"], config["data"]["tile_width"]), dtype=np.float32) |
| _, stitch_coverage = assemble_tiles(test_tile, [[0, config["data"]["tile_height"], 0, config["data"]["tile_width"], 0, |
| config["data"]["tile_height"], 0, config["data"]["tile_width"]]], |
| (config["data"]["tile_height"], config["data"]["tile_width"])) |
| coverage_fraction = float(sum((r[5] - r[4]) * (r[7] - r[6]) for r in records) / |
| (config["data"]["global_height"] * config["data"]["global_width"])) |
| checkpoint_source = str(config["paths"]["checkpoint"]) |
| np.savez_compressed(output, output_kind="sampled_tiles", format_version=FORMAT_VERSION, |
| checkpoint_source=checkpoint_source, sample_count=len(predictions), synthetic=True, |
| coverage_fraction=coverage_fraction, is_complete_global=False, |
| variable_groups=json.dumps({"S": {"channels": [0, 26], "unit": "psu"}, "T": {"channels": [26, 52], "unit": "degC"}, "U": {"channels": [52, 78], "unit": "m s-1"}, "V": {"channels": [78, 104], "unit": "m s-1"}, "SSH": {"channels": [104, 105], "unit": "m"}}), |
| prediction=np.asarray(predictions), truth=np.asarray(truths), initial=np.asarray(initial), |
| latitude_deg=np.asarray(latitudes), depth_mask=np.asarray(masks), tile_origins=np.asarray(origins), |
| tile_records=records, output_shape=data["output_shape"], |
| lead_hours=config["data"]["time_step_hours"] * np.arange(1, config["inference"]["rollout_steps"] + 1)) |
| metadata = {"output_kind": "sampled_tiles", "format_version": FORMAT_VERSION, |
| "checkpoint_source": checkpoint_source, "sample_count": len(predictions), "synthetic": True, |
| "output_shape": data["output_shape"].tolist(), "paper_rollout_steps": config["paper_model"]["rollout_steps"], |
| "executed_rollout_steps": config["inference"]["rollout_steps"], "tile_streaming": True, |
| "tile_origins": np.asarray(origins).tolist(), "coverage_fraction": coverage_fraction, |
| "is_complete_global": False, "stitch_interface_verified": bool(stitch_coverage.all()), |
| "coverage_semantics": "fraction of global cells owned after deterministic overlap crops", |
| "variable_groups": json.loads(str(np.load(output)["variable_groups"]))} |
| metadata_path = ROOT / config["paths"]["inference_metadata"] |
| metadata_path.write_text(json.dumps(metadata, indent=2) + "\n") |
| print(f"predictions={output.relative_to(ROOT)} shape={np.asarray(predictions).shape}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|