Download scripts/inference.py from OneScience-Group/Transolver-Car-Design: direct link, hf CLI and curl.
- Browser
- Download file 8.29 kB
-
https://huggingface.co/OneScience-Group/Transolver-Car-Design/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/Transolver-Car-Design/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/Transolver-Car-Design/resolve/main/scripts/inference.py
8.29 kB
| from __future__ import annotations | |
| import logging | |
| import os | |
| import sys | |
| import time | |
| import importlib.util | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| os.chdir(ROOT) | |
| from model import Transolver3D, Transolver3D_plus | |
| import onescience | |
| from onescience.distributed.manager import DistributedManager | |
| from onescience.utils.YParams import YParams | |
| from onescience.utils.transolver import cal_coefficient, save_prediction_to_vtk, visualize_prediction | |
| def load_shapenet_car_datapipe(): | |
| module_path = Path(onescience.__file__).resolve().parent / "datapipes/cfd/ShapeNetCar.py" | |
| spec = importlib.util.spec_from_file_location("_onescience_shapenetcar", module_path) | |
| if spec is None or spec.loader is None: | |
| raise ImportError(f"Unable to load ShapeNetCarDatapipe from {module_path}") | |
| module = importlib.util.module_from_spec(spec) | |
| spec.loader.exec_module(module) | |
| return module.ShapeNetCarDatapipe | |
| def setup_logging(rank: int) -> logging.Logger: | |
| level = logging.INFO if rank == 0 else logging.WARNING | |
| logging.basicConfig( | |
| level=level, | |
| format="%(asctime)s - %(levelname)s - %(message)s", | |
| datefmt="%Y-%m-%d %H:%M:%S", | |
| ) | |
| logging.getLogger().setLevel(level) | |
| return logging.getLogger(__name__) | |
| def build_model(model_name: str, model_params, device: torch.device) -> torch.nn.Module: | |
| model_cls = { | |
| "Transolver": Transolver3D, | |
| "Transolver_plus": Transolver3D_plus, | |
| }.get(model_name) | |
| if model_cls is None: | |
| raise NotImplementedError(f"Model {model_name} initialization not implemented.") | |
| return model_cls( | |
| n_hidden=model_params.n_hidden, | |
| n_layers=model_params.n_layers, | |
| space_dim=model_params.space_dim, | |
| fun_dim=model_params.fun_dim, | |
| n_head=model_params.n_head, | |
| mlp_ratio=model_params.mlp_ratio, | |
| out_dim=model_params.out_dim, | |
| slice_num=model_params.slice_num, | |
| unified_pos=model_params.unified_pos, | |
| ).to(device) | |
| def resolve_device(gpuid: int) -> torch.device: | |
| if torch.cuda.is_available() and int(gpuid) >= 0: | |
| return torch.device(f"cuda:{gpuid}") | |
| return torch.device("cpu") | |
| def maybe_calculate_coefficient(data_dir: Path, pred_press: np.ndarray, pred_velo: np.ndarray, gt_press: np.ndarray, gt_velo: np.ndarray): | |
| if not (data_dir / "quadpress_smpl.vtk").exists() or not (data_dir / "hexvelo_smpl.vtk").exists(): | |
| return None, None | |
| pred_coef = cal_coefficient(str(data_dir), pred_press[:, None], pred_velo) | |
| gt_coef = cal_coefficient(str(data_dir), gt_press[:, None], gt_velo) | |
| return pred_coef, gt_coef | |
| def main() -> None: | |
| DistributedManager.initialize() | |
| manager = DistributedManager() | |
| logger = setup_logging(manager.rank) | |
| if manager.rank != 0: | |
| logger.warning("Inference should run on a single process; exiting non-zero rank.") | |
| return | |
| config_file_path = str(ROOT / "conf/config.yaml") | |
| cfg = YParams(config_file_path, "model") | |
| cfg_data = YParams(config_file_path, "datapipe") | |
| cfg_train = YParams(config_file_path, "training") | |
| cfg_test = YParams(config_file_path, "inference") | |
| model_name = cfg.name | |
| model_params = cfg.specific_params[model_name] | |
| cfg_data.model_hparams = model_params | |
| device = resolve_device(cfg_test.gpuid) | |
| logger.info("Using device: %s", device) | |
| ShapeNetCarDatapipe = load_shapenet_car_datapipe() | |
| datapipe = ShapeNetCarDatapipe(params=cfg_data, distributed=False) | |
| val_dataset = datapipe.val_dataset | |
| coef_norm = datapipe.coef_norm | |
| val_names = val_dataset.data_list_names | |
| test_loader, _ = datapipe.val_dataloader() | |
| logger.info("Loaded %d validation samples.", len(val_dataset)) | |
| model = build_model(model_name, model_params, device) | |
| checkpoint_path = Path(cfg_train.checkpoint_dir) / f"{model_name}.pth" | |
| if not checkpoint_path.exists(): | |
| raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") | |
| checkpoint = torch.load(checkpoint_path, map_location=device) | |
| model.load_state_dict(checkpoint["model_state_dict"]) | |
| model.eval() | |
| result_root = Path(cfg_test.result_dir) / model_name | |
| npy_dir = result_root / "npy" | |
| vtk_dir = result_root / "vtk" | |
| vis_dir = result_root / "vis" | |
| npy_dir.mkdir(parents=True, exist_ok=True) | |
| if cfg_test.save_vtk: | |
| vtk_dir.mkdir(parents=True, exist_ok=True) | |
| if cfg_test.visualize: | |
| vis_dir.mkdir(parents=True, exist_ok=True) | |
| criterion_func = nn.MSELoss(reduction="none") | |
| l2errs_press, l2errs_velo, mses_press, mses_velo_var, times = [], [], [], [], [] | |
| gt_coef_list, pred_coef_list = [], [] | |
| mean = torch.tensor(coef_norm[2], dtype=torch.float32, device=device) | |
| std = torch.tensor(coef_norm[3], dtype=torch.float32, device=device) | |
| with torch.no_grad(): | |
| for index, data in enumerate(test_loader): | |
| if index >= len(val_names): | |
| break | |
| sample_name = val_names[index] | |
| data = data.to(device) | |
| tic = time.time() | |
| out = model(data) | |
| times.append(time.time() - tic) | |
| targets = data.y | |
| pred_press = out[data.surf, -1] * std[-1] + mean[-1] | |
| gt_press = targets[data.surf, -1] * std[-1] + mean[-1] | |
| pred_velo = out[~data.surf, :-1] * std[:-1] + mean[:-1] | |
| gt_velo = targets[~data.surf, :-1] * std[:-1] + mean[:-1] | |
| out_denorm = out * std + mean | |
| y_denorm = targets * std + mean | |
| safe_name = sample_name.replace("/", "_") | |
| np.save(npy_dir / f"{index}_{safe_name}_pred.npy", out_denorm.cpu().numpy()) | |
| np.save(npy_dir / f"{index}_{safe_name}_gt.npy", y_denorm.cpu().numpy()) | |
| data_dir = ROOT / cfg_data.source.data_dir / sample_name | |
| pred_coef, gt_coef = maybe_calculate_coefficient( | |
| data_dir, | |
| pred_press.cpu().numpy(), | |
| pred_velo.cpu().numpy(), | |
| gt_press.cpu().numpy(), | |
| gt_velo.cpu().numpy(), | |
| ) | |
| if pred_coef is not None and gt_coef is not None: | |
| pred_coef_list.append(pred_coef) | |
| gt_coef_list.append(gt_coef) | |
| l2errs_press.append((torch.norm(pred_press - gt_press) / (torch.norm(gt_press) + 1e-8)).cpu().numpy()) | |
| l2errs_velo.append((torch.norm(pred_velo - gt_velo) / (torch.norm(gt_velo) + 1e-8)).cpu().numpy()) | |
| mses_press.append(criterion_func(out[data.surf, -1], targets[data.surf, -1]).mean().cpu().numpy()) | |
| mses_velo_var.append(criterion_func(out[~data.surf, :-1], targets[~data.surf, :-1]).mean().cpu().numpy()) | |
| if cfg_test.save_vtk and (data_dir / "quadpress_smpl.vtk").exists(): | |
| save_prediction_to_vtk( | |
| out_denorm=out_denorm, | |
| targets=targets, | |
| cfd_data=data, | |
| sample_name=sample_name, | |
| output_dir=str(vtk_dir), | |
| index=index, | |
| data_dir=str(ROOT / cfg_data.source.data_dir), | |
| ) | |
| if cfg_test.visualize and cfg_test.save_vtk: | |
| visualize_prediction(output_dir=str(vtk_dir), vis_dir=str(vis_dir), index=index) | |
| logger.info("Results saved to: %s", result_root) | |
| logger.info("Relative L2 pressure: %.6f", float(np.mean(l2errs_press))) | |
| logger.info("Relative L2 velocity: %.6f", float(np.mean(l2errs_velo))) | |
| logger.info("RMSE pressure: %.6f", float(np.sqrt(np.mean(mses_press)) * coef_norm[3][-1])) | |
| rmse_velo = np.sqrt(np.mean(mses_velo_var, axis=0)) * coef_norm[3][:-1] | |
| logger.info("Combined velocity RMSE: %.6f", float(np.sqrt(np.mean(np.square(rmse_velo))))) | |
| logger.info("Mean inference time (s): %.6f", float(np.mean(times))) | |
| if gt_coef_list: | |
| coef_error = np.mean(np.abs(np.array(pred_coef_list) - np.array(gt_coef_list)) / (np.array(gt_coef_list) + 1e-8)) | |
| logger.info("Mean relative CD error: %.6f", float(coef_error)) | |
| else: | |
| logger.info("Skipped drag coefficient metrics because VTK geometry files were not present.") | |
| if __name__ == "__main__": | |
| main() | |