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()