litert-models / conversion /dynamic_hf_worker.py
unicorn who dev
Complete FireViewer Android learning SDK, source and verification records
dee7f43 verified
Raw History Blame Contribute Delete
6.04 kB
import os,sys,re,json,shutil
from pathlib import Path
os.environ.update(VDS_BATCH2_ROOT='/workspace/vds-litert-batch2',HF_TOKEN_FILE='/workspace/.secrets/hf_token',HF_FIREVIEWER_DEST_REPO='fireviewer/litert-models',TORCH_NUM_THREADS='8',HF_HOME='/workspace/cache/huggingface')
import torch,numpy as np,litert_torch
from ai_edge_litert.interpreter import Interpreter,OpResolverType
from huggingface_hub import HfApi
import dynamic_lowerings
root=Path('/workspace/vds-litert-batch2');sys.path.insert(0,str(root/'package/scripts'))
import convert_batch2 as worker
worker.WORK=root/'dynamic-work'
key=sys.argv[1]
spec=next((s for s in worker.all_specs() if s['key']==key),None)
if spec is None:
import previous_native
old=next(s for s in json.loads((root/'previous-models.json').read_text()) if s['key']==key)
spec={**old,'group':'general_models','route':key,'title':key,'task':'vision embeddings' if key=='dinov2' else 'pose estimation'}
worker.CONV[key]=lambda s,e,o:getattr(previous_native,key)(s,e,o,worker)
original_prefetched=worker.prefetched
worker.prefetched=lambda k:json.loads((root/'previous-prefetch.json').read_text())[k] if k==key else original_prefetched(k)
validation={}
adapter_sha256=worker.sha256(root/'dynamic_lowerings.py')
def dynamic_export(module,args,out):
global validation
module=module.cpu().eval();args=tuple(a.cpu() for a in args)
for sub in module.modules():
if hasattr(sub,'set_attn_implementation'):
try:sub.set_attn_implementation('eager')
except (ValueError,AttributeError):pass
h,w=args[0].shape[-2:]
stride=14 if spec['route'] in {'depth','dinov2'} else 16 if spec['route']=='dinov3_fireviewer' else 32
if h%stride or w%stride:stride=1
height=stride*torch.export.Dim('height_blocks',min=4,max=128) if stride>1 else torch.export.Dim('height',min=128,max=2048)
width=stride*torch.export.Dim('width_blocks',min=4,max=128) if stride>1 else torch.export.Dim('width',min=128,max=2048)
dims=[{} for a in args];dims[0]={2:height,3:width}
masks=[]
for i,a in enumerate(args[1:],1):
if a.ndim==3 and tuple(a.shape[-2:])==(h,w):dims[i]={1:height,2:width};masks.append(i)
shapes=[(h,w),(h-2*stride,w+stride),(h+stride,w-2*stride)]
def inputs(hh,ww):
aa=list(args);aa[0]=torch.from_numpy(np.random.default_rng(42).normal(size=(1,3,hh,ww)).astype(np.float32))
for i in masks:aa[i]=torch.ones((1,hh,ww),dtype=args[i].dtype)
return tuple(aa)
with torch.inference_mode():
for hh,ww in shapes:module(*inputs(hh,ww))
if os.environ.get('REUSE_ARTIFACT'):
shutil.copy2(os.environ['REUSE_ARTIFACT'],out)
else:
converted=litert_torch.convert(module,args,dynamic_shapes=tuple(dims))
out.parent.mkdir(parents=True,exist_ok=True);converted.export(str(out));del converted
interp=Interpreter(model_path=str(out),num_threads=8,experimental_op_resolver_type=OpResolverType.BUILTIN_WITHOUT_DEFAULT_DELEGATES)
details=interp.get_input_details();mapping={}
for j,d in enumerate(details):
match=re.search(r'args_(\d+)',d['name']);index=int(match.group(1)) if match else j
mapping[index]=d
if mapping[0]['shape_signature'].tolist()[2:]!=[-1,-1]:raise RuntimeError('Spatial signature was fixed by converter')
checks=[]
for hh,ww in shapes:
aa=inputs(hh,ww)
for i,a in enumerate(aa):interp.resize_tensor_input(mapping[i]['index'],tuple(a.shape),strict=True)
interp.allocate_tensors()
for i,a in enumerate(aa):interp.set_tensor(mapping[i]['index'],a.numpy())
interp.invoke()
with torch.inference_mode():ref=module(*aa)
if isinstance(ref,torch.Tensor):ref=(ref,)
refs=[v.detach().numpy() for v in ref]
errors=[]
for j,d in enumerate(interp.get_output_details()):
a=interp.get_tensor(d['index']);candidates=[r for r in refs if r.shape==a.shape]
if len(candidates)==1:r=candidates[0]
else:
match=re.search(r'output_(\d+)',d['name'])
if not match:raise RuntimeError(f'Ambiguous output mapping for {d["name"]} {a.shape}')
r=refs[int(match.group(1))]
if r.shape!=a.shape:raise RuntimeError('Output signature mismatch');np.testing.assert_allclose(a,r,rtol=1e-3,atol=1e-3)
if not np.isfinite(a).all():raise RuntimeError('Nonfinite output')
errors.append(float(np.max(np.abs(a-r))))
checks.append({'image_shape':[1,3,hh,ww],'max_abs_errors':errors})
print('DYNAMIC_PARITY_OK',key,hh,ww,errors,flush=True)
validation={'dynamic_validated':True,'checks':checks,'atol':1e-3,'rtol':1e-3,'alignment':stride,'source_resolution_resampled':False,'runtime_requirement':'LiteRT builtin CPU; disable default XNNPACK delegates for these dynamic graphs','android_tested':False,'accuracy_evaluated':False,'conversion_adapter_sha256':adapter_sha256}
return {**worker.inspect_tflite(out),**validation}
base_publish=worker.publish
def publish(api,s,e,out,arts,report):
if not validation.get('dynamic_validated'):raise RuntimeError('Dynamic validation gate did not pass')
report.update(validation)
return base_publish(api,s,e,out,arts,report)
worker.direct_export=dynamic_export;worker.publish=publish
if spec['route']=='rfdetr_medium':
from rfdetr.export._litert.exporter import LiteRTExporter
def patched_convert(self,litert,wrapped_model,input_tensors,output_file):
dynamic_export(wrapped_model,(input_tensors,),Path(output_file))
LiteRTExporter._convert_and_save=patched_convert
if key=='fireviewer_rfdetr_medium_v110' and os.environ.get('RF_NATIVE')=='1':
import rfdetr_native
worker.CONV['rfdetr_medium']=lambda s,e,o:rfdetr_native.convert(s,e,o,worker)
if key=='fireviewer_dfine_m_strict_v1':
import dfine_native
worker.CONV['onnx_existing']=lambda s,e,o:dfine_native.convert(s,e,o,worker)
worker.process(HfApi(token=worker.token()),spec)