karmx's picture
Release TinyQuery 139.7M from scratch with frozen weights, reproducible Mac evaluations and runtime source
b296ad4 verified
Raw
History Blame Contribute Delete
2.87 kB
"""Export only inference weights and public project code; never upload credentials or remote config."""
import argparse
import hashlib
import json
from pathlib import Path
import shutil
import torch
from safetensors.torch import save_file
from tinyquery.model import Config,TinyQuery
def main():
p=argparse.ArgumentParser(); p.add_argument('--checkpoint',required=True); p.add_argument('--data',required=True)
p.add_argument('--out',required=True); args=p.parse_args()
source=Path(args.checkpoint); data=Path(args.data); out=Path(args.out); out.mkdir(parents=True,exist_ok=True)
if source.suffix=='.safetensors':
config=json.loads((source.parent/'config.json').read_text())
shutil.copy2(source,out/'model.safetensors')
info=source.parent/'checkpoint-info.json'
if not info.exists():info=source.parent/'best-info.json'
if info.exists(): shutil.copy2(info,out/'checkpoint-info.json')
else:
checkpoint=torch.load(source,map_location='cpu',weights_only=False)
config=checkpoint['config']
state={k:v.to(torch.bfloat16).contiguous() for k,v in checkpoint['model'].items()}
save_file(state,str(out/'model.safetensors'))
info={k:checkpoint[k] for k in ['step','processed_tokens','response_tokens','training_seconds','random_initialization']}
(out/'checkpoint-info.json').write_text(json.dumps(info,indent=2))
(out/'config.json').write_text(json.dumps(config,indent=2))
shutil.copy2(data/'tokenizer.json',out/'tokenizer.json')
package=Path(__file__).parent
shutil.copytree(package,out/'tinyquery',dirs_exist_ok=True,ignore=shutil.ignore_patterns('__pycache__','*.pyc'))
shutil.copy2(package/'requirements-inference.txt',out/'requirements.txt')
for name in ['tokenization.json','grounding-stats.json','split-audit.json']:
if (data/name).exists(): shutil.copy2(data/name,out/name)
model=TinyQuery(Config(**config)); count=sum(p.numel() for p in model.parameters())
from safetensors.torch import load_file
model.load_state_dict(load_file(str(out/'model.safetensors')),strict=True)
manifest={'parameters':count,'format':'Custom native PyTorch; see tinyquery/model.py, not a Transformers AutoModel checkpoint',
'random_initialization':True,'files':{}}
for file in sorted(out.rglob('*')):
if file.is_file() and file.name!='manifest.json':
digest=hashlib.sha256()
with file.open('rb') as stream:
for block in iter(lambda:stream.read(4*1024*1024),b''): digest.update(block)
manifest['files'][str(file.relative_to(out))]={'bytes':file.stat().st_size,'sha256':digest.hexdigest()}
(out/'manifest.json').write_text(json.dumps(manifest,indent=2)); print(json.dumps({'parameters':count,'files':len(manifest['files'])}))
if __name__=='__main__': main()