Badal023's picture
Update app.py
44eb175 verified
Raw History Blame Contribute Delete
4.25 kB
import os, uuid, subprocess
from flask import Flask, request, jsonify
from faster_whisper import WhisperModel
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
from deep_translator import GoogleTranslator
import torch
app = Flask(__name__)
# ---------------- ERROR HANDLER ----------------
@app.errorhandler(Exception)
def handle_error(e):
return jsonify({"error": str(e)}), 500
# ---------------- LOAD MODELS ----------------
print("Loading models...")
whisper_model = WhisperModel("small", device="cpu", compute_type="int8")
bart_tokenizer = AutoTokenizer.from_pretrained("facebook/bart-base")
bart_model = AutoModelForSeq2SeqLM.from_pretrained("facebook/bart-base")
print("Models loaded successfully")
# ---------------- CLEAN TEXT ----------------
def clean_text(text):
with torch.no_grad():
inputs = bart_tokenizer(
"Extract key medical symptoms: " + text,
return_tensors="pt",
truncation=True
)
outputs = bart_model.generate(**inputs, max_length=50)
return bart_tokenizer.decode(outputs[0], skip_special_tokens=True)
# ---------------- SYMPTOM WEIGHTS ----------------
SYMPTOM_WEIGHTS = {
# 🚨 Critical
"chest pain": 50,
"breathing": 50,
"shortness of breath": 50,
"unconscious": 60,
"stroke": 60,
"bleeding": 50,
# ⚠️ Moderate
"fever": 25,
"vomiting": 25,
"weakness": 20,
"body pain": 15,
"dizziness": 20,
# ✅ Mild
"cold": 5,
"runny nose": 5,
"sneezing": 5,
"cough": 5,
"gas": 5,
"acidity": 5,
"headache": 5
}
# ---------------- SCORING ----------------
def rule_based_score(text):
t = text.lower()
score = 0
matched = []
for symptom, weight in SYMPTOM_WEIGHTS.items():
if symptom in t:
score += weight
matched.append(symptom)
return score, matched
def map_severity(score):
if score >= 70:
return "emergency"
elif score >= 30:
return "urgent"
else:
return "non-urgent"
# ---------------- CORE PIPELINE ----------------
def run_pipeline(text):
# Translate
try:
text_en = GoogleTranslator(source="auto", target="en").translate(text)
except:
text_en = text
# Clean
cleaned = clean_text(text_en)
# Rule scoring
score, matched = rule_based_score(cleaned)
severity = map_severity(score)
return {
"processed_text": cleaned,
"matched_symptoms": matched,
"severity": severity,
"score": min(score, 100)
}
# ---------------- ROUTES ----------------
@app.route("/", methods=["GET"])
def root():
return jsonify({
"status": "running",
"message": "AI Triage API live",
"endpoints": {
"audio": "/process-audio",
"text": "/process-text"
}
})
# 🎤 AUDIO
@app.route("/process-audio", methods=["POST"])
def process_audio():
if "audio" not in request.files:
return jsonify({"error": "No audio file"}), 400
file = request.files["audio"]
uid = str(uuid.uuid4())
input_path = f"{uid}.webm"
wav_path = f"{uid}.wav"
file.save(input_path)
subprocess.run([
"ffmpeg", "-y", "-i", input_path,
"-ac", "1", "-ar", "16000", wav_path
], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
segments, _ = whisper_model.transcribe(wav_path)
text = " ".join([seg.text for seg in segments]).strip()
# cleanup
os.remove(input_path)
if os.path.exists(wav_path):
os.remove(wav_path)
if not text:
return jsonify({"error": "No speech detected"}), 400
result = run_pipeline(text)
return jsonify({
"input_type": "audio",
"original_text": text,
**result
})
# ⌨️ TEXT
@app.route("/process-text", methods=["POST"])
def process_text():
text = request.form.get("text")
if not text:
return jsonify({"error": "No text provided"}), 400
result = run_pipeline(text)
return jsonify({
"input_type": "text",
"original_text": text,
**result
})
# ---------------- RUN ----------------
if __name__ == "__main__":
app.run(host="0.0.0.0", port=7860)