"""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"]):] # Exercise overlap-crop assembly on a bounded canvas without allocating the global field. 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()