from flask import Flask, request, jsonify, render_template_string
from transformers import T5Tokenizer, T5ForConditionalGeneration
from deep_translator import GoogleTranslator
from PIL import Image
from langdetect import detect, DetectorFactory
import torch
import requests
import io
import base64
import os
from dotenv import load_dotenv
# Load environment variables
load_dotenv()
app = Flask(__name__)
# Translation setup
translator = GoogleTranslator(source='auto', target='en')
DetectorFactory.seed = 0
# Hugging Face image generation API
HF_TOKEN = os.getenv("HF_TOKEN")
headers = {"Authorization": f"Bearer {HF_TOKEN}"}
IMAGE_MODEL_API_URL = "https://api-inference.huggingface.co/models/black-forest-labs/FLUX.1-dev"
# Load text generation model
model_id = "google/flan-t5-base"
tokenizer = T5Tokenizer.from_pretrained(model_id, legacy=False)
model = T5ForConditionalGeneration.from_pretrained(model_id).to("cuda" if torch.cuda.is_available() else "cpu")
# Text generation function
def query_text_model(prompt):
structured_prompt = f"Describe the following in bullet points: {prompt}"
inputs = tokenizer(structured_prompt, return_tensors="pt").to(model.device)
outputs = model.generate(**inputs, max_new_tokens=256)
text = tokenizer.decode(outputs[0], skip_special_tokens=True)
return [line.strip("-• ") for line in text.split("\n") if line.strip()]
# Format text for HTML display
def format_text_response(prompt, points):
formatted_response = f"{prompt}:
"
for point in points:
formatted_response += f"- {point}
"
formatted_response += "
"
return formatted_response
# Call Hugging Face API for image generation
def query_image_model(prompt):
try:
response = requests.post(IMAGE_MODEL_API_URL, headers=headers, json={"inputs": prompt})
if response.status_code == 200 and response.headers['Content-Type'].startswith('image/'):
return response.content
else:
raise ValueError("Image generation failed or unexpected content type")
except requests.exceptions.RequestException as e:
print("Image generation error:", str(e))
return None
# Convert image bytes to base64
def image_to_base64(image_bytes):
image = Image.open(io.BytesIO(image_bytes))
buffered = io.BytesIO()
image.save(buffered, format="JPEG")
return base64.b64encode(buffered.getvalue()).decode("utf-8")
# Load HTML page
@app.route('/')
def index():
return render_template_string(open("index.html", "r", encoding="utf-8").read())
# Chat endpoint handling text, image, and translation
@app.route('/chat', methods=['POST'])
def chat():
data = request.json
user_input = data.get('message', '').strip()
response_type = data.get('response_type', 'text')
print(f"[Received] {response_type.upper()} | {user_input}")
if not user_input:
return jsonify({"response": "Please enter something first."})
if response_type == 'image':
image_bytes = query_image_model(user_input)
if image_bytes:
image_base64 = image_to_base64(image_bytes)
return jsonify({"response": f'
'})
return jsonify({"response": "Image generation failed."})
elif response_type == 'translation':
try:
translated = translator.translate(user_input)
return jsonify({"response": translated})
except Exception as e:
print("Translation error:", e)
return jsonify({"response": "Sorry, translation failed."})
else: # Default: text
try:
text_points = query_text_model(user_input)
if text_points:
formatted_response = format_text_response(user_input, text_points)
return jsonify({"response": formatted_response})
return jsonify({"response": "No meaningful response generated."})
except Exception as e:
print("Text generation error:", e)
return jsonify({"response": "Something went wrong while generating a response."})
# Run the app
if __name__ == '__main__':
port = int(os.environ.get("PORT", 7860))
app.run(host='0.0.0.0', port=port)