| |
| from __future__ import annotations |
| import argparse, fcntl, gc, hashlib, importlib.util, json, os, shutil, subprocess, sys, tempfile, time, traceback |
| from pathlib import Path |
| from typing import Any |
|
|
| import numpy as np |
| import torch |
| import torch.nn as nn |
| from huggingface_hub import HfApi, snapshot_download |
|
|
| ROOT = Path(os.environ.get('VDS_BATCH2_ROOT','/workspace/vds-litert-batch2')) |
| PKG = Path(__file__).resolve().parents[1] |
| MANIFEST_PATH = PKG / 'model_manifest.json' |
| TOKEN_FILE = Path(os.environ.get('HF_TOKEN_FILE','/workspace/HF.txt')) |
| DEST_REPO = os.environ.get('HF_DEST_REPO','Charlbi/Lite_rt_prepared_for_android_dataset_builder') |
| FIREVIEWER_DEST_REPO = os.environ.get('HF_FIREVIEWER_DEST_REPO','fireviewer/litert-models') |
| PREFETCH = ROOT/'models'/'upstream' |
| WORK = ROOT/'work' |
| ALLOWED = {'apache-2.0','mit','bsd-2-clause','bsd-3-clause'} |
| THREADS = max(1,int(os.environ.get('TORCH_NUM_THREADS','8'))) |
| torch.set_num_threads(THREADS) |
|
|
|
|
| def read_json(p: Path): return json.loads(p.read_text(encoding='utf-8')) |
| def write_json(p: Path, x: Any): p.parent.mkdir(parents=True,exist_ok=True); p.write_text(json.dumps(x,indent=2,sort_keys=True,default=str)+'\n',encoding='utf-8') |
| def sha256(p: Path): |
| h=hashlib.sha256() |
| with p.open('rb') as f: |
| for c in iter(lambda:f.read(8*1024*1024),b''): h.update(c) |
| return h.hexdigest() |
| def token(): |
| if not TOKEN_FILE.is_file(): raise RuntimeError(f'missing {TOKEN_FILE}') |
| v=TOKEN_FILE.read_text(encoding='utf-8-sig').strip().splitlines()[0].strip() |
| if len(v)<10: raise RuntimeError('HF token file is empty') |
| return v |
|
|
| def run(cmd, cwd=None, env=None): |
| print('+',' '.join(map(str,cmd)),flush=True) |
| subprocess.run(list(map(str,cmd)),cwd=cwd,env=env,check=True) |
|
|
| def inspect_tflite(p: Path): |
| from ai_edge_litert.interpreter import Interpreter |
| with p.open('rb') as f: |
| if f.read(8)[4:8] != b'TFL3': raise RuntimeError(f'bad flatbuffer: {p}') |
| i=Interpreter(model_path=str(p),num_threads=THREADS) |
| if not any(-1 in d.get('shape_signature',[]) for d in i.get_input_details()): i.allocate_tensors() |
| def d(t): return {'name':str(t.get('name','')),'shape':[int(x) for x in t['shape']],'shape_signature':[int(x) for x in t.get('shape_signature',t['shape'])],'dtype':str(np.dtype(t['dtype']))} |
| return {'inputs':[d(x) for x in i.get_input_details()],'outputs':[d(x) for x in i.get_output_details()]} |
|
|
| def direct_export(module: nn.Module, args: tuple[torch.Tensor,...], out: Path): |
| import litert_torch |
| out.parent.mkdir(parents=True,exist_ok=True) |
| module=module.cpu().eval(); args=tuple(x.detach().cpu() for x in args) |
| with torch.inference_mode(): edge=litert_torch.convert(module,args) |
| tmp=out.with_suffix('.partial.tflite'); edge.export(str(tmp)); del edge |
| check=inspect_tflite(tmp); tmp.replace(out); return check |
|
|
| def onnx_to_tflite(onnx: Path, out: Path, spatial_size: int=640): |
| temp=out.parent/(out.stem+'_onnx2tf') |
| shutil.rmtree(temp,ignore_errors=True); temp.mkdir(parents=True,exist_ok=True) |
| run([ROOT/'.venv-onnx/bin/python',PKG/'scripts/convert_onnx.py',str(onnx),str(temp),str(spatial_size)]) |
| cands=sorted(temp.rglob('*float32.tflite')) or sorted(temp.rglob('*.tflite'), key=lambda p:p.stat().st_size, reverse=True) |
| if not cands: raise RuntimeError('onnx2tf produced no tflite') |
| shutil.copy2(cands[0],out); return inspect_tflite(out) |
|
|
| def all_specs(): |
| m=read_json(MANIFEST_PATH) |
| out=[] |
| for group in ('general_models','fireviewer_models'): |
| for s in m[group]: |
| if s['key']=='fireviewer_dfine_m_strict_v1': s={**s,'route':'onnx_existing'} |
| out.append({**s,'group':group}) |
| return out |
|
|
| def resolve_hf(api,s): |
| info=api.model_info(s['repo']) |
| lic=None |
| card=getattr(info,'card_data',None) |
| if card is not None: |
| try: lic=card.get('license') |
| except Exception: lic=getattr(card,'license',None) |
| tags=getattr(info,'tags',[]) or [] |
| if not lic: |
| for x in tags: |
| if str(x).startswith('license:'): lic=str(x).split(':',1)[1] |
| if s['group']=='general_models' and str(lic or s['license']).lower() not in ALLOWED: |
| raise RuntimeError(f"non permissive license for public catalog: {lic}") |
| return {'repo_id':s['repo'],'sha':info.sha,'license':str(lic or s['license'])} |
|
|
| def prefetch(api): |
| ROOT.mkdir(parents=True,exist_ok=True); PREFETCH.mkdir(parents=True,exist_ok=True) |
| state={'created':time.time(),'destination':DEST_REPO,'models':{}} |
| for s in all_specs(): |
| k=s['key']; print(f'[prefetch] {k}',flush=True) |
| try: |
| if s['route']=='rtmdet': |
| ext=PREFETCH/k; ext.mkdir(parents=True,exist_ok=True) |
| ck=ext/'model.pth' |
| if not ck.exists(): run(['curl','-L','--fail','--retry','4','-o',str(ck),s['checkpoint_url']]) |
| state['models'][k]={'local_dir':str(ext),'checkpoint':str(ck),'sha256':sha256(ck),'license':s['license'],'external_url':s['checkpoint_url']} |
| continue |
| up=resolve_hf(api,s); loc=PREFETCH/k |
| snapshot_download(repo_id=s['repo'],revision=up['sha'],local_dir=str(loc),token=token(),ignore_patterns=['*.ckpt','last_ema.pth','checkpoint_best_ema.pth','checkpoints/*']) |
| state['models'][k]={**up,'local_dir':str(loc)} |
| if s['route']=='dinov3_fireviewer': |
| base=PREFETCH/(k+'_base') |
| snapshot_download(repo_id=s['base_repo'],revision=s['base_revision'],local_dir=str(base),token=token()) |
| state['models'][k]['base_local_dir']=str(base) |
| except Exception as e: |
| state['models'][k]={'error':f'{type(e).__name__}: {e}'} |
| print(' !!',state['models'][k]['error'],flush=True) |
| write_json(ROOT/'prefetch_manifest.json',state); return state |
|
|
| def prefetched(k): |
| p=ROOT/'prefetch_manifest.json' |
| if not p.exists(): raise RuntimeError('run setup/prefetch first') |
| e=read_json(p)['models'].get(k,{}) |
| if e.get('error'): raise RuntimeError(e['error']) |
| return e |
|
|
| class TimmWrap(nn.Module): |
| def __init__(self,m): super().__init__(); self.m=m |
| def forward(self,x): |
| f=self.m.forward_features(x); emb=self.m.forward_head(f,pre_logits=True); logits=self.m.forward_head(f) |
| return logits,emb |
| class ImageODWrap(nn.Module): |
| def __init__(self,m): super().__init__(); self.m=m |
| def forward(self,pixel_values,pixel_mask): |
| o=self.m(pixel_values=pixel_values,pixel_mask=pixel_mask,return_dict=True) |
| return o.logits,o.pred_boxes |
| class GroundWrap(nn.Module): |
| def __init__(self,m): super().__init__(); self.m=m |
| def forward(self,pixel_values,pixel_mask,input_ids,attention_mask,token_type_ids): |
| o=self.m(pixel_values=pixel_values,pixel_mask=pixel_mask,input_ids=input_ids.long(),attention_mask=attention_mask.long(),token_type_ids=token_type_ids.long(),return_dict=True) |
| return o.logits,o.pred_boxes |
| class OwlWrap(nn.Module): |
| def __init__(self,m): super().__init__(); self.m=m |
| def forward(self,pixel_values,input_ids,attention_mask): |
| o=self.m(pixel_values=pixel_values,input_ids=input_ids.long(),attention_mask=attention_mask.long(),return_dict=True) |
| return o.logits,o.pred_boxes |
| class DepthWrap(nn.Module): |
| def __init__(self,m): super().__init__(); self.m=m |
| def forward(self,pixel_values): return self.m(pixel_values=pixel_values,return_dict=True).predicted_depth |
| class SegWrap(nn.Module): |
| def __init__(self,m): super().__init__(); self.m=m |
| def forward(self,pixel_values): return self.m(pixel_values=pixel_values,return_dict=True).logits |
| class DictOutWrap(nn.Module): |
| def __init__(self,m): super().__init__(); self.m=m |
| def forward(self,x): |
| o=self.m(x) |
| if isinstance(o,dict): return tuple(o[k] for k in sorted(o)) |
| if isinstance(o,(tuple,list)): return tuple(o) |
| return o |
|
|
| def processor_support_files(local: Path): |
| names={'preprocessor_config.json','processor_config.json','tokenizer.json','tokenizer_config.json','special_tokens_map.json','vocab.json','merges.txt','vocab.txt','sentencepiece.bpe.model','spiece.model','config.json'} |
| return [p for p in local.iterdir() if p.is_file() and p.name in names] |
|
|
| def convert_timm(s,e,out): |
| import timm |
| local=Path(e['local_dir']); cfg=read_json(local/'config.json'); arch=cfg.get('architecture') or cfg.get('model_name') |
| weights=next(iter(local.glob('*.safetensors')),None) or next(iter(local.glob('pytorch_model.bin')),None) |
| if not arch or not weights: raise RuntimeError('missing timm architecture/weights') |
| m=timm.create_model(arch,pretrained=False,checkpoint_path=str(weights)).eval() |
| inp=cfg.get('pretrained_cfg',{}).get('input_size',[3,224,224]); x=torch.zeros(1,int(inp[0]),int(inp[1]),int(inp[2])) |
| p=out/'model.tflite'; chk=direct_export(TimmWrap(m),(x,),p) |
| return [p],{'route':'timm','structural':chk,'outputs':['classification_logits','embedding']} |
|
|
| def convert_object_detection(s,e,out): |
| from transformers import AutoImageProcessor,AutoModelForObjectDetection |
| local=e['local_dir']; proc=AutoImageProcessor.from_pretrained(local,local_files_only=True,trust_remote_code=True) |
| m=AutoModelForObjectDetection.from_pretrained(local,local_files_only=True,trust_remote_code=True).eval() |
| |
| img=np.zeros((640,640,3),dtype=np.uint8); d=proc(images=img,return_tensors='pt') |
| pv=d['pixel_values'].float(); pm=d.get('pixel_mask',torch.ones((pv.shape[0],pv.shape[-2],pv.shape[-1]),dtype=torch.int64)).to(torch.int32) |
| p=out/'model.tflite'; chk=direct_export(ImageODWrap(m),(pv,pm),p) |
| return [p],{'route':'transformers_object_detection','structural':chk,'outputs':['logits','pred_boxes']} |
|
|
| def convert_grounding(s,e,out): |
| from transformers import AutoProcessor,AutoModelForZeroShotObjectDetection |
| local=e['local_dir']; proc=AutoProcessor.from_pretrained(local,local_files_only=True); m=AutoModelForZeroShotObjectDetection.from_pretrained(local,local_files_only=True).eval() |
| d=proc(images=np.zeros((800,800,3),dtype=np.uint8),text='person. car. animal.',padding='max_length',max_length=32,truncation=True,return_tensors='pt') |
| pv=d['pixel_values'].float(); pm=d.get('pixel_mask',torch.ones((1,pv.shape[-2],pv.shape[-1]),dtype=torch.int32)).to(torch.int32) |
| ids=d['input_ids'].to(torch.int32); att=d['attention_mask'].to(torch.int32); tt=d.get('token_type_ids',torch.zeros_like(ids)).to(torch.int32) |
| p=out/'model.tflite'; chk=direct_export(GroundWrap(m),(pv,pm,ids,att,tt),p) |
| return [p],{'route':'grounding_dino','structural':chk,'outputs':['logits','pred_boxes'],'text_max_length':32} |
|
|
| def convert_owl(s,e,out): |
| from transformers import AutoProcessor,AutoModelForZeroShotObjectDetection |
| local=e['local_dir']; proc=AutoProcessor.from_pretrained(local,local_files_only=True); m=AutoModelForZeroShotObjectDetection.from_pretrained(local,local_files_only=True).eval() |
| d=proc(images=np.zeros((960,960,3),dtype=np.uint8),text=[['person','car','animal']],padding='max_length',max_length=16,truncation=True,return_tensors='pt') |
| p=out/'model.tflite'; chk=direct_export(OwlWrap(m),(d['pixel_values'].float(),d['input_ids'].to(torch.int32),d['attention_mask'].to(torch.int32)),p) |
| return [p],{'route':'owlv2','structural':chk,'outputs':['logits','pred_boxes'],'text_max_length':16} |
|
|
| def convert_depth(s,e,out): |
| from transformers import AutoImageProcessor,AutoModelForDepthEstimation |
| local=e['local_dir']; proc=AutoImageProcessor.from_pretrained(local,local_files_only=True); m=AutoModelForDepthEstimation.from_pretrained(local,local_files_only=True).eval() |
| d=proc(images=np.zeros((518,518,3),dtype=np.uint8),return_tensors='pt'); p=out/'model.tflite'; chk=direct_export(DepthWrap(m),(d['pixel_values'].float(),),p) |
| return [p],{'route':'depth','structural':chk,'outputs':['predicted_depth']} |
|
|
| def convert_seg(s,e,out): |
| from transformers import AutoImageProcessor,AutoModelForSemanticSegmentation |
| local=e['local_dir']; proc=AutoImageProcessor.from_pretrained(local,local_files_only=True); m=AutoModelForSemanticSegmentation.from_pretrained(local,local_files_only=True).eval() |
| d=proc(images=np.zeros((512,512,3),dtype=np.uint8),return_tensors='pt'); p=out/'model.tflite'; chk=direct_export(SegWrap(m),(d['pixel_values'].float(),),p) |
| return [p],{'route':'segmentation','structural':chk,'outputs':['segmentation_logits']} |
|
|
| def convert_onnx_existing(s,e,out): |
| local=Path(e['local_dir']); onnx=next(iter(local.rglob('*.onnx')),None) |
| if onnx is None: raise RuntimeError('no ONNX artifact upstream') |
| size=960 if s['key']=='fireviewer_yolo11m_strict_v1' else 704 |
| p=out/'model.tflite'; chk=onnx_to_tflite(onnx,p,spatial_size=size); return [p],{'route':'onnx2tf','structural':chk,'source_onnx':onnx.name,'source_sha256':sha256(onnx),'image_size':size} |
|
|
| def convert_picodet(s,e,out): |
| local=Path(e['local_dir']); model=local/'inference.json'; params=local/'inference.pdiparams' |
| if not model.exists() or not params.exists(): raise RuntimeError('Paddle inference.json/pdiparams missing') |
| onnx=out/'model.onnx' |
| run([ROOT/'.venv-onnx/bin/paddle2onnx','--model_dir',str(local),'--model_filename','inference.json','--params_filename','inference.pdiparams','--save_file',str(onnx),'--opset_version','11','--enable_onnx_checker','True']) |
| p=out/'model.tflite'; chk=onnx_to_tflite(onnx,p); return [p],{'route':'paddle2onnx+onnx2tf','structural':chk} |
|
|
| def convert_rtmdet(s,e,out): |
| repo=ROOT/'vendor'/'mmyolo'; ck=Path(e['checkpoint']); cfg=repo/s['config']; od=out/'onnx'; od.mkdir(parents=True,exist_ok=True) |
| run([ROOT/'.venv-mm/bin/python',str(repo/'projects/easydeploy/tools/export.py'),str(cfg),str(ck),'--work-dir',str(od),'--img-size','640','640','--batch','1','--device','cpu','--opset','11','--model-only'],cwd=repo) |
| ons=sorted(od.rglob('*.onnx'),key=lambda p:p.stat().st_size,reverse=True) |
| if not ons: raise RuntimeError('MMYOLO EasyDeploy produced no ONNX') |
| p=out/'model.tflite'; chk=onnx_to_tflite(ons[0],p); return [p],{'route':'mmyolo-easydeploy+onnx2tf','structural':chk,'postprocess':'external RTMDet decode/NMS'} |
|
|
| def convert_rfdetr_medium(s,e,out): |
| from rfdetr import RFDETRMedium |
| local=Path(e['local_dir']); cands=list(local.rglob('checkpoint_best_total.pth')) or list(local.rglob('*.pth')) |
| if not cands: raise RuntimeError('RF-DETR checkpoint missing') |
| training=read_json(local/'training_config.json') |
| m=RFDETRMedium(pretrain_weights=str(cands[0]),num_classes=int(training['num_classes']),resolution=int(training['model_config']['resolution']),device='cpu') |
| res=m.export(format='litert',output_dir=str(out),quantization=None) |
| p=Path(res) if isinstance(res,(str,Path)) else None |
| if p is None or not p.is_file(): |
| cs=list(out.rglob('*.tflite')); p=cs[0] if len(cs)==1 else None |
| if p is None: raise RuntimeError('RF-DETR produced no unique LiteRT file') |
| final=out/'model.tflite'; |
| if p.resolve()!=final.resolve(): shutil.copy2(p,final) |
| return [final],{'route':'rfdetr-direct-litert','structural':inspect_tflite(final)} |
|
|
| def convert_dinov3_fv(s,e,out): |
| local=Path(e['local_dir']); adapter=local/'dinov3_adapter.py'; weights=local/'model.safetensors' |
| if not adapter.exists() or not weights.exists(): raise RuntimeError('FireViewer DINOv3 adapter/safetensors missing') |
| spec=importlib.util.spec_from_file_location('fv_dinov3_adapter',adapter); mod=importlib.util.module_from_spec(spec); assert spec and spec.loader |
| sys.modules[spec.name]=mod; spec.loader.exec_module(mod) |
| cls=getattr(mod,'DinoV3MultiTaskModel'); base=e.get('base_local_dir') or s['base_repo'] |
| |
| import safetensors.torch as st |
| original_load=st.load_file |
| def load_compatible(*args,**kwargs): |
| state=original_load(*args,**kwargs) |
| return {(k.replace('backbone.layer.','backbone.model.layer.',1) if k.startswith('backbone.layer.') else k):v for k,v in state.items()} |
| st.load_file=load_compatible |
| try:obj=cls.from_safetensors(str(weights),model_id=str(base),revision=None,image_size=448,token=False) |
| finally:st.load_file=original_load |
| net=getattr(obj,'network',obj).eval(); p=out/'model.tflite'; chk=direct_export(DictOutWrap(net),(torch.zeros(1,3,448,448),),p) |
| return [p],{'route':'fireviewer-dinov3-adapter','structural':chk,'warning':'pilot; no public-release approval inherited'} |
|
|
| CONV={'timm':convert_timm,'object_detection':convert_object_detection,'grounding_dino':convert_grounding,'owlv2':convert_owl,'depth':convert_depth,'segmentation':convert_seg,'onnx_existing':convert_onnx_existing,'picodet':convert_picodet,'rtmdet':convert_rtmdet,'rfdetr_medium':convert_rfdetr_medium,'dinov3_fireviewer':convert_dinov3_fv} |
|
|
| def support_files(s,e): |
| if not e.get('local_dir'): return [] |
| local=Path(e['local_dir']) |
| notices=[p for p in local.iterdir() if p.is_file() and (p.name.lower().startswith(('license','copying','notice')) or p.name in {'CHECKPOINT_LICENSE.md','CHECKPOINT_RIGHTS.md','CITATION.md','COMMERCIAL_USE.md','MODEL_RIGHTS_POLICY.md','RIGHTS_AND_ATTRIBUTION.md','THIRD_PARTY_NOTICES.md','labels.json'})] |
| return list(dict.fromkeys(processor_support_files(local)+notices)) |
|
|
| def publish(api,s,e,out,arts,report): |
| group='models' |
| dest=DEST_REPO if s['group']=='general_models' else FIREVIEWER_DEST_REPO |
| if s['group']=='fireviewer_models' and not s.get('publish_default',False) and os.environ.get('PUBLISH_RESTRICTED_DINOV3','0')!='1': |
| return 'converted_local_only_rights_gate' |
| stage=out/'publish'; shutil.rmtree(stage,ignore_errors=True); stage.mkdir(parents=True) |
| for a in arts: shutil.copy2(a,stage/a.name) |
| for sf in support_files(s,e): shutil.copy2(sf,stage/sf.name) |
| source_readme=Path(e['local_dir'])/'README.md' |
| if source_readme.is_file(): shutil.copy2(source_readme,stage/'UPSTREAM_README.md') |
| contract={'schema':1,'key':s['key'],'title':s['title'],'task':s['task'],'route':report.get('route',s['route']),'upstream':s.get('repo') or s.get('checkpoint_url'),'upstream_revision':e.get('sha'),'license':e.get('license',s['license']),'tflite':[a.name for a in arts],'structural_only':not report.get('dynamic_validated',False),'benchmark':False,'report':report} |
| write_json(stage/'runtime_contract.json',contract) |
| write_json(stage/'artifact_manifest.json',{p.name:{'sha256':sha256(p),'bytes':p.stat().st_size} for p in stage.iterdir() if p.is_file()}) |
| (stage/'README.md').write_text(f"---\nlicense: {str(e.get('license',s['license'])).lower()}\ntags:\n- litert\n- tflite\n- android\n- dataset-preparation\n---\n\n# {s['title']} — LiteRT conversion\n\nSource: `{s.get('repo') or s.get('checkpoint_url')}`.\n\nTask: {s['task']}.\n\nThis conversion is prepared for human-reviewed dataset preprocessing. No accuracy or speed benchmark is claimed by this batch. The serialized LiteRT files were only structurally loaded and tensor allocation was checked.\n\nConversion route: `{s['route']}`.\n",encoding='utf-8') |
| if report.get('dynamic_validated'): |
| (stage/'README.md').write_text(f"# {s['title']} — dynamic LiteRT conversion\n\nSource: `{s.get('repo')}`, revision `{e.get('sha')}`.\n\nThe image height and width are dynamic. No image resampling is embedded in this graph. The supported bounds and any alignment constraints are recorded in `runtime_contract.json`. Batch size is one. Inputs must be normalized as required by the upstream configuration.\n\nValidation includes real LiteRT inference at several square and rectangular dimensions and numerical comparison against the source model. See `validation.json` for exact shapes, tolerances and measured errors. These are conversion checks, not accuracy or performance benchmarks. Android execution has not been tested. The required CPU interpreter/delegate configuration, when applicable, is recorded in `runtime_contract.json`.\n\nPreserved upstream license and attribution files apply.\n",encoding='utf-8') |
| write_json(stage/'validation.json',report) |
| write_json(stage/'artifact_manifest.json',{p.name:{'sha256':sha256(p),'bytes':p.stat().st_size} for p in stage.iterdir() if p.is_file() and p.name!='artifact_manifest.json'}) |
| with (ROOT/'.publish.lock').open('a') as lock: |
| fcntl.flock(lock,fcntl.LOCK_EX) |
| parent=api.repo_info(repo_id=dest,repo_type='model',revision='main').sha |
| commit=api.upload_folder(repo_id=dest,repo_type='model',folder_path=str(stage),path_in_repo=f'{group}/{s["key"]}',revision='main',parent_commit=parent,commit_message=f'Add LiteRT conversion: {s["title"]}') |
| paths=[f'{group}/{s["key"]}/{a.name}' for a in arts] |
| remote=api.get_paths_info(repo_id=dest,paths=paths,revision=commit.oid,repo_type='model') |
| indexed={p.path:p for p in remote} |
| for a,path in zip(arts,paths): |
| entry=indexed[path] |
| if entry.size!=a.stat().st_size or (entry.lfs and entry.lfs.sha256!=sha256(a)): |
| raise RuntimeError(f'Remote artifact mismatch: {path}') |
| write_json(out/'publication_receipt.json',{'repo':dest,'commit':commit.oid,'verified_paths':paths}) |
| return f'{dest}/tree/{commit.oid}/{group}/{s["key"]}' |
|
|
| def process(api,s): |
| k=s['key']; base=WORK/k; base.mkdir(parents=True,exist_ok=True); attempt=Path(tempfile.mkdtemp(prefix='attempt-',dir=base)); st={'key':k,'title':s['title'],'phase':'starting'} |
| try: |
| e=prefetched(k); st['upstream']=e; st['phase']='converting'; write_json(base/'status.json',st) |
| arts,rep=CONV[s['route']](s,e,attempt); rep['artifacts']=[{'name':p.name,'bytes':p.stat().st_size,'sha256':sha256(p),'tensors':inspect_tflite(p)} for p in arts] |
| st['report']=rep; st['phase']='publishing'; write_json(base/'status.json',st) |
| st['destination']=publish(api,s,e,attempt,arts,rep); st['phase']='pushed' if '/tree/' in st['destination'] else 'converted_local_only' |
| except Exception as ex: |
| st.update(phase='failed',error=f'{type(ex).__name__}: {ex}',traceback=traceback.format_exc()) |
| print(st['traceback'],flush=True) |
| finally: |
| write_json(base/'status.json',st); gc.collect(); |
| if torch.cuda.is_available(): |
| try: torch.cuda.empty_cache() |
| except Exception: pass |
| return st |
|
|
| def main(): |
| ap=argparse.ArgumentParser(); ap.add_argument('--prefetch-only',action='store_true'); ap.add_argument('--general-only',action='store_true'); ap.add_argument('--fireviewer-only',action='store_true'); ap.add_argument('models',nargs='*'); a=ap.parse_args() |
| api=HfApi(token=token()) |
| if a.prefetch_only: |
| api.repo_info(repo_id=DEST_REPO,repo_type='model'); prefetch(api); return |
| specs=all_specs() |
| if a.general_only: specs=[s for s in specs if s['group']=='general_models'] |
| if a.fireviewer_only: specs=[s for s in specs if s['group']=='fireviewer_models'] |
| if a.models: specs=[s for s in specs if s['key'] in set(a.models)] |
| results={} |
| for i,s in enumerate(specs,1): |
| print(f'\n[{i}/{len(specs)}] {s["key"]}',flush=True); results[s['key']]=process(api,s); print(' ->',results[s['key']]['phase'],flush=True) |
| write_json(WORK/'summary.json',results) |
| print('\nPUSHED:',', '.join(k for k,v in results.items() if v['phase']=='pushed') or 'none') |
| print('FAILED:',', '.join(k for k,v in results.items() if v['phase']=='failed') or 'none') |
| print('LOCAL ONLY:',', '.join(k for k,v in results.items() if v['phase']=='converted_local_only') or 'none') |
|
|
| if __name__=='__main__': main() |
|
|