bghira's picture
Squash history while preserving current repository contents
7aa9202
Raw History Blame Contribute Delete
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()