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}:" 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)