Image21-INT8 / scripts /benchmark_edits_v2.py
ixim's picture
Release verified Image21-INT8 conversion (part 2)
e9190ff verified
Raw History Blame Contribute Delete
7.89 kB
"""Independent-seed editing evaluation; run each precision in a fresh process.
Historical benchmark_edits.py is retained byte-for-byte for reproduction.
"""
import argparse
import gc
import inspect
import json
import os
import time
from pathlib import Path
import torch
from scripts.benchmark import environment, load_pipeline, memory_status
from scripts.editing_protocol import (PROTOCOL, DEFAULT_SEEDS, validate_seeds, load_input,
edit_dimensions, image_diagnostics, save_diagnostics)
from scripts.integrity import sha256
from scripts.provenance import model_identity
def main():
ap = argparse.ArgumentParser(__doc__)
ap.add_argument('--model', required=True)
ap.add_argument('--inputs', required=True)
ap.add_argument('--output', required=True)
ap.add_argument('--cases', default='benchmarks/editing-v2.json')
ap.add_argument('--case', nargs='+')
ap.add_argument('--seeds', type=int, nargs='+', default=list(DEFAULT_SEEDS))
ap.add_argument('--resolution', type=int, default=1024, help='Square root of target area; preserve input aspect ratio')
ap.add_argument('--steps', type=int, default=40)
ap.add_argument('--vae-tiling', action=argparse.BooleanOptionalAction, default=False)
ap.add_argument('--kv-cache', action=argparse.BooleanOptionalAction, default=True)
ap.add_argument('--no-warmup', action='store_true', help='Diagnostic runs only; excludes latency comparisons')
ap.add_argument('--save-latents', action='store_true', help='Retain decoder input for one-variable VAE diagnosis')
args = ap.parse_args()
if args.steps <= 0:
ap.error('Positive steps required')
cases = json.loads(Path(args.cases).read_text(encoding='utf-8'))
if args.case:
if set(args.case) - {c['id'] for c in cases}:
ap.error('Unknown case')
cases = [c for c in cases if c['id'] in args.case]
if not cases or len({c['id'] for c in cases}) != len(cases):
ap.error('Nonempty distinct cases required')
validate_seeds(cases, args.seeds)
# Validate all inputs before allocating CUDA or creating an output directory.
inputs = {}
for case in cases:
name = case['input_file']
if Path(name).name != name or Path(case['id']).name != case['id']:
ap.error('Case IDs and input files must be plain names')
path = Path(args.inputs)/name
if sha256(path) != case['input_sha256']:
raise ValueError('Input hash mismatch: '+case['id'])
inputs[case['id']] = load_input(path)
edit_dimensions(inputs[case['id']].size, args.resolution)
out = Path(args.output)
out.mkdir(parents=True, exist_ok=False)
(out/'cases.json').write_text(json.dumps(cases, ensure_ascii=False, indent=2), encoding='utf-8')
env = environment()
env.update(protocol=PROTOCOL, pid=os.getpid(), before_load_memory=memory_status(),
model_identity=model_identity(args.model), offload='model', offload_aux_fix=True,
warmup=not args.no_warmup, generator_device='cpu', seeds=args.seeds,
resolution=args.resolution, steps=args.steps, kv_cache=args.kv_cache,
vae_tiling=args.vae_tiling, cases_sha256=sha256(out/'cases.json'),
source_sha256={name: sha256(Path(__file__).with_name(name)) for name in
('benchmark_edits_v2.py', 'editing_protocol.py', 'benchmark.py', 'runtime.py')})
if env['before_load_memory']['allocated_bytes'] != 0:
raise ValueError('Expected fresh CUDA allocation baseline')
pipe = load_pipeline(args.model)
if args.vae_tiling:
pipe.vae.enable_tiling()
env.update(scheduler_config=dict(pipe.scheduler.config), vae_dtype=str(pipe.vae.dtype),
upstream_source_sha256={type(component).__name__: sha256(inspect.getfile(type(component)))
for component in (pipe, pipe.vae, pipe.transformer, pipe.scheduler)},
vae_tiles={n: getattr(pipe.vae, n) for n in
('tile_sample_min_height', 'tile_sample_min_width',
'tile_sample_stride_height', 'tile_sample_stride_width')})
(out/'environment.json').write_text(json.dumps(env, indent=2), encoding='utf-8')
def call(case, seed):
image = inputs[case['id']]
width, height = edit_dimensions(image.size, args.resolution)
with torch.inference_mode():
result = pipe(prompt=case['prompt'], image=image, width=width, height=height,
output_resolution=args.resolution, num_inference_steps=args.steps,
true_cfg_scale=1., use_kv_cache=args.kv_cache,
generator=torch.Generator('cpu').manual_seed(seed)).images[0]
if result.size != (width, height):
raise ValueError('Unexpected output dimensions')
return result
if not args.no_warmup:
warmup_seed = 1000000
while warmup_seed in args.seeds or warmup_seed in {c.get('source_seed') for c in cases}:
warmup_seed += 1
print('Full-settings warmup, excluded', flush=True)
call(cases[0], warmup_seed)
original_decode = pipe.vae.decode
for case in cases:
for seed in args.seeds:
stem = f'{case["id"]}-s{seed}'
print('Generating '+stem, flush=True)
latent_path = out/(stem+'-decode.pt')
if args.save_latents:
def decode(latents, *a, **kw):
torch.save(latents.detach().cpu(), latent_path)
return original_decode(latents, *a, **kw)
pipe.vae.decode = decode
gc.collect()
torch.cuda.empty_cache()
before = memory_status()
torch.cuda.reset_peak_memory_stats()
started = time.perf_counter()
try:
image = call(case, seed)
finally:
pipe.vae.decode = original_decode
torch.cuda.synchronize()
seconds = time.perf_counter()-started
image.save(out/(stem+'.png'))
width, height = image.size
row = dict(protocol=PROTOCOL, case_id=case['id'], category=case['category'],
prompt=case['prompt'], seed=seed, source_seed=case.get('source_seed'),
input_file=case['input_file'], input_sha256=case['input_sha256'],
input_diagnostics=image_diagnostics(inputs[case['id']]),
width=width, height=height, output_resolution=args.resolution,
steps=args.steps, cfg=1., kv_cache=args.kv_cache, vae_tiling=args.vae_tiling,
offload='model', seconds=seconds, before_memory=before,
peak_allocated_bytes=torch.cuda.max_memory_allocated(),
peak_reserved_bytes=torch.cuda.max_memory_reserved(),
image=stem+'.png', image_sha256=sha256(out/(stem+'.png')),
sigmas=pipe.scheduler.sigmas.cpu().tolist(),
latency_comparable=not args.no_warmup and not args.save_latents,
**image_diagnostics(image))
if args.save_latents:
row.update(decode_latents=latent_path.name, decode_latents_sha256=sha256(latent_path))
save_diagnostics(image, out/'diagnostics', stem)
with (out/'records.jsonl').open('a', encoding='utf-8') as stream:
stream.write(json.dumps(row, ensure_ascii=False)+'\n')
print(f'{stem}: {seconds:.2f}s', flush=True)
(out/'COMPLETE.json').write_text(json.dumps(dict(records=len(cases)*len(args.seeds),
records_sha256=sha256(out/'records.jsonl'))))
if __name__ == '__main__':
main()