Download evaluation/render_fixed_vae.py from SimpleTuner/Qwen-Image-2.1-LoRA-experiments: direct link, hf CLI and curl.
- Browser
- Download file 8.08 kB
-
https://huggingface.co/SimpleTuner/Qwen-Image-2.1-LoRA-experiments/resolve/main/evaluation/render_fixed_vae.py
- Command line
-
hf download hf://SimpleTuner/Qwen-Image-2.1-LoRA-experiments/evaluation/render_fixed_vae.py
-
curl -L -o render_fixed_vae.py https://huggingface.co/SimpleTuner/Qwen-Image-2.1-LoRA-experiments/resolve/main/evaluation/render_fixed_vae.py
8.08 kB
| 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() | |