Spaces:
Sleeping
Sleeping
Download app.py from Badal023/Symptom_Classifier: direct link, hf CLI and curl.
- Browser
- Download file 4.25 kB
-
https://huggingface.co/spaces/Badal023/Symptom_Classifier/resolve/main/app.py
- Command line
-
hf download hf://spaces/Badal023/Symptom_Classifier/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Badal023/Symptom_Classifier/resolve/main/app.py
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 ---------------- | |
| 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 ---------------- | |
| def root(): | |
| return jsonify({ | |
| "status": "running", | |
| "message": "AI Triage API live", | |
| "endpoints": { | |
| "audio": "/process-audio", | |
| "text": "/process-text" | |
| } | |
| }) | |
| # 🎤 AUDIO | |
| 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 | |
| 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) |