Text-to-Image
Diffusers
Safetensors
sdxs-1b / generate.py
recoilme's picture
1.6b
5ae8af8
Raw
History Blame Contribute Delete
4.37 kB
"""
SDXS-1B image generation from prompts.txt.
Usage:
python generate.py # generates grid from prompts.txt
python generate.py --prompts my_list.txt # custom prompts file
python generate.py --prompt "your text" # single prompt (без файла)
python generate.py --index 0 # только промпт #0 из prompts.txt
"""
import argparse
import os
import math
import textwrap
import torch
import numpy as np
from pipeline_sdxs import SdxsPipeline
NEGATIVE_PROMPT = (
"bad quality grainy image with low details, incomplete text, "
"despite numerous technical flaws and distorted figures"
)
DEFAULT_PROMPTS_FILE = "prompts.txt"
GRID_OUTPUT = "media/result_grid.jpg"
def main():
parser = argparse.ArgumentParser(description="SDXS-1B generation")
parser.add_argument("--prompts", type=str, default=None, help="prompts file")
parser.add_argument("--prompt", type=str, default=None, help="single prompt")
parser.add_argument("--index", type=int, default=None, help="generate only prompt #N from file")
parser.add_argument("--steps", type=int, default=40)
parser.add_argument("--guidance", type=float, default=5.0)
parser.add_argument("--seed", type=int, default=0)
parser.add_argument("--width", type=int, default=1024)
parser.add_argument("--height", type=int, default=1408)
parser.add_argument("--cols", type=int, default=4, help="колонки в гриде")
parser.add_argument("--output", type=str, default=None,
help="папка для отдельных картинок (по умолчанию не сохраняются)")
args = parser.parse_args()
# ---------- сбор промптов ----------
prompts = None
if args.prompt:
prompts = [args.prompt]
else:
path = args.prompts or DEFAULT_PROMPTS_FILE
if not os.path.exists(path):
print(f"Файл промптов не найден: {path}")
return
with open(path, encoding="utf-8") as f:
prompts = [l.strip() for l in f if l.strip()]
if args.index is not None:
if args.index >= len(prompts):
print(f"Нет промпта #{args.index} (всего {len(prompts)})")
return
prompts = [prompts[args.index]]
if not prompts:
print("Нет промптов")
return
print(f"Промптов: {len(prompts)}")
# ---------- пайплайн ----------
device = "cuda" if torch.cuda.is_available() else "cpu"
pipe = SdxsPipeline.from_pretrained(
os.path.dirname(os.path.abspath(__file__)),
torch_dtype=torch.float16,
).to(device)
# ---------- генерация ----------
images = []
for i, prompt in enumerate(prompts):
print(f"[{i+1}/{len(prompts)}] {prompt[:80]}...")
image = pipe(
prompt=prompt,
negative_prompt=NEGATIVE_PROMPT,
guidance_scale=args.guidance,
width=args.width,
height=args.height,
seed=args.seed,
num_inference_steps=args.steps,
)[0][0]
images.append(image)
if args.output:
os.makedirs(args.output, exist_ok=True)
path = os.path.join(args.output, f"{i+1:03d}.png")
image.save(path)
print(f" saved {path}")
# ---------- грид ----------
os.makedirs(os.path.dirname(GRID_OUTPUT), exist_ok=True)
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
cols = min(args.cols, len(images))
rows = math.ceil(len(images) / cols)
fig, axes = plt.subplots(rows, cols, figsize=(cols * 4, rows * 4.5), constrained_layout=True)
axes = list(np.array(axes).flatten())[:len(images)]
for i, (img, prompt) in enumerate(zip(images, prompts)):
ax = axes[i]
ax.imshow(img)
ax.axis("off")
ax.set_aspect("equal")
text = (prompt[:200] + "…") if len(prompt) > 200 else prompt
lines = textwrap.wrap(text, width=35)
while len(lines) < 4:
lines.append("")
ax.set_title("\n".join(lines), fontsize=9, pad=8)
plt.savefig(GRID_OUTPUT, bbox_inches="tight", dpi=150, format="jpeg")
print(f"grid -> {GRID_OUTPUT}")
if __name__ == "__main__":
main()