"""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()