File size: 4,247 Bytes
de1fd83
8de178b
de1fd83
44eb175
de1fd83
8de178b
de1fd83
 
 
 
6b69b85
8de178b
 
 
 
de1fd83
 
 
 
 
0913f8b
 
de1fd83
0913f8b
 
44eb175
0913f8b
 
 
 
 
 
 
 
 
de1fd83
44eb175
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
de1fd83
 
44eb175
 
6b69b85
44eb175
 
6b69b85
44eb175
 
 
 
6b69b85
44eb175
6b69b85
44eb175
 
 
 
 
 
 
6b69b85
b090625
 
 
6b69b85
b090625
 
 
 
 
6b69b85
b090625
6b69b85
44eb175
 
 
b090625
 
 
44eb175
 
 
b090625
 
de1fd83
b090625
8de178b
 
 
 
6b69b85
8de178b
 
 
 
 
de1fd83
b090625
 
 
de1fd83
0913f8b
8de178b
0913f8b
de1fd83
 
 
 
 
 
 
 
6b69b85
 
 
 
de1fd83
8de178b
0913f8b
 
8de178b
 
 
 
 
0913f8b
8de178b
de1fd83
b090625
de1fd83
b090625
 
 
 
 
de1fd83
b090625
 
 
de1fd83
b090625
de1fd83
b090625
8de178b
de1fd83
b090625
de1fd83
 
b090625
0913f8b
b090625
de1fd83
 
0913f8b
de1fd83
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
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)