Z-Image-Turbo-W4A8 / test_generate.py
CaoTurkey's picture
Upload 3 files
2213223 verified
Raw
History Blame Contribute Delete
5.02 kB
#!/usr/bin/env python3
"""End-to-end Z-Image-Turbo txt2img comparison across diffusion model variants.
Runs the standard txt2img flow (CLIP encode -> KSampler -> VAE decode) once per
given diffusion checkpoint, using the same prompt/seed/size, and reports load,
sample and decode timings per variant. Outputs PNGs + a timing log.
Usage:
python test_generate.py --te qwen_3_4b.safetensors --vae qwen_image_vae.safetensors \
--diffusion bf16.safetensors --diffusion w4a8.safetensors --diffusion int8_convrot.safetensors \
--prompt "..." --outdir out
"""
import argparse
import json
import os
import sys
import time
import torch
REPO = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", "ComfyUI")
if REPO not in sys.path:
sys.path.insert(0, REPO)
import comfy.sample
import comfy.sd
import comfy.utils
def load_clip_cond(te_path, prompt, negative):
clip = comfy.sd.load_clip([te_path], clip_type=comfy.sd.CLIPType.QWEN_IMAGE)
positive = clip.encode_from_tokens_scheduled(clip.tokenize(prompt))
negative = clip.encode_from_tokens_scheduled(clip.tokenize(negative))
return clip, positive, negative
def save_png(tensor, path):
img = torch.clamp(tensor, 0.0, 1.0).cpu().numpy()
img = (img * 255.0).astype("uint8")
try:
import torchvision.transforms.functional as F
F.to_pil_image(torch.from_numpy(img).permute(2, 0, 1)).save(path)
except ImportError:
from PIL import Image
Image.fromarray(img).save(path)
def generate(args, label, dm_path, positive, negative, vae):
t0 = time.time()
patcher = comfy.sd.load_diffusion_model(dm_path)
t_load = time.time() - t0
latent_format = patcher.get_model_object("latent_format")
batch = 1
latent = torch.zeros([batch, latent_format.latent_channels,
args.height // 8, args.width // 8], dtype=torch.float32)
noise = comfy.sample.prepare_noise(latent, args.seed)
t0 = time.time()
samples = comfy.sample.sample(
patcher, noise, args.steps, args.cfg, args.sampler, args.scheduler,
positive, negative, latent, denoise=1.0,
disable_pbar=not args.pbar, seed=args.seed)
t_sample = time.time() - t0
t0 = time.time()
images = vae.decode(samples)
t_decode = time.time() - t0
out_path = os.path.join(args.outdir, f"{label}.png")
save_png(images[0], out_path)
peak = 0
if torch.cuda.is_available():
peak = torch.cuda.max_memory_allocated() / 2**30
return {
"label": label,
"output": out_path,
"load_s": round(t_load, 2),
"sample_s": round(t_sample, 2),
"decode_s": round(t_decode, 2),
"total_s": round(t_load + t_sample + t_decode, 2),
"peak_vram_gb": round(peak, 2),
}
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--te", required=True, help="text encoder safetensors (BF16)")
ap.add_argument("--vae", required=True, help="VAE safetensors")
ap.add_argument("--diffusion", action="append", required=True, help="diffusion model path (repeatable)")
ap.add_argument("--prompt", default="A cute corgi sitting on a mossy rock in a forest, soft sunlight, detailed fur, photographic")
ap.add_argument("--negative", default="")
ap.add_argument("--steps", type=int, default=8)
ap.add_argument("--cfg", type=float, default=1.0)
ap.add_argument("--sampler", default="euler")
ap.add_argument("--scheduler", default="beta")
ap.add_argument("--width", type=int, default=1024)
ap.add_argument("--height", type=int, default=1024)
ap.add_argument("--seed", type=int, default=42)
ap.add_argument("--outdir", default="zimage_out")
ap.add_argument("--pbar", action="store_true")
args = ap.parse_args()
os.makedirs(args.outdir, exist_ok=True)
clip, positive, negative = load_clip_cond(args.te, args.prompt, args.negative)
sd, _ = comfy.utils.load_torch_file(args.vae, return_metadata=True)
vae = comfy.sd.VAE(sd=sd)
vae.throw_exception_if_invalid()
results = []
for i, dm in enumerate(args.diffusion):
label = os.path.splitext(os.path.basename(dm))[0]
print(f"\n=== [{i + 1}/{len(args.diffusion)}] {label} ===", flush=True)
try:
r = generate(args, label, dm, positive, negative, vae)
results.append(r)
print(json.dumps(r, ensure_ascii=False, indent=2), flush=True)
except Exception as e:
import traceback
traceback.print_exc()
results.append({"label": label, "error": str(e)})
if torch.cuda.is_available():
torch.cuda.empty_cache()
log_path = os.path.join(args.outdir, "timings.json")
with open(log_path, "w") as f:
json.dump({"settings": vars(args), "results": results}, f, ensure_ascii=False, indent=2)
print(f"\nWrote {log_path}")
for r in results:
print(json.dumps(r, ensure_ascii=False))
if __name__ == "__main__":
main()