File size: 5,088 Bytes
4f03424
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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()