Badal023 commited on
Commit
6b69b85
·
verified ·
1 Parent(s): 8de178b

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +39 -15
app.py CHANGED
@@ -9,7 +9,7 @@ import torch.nn.functional as F
9
 
10
  app = Flask(__name__)
11
 
12
- # ---------------- GLOBAL ERROR HANDLER ----------------
13
  @app.errorhandler(Exception)
14
  def handle_error(e):
15
  return jsonify({"error": str(e)}), 500
@@ -44,23 +44,46 @@ def clean_text(text):
44
  outputs = bart_model.generate(**inputs, max_length=50)
45
  return bart_tokenizer.decode(outputs[0], skip_special_tokens=True)
46
 
 
47
  ref_texts = {
48
- "emergency": "severe chest pain unconscious breathing failure stroke heart attack",
49
- "urgent": "fever infection moderate pain vomiting weakness",
50
- "non-urgent": "mild headache cold cough stable"
51
  }
52
 
53
  ref_embeddings = {k: get_embedding(v) for k, v in ref_texts.items()}
54
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
55
  # ---------------- CORE PIPELINE ----------------
56
  def run_pipeline(text):
57
 
 
58
  try:
59
  text_en = GoogleTranslator(source="auto", target="en").translate(text)
60
  except:
61
  text_en = text
62
 
 
63
  cleaned = clean_text(text_en)
 
 
64
  emb = get_embedding(cleaned)
65
 
66
  scores = {
@@ -71,8 +94,13 @@ def run_pipeline(text):
71
  label = max(scores, key=scores.get)
72
  confidence = scores[label]
73
 
74
- score_map = {"emergency": 90, "urgent": 60, "non-urgent": 20}
75
- score = int(score_map[label] * max(confidence, 0.5))
 
 
 
 
 
76
 
77
  return {
78
  "processed_text": cleaned,
@@ -83,12 +111,11 @@ def run_pipeline(text):
83
 
84
  # ---------------- ROUTES ----------------
85
 
86
- # ✅ FIXED ROOT (NO HTML)
87
  @app.route("/", methods=["GET"])
88
  def root():
89
  return jsonify({
90
  "status": "running",
91
- "message": "AI Triage API is live",
92
  "endpoints": {
93
  "audio": "/process-audio",
94
  "text": "/process-text"
@@ -110,13 +137,10 @@ def process_audio():
110
 
111
  file.save(input_path)
112
 
113
- try:
114
- subprocess.run([
115
- "ffmpeg", "-y", "-i", input_path,
116
- "-ac", "1", "-ar", "16000", wav_path
117
- ], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
118
- except:
119
- return jsonify({"error": "FFmpeg failed"}), 500
120
 
121
  segments, _ = whisper_model.transcribe(wav_path)
122
  text = " ".join([seg.text for seg in segments]).strip()
 
9
 
10
  app = Flask(__name__)
11
 
12
+ # ---------------- ERROR HANDLER ----------------
13
  @app.errorhandler(Exception)
14
  def handle_error(e):
15
  return jsonify({"error": str(e)}), 500
 
44
  outputs = bart_model.generate(**inputs, max_length=50)
45
  return bart_tokenizer.decode(outputs[0], skip_special_tokens=True)
46
 
47
+ # ---------------- BETTER REFERENCE ----------------
48
  ref_texts = {
49
+ "emergency": "severe chest pain heart attack breathing difficulty unconscious stroke heavy bleeding",
50
+ "urgent": "high fever infection vomiting dehydration moderate pain weakness",
51
+ "non-urgent": "common cold runny nose mild headache sneezing cough minor symptoms"
52
  }
53
 
54
  ref_embeddings = {k: get_embedding(v) for k, v in ref_texts.items()}
55
 
56
+ # ---------------- RULE ENGINE (KEY FIX) ----------------
57
+ def apply_rules(text, label, score):
58
+ t = text.lower()
59
+
60
+ # 🚨 Emergency overrides
61
+ if any(x in t for x in ["chest pain", "heart", "breathing", "unconscious", "stroke", "bleeding"]):
62
+ return "emergency", max(score, 90)
63
+
64
+ # ⚠️ Urgent
65
+ if any(x in t for x in ["fever", "vomiting", "infection", "high temperature"]):
66
+ return "urgent", max(score, 60)
67
+
68
+ # ✅ Non-urgent
69
+ if any(x in t for x in ["cold", "runny nose", "sneezing", "mild headache"]):
70
+ return "non-urgent", min(score, 25)
71
+
72
+ return label, score
73
+
74
  # ---------------- CORE PIPELINE ----------------
75
  def run_pipeline(text):
76
 
77
+ # Translate
78
  try:
79
  text_en = GoogleTranslator(source="auto", target="en").translate(text)
80
  except:
81
  text_en = text
82
 
83
+ # Clean
84
  cleaned = clean_text(text_en)
85
+
86
+ # Embedding
87
  emb = get_embedding(cleaned)
88
 
89
  scores = {
 
94
  label = max(scores, key=scores.get)
95
  confidence = scores[label]
96
 
97
+ base_map = {"emergency": 90, "urgent": 60, "non-urgent": 20}
98
+
99
+ # 🔥 Improved scoring
100
+ score = int((base_map[label] * 0.7) + (confidence * 30))
101
+
102
+ # 🔥 Apply rules (CRITICAL)
103
+ label, score = apply_rules(cleaned, label, score)
104
 
105
  return {
106
  "processed_text": cleaned,
 
111
 
112
  # ---------------- ROUTES ----------------
113
 
 
114
  @app.route("/", methods=["GET"])
115
  def root():
116
  return jsonify({
117
  "status": "running",
118
+ "message": "AI Triage API live",
119
  "endpoints": {
120
  "audio": "/process-audio",
121
  "text": "/process-text"
 
137
 
138
  file.save(input_path)
139
 
140
+ subprocess.run([
141
+ "ffmpeg", "-y", "-i", input_path,
142
+ "-ac", "1", "-ar", "16000", wav_path
143
+ ], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL)
 
 
 
144
 
145
  segments, _ = whisper_model.transcribe(wav_path)
146
  text = " ".join([seg.text for seg in segments]).strip()