import os import streamlit as st import requests from groq import Groq from PIL import Image import io # ------------------------- # PAGE CONFIG # ------------------------- st.set_page_config(page_title="AI Image Generator Pro", layout="centered") st.title("🎨 AI Image Generator (Groq + SDXL)") st.markdown("Generate high-quality AI images with optional AI prompt enhancement.") # ------------------------- # LOAD API KEYS # ------------------------- HF_TOKEN = os.getenv("HF_TOKEN") GROQ_API_KEY = os.getenv("GROQ_API_KEY") if not HF_TOKEN: st.error("HF_TOKEN not found in Secrets") st.stop() if not GROQ_API_KEY: st.error("GROQ_API_KEY not found in Secrets") st.stop() # ------------------------- # INIT GROQ CLIENT # ------------------------- groq_client = Groq(api_key=GROQ_API_KEY) # ------------------------- # SDXL MODEL ENDPOINT # ------------------------- API_URL = "https://router.huggingface.co/hf-inference/models/stabilityai/stable-diffusion-xl-base-1.0" headers = { "Authorization": f"Bearer {HF_TOKEN}" } # ------------------------- # USER INPUT # ------------------------- prompt = st.text_area("Enter your prompt:", "A futuristic AI robot in cyberpunk city") enhance = st.checkbox("Enhance prompt using AI") generate_button = st.button("Generate Image") # ------------------------- # PROMPT ENHANCEMENT FUNCTION # ------------------------- def enhance_prompt(user_prompt): response = groq_client.chat.completions.create( model="llama-3.1-8b-instant", messages=[ { "role": "system", "content": "You are a professional AI prompt engineer. Rewrite user prompts into detailed cinematic image generation prompts." }, { "role": "user", "content": user_prompt } ], temperature=0.8, max_tokens=150 ) return response.choices[0].message.content # ------------------------- # IMAGE GENERATION FUNCTION # ------------------------- def generate_image(final_prompt): payload = { "inputs": final_prompt } response = requests.post(API_URL, headers=headers, json=payload) if response.status_code == 200: return Image.open(io.BytesIO(response.content)) else: st.error(f"Image generation error {response.status_code}: {response.text}") return None # ------------------------- # GENERATE BUTTON ACTION # ------------------------- if generate_button: if not prompt.strip(): st.warning("Please enter a prompt.") st.stop() with st.spinner("Processing..."): final_prompt = prompt # Enhance Prompt if enhance: st.write("🧠 Enhancing prompt with Groq...") final_prompt = enhance_prompt(prompt) st.write("✨ Enhanced Prompt:") st.write(final_prompt) # Generate Image st.write("🎨 Generating Image...") image = generate_image(final_prompt) if image: # Convert image to bytes for download img_bytes = io.BytesIO() image.save(img_bytes, format="PNG") img_bytes.seek(0) # Create safe filename safe_filename = prompt[:25].replace(" ", "_").replace("/", "") if not safe_filename: safe_filename = "generated_image" # DOWNLOAD BUTTON (Above Image) st.download_button( label="⬇️ Download Image", data=img_bytes, file_name=f"{safe_filename}.png", mime="image/png" ) # Show Image Below Button st.image(image, caption="Generated Image", use_column_width=True) st.success("Image generated successfully!")