PBL6_AI / app.py
phuongsuga's picture
Update app.py
d930fec verified
Raw
History Blame Contribute Delete
6.48 kB
import os
import re
import json
import torch
import torch.nn.functional as F
from flask import Flask, request, jsonify
from transformers import AutoTokenizer, AutoModelForSequenceClassification
from pyvi import ViTokenizer
from PIL import Image
from torchvision import transforms
import timm
from flask_cors import CORS
from huggingface_hub import hf_hub_download # <--- Dùng cái này để tải file config an toàn
# =====================
# CONFIG
# =====================
os.environ["TRANSFORMERS_CACHE"] = "/tmp/hf_cache"
TEXT_MODEL_REPO = "phuongsuga/PBL6_AI_Model_Text_Image"
IMAGE_MODEL_REPO = "phuongsuga/PBL6_AI_Model_Image"
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
THRESHOLD = 0.65
app = Flask(__name__)
CORS(app)
# =====================
# LOAD TEXT MODEL
# =====================
print("🔹 Downloading text model...")
# 1. Tải tokenizer
tokenizer = AutoTokenizer.from_pretrained(
TEXT_MODEL_REPO,
subfolder="text_model",
use_fast=False
)
# 2. Tải model
text_model = AutoModelForSequenceClassification.from_pretrained(
TEXT_MODEL_REPO,
subfolder="text_model/checkpoint-3390"
)
text_model.to(DEVICE).eval()
# 3. Tải label2id.json AN TOÀN bằng hf_hub_download
try:
print("🔹 Loading label2id.json...")
label2id_path = hf_hub_download(repo_id=TEXT_MODEL_REPO, filename="text_model/label2id.json")
with open(label2id_path, "r", encoding="utf-8") as f:
label2id = json.load(f)
id2label_text = {int(v): k for k, v in label2id.items()} # Đảm bảo key là int
print("✅ Loaded label2id successfully.")
except Exception as e:
print(f"⚠️ Warning: Could not load label2id.json from Hub. Error: {e}")
# Fallback nếu file không tồn tại: Lấy từ config của model
print("🔹 Attempting to use model config instead...")
id2label_text = text_model.config.id2label
label2id = text_model.config.label2id
# =====================
# LOAD IMAGE MODEL
# =====================
print("🔹 Downloading image model...")
class_names = ["an_toan", "bao_luc", "khieu_dam_doi_truy", "nhay_cam_chinh_tri"]
def build_model(num_classes=4):
return timm.create_model("efficientnet_b3", pretrained=False, num_classes=num_classes)
image_model = build_model()
image_model_path = "/tmp/efficientnet_b3.pth"
# Tải weight ảnh nếu chưa có
if not os.path.exists(image_model_path):
try:
url = f"https://huggingface.co/{IMAGE_MODEL_REPO}/resolve/main/image_model/efficientnet_b3.pth"
torch.hub.download_url_to_file(url, image_model_path)
except Exception as e:
print(f"❌ Error downloading image model: {e}")
# Load state dict
if os.path.exists(image_model_path):
state_dict = torch.load(image_model_path, map_location=DEVICE)
image_model.load_state_dict(state_dict)
image_model.to(DEVICE).eval()
else:
print("❌ Image model weight file not found!")
val_transforms = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
# =====================
# UTILS
# =====================
def seg_pyvi(text: str) -> str:
try:
return ViTokenizer.tokenize(text).replace(" ", "_")
except Exception:
return text
def split_sentences(text: str):
return [s for s in re.split(r'(?<=[.!?])\s+|\n+', text.strip()) if s.strip()]
def predict_text(text: str):
sentences = split_sentences(text)
results = []
# Kiểm tra xem label "an_toan" có trong dict không, nếu không lấy key đầu tiên làm safe label
safe_label = "an_toan"
safe_id = label2id.get(safe_label, 0)
for sent in sentences:
seg = seg_pyvi(sent)
inputs = tokenizer(seg, truncation=True, padding="max_length",
max_length=128, return_tensors="pt").to(DEVICE)
with torch.no_grad():
logits = text_model(**inputs).logits
probs = F.softmax(logits, dim=-1).cpu().numpy()[0]
pred_id = probs.argmax()
label = id2label_text.get(int(pred_id), "Unknown")
prob = float(probs[pred_id])
if label != safe_label and prob >= THRESHOLD:
results.append({"sentence": sent, "label": label, "confidence": prob})
else:
# Chỉ append nếu bạn muốn log cả câu an toàn
results.append({
"sentence": sent,
"label": safe_label,
"confidence": float(probs[safe_id]) if safe_id < len(probs) else 0.0
})
return results
def predict_image(pil_image: Image.Image):
img = val_transforms(pil_image).unsqueeze(0).to(DEVICE)
with torch.no_grad():
outputs = image_model(img)
probs = F.softmax(outputs, dim=1)[0].cpu().numpy()
pred_id = probs.argmax()
label = class_names[pred_id]
prob = float(probs[pred_id])
if label != "an_toan" and prob >= THRESHOLD:
return {"label": label, "confidence": prob}
else:
return {
"label": "an_toan",
"confidence": float(probs[class_names.index("an_toan")])
}
# =====================
# ROUTES
# =====================
@app.route("/")
def home():
return jsonify({"message": "✅ AI moderation API is running!"})
@app.route("/analyze", methods=["POST"])
def analyze():
result = {"text_result": [], "image_result": []}
try:
# Xử lý Text
if "content" in request.form:
text_input = request.form["content"]
if text_input:
result["text_result"] = predict_text(text_input)
# Xử lý Image
if "image" in request.files:
image_files = request.files.getlist("image")
for image_file in image_files:
try:
image = Image.open(image_file.stream).convert("RGB")
image_result = predict_image(image)
result["image_result"].append(image_result)
except Exception as img_err:
print(f"❌ Error processing image: {img_err}")
result["image_result"].append({"error": "Invalid image file"})
return jsonify(result)
except Exception as e:
print(f"❌ Server Error: {str(e)}")
return jsonify({"error": str(e)}), 500
if __name__ == "__main__":
port = int(os.environ.get("PORT", 7860))
app.run(host="0.0.0.0", port=port)