File size: 7,888 Bytes
e9190ff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
"""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()