Team Ai
Apppublic

M-A-Z/Text_To_Image

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
app.py156 linesDownload Raw Back to root
1# app.py2 3import torch4from diffusers import StableDiffusionPipeline5import gradio as gr6import os7 8# Global variable to cache the pipeline9pipe = None10 11def load_model():12    """13    Loads the Stable Diffusion model for CPU. This function will be called once14    when the Gradio app starts.15    """16    global pipe17    if pipe is None:18        print("Loading Stable Diffusion model for CPU... This will take a moment.")19        # We recommend "runwayml/stable-diffusion-v1-5" for CPU as it's lighter.20        # Avoid larger models like SDXL on CPU.21        model_id = "runwayml/stable-diffusion-v1-5"22 23        # Always use float32 for CPU for compatibility and stability.24        torch_dtype = torch.float3225 26        try:27            # Load the pipeline from Hugging Face Hub28            pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=torch_dtype)29            # Explicitly move the model to CPU (usually default if no GPU)30            pipe = pipe.to("cpu")31            print("Stable Diffusion model loaded successfully on CPU.")32            print("WARNING: Image generation on CPU will be significantly slower.")33        except Exception as e:34            print(f"Error loading model: {e}")35            print("Please ensure you have an active internet connection to download the model.")36            print("If running in a limited environment, consider pre-downloading the model.")37            pipe = None # Ensure pipe is None if loading failed38 39    return pipe40 41def generate_image(prompt: str, negative_prompt: str = "", num_inference_steps: int = 30, guidance_scale: float = 7.5):42    """43    Generates an image from a text prompt using the loaded Stable Diffusion model on CPU.44 45    Args:46        prompt (str): The positive text prompt to guide image generation.47        negative_prompt (str): The negative text prompt to guide what NOT to include.48        num_inference_steps (int): Number of denoising steps (fewer steps recommended for CPU).49        guidance_scale (float): Controls how strongly the prompt is followed.50 51    Returns:52        PIL.Image: The generated image.53    """54    global pipe55    if pipe is None:56        gr.Warning("Model not loaded. Attempting to load now...")57        load_model()58        if pipe is None:59            gr.Error("Failed to load model. Cannot generate image.")60            return None # Return None if model failed to load61 62    if not prompt:63        gr.Warning("Please enter a text prompt to generate an image.")64        return None65 66    # Use no_grad for inference to save memory and speed up67    with torch.no_grad():68        try:69            image = pipe(70                prompt=prompt,71                negative_prompt=negative_prompt,72                num_inference_steps=num_inference_steps,73                guidance_scale=guidance_scale74            ).images[0]75            return image76        except Exception as e:77            gr.Error(f"An error occurred during image generation: {e}")78            return None79 80# Define the Gradio interface81with gr.Blocks(title="Hugging Face Text-to-Image Generator (CPU)") as demo:82    gr.Markdown(83        """84        # 🎨 Text-to-Image Generator with Hugging Face Diffusers (CPU)85        **⚠️ Important Note for CPU Users:**86        Image generation on a CPU is **very slow** (can take several minutes per image).87        For faster results, a GPU is highly recommended.88        """89    )90 91    with gr.Row():92        with gr.Column(scale=2):93            prompt_input = gr.Textbox(94                label="Text Prompt",95                placeholder="A high-quality photo of an astronaut riding a horse on Mars, cinematic, realistic",96                lines=397            )98            negative_prompt_input = gr.Textbox(99                label="Negative Prompt (Optional)",100                placeholder="blurry, low resolution, ugly, deformed, text, watermark",101                lines=2102            )103            generate_button = gr.Button("Generate Image")104        with gr.Column(scale=1):105            num_inference_steps_slider = gr.Slider(106                minimum=10,107                maximum=50, # Reduced max steps for CPU to manage generation time108                step=5,109                value=30, # Default to fewer steps for CPU110                label="Inference Steps",111                info="More steps can improve quality but will significantly increase generation time on CPU."112            )113            guidance_scale_slider = gr.Slider(114                minimum=1.0,115                maximum=15.0, # Slightly reduced max guidance scale for CPU116                step=0.5,117                value=7.5,118                label="Guidance Scale",119                info="Higher values make the image adhere more to the prompt."120            )121    122    output_image = gr.Image(type="pil", label="Generated Image")123 124    # Connect the button click to the generation function125    generate_button.click(126        fn=generate_image,127        inputs=[128            prompt_input,129            negative_prompt_input,130            num_inference_steps_slider,131            guidance_scale_slider132        ],133        outputs=output_image134    )135 136    # Examples for quick testing137    gr.Examples(138        examples=[139            ["A simple drawing of a house, cartoon style"],140            ["A red apple on a wooden table"],141            ["A happy golden retriever puppy playing in a field"]142        ],143        inputs=prompt_input,144        outputs=output_image,145        fn=generate_image,146        cache_examples=False # Do not cache examples for CPU as they are slow to generate147    )148 149# Load the model when the app starts. This is outside the Gradio blocks context150# so it runs once when the script is executed.151load_model()152 153# Launch the Gradio app154# Use `share=True` to get a public URL (useful for Colab or sharing)155# Use `debug=True` for more detailed logging in your console156demo.launch(debug=True)