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)