img / generate.py
recoilme's picture
Upload folder using huggingface_hub
f2ff3b1 verified
Raw
History Blame Contribute Delete
3.07 kB
import argparse
import torch
from diffusers import ZImagePipeline, FlowMatchEulerDiscreteScheduler
NEGATIVE_PROMPT = (
"low quality, bad quality, blurry, sketch, sepia, text, "
"bad anatomy, bad proportions, bad hands, missing fingers, child drawing"
)
LORA = "out/lora/realistic_aesthetic_v7.safetensors"
PROMPTS_FILE = "out/prompts.txt"
OUT_DIR = "out/img"
SHIFT = 1.0
def load_pipe(shift=SHIFT):
pipe = ZImagePipeline.from_pretrained(
"Tongyi-MAI/Z-Image",
torch_dtype=torch.bfloat16,
)
pipe.to("cuda:0")
pipe.scheduler = FlowMatchEulerDiscreteScheduler.from_config(
pipe.scheduler.config, shift=shift
)
return pipe
def generate(pipe, prompt, path, lora_path=None, lora_scale=0.0):
if lora_path:
pipe.load_lora_weights(lora_path)
pipe.fuse_lora(lora_scale=lora_scale)
else:
# ensure no LoRA is active
if hasattr(pipe, 'unfuse_lora'):
try: pipe.unfuse_lora()
except: pass
if hasattr(pipe, 'unload_lora_weights'):
try: pipe.unload_lora_weights()
except: pass
image = pipe(
prompt=prompt,
negative_prompt=NEGATIVE_PROMPT,
height=1280,
width=832,
num_inference_steps=50,
guidance_scale=4.0,
cfg_normalization=True,
generator=torch.Generator("cuda:0").manual_seed(42),
).images[0]
if lora_path:
pipe.unfuse_lora()
image.save(path, quality=97)
print(f" -> {path} [{image.size[0]}x{image.size[1]}]")
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--lora", type=str, default=LORA)
parser.add_argument("--no-lora", action="store_true")
parser.add_argument("--scales", type=str, default=None,
help="comma-separated, e.g. 0.4,0.6,0.8")
parser.add_argument("--prompt-idx", type=int, default=None,
help="single prompt index (0-based)")
parser.add_argument("--shift", type=float, default=SHIFT,
help="scheduler shift (default 1.0)")
args = parser.parse_args()
with open(PROMPTS_FILE) as f:
prompts = [line.strip() for line in f if line.strip()]
pipe = load_pipe(shift=args.shift)
lora = None if args.no_lora else args.lora
indices = [args.prompt_idx] if args.prompt_idx is not None else range(len(prompts))
scales = [float(s) for s in args.scales.split(",")] if args.scales else [None]
for i in indices:
# Always generate non-lora first
path = f"{OUT_DIR}/{i+1:03d}.jpg"
print(f"[{i+1}/{len(prompts)}] no lora — {prompts[i][:60]}...")
generate(pipe, prompts[i], path)
for s in scales:
if s is None:
continue
path = f"{OUT_DIR}/{i+1:03d}_lora_{int(s*100):.0f}.jpg"
print(f"[{i+1}/{len(prompts)}] lora {int(s*100)}% — {prompts[i][:60]}...")
generate(pipe, prompts[i], path, lora_path=lora, lora_scale=s)
print("Done.")
if __name__ == "__main__":
main()