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