patdev's picture
Fix TensorRT skin decoder runtime dtype bindings
1a4cfbd verified
Raw
History Blame Contribute Delete
6.6 kB
from __future__ import annotations
import torch
try:
from .trt_torch import TorchTensorRTEngine
except ImportError:
from trt_torch import TorchTensorRTEngine
class SSFlowTRT(torch.nn.Module):
resolution=16; in_channels=8; in_channels_skl=4; out_channels=8; out_channels_skl=4
z_is_global=False; z_skl_is_global=False; global_token_num=1024; global_token_num_skl=1024
def __init__(self,plan): super().__init__(); self.engine=TorchTensorRTEngine(plan); self.device=torch.device('cuda'); self._dummy=torch.nn.Parameter(torch.empty(0,device='cuda'),requires_grad=False)
def forward(self,x,x_skl,t,cond,**kwargs):
o=self.engine.run({'x':x,'x_skl':x_skl,'timestep':t,'cond':cond});return o['velocity'],o['velocity_skl']
def to(self,*a,**k):return self
def eval(self):return self
class SSDecoderTRT(torch.nn.Module):
def __init__(self,plan):super().__init__();self.engine=TorchTensorRTEngine(plan);self.device=torch.device('cuda');self._dummy=torch.nn.Parameter(torch.empty(0,device='cuda'),requires_grad=False)
def forward(self,z,z_skl):
o=self.engine.run({'z':z,'z_skl':z_skl});return o['occupancy'],o['occupancy_skl']
def to(self,*a,**k):return self
def eval(self):return self
class SkinDecoderTRT(torch.nn.Module):
skin_feat_channels=4
def __init__(self,plan):
super().__init__();self.engine=TorchTensorRTEngine(plan);self.device=torch.device('cuda');self._dummy=torch.nn.Parameter(torch.empty(0,device='cuda'),requires_grad=False)
def _cast_for_binding(self,name,t):
dt=self.engine._torch_dtype(self.engine.engine.get_tensor_dtype(name))
return t.to(device='cuda',dtype=dt).contiguous()
def forward(self,vertex_features,joint_features,parents):
e=self.engine;ctx=e.context
vf=self._cast_for_binding('vertex_features',vertex_features)
jf=self._cast_for_binding('joint_features',joint_features)
pp=self._cast_for_binding('parents',parents)
feeds={'vertex_features':vf,'joint_features':jf,'parents':pp}
for n,t in feeds.items():
ctx.set_input_shape(n,tuple(t.shape));ctx.set_tensor_address(n,int(t.data_ptr()))
if hasattr(ctx,'infer_shapes'):
try:ctx.infer_shapes()
except Exception:pass
out_name='skin_weights';dt=e._torch_dtype(e.engine.get_tensor_dtype(out_name))
shape=tuple(ctx.get_tensor_shape(out_name))
if any(int(x)<0 for x in shape):shape=(vf.shape[0],vf.shape[1],jf.shape[1])
y=torch.empty(shape,device='cuda',dtype=dt);ctx.set_tensor_address(out_name,int(y.data_ptr()))
ok=ctx.execute_async_v3(stream_handle=torch.cuda.current_stream().cuda_stream)
if not ok:raise RuntimeError(f'TensorRT skin decoder execute failed; vf={tuple(vf.shape)}/{vf.dtype} jf={tuple(jf.shape)}/{jf.dtype} parents={tuple(pp.shape)}/{pp.dtype} out={shape}/{dt}')
return y
def to(self,*a,**k):return self
def eval(self):return self
class SLatFlowCoreTRTAdapter(torch.nn.Module):
def __init__(self,shell_model,plan,max_geo=16384,max_skl=8192,fallback_factory=None):
super().__init__();self.native=shell_model;self.engine=TorchTensorRTEngine(plan);self.max_geo=max_geo;self.max_skl=max_skl;self.fallback_factory=fallback_factory;self._fallback_model=None
for name in ['resolution','in_channels','in_channels_vert_skin','in_channels_skl','out_channels','out_channels_vert_skin','out_channels_skl','dtype','pe_mode','share_mod','z_is_global','z_skl_is_global','global_token_num','global_token_num_skl','predict_x0','t_scale','t_eps','use_joint_num_cond']:
if hasattr(shell_model,name):setattr(self,name,getattr(shell_model,name))
@property
def device(self):return self.native.device
def _fallback(self,x,x_skl,t,cond,**kwargs):
if self._fallback_model is None:
if self.fallback_factory is None:raise RuntimeError('SLat TRT outside profile and no fallback')
self._fallback_model=self.fallback_factory()
return self._fallback_model(x,x_skl,t,cond,**kwargs)
def forward(self,x,x_skl,t,cond,joints_num=None,**kwargs):
m=self.native;n_geo=int(x.feats.shape[0]);n_skl=int(x_skl.feats.shape[0])
if x.shape[0]!=1 or x_skl.shape[0]!=1 or n_geo>self.max_geo or n_skl>self.max_skl:return self._fallback(x,x_skl,t,cond,joints_num=joints_num,**kwargs)
cond=cond.type(m.dtype);feats,feats_skin=x.feats[:,:m.in_channels],x.feats[:,m.in_channels:];x_geo,x_skin=x.replace(feats),x.replace(feats_skin)
if getattr(m,'predict_x0',False):xt_skin,xt_skl=feats_skin.clone(),x_skl.feats.clone()
joint_vs=joint_skl=None
if getattr(m,'use_joint_num_cond',False):
joint_vs,joint_skl=m._get_joint_num_emb(joints_num,x.shape[0],x.device);joint_vs=joint_vs.type(m.dtype);joint_skl=joint_skl.type(m.dtype)
d={'':x_geo,'_vert_skin':x_skin,'_skl':x_skl};ce={'':None,'_vert_skin':joint_vs,'_skl':joint_skl}
for postfix in ['', '_vert_skin','_skl']:
d[postfix],d[f't_emb{postfix}'],d[f't_mod{postfix}'],d[f'skips{postfix}']=m.forward_stage(d[postfix],t,postfix,stage='in',cond_emb=ce[postfix])
try:
o=self.engine.run({'geo':d[''].feats.unsqueeze(0).contiguous(),'skin':d['_vert_skin'].feats.unsqueeze(0).contiguous(),'skl':d['_skl'].feats.unsqueeze(0).contiguous(),'mod_geo':d['t_mod'].contiguous(),'mod_skin':d['t_mod_vert_skin'].contiguous(),'mod_skl':d['t_mod_skl'].contiguous(),'cond':cond.contiguous()})
d['']=d[''].replace(o['geo_out'][0]);d['_vert_skin']=d['_vert_skin'].replace(o['skin_out'][0]);d['_skl']=d['_skl'].replace(o['skl_out'][0])
except Exception as exc:
print(f'[companion-runtime] SLat TRT execution fallback: {exc}',flush=True);return self._fallback(x,x_skl,t,cond,joints_num=joints_num,**kwargs)
for postfix in ['', '_vert_skin','_skl']:
d[postfix]=m.forward_stage(d[postfix],t,postfix,stage='out',t_emb=d[f't_emb{postfix}'],skips=d[f'skips{postfix}'],original_dtype=x.dtype)
if getattr(m,'predict_x0',False):
tn=t/m.t_scale;factor=(1/tn.clamp_min(m.t_eps))[:,None]
d['_vert_skin']=d['_vert_skin'].replace((d['_vert_skin'].feats-xt_skin)*factor[d['_vert_skin'].coords[:,0]])
d['_skl']=d['_skl'].replace((d['_skl'].feats-xt_skl)*factor[d['_skl'].coords[:,0]])
return x.replace(torch.cat([d[''].feats,d['_vert_skin'].feats],dim=1)),x_skl.replace(d['_skl'].feats)
def to(self,*a,**k):self.native.to(*a,**k);return self
def eval(self):self.native.eval();return self