litert-models / conversion /package /scripts /convert_batch2.py
Charlbi's picture
Complete FireViewer Android learning SDK, source and verification records
dee7f43 verified
Raw
History Blame Contribute Delete
23.7 kB
#!/usr/bin/env python3
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()
# Use a static 640 square unless processor fixes another size.
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']
# Transformers 5.17 nests the existing ViT encoder under model; keep strict load.
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()