BENO / scripts /inference.py
OneScience's picture
Upload folder using huggingface_hub
7087ec2 verified
Raw
History Blame Contribute Delete
7.73 kB
import json
import logging
import os
import sys
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
PROJECT_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(PROJECT_ROOT))
from model import HeteroGNS
from onescience.datapipes.cfd import BENODatapipe
from onescience.distributed.manager import DistributedManager
from onescience.utils.YParams import YParams
from onescience.utils.beno.util import record_data, to_np_array
from onescience.utils.beno.utilities import LpLoss, plot_data
logging.basicConfig(level=logging.INFO)
LOGGER = logging.getLogger("BENO_Inference")
def set_default_data_env():
os.environ.setdefault("ONESCIENCE_BENO_DATA_DIR", str(PROJECT_ROOT / "data"))
def resolve_path(path_value):
path = Path(path_value)
return path if path.is_absolute() else PROJECT_ROOT / path
def load_config():
set_default_data_env()
cfg = YParams(str(PROJECT_ROOT / "conf" / "config.yaml"), "root")
cfg.datapipe.source.data_dir = str(resolve_path(cfg.datapipe.source.data_dir))
cfg.datapipe.source.cache_dir = str(resolve_path(cfg.datapipe.source.cache_dir))
cfg.inference.checkpoint_dir = str(resolve_path(cfg.inference.checkpoint_dir))
cfg.inference.result_dir = str(resolve_path(cfg.inference.result_dir))
cfg.inference.picture_dir = str(resolve_path(cfg.inference.picture_dir))
cfg.inference.metrics_path = str(resolve_path(cfg.inference.metrics_path))
if cfg.inference.checkpoint_path:
cfg.inference.checkpoint_path = str(resolve_path(cfg.inference.checkpoint_path))
return cfg
def activation_from_name(name):
name = str(name).lower()
if name == "relu":
return nn.ReLU
if name == "elu":
return nn.ELU
if name == "leakyrelu":
return nn.LeakyReLU
return nn.SiLU
def build_model(model_cfg):
return HeteroGNS(
nnode_in_features=model_cfg.nnode_in_features,
nnode_out_features=model_cfg.nnode_out_features,
nedge_in_features=model_cfg.nedge_in_features,
latent_dim=model_cfg.get("latent_dim", model_cfg.get("width", 128)),
nmessage_passing_steps=model_cfg.get("nmessage_passing_steps", 10),
nmlp_layers=model_cfg.nmlp_layers,
mlp_hidden_dim=model_cfg.get("mlp_hidden_dim", model_cfg.get("width", 128)),
activation=activation_from_name(model_cfg.act),
boundary_dim=model_cfg.boundary_dim,
trans_layer=model_cfg.trans_layer,
)
def select_device(device_name, dist):
device_name = str(device_name).lower()
if device_name == "cpu":
return torch.device("cpu")
if device_name == "cuda":
if not torch.cuda.is_available():
raise RuntimeError("inference.device is cuda, but CUDA is not available.")
return torch.device(f"cuda:{dist.local_rank}")
return dist.device
def find_checkpoint(cfg):
if cfg.inference.checkpoint_path:
path = Path(cfg.inference.checkpoint_path)
return path if path.exists() else None
checkpoints = sorted(Path(cfg.inference.checkpoint_dir).glob("model_epoch_*.pt"))
return checkpoints[-1] if checkpoints else None
def main():
cfg = load_config()
DistributedManager.initialize()
dist = DistributedManager()
device = select_device(cfg.inference.get("device", "auto"), dist)
datapipe = BENODatapipe(cfg, distributed=(dist.world_size > 1))
test_loader, _ = datapipe.test_dataloader()
u_normalizer = datapipe.u_normalizer.to(device)
a_normalizer = datapipe.a_normalizer.to(device)
resolution = int(cfg.datapipe.data.resolution)
model = build_model(cfg.model).to(device)
checkpoint = find_checkpoint(cfg)
if checkpoint is not None:
LOGGER.info("Loading checkpoint: %s", checkpoint)
state_dict = torch.load(checkpoint, map_location=device)
model.load_state_dict(state_dict)
else:
LOGGER.warning("No checkpoint found. Running with random weights.")
model.eval()
result_dir = Path(cfg.inference.result_dir)
picture_dir = Path(cfg.inference.picture_dir)
result_dir.mkdir(parents=True, exist_ok=True)
picture_dir.mkdir(parents=True, exist_ok=True)
myloss = LpLoss(size_average=False)
analysis_record = {}
out_all = np.array([])
label_all = np.array([])
a_ori_all = np.array([])
mask_all = np.array([])
grid_all = np.array([])
with torch.no_grad():
for sample_id, data in enumerate(test_loader):
data = data.to(device)
out_indomain = model(data)
full_out = torch.zeros((resolution * resolution, 1), device=device)
full_label = torch.zeros((resolution * resolution), device=device)
full_input_a = torch.zeros((resolution * resolution, 10), device=device)
full_grid = torch.zeros((resolution * resolution, 2), device=device)
indices = data["G1"].sample_idx
full_out[indices] = out_indomain
full_label[indices] = data["G1+2"].y
full_input_a[indices, :] = data["G1"].x
full_grid[indices, :] = data["G1"].x[:, :2]
pred_decoded = u_normalizer.decode(full_out.view(1, -1))
a_decoded = a_normalizer.decode(full_input_a[:, 2].view(1, -1))
label_reshaped = full_label.view(1, -1)
cell_state_full = torch.zeros((1, resolution * resolution), device=device)
cell_state_full[0, :] = data["G1"].cell_state
l2_item = myloss(pred_decoded, label_reshaped).item()
mae_item = nn.L1Loss()(pred_decoded, label_reshaped).item()
record_data(analysis_record, [l2_item, mae_item], ["L2", "MAE"])
np.savez(
result_dir / f"sample_{sample_id:04d}.npz",
predict=to_np_array(pred_decoded).reshape(resolution, resolution),
label=to_np_array(label_reshaped).reshape(resolution, resolution),
forcing=to_np_array(a_decoded).reshape(resolution, resolution),
mask=to_np_array(cell_state_full).reshape(resolution, resolution),
)
out_all = np.append(out_all, to_np_array(pred_decoded))
label_all = np.append(label_all, to_np_array(label_reshaped))
a_ori_all = np.append(a_ori_all, to_np_array(a_decoded))
mask_all = np.append(mask_all, to_np_array(cell_state_full))
grid_all = np.append(grid_all, to_np_array(full_grid.unsqueeze(0)))
metrics = {
"mean_l2": float(np.mean(analysis_record["L2"])),
"std_l2": float(np.std(analysis_record["L2"])),
"mean_mae": float(np.mean(analysis_record["MAE"])),
"num_samples": int(len(analysis_record["L2"])),
}
metrics_path = Path(cfg.inference.metrics_path)
metrics_path.parent.mkdir(parents=True, exist_ok=True)
metrics_path.write_text(json.dumps(metrics, indent=2), encoding="utf-8")
LOGGER.info("Metrics: %s", metrics)
plot_samples = min(int(cfg.inference.num_visualize), metrics["num_samples"])
if plot_samples >= 2:
plot_path = picture_dir / "forcing_solution_comparison.png"
plot_data(
predict_term=out_all,
true_term=label_all,
forcing_term=a_ori_all,
forcing_mask=mask_all,
grid_info=grid_all,
resolution=resolution,
num_samples=plot_samples,
interpolation="bilinear",
save_path=str(plot_path),
)
LOGGER.info("Plot saved to %s", plot_path)
elif plot_samples == 1:
LOGGER.info("Skipping plot_data for one sample because the runtime helper expects a 2D axes grid.")
dist.cleanup()
if __name__ == "__main__":
main()