FuXi-Ocean / scripts /inference.py
zhangrenchao's picture
Publish FuXi-Ocean engineering reproduction
d3e46b7 verified
Raw
History Blame Contribute Delete
5.87 kB
"""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()