Text_to_Image / app.py
santosh0223's picture
Update app.py
4ffa91e verified
Raw
History Blame Contribute Delete
3.82 kB
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!")