Download evaluation/render.py from SimpleTuner/Qwen-Image-2.1-LoRA-experiments: direct link, hf CLI and curl.
- Browser
- Download file 3.98 kB
-
https://huggingface.co/SimpleTuner/Qwen-Image-2.1-LoRA-experiments/resolve/main/evaluation/render.py
- Command line
-
hf download hf://SimpleTuner/Qwen-Image-2.1-LoRA-experiments/evaluation/render.py
-
curl -L -o render.py https://huggingface.co/SimpleTuner/Qwen-Image-2.1-LoRA-experiments/resolve/main/evaluation/render.py
3.98 kB
| import argparse,gc,hashlib,json,time | |
| from pathlib import Path | |
| import torch | |
| from diffusers import FlowMatchEulerDiscreteScheduler | |
| from transformers import Qwen3VLForConditionalGeneration,Qwen3VLProcessor | |
| from simpletuner.helpers.models.qwen_image.autoencoder_21 import AutoencoderKLQwenImage21 | |
| from simpletuner.helpers.models.qwen_image.pipeline_21 import QwenImage21Pipeline | |
| from simpletuner.helpers.models.qwen_image.transformer_21 import QwenImage21Transformer2DModel | |
| p=argparse.ArgumentParser();p.add_argument('--checkpoint',type=Path);p.add_argument('--output',type=Path,required=True);args=p.parse_args() | |
| assert not args.output.exists();args.output.mkdir(parents=True) | |
| base='Qwen/Qwen-Image-2.1';revision='790c92633540aa0cb11d9abf19eb46d861714758' | |
| prompts=json.loads(Path('config/experiment/prompts.json').read_text());all_prompts=dict(prompts,negative='') | |
| cache=Path('cache/matrix-render-prompts.pt') | |
| if not cache.exists(): | |
| encoder=QwenImage21Pipeline(scheduler=FlowMatchEulerDiscreteScheduler.from_pretrained(base,subfolder='scheduler',revision=revision),vae=None,transformer=None,text_encoder=Qwen3VLForConditionalGeneration.from_pretrained(base,subfolder='text_encoder',revision=revision,dtype=torch.bfloat16).to('cuda').eval(),processor=Qwen3VLProcessor.from_pretrained(base,subfolder='processor',revision=revision)) | |
| embeds={} | |
| with torch.inference_mode(): | |
| for name,prompt in all_prompts.items(): | |
| values=encoder.encode_prompt(prompt,device=torch.device('cuda'));embeds[name]=[v.cpu() if v is not None else None for v in values] | |
| cache.parent.mkdir(parents=True,exist_ok=True);torch.save(dict(prompts=all_prompts,embeddings=embeds),cache) | |
| del encoder,embeds,values;gc.collect();torch.cuda.empty_cache() | |
| cached=torch.load(cache,map_location='cpu',weights_only=True);assert cached['prompts']==all_prompts | |
| pipe=QwenImage21Pipeline(scheduler=FlowMatchEulerDiscreteScheduler.from_pretrained(base,subfolder='scheduler',revision=revision),vae=AutoencoderKLQwenImage21.from_pretrained(base,subfolder='vae',revision=revision,torch_dtype=torch.bfloat16),transformer=QwenImage21Transformer2DModel.from_pretrained(base,subfolder='transformer',revision=revision,torch_dtype=torch.bfloat16),text_encoder=None,processor=None).to('cuda') | |
| pipe.vae.disable_tiling();pipe.vae.enable_slicing(); | |
| if args.checkpoint is not None: pipe.load_lora_weights(str(args.checkpoint),weight_name='pytorch_lora_weights.safetensors') | |
| pipe.transformer.eval() | |
| # Eager inference makes the pass-count audit independent of graph compilation. | |
| passes=[0] | |
| def count_pass(module,inputs):passes[0]+=1 | |
| handle=pipe.transformer.register_forward_pre_hook(count_pass) | |
| records=[] | |
| with torch.inference_mode(): | |
| for cfg,seed in [(1.0,42),(4.0,42),(1.0,123),(4.0,123)]: | |
| for name,prompt in prompts.items(): | |
| embeds,mask,_=[v.to('cuda') if v is not None else None for v in cached['embeddings'][name]] | |
| negative,negative_mask,_=[v.to('cuda') if v is not None else None for v in cached['embeddings']['negative']] | |
| kwargs=dict(prompt_embeds=embeds,prompt_embeds_mask=mask,height=512,width=512,num_inference_steps=40,true_cfg_scale=cfg,generator=torch.Generator(device='cuda').manual_seed(seed)) | |
| if cfg>1:kwargs.update(negative_prompt_embeds=negative,negative_prompt_embeds_mask=negative_mask) | |
| passes[0]=0;started=time.monotonic();result=pipe(**kwargs).images[0] | |
| assert passes[0]==40*(2 if cfg>1 else 1),passes[0] | |
| name_out=f'{name}-cfg{int(cfg)}-seed{seed}.png';result.save(args.output/name_out) | |
| row=dict(file=name_out,prompt=prompt,cfg=cfg,seed=seed,transformer_passes=passes[0],seconds=time.monotonic()-started);records.append(row);print(json.dumps(row),flush=True) | |
| (args.output/'receipt.json').write_text(json.dumps(dict(status='completed' if len(records)==4*len(prompts) else 'running',checkpoint_sha256=hashlib.sha256((args.checkpoint/'pytorch_lora_weights.safetensors').read_bytes()).hexdigest() if args.checkpoint is not None else None,images=records),indent=2)+'\n') | |
| handle.remove() | |