ace-mini-1 / inference.py
ConnorKapoor's picture
Release V003 research weights with verified standalone inference
278e605 verified
Raw History Blame Contribute Delete
2.27 kB
"""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