"""Load Ace-Mini safetensors and evaluate prepared casting-frame query points.""" import json from pathlib import Path import numpy as np import torch from safetensors.torch import load_file from stream_surrogate import StreamNet,features from warp_surrogate import WarpNet class AceMini: def __init__(self,directory=None,device='cpu'): root=Path(directory or Path(__file__).parent) self.config=json.loads((root/'config.json').read_text()) self.device=torch.device(device) self.model=StreamNet(self.config['query_dim']).to(self.device).eval() self.model.load_state_dict(load_file(str(root/'model.safetensors'),device=str(self.device))) self.warp=WarpNet(self.config['query_dim']).to(self.device).eval() self.warp.load_state_dict(load_file(str(root/'deformation.safetensors'),device=str(self.device))) @torch.inference_mode() def predict(self,spec,data,time_s=0.,deformation_phase=1.): """All geometry must already be in a frame with gravity along -Z. Returns de-normalized SI fields (defect percentages as percentage points). Deformation phase is normalized progress, not physical time. """ if not np.isfinite(time_s) or time_s<0:raise ValueError('time_s must be finite and nonnegative') if not 0<=deformation_phase<=1:raise ValueError('deformation_phase must be in [0,1]') cloud,q=features(spec,data) if not np.isfinite(cloud).all() or not np.isfinite(q).all():raise ValueError('Inputs must be finite') cloud=torch.from_numpy(cloud).to(self.device);q=torch.from_numpy(q).to(self.device) code=self.model.encode(cloud);result={} for task,stats in self.config['stats'].items(): t=0. if task in ('maps','mechanical') else float(time_s) y=self.model(cloud,q,torch.full((len(q),),t,device=self.device),task,code).cpu().numpy() result[task]=y*np.array(stats['std'])+np.array(stats['mean']) final_mm=torch.tensor(result['mechanical'][:,:3]*1000,dtype=torch.float32,device=self.device) phase=torch.full((len(q),),float(deformation_phase),device=self.device) result['deformation_m']=self.warp(q,code,phase,final_mm).cpu().numpy()/1000 return result