DALL-E-mini / app.py
RexTRO111's picture
Create app.py
8b4f2e7 verified
Raw
History Blame Contribute Delete
2.18 kB
import gradio as gr
import torch
from min_dalle import MinDalle
# Initialize the DALL-E Mini/Mega model
# Use float16 on GPU to save memory and increase generation speed
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.float16 if device == "cuda" else torch.float32
model = MinDalle(
is_mega=False,
models_root="pretrained",
is_reusable=True,
dtype=dtype,
device=device
)
def generate_27_images(prompt, seed, progress=gr.Progress(track_tqdm=True)):
if not prompt or not prompt.strip():
return []
images = []
# Generate 27 images iteratively
for i in range(27):
current_seed = seed + i if seed != -1 else -1
img = model.generate_image(
text=prompt,
seed=current_seed,
grid_size=1,
is_seamless=False,
temperature=1.0,
top_k=256,
supercondition_factor=16.0
)
images.append(img)
return images
# Gradio Interface
with gr.Blocks(title="DALL-E Mini 27-Image Grid") as demo:
gr.Markdown("# 🎨 DALL·E Mini - 27 Image Generator")
gr.Markdown("Enter a text prompt to generate 27 distinct image variants rendered in a responsive grid.")
with gr.Row():
with gr.Column(scale=1):
prompt_input = gr.Textbox(
label="Prompt",
placeholder="An astronaut riding a green horse on Mars...",
lines=3
)
seed_input = gr.Number(
label="Seed (-1 for random)",
value=-1,
precision=0
)
generate_btn = gr.Button("Generate 27 Images", variant="primary")
with gr.Column(scale=2):
gallery_output = gr.Gallery(
label="Generated Images (27)",
columns=[3, 4, 6],
rows=[9, 7, 5],
object_fit="contain",
height="auto"
)
generate_btn.click(
fn=generate_27_images,
inputs=[prompt_input, seed_input],
outputs=gallery_output
)
if __name__ == "__main__":
demo.launch()