rashed2/image-captioning-api
0
1# -*- coding: utf-8 -*-2"""Untitled5.ipynb3 4Automatically generated by Colab.5 6Original file is located at7 https://colab.research.google.com/drive/16Ic_VB5tHJDr6wgzY_t49TdUg9ZE5g3R8"""9 10import os, torch, re, io11from PIL import Image12from flask import Flask, request, jsonify13from transformers import BlipProcessor, BlipForConditionalGeneration, AutoTokenizer, AutoModelForSeq2SeqLM14 15# تحميل النماذج مرة واحدة عند بدء التشغيل16blip_model_name = "Salesforce/blip-image-captioning-base"17blip_processor = BlipProcessor.from_pretrained(blip_model_name)18blip_model = BlipForConditionalGeneration.from_pretrained(19 blip_model_name, torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32, device_map="auto"20)21 22trans_model_name = "facebook/m2m100_418M"23trans_tokenizer = AutoTokenizer.from_pretrained(trans_model_name)24trans_model = AutoModelForSeq2SeqLM.from_pretrained(trans_model_name)25 26def load_and_compress_image(file, max_side=512, quality=80):27 img = Image.open(file).convert("RGB")28 w, h = img.size29 scale = min(1.0, max_side / max(w, h))30 if scale < 1.0:31 img = img.resize((int(w * scale), int(h * scale)), Image.LANCZOS)32 buf = io.BytesIO()33 img.save(buf, format="JPEG", quality=quality, optimize=True)34 buf.seek(0)35 return Image.open(buf).convert("RGB")36 37def dedup_keywords(kws):38 cleaned, seen = [], set()39 for k in kws:40 k = k.strip().strip("،,;.- ").replace(" ", " ")41 if k and k.lower() not in seen:42 seen.add(k.lower())43 cleaned.append(k)44 return cleaned45 46def translate_to_arabic(text_en):47 inputs = trans_tokenizer(text_en, return_tensors="pt", truncation=True)48 forced_bos_token_id = trans_tokenizer.get_lang_id("ar")49 with torch.no_grad():50 outputs = trans_model.generate(51 **inputs,52 forced_bos_token_id=forced_bos_token_id,53 max_new_tokens=8054 )55 return trans_tokenizer.decode(outputs[0], skip_special_tokens=True)56 57def refine_arabic_text(text_ar):58 text_ar = re.sub(r"\s+", " ", text_ar)59 text_ar = text_ar.replace(" ,", ",").replace(" .", ".")60 return text_ar.strip()61 62def analyze_image_bilingual(file):63 image = load_and_compress_image(file)64 inputs = blip_processor(images=image, return_tensors="pt").to(blip_model.device)65 out_ids = blip_model.generate(**inputs, max_new_tokens=50)66 caption_en = blip_processor.decode(out_ids[0], skip_special_tokens=True)67 caption_ar = refine_arabic_text(translate_to_arabic(caption_en))68 return {69 "arabic_description": caption_ar,70 "arabic_keywords": dedup_keywords(caption_ar.split()),71 "english_keywords": dedup_keywords(caption_en.split())72 }73 74app = Flask(__name__)75 76@app.route("/analyze", methods=["POST"])77def analyze():78 if "image" not in request.files:79 return jsonify({"error": "No image uploaded"}), 40080 file = request.files["image"]81 try:82 result = analyze_image_bilingual(file)83 return jsonify(result)84 except Exception as e:85 return jsonify({"error": str(e)}), 50086 87@app.route("/", methods=["GET"])88def home():89 return jsonify({"message": "Image Captioning API is running"})90 91if __name__ == "__main__":92 port = int(os.environ.get("PORT", 7860))93 app.run(host="0.0.0.0", port=port)