Spaces:
Sleeping
Sleeping
| 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!") |