Image21-MLX-8bit / scripts /benchmark.py
ixim's picture
Add files using upload-large-folder tool
4f03424 verified
Raw History Blame Contribute Delete
5.09 kB
"""Paired 1024px evaluation, separate processes per precision, unchanged RGBA samples."""
import argparse
import gc
import importlib.metadata
import json
import platform
import time
from pathlib import Path
import mlx.core as mx
import numpy as np
import psutil
from scripts.common import sha256, write_json
from scripts.runtime import load, generate, enable_progress
def main():
ap=argparse.ArgumentParser()
ap.add_argument('--model',type=Path,required=True)
ap.add_argument('--output',type=Path,required=True)
ap.add_argument('--cases',type=Path,default=Path('benchmarks/cases.json'))
ap.add_argument('--only',nargs='+')
ap.add_argument('--seeds',nargs='+',type=int,default=[42])
ap.add_argument('--steps',type=int,default=40)
ap.add_argument('--size',type=int,default=1024)
ap.add_argument('--edit-input',type=Path)
ap.add_argument('--no-warmup',action='store_true')
ap.add_argument('--resident',action='store_true',help='Diagnostic mode: keep all components resident')
args=ap.parse_args()
if args.output.exists(): raise FileExistsError(args.output)
cases=json.loads(args.cases.read_text())
if args.only: cases=[x for x in cases if x['id'] in args.only]
if not cases: raise ValueError('No cases selected')
args.output.mkdir(parents=True)
env=dict(model=str(args.model.resolve()),device=mx.device_info(),platform=platform.platform(),
total_memory_gib=psutil.virtual_memory().total/2**30,
packages={p:importlib.metadata.version(p) for p in ('mlx','mlx-vlm','numpy','transformers')},
cases_sha256=sha256(args.cases),warmup=not args.no_warmup,
phase_offload=not args.resident,
note='MLX allocated peak is not whole-system memory or a proven minimum RAM requirement.')
if (args.model/'conversion.json').exists(): env['conversion']=json.loads((args.model/'conversion.json').read_text())
write_json(args.output/'environment.json',env)
enable_progress()
start=time.perf_counter(); pipe=load(args.model,phase_offload=not args.resident)
if args.resident:
mx.eval(pipe.transformer.parameters(),pipe.text_encoder.model.parameters(),pipe.vae.parameters())
load_seconds=time.perf_counter()-start
print(f'Pipeline ready in {load_seconds:.2f}s; phase loads are included in image timings',flush=True)
if not args.no_warmup:
print('Full untimed warm-up',flush=True)
generate(pipe,cases[0]['prompt'],seed=20260926,steps=args.steps,width=args.size,height=args.size)
mx.synchronize(); gc.collect(); mx.clear_cache()
for case in cases:
for generation_seed in args.seeds:
is_edit=case['id']=='edit'
seed=1_000_000+generation_seed if is_edit else generation_seed
inputs=None
if is_edit:
ref=args.edit_input or args.output/f'portrait-s{generation_seed}.png'
if not ref.exists(): raise FileNotFoundError(f'Provide the shared BF16 portrait: {ref}')
inputs=[str(ref)]
gc.collect(); mx.clear_cache(); mx.reset_peak_memory()
swap_before=psutil.swap_memory().used/2**30
start=time.perf_counter()
print(f'Running {case["id"]} seed={seed}',flush=True)
img=generate(pipe,case['prompt'],seed=seed,steps=args.steps,width=args.size,height=args.size,
inputs=inputs,source_seed=generation_seed if is_edit else None,resolution=args.size)
mx.synchronize(); seconds=time.perf_counter()-start
name=f'{case["id"]}-s{seed}.png'; img.save(args.output/name)
pixels=np.asarray(img); alpha=pixels[:,:,3]
row=dict(case_id=case['id'],prompt=case['prompt'],seed=seed,width=args.size,height=args.size,
steps=args.steps,cfg=1.0,vae_tiling=False,kv_cache=True,phase_offload=not args.resident,
input_sha256=sha256(inputs[0]) if inputs else None,pipeline_init_seconds=load_seconds,
seconds=seconds,mlx_peak_gib=mx.get_peak_memory()/2**30,
rss_gib=psutil.Process().memory_info().rss/2**30,
system_available_gib=psutil.virtual_memory().available/2**30,
swap_used_gib=psutil.swap_memory().used/2**30,
swap_before_gib=swap_before,swap_delta_gib=psutil.swap_memory().used/2**30-swap_before,
output=name,output_sha256=sha256(args.output/name),mode=img.mode,
alpha_min=int(alpha.min()),alpha_max=int(alpha.max()),
alpha_fraction_below_250=float(np.mean(alpha<250)),
rgb_std=float(pixels[:,:,:3].astype(np.float32).std()),warmup=not args.no_warmup)
with (args.output/'results.jsonl').open('a') as f: f.write(json.dumps(row,ensure_ascii=False)+'\n')
print(json.dumps(row,ensure_ascii=False),flush=True)
write_json(args.output/'COMPLETE.json',dict(cases=len(cases),seeds=args.seeds,steps=args.steps,size=args.size))
if __name__=='__main__': main()