Download inference.py from DigitalMetal/ace-mini-1: direct link, hf CLI and curl.
- Browser
- Download file 2.27 kB
-
https://huggingface.co/DigitalMetal/ace-mini-1/resolve/main/inference.py
- Command line
-
hf download hf://DigitalMetal/ace-mini-1/inference.py
-
curl -L -o inference.py https://huggingface.co/DigitalMetal/ace-mini-1/resolve/main/inference.py
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))) | |
| 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 | |