medtechdev's picture
Upload app.py
1292d2d verified
Raw
History Blame Contribute Delete
16.9 kB
"""
NON-INVASIVE ANEMIA SCREENING - Palpebral & Buccal
===================================================
Palpebral conjunctiva (required) + Buccal mucosa (optional)
"""
import gradio as gr
import numpy as np
import pandas as pd
import pickle
import cv2
import uuid
from datetime import datetime
from pathlib import Path
# ============================================================
# CONFIG
# ============================================================
MODEL_PATH = "model.pkl"
YOLO_PALP = "best.pt"
YOLO_BUCC = "buccal_best.pt"
DATA_LOG = "collected_data.csv"
# ============================================================
# LOAD MODELS
# ============================================================
def load_models():
global ml_model, feature_names, yolo_palp, yolo_bucc
with open(MODEL_PATH, 'rb') as f:
data = pickle.load(f)
ml_model = data['model']
feature_names = data['features']
print(f"ML model: {data.get('name', 'Unknown')}, {len(feature_names)} features")
from ultralytics import YOLO
yolo_palp = YOLO(YOLO_PALP)
print("Palpebral YOLO loaded")
try:
yolo_bucc = YOLO(YOLO_BUCC)
print("Buccal YOLO loaded")
except:
yolo_bucc = None
print("Buccal YOLO not found")
# ============================================================
# IMAGE PROCESSING
# ============================================================
def calibrate_image(img):
try:
if img is None:
return None
if len(img.shape) == 2:
img_rgb = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)
elif img.shape[2] == 4:
img_rgb = cv2.cvtColor(img, cv2.COLOR_RGBA2RGB)
else:
img_rgb = img.copy()
gray = cv2.cvtColor(img_rgb, cv2.COLOR_RGB2GRAY)
_, thresh = cv2.threshold(gray, 180, 255, cv2.THRESH_BINARY)
contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
if not contours:
return img_rgb
c = max(contours, key=cv2.contourArea)
x, y, w, h = cv2.boundingRect(c)
if w * h < 1000:
return img_rgb
patch = img_rgb[y+h//4:y+3*h//4, x+w//4:x+3*w//4]
if patch.size == 0:
return img_rgb
avg = np.mean(patch, axis=(0, 1))
if any(v < 10 for v in avg):
return img_rgb
scale = np.array([128, 128, 128]) / avg
return np.clip(img_rgb * scale, 0, 255).astype(np.uint8)
except:
return img
def extract_features(img, yolo, prefix):
if img is None or yolo is None:
return None
try:
img_cal = calibrate_image(img)
if img_cal is None:
return None
results = yolo(img_cal, verbose=False)
if len(results) > 0 and results[0].masks is not None:
masks = results[0].masks.data.cpu().numpy()
confs = results[0].boxes.conf.cpu().numpy()
if len(confs) > 0:
m = masks[np.argmax(confs)]
m = cv2.resize(m, (img_cal.shape[1], img_cal.shape[0]))
mask = (m > 0.5).astype(np.uint8) * 255
else:
h, w = img_cal.shape[:2]
mask = np.zeros((h, w), dtype=np.uint8)
mask[h//4:3*h//4, w//4:3*w//4] = 255
else:
h, w = img_cal.shape[:2]
mask = np.zeros((h, w), dtype=np.uint8)
mask[h//4:3*h//4, w//4:3*w//4] = 255
px = img_cal[mask > 0]
if len(px) < 100:
return None
R, G, B = [np.mean(px[:, i]) for i in range(3)]
hsv = cv2.cvtColor(img_cal, cv2.COLOR_RGB2HSV)
H, S, V = [np.mean(hsv[mask > 0][:, i]) for i in range(3)]
lab = cv2.cvtColor(img_cal, cv2.COLOR_RGB2LAB)
L, A, Bc = [np.mean(lab[mask > 0][:, i]) for i in range(3)]
return {
f'{prefix}_R': R, f'{prefix}_G': G, f'{prefix}_B': B,
f'{prefix}_H': H, f'{prefix}_S': S, f'{prefix}_V': V,
f'{prefix}_L_color': L, f'{prefix}_A_color': A, f'{prefix}_B_color': Bc
}
except Exception as e:
print(f"Error {prefix}: {e}")
return None
# ============================================================
# FORM TO FEATURES
# ============================================================
def form_to_features(age, income_range, family_size, diet,
food_grains, food_leafy, food_cereals, food_pulses,
food_fish, food_eggs, food_meat, food_jaggery, food_nuts,
fruit_freq, veg_freq, tea_coffee,
comorb_dm, comorb_htn, comorb_thyroid,
prev_anemia, family_anemia, hookworm, bleeding_disorder, deworming,
menstrual, pads_per_day,
gravida, para, live_births, abortions, stillbirths,
family_dm, family_htn):
income_map = {"Below ₹5,000": 2500, "₹5,000 - ₹10,000": 7500, "₹10,000 - ₹15,000": 12500,
"₹15,000 - ₹25,000": 20000, "₹25,000 - ₹50,000": 37500, "Above ₹50,000": 75000}
freq_map = {"Daily": 4, "Twice a week": 3, "Once a week": 2, "Rarely": 1, "Never": 0}
pads_map = {"1-2": 1.5, "2-3": 2.5, "3-4": 3.5, "4-5": 4.5, "More than 5": 6, "Not applicable": 2.5}
income = income_map.get(income_range, 20000)
return {
'age': age, 'income': income, 'family_size': family_size,
'income_per_capita': income / max(family_size, 1),
'diet_veg': 1 if diet == "Vegetarian" else 0,
'diet_nonveg': 1 if diet == "Non-vegetarian" else 0,
'diet_mixed': 1 if diet == "Mixed" else 0,
'food_whole_grains': int(food_grains), 'food_green_leafy_vegetables': int(food_leafy),
'food_cereals': int(food_cereals), 'food_pulses': int(food_pulses),
'food_fish': int(food_fish), 'food_eggs': int(food_eggs),
'food_meat': int(food_meat), 'food_jaggery': int(food_jaggery), 'food_nuts': int(food_nuts),
'food_iron_count': sum([food_grains, food_leafy, food_cereals, food_pulses, food_fish, food_eggs, food_meat, food_jaggery, food_nuts]),
'fruit_freq': freq_map.get(fruit_freq, 2), 'veg_freq': freq_map.get(veg_freq, 2),
'tea_coffee_daily': 1 if tea_coffee == "Yes" else 0, 'tea_coffee_daily_missing': 0,
'comorb_dm': int(comorb_dm), 'comorb_htn': int(comorb_htn), 'comorb_thyroid': int(comorb_thyroid),
'comorb_none': 1 if not any([comorb_dm, comorb_htn, comorb_thyroid]) else 0,
'comorb_count': sum([int(comorb_dm), int(comorb_htn), int(comorb_thyroid)]),
'prev_anemia': 1 if prev_anemia == "Yes" else 0, 'prev_anemia_missing': 0,
'family_anemia': 1 if family_anemia == "Yes" else 0, 'family_anemia_missing': 0,
'hookworm_history': 1 if hookworm == "Yes" else 0, 'hookworm_history_missing': 0,
'bleeding_disorder': 1 if bleeding_disorder == "Yes" else 0, 'bleeding_disorder_missing': 0,
'deworming_done': 1 if deworming == "Yes" else 0, 'deworming_done_missing': 0,
'menstrual_regular': 1 if menstrual == "Regular" else 0,
'menstrual_irregular': 1 if menstrual == "Irregular" else 0,
'pads_per_day': pads_map.get(pads_per_day, 2.5), 'pads_missing': 0,
'gravida': gravida, 'para': para, 'live_births': live_births,
'abortions': abortions, 'stillbirths': stillbirths,
'has_obstetric_history': 1 if gravida > 0 else 0,
'pregnancy_loss': abortions + stillbirths,
'parity_ratio': para / gravida if gravida > 0 else 0,
'family_dm': 1 if family_dm == "Yes" else 0, 'family_dm_missing': 0,
'family_htn': 1 if family_htn == "Yes" else 0, 'family_htn_missing': 0,
}
# ============================================================
# PREDICTION
# ============================================================
def predict_hb(img_palp, img_bucc,
age, income_range, family_size, diet,
food_grains, food_leafy, food_cereals, food_pulses,
food_fish, food_eggs, food_meat, food_jaggery, food_nuts,
fruit_freq, veg_freq, tea_coffee,
comorb_dm, comorb_htn, comorb_thyroid,
prev_anemia, family_anemia, hookworm, bleeding_disorder, deworming,
menstrual, pads_per_day,
gravida, para, live_births, abortions, stillbirths,
family_dm, family_htn, actual_hb):
session_id = str(uuid.uuid4())[:8]
sites = []
if img_palp is None:
return "❌ Palpebral image required.", "", ""
all_features = {}
# Palpebral (required)
palp_feat = extract_features(img_palp, yolo_palp, 'palpebral')
if palp_feat is None:
return "❌ Could not process palpebral image.", "", ""
all_features.update(palp_feat)
sites.append('palpebral')
# Buccal (optional)
if img_bucc is not None and yolo_bucc:
bucc_feat = extract_features(img_bucc, yolo_bucc, 'buccal')
if bucc_feat:
all_features.update(bucc_feat)
sites.append('buccal')
# Tabular
tab_feat = form_to_features(age, income_range, family_size, diet,
food_grains, food_leafy, food_cereals, food_pulses,
food_fish, food_eggs, food_meat, food_jaggery, food_nuts,
fruit_freq, veg_freq, tea_coffee,
comorb_dm, comorb_htn, comorb_thyroid,
prev_anemia, family_anemia, hookworm, bleeding_disorder, deworming,
menstrual, pads_per_day,
gravida, para, live_births, abortions, stillbirths,
family_dm, family_htn)
all_features.update(tab_feat)
# Default values for missing image features
defaults = {
'palpebral_R': 150, 'palpebral_G': 100, 'palpebral_B': 100,
'palpebral_H': 10, 'palpebral_S': 100, 'palpebral_V': 150,
'palpebral_L_color': 130, 'palpebral_A_color': 150, 'palpebral_B_color': 140,
'buccal_R': 160, 'buccal_G': 110, 'buccal_B': 110,
'buccal_H': 8, 'buccal_S': 90, 'buccal_V': 160,
'buccal_L_color': 140, 'buccal_A_color': 145, 'buccal_B_color': 135,
}
for f in feature_names:
if f not in all_features:
all_features[f] = defaults.get(f, 0)
X = pd.DataFrame([all_features])[feature_names]
hb = round(float(ml_model.predict(X)[0]), 1)
if hb >= 12:
sev, col, rec = "Normal", "🟢", "No anemia detected."
elif hb >= 11:
sev, col, rec = "Mild Anemia", "🟡", "Consider iron-rich foods."
elif hb >= 8:
sev, col, rec = "Moderate Anemia", "🟠", "Consult healthcare provider."
else:
sev, col, rec = "Severe Anemia", "🔴", "Seek immediate medical attention."
# Log
try:
record = {'session_id': session_id, 'timestamp': datetime.now().isoformat(),
'predicted_hb': hb, 'actual_hb': actual_hb if actual_hb > 0 else None,
'sites': ','.join(sites), **all_features}
df = pd.DataFrame([record])
if Path(DATA_LOG).exists():
df.to_csv(DATA_LOG, mode='a', header=False, index=False)
else:
df.to_csv(DATA_LOG, index=False)
except:
pass
result = f"## {col} Estimated Hb: {hb} g/dL\n### {sev}\n{rec}\n\n---\n**Sites:** {', '.join(sites)} | **Session:** {session_id}"
return result, session_id, f"{hb} g/dL"
# ============================================================
# INTERFACE
# ============================================================
def create_app():
with gr.Blocks(title="Anemia Screening", theme=gr.themes.Soft()) as app:
gr.Markdown("# 🩺 Non-Invasive Anemia Screening\n**Palpebral conjunctiva required. Buccal mucosa optional.**")
with gr.Row():
with gr.Column(scale=1):
img_palp = gr.Image(label="Palpebral Conjunctiva (Required)", type="numpy")
img_bucc = gr.Image(label="Buccal Mucosa (Optional)", type="numpy")
with gr.Column(scale=2):
with gr.Accordion("Demographics", open=True):
with gr.Row():
age = gr.Slider(10, 80, 25, step=1, label="Age")
family_size = gr.Slider(1, 15, 4, step=1, label="Family Size")
income_range = gr.Dropdown(["Below ₹5,000", "₹5,000 - ₹10,000", "₹10,000 - ₹15,000", "₹15,000 - ₹25,000", "₹25,000 - ₹50,000", "Above ₹50,000"], value="₹10,000 - ₹15,000", label="Income")
with gr.Accordion("Diet", open=True):
diet = gr.Radio(["Vegetarian", "Non-vegetarian", "Mixed"], value="Mixed", label="Diet")
with gr.Row():
food_grains = gr.Checkbox(label="Whole Grains", value=True)
food_leafy = gr.Checkbox(label="Leafy Veg", value=True)
food_cereals = gr.Checkbox(label="Cereals", value=True)
with gr.Row():
food_pulses = gr.Checkbox(label="Pulses", value=True)
food_fish = gr.Checkbox(label="Fish")
food_eggs = gr.Checkbox(label="Eggs")
with gr.Row():
food_meat = gr.Checkbox(label="Meat")
food_jaggery = gr.Checkbox(label="Jaggery")
food_nuts = gr.Checkbox(label="Nuts")
with gr.Row():
fruit_freq = gr.Dropdown(["Daily", "Twice a week", "Once a week", "Rarely", "Never"], value="Twice a week", label="Fruits")
veg_freq = gr.Dropdown(["Daily", "Twice a week", "Once a week", "Rarely", "Never"], value="Daily", label="Vegetables")
tea_coffee = gr.Radio(["Yes", "No"], value="Yes", label="Daily Tea/Coffee")
with gr.Accordion("Medical History", open=False):
with gr.Row():
comorb_dm = gr.Checkbox(label="Diabetes")
comorb_htn = gr.Checkbox(label="Hypertension")
comorb_thyroid = gr.Checkbox(label="Thyroid")
with gr.Row():
prev_anemia = gr.Radio(["Yes", "No"], value="No", label="Anemia (past 12mo)")
family_anemia = gr.Radio(["Yes", "No"], value="No", label="Family Anemia")
with gr.Row():
hookworm = gr.Radio(["Yes", "No"], value="No", label="Hookworm")
bleeding_disorder = gr.Radio(["Yes", "No"], value="No", label="Bleeding Disorder")
deworming = gr.Radio(["Yes", "No"], value="Yes", label="Deworming Done")
with gr.Row():
family_dm = gr.Radio(["Yes", "No"], value="No", label="Family Diabetes")
family_htn = gr.Radio(["Yes", "No"], value="No", label="Family Hypertension")
with gr.Accordion("Menstrual/Obstetric", open=False):
with gr.Row():
menstrual = gr.Radio(["Regular", "Irregular", "Not applicable"], value="Regular", label="Menstrual")
pads_per_day = gr.Dropdown(["1-2", "2-3", "3-4", "4-5", "More than 5", "Not applicable"], value="2-3", label="Pads/day")
with gr.Row():
gravida = gr.Slider(0, 10, 0, step=1, label="Gravida")
para = gr.Slider(0, 10, 0, step=1, label="Para")
live_births = gr.Slider(0, 10, 0, step=1, label="Live Births")
with gr.Row():
abortions = gr.Slider(0, 5, 0, step=1, label="Abortions")
stillbirths = gr.Slider(0, 5, 0, step=1, label="Stillbirths")
with gr.Row():
actual_hb = gr.Number(label="Actual Hb (optional)", value=0, minimum=0, maximum=25)
submit = gr.Button("🔬 Estimate Hb", variant="primary", size="lg")
result = gr.Markdown()
with gr.Row():
session = gr.Textbox(label="Session ID", interactive=False)
hb_out = gr.Textbox(label="Estimated Hb", interactive=False)
gr.Markdown("---\n🔒 Anonymous data collection. No PII stored.")
submit.click(predict_hb,
[img_palp, img_bucc, age, income_range, family_size, diet,
food_grains, food_leafy, food_cereals, food_pulses, food_fish, food_eggs, food_meat, food_jaggery, food_nuts,
fruit_freq, veg_freq, tea_coffee, comorb_dm, comorb_htn, comorb_thyroid,
prev_anemia, family_anemia, hookworm, bleeding_disorder, deworming,
menstrual, pads_per_day, gravida, para, live_births, abortions, stillbirths,
family_dm, family_htn, actual_hb],
[result, session, hb_out])
return app
if __name__ == "__main__":
load_models()
create_app().launch()