import argparse from collections import defaultdict import hashlib import importlib.metadata import json from pathlib import Path import time import torch from diffusers import AutoencoderKLQwenImage21, QwenImage21Pipeline from huggingface_hub import hf_hub_download from peft import get_peft_model_state_dict from PIL.PngImagePlugin import PngInfo from safetensors.torch import load_file def sha256(path): with path.open('rb') as handle: return hashlib.file_digest(handle, 'sha256').hexdigest() def groups(jobs): result = defaultdict(list) for job in jobs: key = tuple(job[k] for k in ['checkpoint_repo', 'checkpoint_revision', 'checkpoint_file', 'checkpoint_sha256', 'width', 'height', 'seed', 'generator_device', 'true_cfg_scale', 'num_inference_steps']) result[key].append(job) return list(result.values()) def main(): parser = argparse.ArgumentParser() parser.add_argument('--manifest', type=Path, required=True) parser.add_argument('--output', type=Path, required=True) parser.add_argument('--smoke', action='store_true') args = parser.parse_args() manifest = json.loads(args.manifest.read_text()) args.output.mkdir(parents=True, exist_ok=False) assert torch.cuda.is_available() assert 'H100' in torch.cuda.get_device_name() assert manifest['batch_size'] == 4 jobs = manifest['jobs'] grouped = groups(jobs) if args.smoke: base = next(g for g in grouped if g[0]['checkpoint_file'] is None and g[0]['width'] == 1024 and g[0]['true_cfg_scale'] == 1 and g[0]['generator_device'] == 'cpu' and len(g) >= 4) concept = next(g for g in grouped if g[0]['checkpoint_file'] == 'checkpoints/assisted-v2-multiscale-repa-autoshift/step-10000/pytorch_lora_weights.safetensors' and g[0]['width'] == 1024 and len(g) >= 4) grouped = [base[:4], concept[:4]] total = sum(len(g) for g in grouped) runtime = {package: importlib.metadata.version(package) for package in ['torch', 'torchvision', 'diffusers', 'transformers', 'accelerate', 'peft', 'huggingface-hub', 'safetensors']} runtime.update(gpu=torch.cuda.get_device_name(), cuda=torch.version.cuda, manifest_sha256=sha256(args.manifest), renderer_sha256=sha256(Path(__file__)), diffusers_revision=manifest['diffusers_revision']) (args.output / 'runtime.json').write_text(json.dumps(runtime, indent=2) + '\n') print(json.dumps({'status': 'loading', 'jobs': total, 'groups': len(grouped), 'runtime': runtime}), flush=True) vae = AutoencoderKLQwenImage21.from_pretrained(manifest['vae_model'], revision=manifest['vae_revision'], dtype=torch.bfloat16) pipe = QwenImage21Pipeline.from_pretrained(manifest['base_model'], revision=manifest['base_revision'], vae=vae, dtype=torch.bfloat16).to('cuda') pipe.vae.disable_tiling() pipe.vae.enable_slicing() pipe.transformer.eval() pipe.set_progress_bar_config(disable=True) current = None completed = [] audit = {} initial_noise = {} def inspect_forward(module, positional, kwargs): audit['passes'] += 1 if audit['passes'] == 1: hidden = kwargs['hidden_states'] assert hidden.shape[0] == audit['batch_size'] assert all(torch.equal(hidden[0], hidden[i]) for i in range(1, hidden.shape[0])) audit['initial_noise_sha256'] = hashlib.sha256(hidden[0].contiguous().view(torch.uint8).cpu().numpy().tobytes()).hexdigest() hook = pipe.transformer.register_forward_pre_hook(inspect_forward, with_kwargs=True) started = time.monotonic() with torch.inference_mode(): for group in grouped: spec = group[0] selected = spec['checkpoint_sha256'] if selected != current: if current is not None: pipe.unload_lora_weights() if selected is not None: path = Path(hf_hub_download(spec['checkpoint_repo'], spec['checkpoint_file'], revision=spec['checkpoint_revision'])) assert sha256(path) == selected pipe.load_lora_weights(str(path.parent), weight_name=path.name, adapter_name='experiment') pipe.set_adapters('experiment', adapter_weights=1.0) expected = {key.removeprefix('transformer.'): value for key, value in load_file(path).items()} loaded = get_peft_model_state_dict(pipe.transformer, adapter_name='experiment') assert expected.keys() == loaded.keys() assert all(torch.equal(expected[key], loaded[key].detach().cpu()) for key in expected) assert pipe.get_active_adapters() == ['experiment'] del expected, loaded current = selected for offset in range(0, len(group), manifest['batch_size']): batch = group[offset:offset + manifest['batch_size']] generators = [torch.Generator(device=spec['generator_device']).manual_seed(spec['seed']) for _ in batch] assert len({id(g) for g in generators}) == len(batch) assert all(g.initial_seed() == spec['seed'] for g in generators) assert all(torch.equal(g.get_state(), generators[0].get_state()) for g in generators) audit.clear() audit.update(passes=0, batch_size=len(batch)) batch_started = time.monotonic() images = pipe( prompt=[job['prompt'] for job in batch], negative_prompt=[''] * len(batch) if spec['true_cfg_scale'] > 1 else None, width=spec['width'], height=spec['height'], num_inference_steps=spec['num_inference_steps'], true_cfg_scale=spec['true_cfg_scale'], num_images_per_prompt=1, generator=generators, ).images assert len(images) == len(batch) assert audit['passes'] == spec['num_inference_steps'] * (2 if spec['true_cfg_scale'] > 1 else 1) noise_key = (spec['width'], spec['height'], spec['seed'], spec['generator_device']) if noise_key in initial_noise: assert initial_noise[noise_key] == audit['initial_noise_sha256'] initial_noise[noise_key] = audit['initial_noise_sha256'] seconds = time.monotonic() - batch_started for image, job in zip(images, batch): assert image.size == (job['width'], job['height']) output = args.output / job['output'] output.parent.mkdir(parents=True, exist_ok=True) parameters = dict(job, base_model=manifest['base_model'], base_revision=manifest['base_revision'], vae_model=manifest['vae_model'], vae_revision=manifest['vae_revision'], diffusers_revision=manifest['diffusers_revision'], batch_size=len(batch)) metadata = PngInfo() metadata.add_text('parameters', json.dumps(parameters, ensure_ascii=False, sort_keys=True)) image.save(output, pnginfo=metadata) completed.append(dict(**parameters, sha256=sha256(output), initial_noise_sha256=audit['initial_noise_sha256'], transformer_passes=audit['passes'], batch_seconds=seconds)) receipt = dict(status='completed' if len(completed) == total else 'running', expected=total, completed=len(completed), elapsed_seconds=time.monotonic() - started, images=completed) temp = args.output / 'receipt.pending.json' temp.write_text(json.dumps(receipt, indent=2, ensure_ascii=False) + '\n') temp.replace(args.output / 'receipt.json') print(json.dumps({'completed': len(completed), 'total': total, 'batch_size': len(batch), 'batch_seconds': round(seconds, 2), 'checkpoint': spec['checkpoint_file'], 'resolution': spec['width'], 'cfg': spec['true_cfg_scale']}), flush=True) hook.remove() print('RENDER_COMPLETE', flush=True) if __name__ == '__main__': main()