File size: 23,655 Bytes
dee7f43 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 | #!/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()
|