diff --git a/Dockerfile b/Dockerfile
index 257db7661393709c3f5ec4052b35ce2fee0e8051..31df3983c59922d719290478fe78fb20a00dd632 100644
--- a/Dockerfile
+++ b/Dockerfile
@@ -7,26 +7,13 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
libglib2.0-0 libsm6 libxext6 libxrender-dev libgl1 \
&& rm -rf /var/lib/apt/lists/*
-# Copy and install Python deps first (layer cache)
+# Copy and install backend runtime deps first.
COPY backend/requirements.txt .
-RUN pip install --no-cache-dir -r requirements.txt
+RUN pip install --no-cache-dir --upgrade pip && pip install --no-cache-dir -r requirements.txt
# Copy backend source
COPY backend/ .
-# Lock PyTorch to single thread for memory-constrained deployments
-ENV OMP_NUM_THREADS=1
-ENV MKL_NUM_THREADS=1
-ENV PYTHONUNBUFFERED=1
+EXPOSE 5000
-EXPOSE 7860
-
-# Production: uvicorn with FastAPI (single worker โ memory constrained)
-CMD ["uvicorn", "app.main:app", \
- "--host", "0.0.0.0", \
- "--port", "7860", \
- "--workers", "1", \
- "--timeout-keep-alive", "30", \
- "--limit-concurrency", "4", \
- "--proxy-headers", \
- "--forwarded-allow-ips=*"]
+CMD ["python", "-m", "uvicorn", "app.main:app", "--host", "0.0.0.0", "--port", "5000"]
diff --git a/README.md b/README.md
index 6db5dcd2e7e295b964dbca84f9de78561d30fcc8..c5ce67eb3040ce4feaee32180da3dd8909b04de7 100644
--- a/README.md
+++ b/README.md
@@ -4,7 +4,7 @@ emoji: ๐ฉธ
colorFrom: red
colorTo: gray
sdk: docker
-app_port: 7860
+app_port: 5000
base_path: /docs
pinned: false
license: mit
diff --git a/backend/.env.example b/backend/.env.example
new file mode 100644
index 0000000000000000000000000000000000000000..e5f0f33780edcd63d1c589a558e8c8a69de747a4
--- /dev/null
+++ b/backend/.env.example
@@ -0,0 +1,61 @@
+# AnemiaLens backend environment variables
+# Copy this file to .env and fill in the secrets.
+# NEVER commit .env to version control.
+
+# -- Mistral AI guidance --
+ANEMIALENS_MISTRAL_API_KEY=
+ANEMIALENS_MISTRAL_ENABLED=true
+ANEMIALENS_MISTRAL_MODEL=mistral-small-latest
+ANEMIALENS_GUIDANCE_TIMEOUT=20
+
+# -- Server --
+ANEMIALENS_LOG_LEVEL=INFO
+ANEMIALENS_CORS_ORIGINS=["http://localhost:5173","http://127.0.0.1:5173"]
+
+# -- Database --
+DATABASE_URL=sqlite+aiosqlite:///./anemialens.db
+
+# -- Auth (JWT) --
+JWT_SECRET_KEY=change-me-to-a-random-64-char-string
+JWT_ALGORITHM=HS256
+JWT_ACCESS_TOKEN_EXPIRE_MINUTES=60
+JWT_REFRESH_TOKEN_EXPIRE_DAYS=30
+
+# -- Rate limiting --
+ANEMIALENS_RATE_LIMIT_ANALYZE=10
+ANEMIALENS_RATE_LIMIT_QUALITY=30
+
+# -- Hosted email reports (recommended on Hugging Face Spaces) --
+# Gmail API works over HTTPS and is the current hosted delivery path.
+ANEMIALENS_EMAIL_PROVIDER=gmail_api
+ANEMIALENS_GMAIL_CLIENT_ID=
+ANEMIALENS_GMAIL_CLIENT_SECRET=
+ANEMIALENS_GMAIL_REFRESH_TOKEN=
+ANEMIALENS_EMAIL_FROM_NAME=AnemiaLens
+ANEMIALENS_EMAIL_FROM_EMAIL=
+ANEMIALENS_EMAIL_REPLY_TO=
+
+# -- SMTP fallback (local or non-restricted hosts) --
+# ANEMIALENS_EMAIL_PROVIDER=smtp
+# ANEMIALENS_SMTP_HOST=smtp.gmail.com
+# ANEMIALENS_SMTP_PORT=465
+# ANEMIALENS_SMTP_USE_SSL=true
+# ANEMIALENS_SMTP_USE_STARTTLS=false
+# ANEMIALENS_SMTP_USERNAME=your.gmail@gmail.com
+# ANEMIALENS_SMTP_PASSWORD=your-16-char-app-password
+# ANEMIALENS_SMTP_TIMEOUT=20
+
+# -- HTTP API alternatives --
+# Resend:
+# ANEMIALENS_EMAIL_PROVIDER=resend
+# ANEMIALENS_RESEND_API_KEY=
+# ANEMIALENS_EMAIL_FROM_NAME=AnemiaLens
+# ANEMIALENS_EMAIL_FROM_EMAIL=onboarding@resend.dev
+# ANEMIALENS_EMAIL_REPLY_TO=your@email.com
+#
+# SendGrid:
+# ANEMIALENS_EMAIL_PROVIDER=sendgrid
+# ANEMIALENS_SENDGRID_API_KEY=
+# ANEMIALENS_EMAIL_FROM_NAME=AnemiaLens
+# ANEMIALENS_EMAIL_FROM_EMAIL=your_verified_sender@gmail.com
+# ANEMIALENS_EMAIL_REPLY_TO=your_verified_sender@gmail.com
diff --git a/backend/app/api/history.py b/backend/app/api/history.py
index 6e93ce12d6e8b4e3c0b206de980ea73bb613b3af..9a4f4e8b46632c3258967f070520053302d144be 100644
--- a/backend/app/api/history.py
+++ b/backend/app/api/history.py
@@ -17,6 +17,8 @@ from app.database import get_db
from app.dependencies import get_current_user
from app.models.screening import Screening
from app.models.user import User
+from app.schemas import AnalyzeResponse
+from app.services.screening_store import persist_screening_result
log = logging.getLogger("anemialens.history")
@@ -67,6 +69,16 @@ class DeleteResponse(BaseModel):
uid: str
+class SaveScreeningRequest(BaseModel):
+ analysis: AnalyzeResponse
+
+
+class SaveScreeningResponse(BaseModel):
+ saved: bool
+ uid: str
+ message: str
+
+
# ---------------------------------------------------------------------------
# Routes
# ---------------------------------------------------------------------------
@@ -238,6 +250,30 @@ async def delete_screening(
return DeleteResponse(deleted=True, uid=screening_uid)
+@router.post(
+ "/save-current",
+ response_model=SaveScreeningResponse,
+ summary="Save the current screening result to the authenticated account",
+)
+async def save_current_screening(
+ body: SaveScreeningRequest,
+ user: Annotated[User, Depends(get_current_user)],
+) -> SaveScreeningResponse:
+ analysis = body.analysis
+ screening = await persist_screening_result(
+ request_id=analysis.analysis_meta.request_id,
+ analysis=analysis,
+ user_id=user.id,
+ processing_time_ms=analysis.analysis_meta.processing_time_ms,
+ )
+ log.info("Screening saved to account: %s by user %s", screening.uid, user.uid)
+ return SaveScreeningResponse(
+ saved=True,
+ uid=screening.uid,
+ message="Screening saved to your account history.",
+ )
+
+
# ---------------------------------------------------------------------------
# CSV Export (Pro only)
# ---------------------------------------------------------------------------
diff --git a/backend/app/config.py b/backend/app/config.py
index 28d3e325150edff0f090e0682cadbcf8552f8dd0..9ba927597dd2ce69dbeeab4051c049efb97c9d41 100644
--- a/backend/app/config.py
+++ b/backend/app/config.py
@@ -30,10 +30,14 @@ DEFAULT_ENSEMBLE_PATH = MODELS_DIR / "ensemble_model.json"
DEFAULT_DEEP_STACK_PATH = MODELS_DIR / "deep_stack_model.joblib"
DEFAULT_ARCHIVE_MODEL_PATH = MODELS_DIR / "archive_screening_model.joblib"
DEFAULT_EFFICIENTNET_MODEL_PATH = MODELS_DIR / "efficientnet_anemia.pth"
-DEFAULT_EFFICIENTNET_REPORT_PATH = MODELS_DIR / "efficientnet_report.json"
-DEFAULT_RUNTIME_STACK_REPORT_PATH = MODELS_DIR / "runtime_stack_report.json"
-DEFAULT_DEPLOYED_SCREENING_REPORT_PATH = MODELS_DIR / "deployed_screening_report.json"
-DEFAULT_TRAINING_REPORT_PATH = MODELS_DIR / "training_report.json"
+DEFAULT_EFFICIENTNET_REPORT_PATH = MODELS_DIR / "efficientnet_report.json"
+DEFAULT_RUNTIME_STACK_REPORT_PATH = MODELS_DIR / "runtime_stack_report.json"
+DEFAULT_RUNTIME_CALIBRATOR_PATH = MODELS_DIR / "runtime_risk_calibrator.pkl"
+DEFAULT_RUNTIME_CALIBRATION_REPORT_PATH = MODELS_DIR / "runtime_calibration_report.json"
+DEFAULT_RUNTIME_REFINER_PATH = MODELS_DIR / "runtime_screening_refiner.pkl"
+DEFAULT_RUNTIME_REFINEMENT_REPORT_PATH = MODELS_DIR / "runtime_refinement_report.json"
+DEFAULT_DEPLOYED_SCREENING_REPORT_PATH = MODELS_DIR / "deployed_screening_report.json"
+DEFAULT_TRAINING_REPORT_PATH = MODELS_DIR / "training_report.json"
# ---------------------------------------------------------------------------
@@ -83,17 +87,24 @@ class Settings(BaseSettings):
"PRELOAD_MODELS_ON_STARTUP",
),
)
- warmup_models_on_startup: bool = Field(
- default=False,
- validation_alias=AliasChoices(
- "ANEMIALENS_WARMUP_MODELS_ON_STARTUP",
- "WARMUP_MODELS_ON_STARTUP",
- ),
- )
- cors_origins: list[str] = Field(
- default=[
- "http://localhost:5173",
- "http://127.0.0.1:5173",
+ warmup_models_on_startup: bool = Field(
+ default=False,
+ validation_alias=AliasChoices(
+ "ANEMIALENS_WARMUP_MODELS_ON_STARTUP",
+ "WARMUP_MODELS_ON_STARTUP",
+ ),
+ )
+ enable_efficientnet_fallback: bool = Field(
+ default=False,
+ validation_alias=AliasChoices(
+ "ANEMIALENS_ENABLE_EFFICIENTNET_FALLBACK",
+ "ENABLE_EFFICIENTNET_FALLBACK",
+ ),
+ )
+ cors_origins: list[str] = Field(
+ default=[
+ "http://localhost:5173",
+ "http://127.0.0.1:5173",
"http://localhost:5174",
"http://127.0.0.1:5174",
]
diff --git a/backend/app/database.py b/backend/app/database.py
index e45efd5bcc8efb39c663ce585c07122e661f2c74..e9c20756b230fb04cefc1ea487c3bfb2bb466813 100644
--- a/backend/app/database.py
+++ b/backend/app/database.py
@@ -7,13 +7,18 @@ Supports SQLite (dev) and PostgreSQL (production) via DATABASE_URL.
from __future__ import annotations
import os
+from pathlib import Path
+from dotenv import load_dotenv
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from sqlalchemy.orm import DeclarativeBase
+BACKEND_ROOT = Path(__file__).resolve().parents[1]
+load_dotenv(BACKEND_ROOT / ".env")
+
DATABASE_URL = os.getenv("DATABASE_URL", "sqlite+aiosqlite:///./anemialens.db").strip()
-# For PostgreSQL on Render, the URL starts with postgres:// but SQLAlchemy needs postgresql+asyncpg://
+# For managed PostgreSQL providers, postgres:// must be normalized to postgresql+asyncpg://
if DATABASE_URL.startswith("postgres://"):
DATABASE_URL = DATABASE_URL.replace("postgres://", "postgresql+asyncpg://", 1)
elif DATABASE_URL.startswith("postgresql://") and "+asyncpg" not in DATABASE_URL:
diff --git a/backend/app/main.py b/backend/app/main.py
index 1f01df69cfae03a44fabe5f183772d382f88f680..eaf8eae8e155c8fd170abe8277841819790b8418 100644
--- a/backend/app/main.py
+++ b/backend/app/main.py
@@ -33,7 +33,7 @@ from typing import Annotated
from dotenv import load_dotenv
from fastapi import FastAPI, File, Form, Request, UploadFile, status, Depends
from fastapi.middleware.cors import CORSMiddleware
-from fastapi.responses import JSONResponse, RedirectResponse
+from fastapi.responses import JSONResponse, RedirectResponse
from PIL import UnidentifiedImageError
from app.config import BACKEND_ROOT, settings
@@ -43,17 +43,20 @@ from app.services.analysis_meta import build_analysis_meta
from app.services.case_insight import CaseInsightService
from app.services.clinical_brief import ClinicalBriefService
from app.services.decision_audit import build_decision_audit
-from app.services.guidance import GuidanceService
-from app.services.handoff import HandoffSummaryService
-from app.services.image_quality import ImageQualityService
-from app.services.prediction import ScreeningPredictor
-from app.services.request_parsing import (
- InvalidRequestPayload,
- normalize_optional_text,
- parse_symptoms,
-)
-from app.services.runtime_status import build_runtime_status
-from app.services.triage import TriageService
+from app.services.guidance import GuidanceService
+from app.services.handoff import HandoffSummaryService
+from app.services.image_quality import ImageQualityService
+from app.services.patient_case import PatientCaseService
+from app.services.prediction import ScreeningPredictor
+from app.services.request_parsing import (
+ InvalidRequestPayload,
+ normalize_optional_text,
+ parse_patient_profile,
+ parse_symptoms,
+)
+from app.services.runtime_status import build_runtime_status
+from app.services.screening_store import persist_screening_result
+from app.services.triage import TriageService
load_dotenv(BACKEND_ROOT / ".env")
@@ -111,10 +114,11 @@ async def lifespan(app: FastAPI):
app.state.predictor = ScreeningPredictor()
app.state.triage_service = TriageService()
app.state.guidance_service = GuidanceService()
- app.state.case_insight_service = CaseInsightService()
- app.state.clinical_brief_service = ClinicalBriefService()
- app.state.handoff_service = HandoffSummaryService()
- log.info("All ML services initialised.")
+ app.state.case_insight_service = CaseInsightService()
+ app.state.clinical_brief_service = ClinicalBriefService()
+ app.state.handoff_service = HandoffSummaryService()
+ app.state.patient_case_service = PatientCaseService()
+ log.info("All ML services initialised.")
# ---------- Model warm-up ----------
if app.state.predictor.is_ready():
@@ -144,7 +148,7 @@ async def lifespan(app: FastAPI):
# App
# ---------------------------------------------------------------------------
-app = FastAPI(
+app = FastAPI(
title="AnemiaLens API",
version="1.0.0",
description=(
@@ -155,7 +159,7 @@ app = FastAPI(
lifespan=lifespan,
docs_url="/docs",
redoc_url="/redoc",
-)
+)
# ---------------------------------------------------------------------------
# Middleware stack (order matters โ outermost first)
@@ -199,57 +203,57 @@ app.add_middleware(MemoryGuardMiddleware)
# ---------------------------------------------------------------------------
@app.middleware("http")
-async def request_id_middleware(request: Request, call_next):
- request_id = str(uuid.uuid4())[:8]
- request.state.request_id = request_id
- request.state.started_at = time.perf_counter()
-
- try:
- response = await call_next(request)
- except Exception:
- elapsed_ms = (time.perf_counter() - request.state.started_at) * 1000
- log.exception(
- "%s %s -> %d (%.1fms) [%s]",
- request.method,
- request.url.path,
- status.HTTP_500_INTERNAL_SERVER_ERROR,
- elapsed_ms,
- request_id,
- extra={"request_id": request_id},
- )
- raise
-
- elapsed_ms = (time.perf_counter() - request.state.started_at) * 1000
- response.headers["X-Request-ID"] = request_id
- response.headers["X-Response-Time"] = f"{elapsed_ms:.1f}ms"
-
- log.info(
- "%s %s -> %d (%.1fms) [%s]",
- request.method,
- request.url.path,
- response.status_code,
- elapsed_ms,
- request_id,
- extra={"request_id": request_id},
- )
- return response
+async def request_id_middleware(request: Request, call_next):
+ request_id = str(uuid.uuid4())[:8]
+ request.state.request_id = request_id
+ request.state.started_at = time.perf_counter()
+
+ try:
+ response = await call_next(request)
+ except Exception:
+ elapsed_ms = (time.perf_counter() - request.state.started_at) * 1000
+ log.exception(
+ "%s %s -> %d (%.1fms) [%s]",
+ request.method,
+ request.url.path,
+ status.HTTP_500_INTERNAL_SERVER_ERROR,
+ elapsed_ms,
+ request_id,
+ extra={"request_id": request_id},
+ )
+ raise
+
+ elapsed_ms = (time.perf_counter() - request.state.started_at) * 1000
+ response.headers["X-Request-ID"] = request_id
+ response.headers["X-Response-Time"] = f"{elapsed_ms:.1f}ms"
+
+ log.info(
+ "%s %s -> %d (%.1fms) [%s]",
+ request.method,
+ request.url.path,
+ response.status_code,
+ elapsed_ms,
+ request_id,
+ extra={"request_id": request_id},
+ )
+ return response
# ---------------------------------------------------------------------------
# Include API route modules (Phase 2 & 3)
# ---------------------------------------------------------------------------
-from app.api.auth import router as auth_router
-from app.api.history import router as history_router
-from app.api.admin import router as admin_router
-from app.api.billing import router as billing_router
-from app.api.email_report import router as email_report_router
-
-app.include_router(auth_router)
-app.include_router(history_router)
-app.include_router(admin_router)
-app.include_router(billing_router)
-app.include_router(email_report_router)
+from app.api.auth import router as auth_router
+from app.api.history import router as history_router
+from app.api.admin import router as admin_router
+from app.api.billing import router as billing_router
+from app.api.email_report import router as email_report_router
+
+app.include_router(auth_router)
+app.include_router(history_router)
+app.include_router(admin_router)
+app.include_router(billing_router)
+app.include_router(email_report_router)
# ---------------------------------------------------------------------------
@@ -277,12 +281,12 @@ def _too_large_response(request_id: str, max_mb: float) -> JSONResponse:
)
-def _attempt_raw_frame_rescue(services, image_bytes: bytes, quality, symptom_score: float = 0.0):
+def _attempt_raw_frame_rescue(services, image_bytes: bytes, quality):
if quality.passed or not services.quality_service.allows_raw_frame_rescue(quality):
return quality, None, False
raw_image = load_image_bytes(image_bytes).convert("RGB")
- raw_prediction = services.predictor.predict(raw_image, quality, symptom_score=symptom_score)
+ raw_prediction = services.predictor.predict(raw_image, quality)
if not services.predictor.should_accept_raw_frame_rescue(raw_prediction):
return quality, None, False
@@ -294,73 +298,37 @@ def _attempt_raw_frame_rescue(services, image_bytes: bytes, quality, symptom_sco
# Screening persistence helper (Phase 2)
# ---------------------------------------------------------------------------
-async def _persist_screening(
- request_id: str,
- analysis: AnalyzeResponse,
- user_id: int | None,
- processing_time_ms: float,
-) -> None:
- """Save the screening result to the database."""
- try:
- from app.database import async_session_factory
- from app.models.screening import Screening
-
- screening = Screening(
- request_id=request_id,
- user_id=user_id,
- triage_band=analysis.triage.band,
- triage_score=analysis.triage.score,
- triage_label=analysis.triage.label,
- anemia_risk=analysis.prediction.anemia_risk if analysis.prediction else None,
- predicted_hemoglobin=analysis.prediction.predicted_hemoglobin if analysis.prediction else None,
- confidence=analysis.prediction.confidence if analysis.prediction else None,
- uncertainty=analysis.prediction.uncertainty if analysis.prediction else None,
- screening_label=analysis.prediction.screening_label if analysis.prediction else None,
- model_source=analysis.prediction.model_source if analysis.prediction else None,
- quality_passed=analysis.quality.passed,
- blocked=analysis.blocked,
- processing_path=analysis.decision_audit.processing_path,
- guidance_source=analysis.guidance.source,
- symptoms_json=json.dumps(analysis.symptoms.model_dump()),
- full_response_json=json.dumps(analysis.model_dump(), default=str),
- share_text=analysis.handoff_summary.share_text,
- urgency_label=analysis.handoff_summary.urgency_label,
- headline=analysis.handoff_summary.headline,
- processing_time_ms=processing_time_ms,
- language=analysis.language,
- region=analysis.region,
- )
+async def _persist_screening(
+ request_id: str,
+ analysis: AnalyzeResponse,
+ user_id: int | None,
+ processing_time_ms: float,
+) -> None:
+ """Save the screening result to the database."""
+ try:
+ await persist_screening_result(
+ request_id=request_id,
+ analysis=analysis,
+ user_id=user_id,
+ processing_time_ms=processing_time_ms,
+ )
+
+ except Exception as exc:
+ log.warning("Failed to persist screening (non-fatal): %s", exc)
- async with async_session_factory() as session:
- session.add(screening)
- await session.commit()
- # Update user scan count
- if user_id is not None:
- from app.models.user import User
- from sqlalchemy import select
- result = await session.execute(select(User).where(User.id == user_id))
- user = result.scalar_one_or_none()
- if user:
- user.scan_count += 1
- await session.commit()
+# ---------------------------------------------------------------------------
+# Routes โ Health / Meta
+# ---------------------------------------------------------------------------
- except Exception as exc:
- log.warning("Failed to persist screening (non-fatal): %s", exc)
+@app.get("/", include_in_schema=False)
+async def root() -> RedirectResponse:
+ """Redirect the Space root to Swagger UI so Docker Space routing has a valid landing page."""
+ return RedirectResponse(url="/docs", status_code=status.HTTP_307_TEMPORARY_REDIRECT)
-# ---------------------------------------------------------------------------
-# Routes โ Health / Meta
-# ---------------------------------------------------------------------------
-
-@app.get("/", include_in_schema=False)
-async def root() -> RedirectResponse:
- """Redirect the Space root to Swagger UI so Docker Space routing has a valid landing page."""
- return RedirectResponse(url="/docs", status_code=status.HTTP_307_TEMPORARY_REDIRECT)
-
-
-@app.get("/health", tags=["meta"], summary="Liveness probe")
-async def health(request: Request) -> dict[str, object]:
+@app.get("/health", tags=["meta"], summary="Liveness probe")
+async def health(request: Request) -> dict[str, object]:
"""Returns 200 OK when the server is alive."""
guidance_status = request.app.state.guidance_service.runtime_status()
return {
@@ -436,13 +404,14 @@ async def quality_check(
summary="Full conjunctiva screening pipeline",
status_code=status.HTTP_200_OK,
)
-async def analyze(
- request: Request,
- image: Annotated[UploadFile, File(description="Eye photo (JPEG or PNG).")],
- symptoms: Annotated[str | None, Form(description="JSON-encoded symptom flags.")] = None,
- language: Annotated[str | None, Form(description="Preferred language for guidance.")] = None,
- region: Annotated[str | None, Form(description="Geographic region for localised guidance.")] = None,
-) -> AnalyzeResponse | JSONResponse:
+async def analyze(
+ request: Request,
+ image: Annotated[UploadFile, File(description="Eye photo (JPEG or PNG).")],
+ symptoms: Annotated[str | None, Form(description="JSON-encoded symptom flags.")] = None,
+ patient_profile: Annotated[str | None, Form(description="JSON-encoded intake profile.")] = None,
+ language: Annotated[str | None, Form(description="Preferred language for guidance.")] = None,
+ region: Annotated[str | None, Form(description="Geographic region for localised guidance.")] = None,
+) -> AnalyzeResponse | JSONResponse:
"""
Full pipeline: quality gate โ ML inference โ triage โ guidance โ insight packs.
Works for both authenticated and anonymous users.
@@ -491,13 +460,14 @@ async def analyze(
)
# --- Input validation --------------------------------------------------
- try:
- symptom_input = parse_symptoms(symptoms)
- language = normalize_optional_text(language, field_name="language")
- region = normalize_optional_text(region, field_name="region")
- except InvalidRequestPayload as exc:
- return JSONResponse(
- status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
+ try:
+ symptom_input = parse_symptoms(symptoms)
+ patient_profile_input = parse_patient_profile(patient_profile)
+ language = normalize_optional_text(language, field_name="language")
+ region = normalize_optional_text(region, field_name="region")
+ except InvalidRequestPayload as exc:
+ return JSONResponse(
+ status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
content={"error": str(exc), "request_id": rid},
)
@@ -513,12 +483,10 @@ async def analyze(
return _image_error_response(rid)
# --- Inference (skipped on quality failure) -----------------------------
- # Compute symptom_score early so it can influence ML prediction
- symptom_score = svc.triage_service._symptom_score(symptom_input)
- prediction = svc.predictor.predict(rgb, quality, symptom_score=symptom_score) if quality.passed else None
- used_raw_frame_rescue = False
- if prediction is None:
- quality, prediction, used_raw_frame_rescue = _attempt_raw_frame_rescue(svc, image_bytes, quality, symptom_score=symptom_score)
+ prediction = svc.predictor.predict(rgb, quality) if quality.passed else None
+ used_raw_frame_rescue = False
+ if prediction is None:
+ quality, prediction, used_raw_frame_rescue = _attempt_raw_frame_rescue(svc, image_bytes, quality)
# --- Triage + guidance -------------------------------------------------
signal_breakdown = svc.triage_service.build_signal_breakdown(quality, prediction, symptom_input)
@@ -568,31 +536,55 @@ async def analyze(
processing_time_ms = (time.perf_counter() - request.state.started_at) * 1000
- analysis_meta = build_analysis_meta(
- request_id=rid,
- api_version=app.version,
- processing_time_ms=processing_time_ms,
+ analysis_meta = build_analysis_meta(
+ request_id=rid,
+ api_version=app.version,
+ processing_time_ms=processing_time_ms,
quality=quality,
decision_audit=decision_audit,
- guidance=guidance,
- used_raw_frame_rescue=used_raw_frame_rescue,
- )
-
- response = AnalyzeResponse(
- blocked=not quality.passed,
- quality=quality,
+ guidance=guidance,
+ used_raw_frame_rescue=used_raw_frame_rescue,
+ )
+ patient_profile_result = svc.patient_case_service.build_profile(
+ rid,
+ patient_profile_input,
+ symptom_input,
+ )
+ workflow_stages = svc.patient_case_service.build_workflow_stages(
+ quality,
+ prediction,
+ triage,
+ guidance,
+ symptom_input,
+ )
+ structured_case = svc.patient_case_service.build_structured_case(
+ rid,
+ patient_profile_result,
+ quality,
+ prediction,
+ triage,
+ guidance,
+ symptom_input,
+ )
+
+ response = AnalyzeResponse(
+ blocked=not quality.passed,
+ quality=quality,
prediction=prediction,
decision_audit=decision_audit,
triage=triage,
guidance=guidance,
insight_pack=insight_pack,
- clinical_brief=clinical_brief,
- handoff_summary=handoff_summary,
- analysis_meta=analysis_meta,
- symptoms=symptom_input,
- language=language,
- region=region,
- )
+ clinical_brief=clinical_brief,
+ handoff_summary=handoff_summary,
+ analysis_meta=analysis_meta,
+ patient_profile=patient_profile_result,
+ workflow_stages=workflow_stages,
+ structured_case=structured_case,
+ symptoms=symptom_input,
+ language=language,
+ region=region,
+ )
# --- Persist to database (async, non-blocking) -------------------------
import asyncio
diff --git a/backend/app/ml/efficientnet_model.py b/backend/app/ml/efficientnet_model.py
index 2cd68a615509cbb08695bca3e55efdbafd6076d3..f246bfa70dfd6434866e260c258d68dbe2e9cfbc 100644
--- a/backend/app/ml/efficientnet_model.py
+++ b/backend/app/ml/efficientnet_model.py
@@ -1,15 +1,11 @@
from __future__ import annotations
from pathlib import Path
-from typing import TYPE_CHECKING, Any
+from typing import Any
import numpy as np
from PIL import Image
-if TYPE_CHECKING:
- import torch
- from torch import nn
-
EFFICIENTNET_VERSION = "efficientnet-b0-ft-v2"
IMAGE_SIZE = 224
@@ -21,56 +17,30 @@ def clamp(value: float, lower: float = 0.0, upper: float = 1.0) -> float:
return max(lower, min(upper, value))
-def build_efficientnet_model(*, pretrained: bool = True) -> "nn.Module":
+def build_efficientnet_model(*, pretrained: bool = True):
from torch import nn
from torchvision.models import EfficientNet_B0_Weights, efficientnet_b0
- class SpatialAttention(nn.Module):
- """
- Focuses the model on the most informative spatial regions (like the conjunctiva area).
- """
-
- def __init__(self, kernel_size: int = 7) -> None:
- super().__init__()
- self.conv = nn.Conv2d(2, 1, kernel_size=kernel_size, padding=kernel_size // 2, bias=False)
- self.sigmoid = nn.Sigmoid()
-
- def forward(self, x):
- import torch
-
- avg_out = torch.mean(x, dim=1, keepdim=True)
- max_out, _ = torch.max(x, dim=1, keepdim=True)
- combined = torch.cat([avg_out, max_out], dim=1)
- scale = self.sigmoid(self.conv(combined))
- return x * scale
-
weights = EfficientNet_B0_Weights.IMAGENET1K_V1 if pretrained else None
model = efficientnet_b0(weights=weights)
-
- model.features.add_module("spatial_attention", SpatialAttention())
model.classifier = nn.Sequential(
nn.Dropout(0.35),
nn.Linear(1280, 512),
nn.GELU(),
- nn.BatchNorm1d(512),
nn.Dropout(0.25),
nn.Linear(512, 128),
nn.GELU(),
- nn.BatchNorm1d(128),
nn.Dropout(0.15),
nn.Linear(128, 2),
)
for param in model.features.parameters():
param.requires_grad = False
-
for name, param in model.features.named_parameters():
- if name.startswith(("4", "5", "6", "7", "8", "spatial_attention")):
+ if name.startswith(("4", "5", "6", "7", "8")):
param.requires_grad = True
-
for param in model.classifier.parameters():
param.requires_grad = True
-
return model
@@ -81,14 +51,27 @@ def build_train_transform():
[
transforms.RandomHorizontalFlip(),
transforms.RandomVerticalFlip(p=0.15),
- transforms.ColorJitter(brightness=0.4, contrast=0.4, saturation=0.3, hue=0.05),
+ transforms.ColorJitter(
+ brightness=0.4,
+ contrast=0.4,
+ saturation=0.3,
+ hue=0.05,
+ ),
transforms.RandomRotation(20),
- transforms.RandomAffine(degrees=0, translate=(0.12, 0.12), scale=(0.88, 1.12)),
+ transforms.RandomAffine(
+ degrees=0,
+ translate=(0.12, 0.12),
+ scale=(0.88, 1.12),
+ ),
transforms.RandomPerspective(distortion_scale=0.15, p=0.3),
transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),
transforms.ToTensor(),
transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),
- transforms.RandomErasing(p=0.25, scale=(0.02, 0.12), ratio=(0.3, 3.3)),
+ transforms.RandomErasing(
+ p=0.25,
+ scale=(0.02, 0.12),
+ ratio=(0.3, 3.3),
+ ),
]
)
@@ -108,14 +91,14 @@ def build_val_transform():
def load_efficientnet_checkpoint(
path: str | Path,
*,
- map_location: str | "torch.device" = "cpu",
+ map_location: str = "cpu",
) -> dict[str, Any]:
import torch
checkpoint = torch.load(path, map_location=map_location)
model = build_efficientnet_model(pretrained=False)
state_dict = checkpoint["state_dict"] if "state_dict" in checkpoint else checkpoint
- model.load_state_dict(state_dict, strict=False)
+ model.load_state_dict(state_dict)
device = torch.device(map_location)
model.to(device)
model.eval()
@@ -164,18 +147,15 @@ def predict_with_efficientnet_model(
_enable_dropout(model)
output = model(tensor)
probabilities.append(float(torch.sigmoid(output[:, 0]).item()))
- hemoglobin_values.append(float((output[:, 1].item() * hb_std_scale) + hb_mean))
+ hemoglobin_values.append(
+ float((output[:, 1].item() * hb_std_scale) + hb_mean)
+ )
mean_probability = float(np.mean(probabilities))
mean_hemoglobin = float(np.mean(hemoglobin_values))
probability_std = float(np.std(probabilities))
hemoglobin_std = float(np.std(hemoglobin_values))
- hb_mean_val = float(bundle.get("hb_mean", 12.8))
- hb_spread_factor = float(bundle.get("hb_spread_factor", 1.30))
- deviation = mean_hemoglobin - hb_mean_val
- mean_hemoglobin = float(np.clip(hb_mean_val + deviation * hb_spread_factor, 5.0, 20.0))
-
margin_uncertainty = 1.0 - min(1.0, abs(mean_probability - 0.5) * 2.5)
uncertainty = clamp(
(probability_std * 2.2)
@@ -196,7 +176,7 @@ def predict_with_efficientnet_model(
}
-def _enable_dropout(model: "nn.Module") -> None:
+def _enable_dropout(model) -> None:
from torch import nn
for module in model.modules():
diff --git a/backend/app/ml/lightweight_model.py b/backend/app/ml/lightweight_model.py
index 119afaa4126bc116d04f17fe63ed534acda27f0d..d25384e17aec75ee5ab1755fb7bd27d42164836e 100644
--- a/backend/app/ml/lightweight_model.py
+++ b/backend/app/ml/lightweight_model.py
@@ -2,7 +2,7 @@
Lightweight fallback model for AnemiaLens.
Used when:
-- Available RAM < 512MB (Render free tier)
+- Available RAM < 512MB on constrained free hosts
- Inference time budget is tight
- Primary model artifacts are unavailable
diff --git a/backend/app/ml/runtime_calibration.py b/backend/app/ml/runtime_calibration.py
new file mode 100644
index 0000000000000000000000000000000000000000..7f01096dcdc32a7119ca5842af2552015c81925f
--- /dev/null
+++ b/backend/app/ml/runtime_calibration.py
@@ -0,0 +1,50 @@
+from __future__ import annotations
+
+import pickle
+from dataclasses import dataclass, field
+from pathlib import Path
+from typing import Literal
+
+from app.ml.archive_model import clamp
+from app.ml.calibration import CompositeCalibrator
+
+SourceHint = Literal["roi_original", "palpebral", "forniceal_palpebral"]
+
+
+@dataclass
+class RuntimeRiskCalibrator:
+ version: str = "runtime-risk-calibrator-v1"
+ method: str = "temperature"
+ calibrator: CompositeCalibrator = field(
+ default_factory=lambda: CompositeCalibrator(method="temperature")
+ )
+ source_thresholds: dict[str, float] = field(default_factory=dict)
+ report: dict[str, object] = field(default_factory=dict)
+
+ def calibrate(
+ self,
+ probability: float,
+ *,
+ source_hint: SourceHint = "roi_original",
+ ) -> float:
+ _ = source_hint
+ return clamp(float(self.calibrator.calibrate(probability)), 0.0, 1.0)
+
+ def threshold_for_source(
+ self,
+ source_hint: SourceHint,
+ *,
+ fallback: float,
+ ) -> float:
+ return float(self.source_thresholds.get(source_hint, fallback))
+
+ def save(self, path: str | Path) -> None:
+ path = Path(path)
+ path.parent.mkdir(parents=True, exist_ok=True)
+ with path.open("wb") as handle:
+ pickle.dump(self, handle)
+
+ @classmethod
+ def load(cls, path: str | Path) -> "RuntimeRiskCalibrator":
+ with Path(path).open("rb") as handle:
+ return pickle.load(handle)
diff --git a/backend/app/ml/runtime_refinement.py b/backend/app/ml/runtime_refinement.py
new file mode 100644
index 0000000000000000000000000000000000000000..72053394541e38dc38cca6291ace666da458ee39
--- /dev/null
+++ b/backend/app/ml/runtime_refinement.py
@@ -0,0 +1,111 @@
+from __future__ import annotations
+
+import pickle
+from dataclasses import dataclass, field
+from pathlib import Path
+from typing import Any
+
+import numpy as np
+
+from app.ml.archive_model import clamp
+
+FEATURE_ORDER: tuple[str, ...] = (
+ "base_anemia_risk",
+ "uncertainty",
+ "predicted_hemoglobin",
+ "predicted_hemoglobin_missing",
+ "brightness_score",
+ "contrast_score",
+ "blur_score",
+ "framing_score",
+ "lighting_score",
+ "glare_risk",
+ "shadow_risk",
+ "lighting_balanced",
+ "lighting_overexposed",
+ "lighting_glare_heavy",
+ "lighting_shadow_heavy",
+ "lighting_flat_contrast",
+ "lighting_dim",
+ "base_likely",
+)
+
+
+@dataclass
+class RuntimeScreeningRefiner:
+ version: str = "runtime-screening-refiner-v1"
+ method: str = "logistic-regression"
+ threshold: float = 0.53
+ feature_order: tuple[str, ...] = FEATURE_ORDER
+ model: Any = None
+ report: dict[str, object] = field(default_factory=dict)
+
+ def _feature_vector(
+ self,
+ *,
+ base_anemia_risk: float,
+ uncertainty: float,
+ predicted_hemoglobin: float | None,
+ quality,
+ base_likely: bool,
+ ) -> list[float]:
+ hb_missing = predicted_hemoglobin is None
+ hb_value = 13.5 if predicted_hemoglobin is None else float(predicted_hemoglobin)
+ lighting = str(getattr(quality, "lighting_condition", "balanced"))
+ return [
+ float(base_anemia_risk),
+ float(uncertainty),
+ hb_value,
+ float(hb_missing),
+ float(getattr(quality, "brightness_score", 0.0)),
+ float(getattr(quality, "contrast_score", 0.0)),
+ float(getattr(quality, "blur_score", 0.0)),
+ float(getattr(quality, "framing_score", 0.0)),
+ float(getattr(quality, "lighting_score", 0.0)),
+ float(getattr(quality, "glare_risk", 0.0)),
+ float(getattr(quality, "shadow_risk", 0.0)),
+ float(lighting == "balanced"),
+ float(lighting == "overexposed"),
+ float(lighting == "glare_heavy"),
+ float(lighting == "shadow_heavy"),
+ float(lighting == "flat_contrast"),
+ float(lighting == "dim"),
+ float(base_likely),
+ ]
+
+ def refine(
+ self,
+ *,
+ base_anemia_risk: float,
+ uncertainty: float,
+ predicted_hemoglobin: float | None,
+ quality,
+ base_likely: bool,
+ ) -> float:
+ if self.model is None:
+ return clamp(float(base_anemia_risk), 0.0, 1.0)
+ vector = np.asarray(
+ [
+ self._feature_vector(
+ base_anemia_risk=base_anemia_risk,
+ uncertainty=uncertainty,
+ predicted_hemoglobin=predicted_hemoglobin,
+ quality=quality,
+ base_likely=base_likely,
+ )
+ ],
+ dtype=np.float32,
+ )
+ probability = float(self.model.predict_proba(vector)[0, 1])
+ return clamp(probability, 0.0, 1.0)
+
+ def save(self, path: str | Path) -> None:
+ path = Path(path)
+ path.parent.mkdir(parents=True, exist_ok=True)
+ with path.open("wb") as handle:
+ pickle.dump(self, handle)
+
+ @classmethod
+ def load(cls, path: str | Path) -> "RuntimeScreeningRefiner":
+ with Path(path).open("rb") as handle:
+ return pickle.load(handle)
diff --git a/backend/app/ml/runtime_stack.py b/backend/app/ml/runtime_stack.py
index 0615124099f126bf268dec42935826317ede057c..5d11bf4996ec42e850a22cda3b5d4963e203497d 100644
--- a/backend/app/ml/runtime_stack.py
+++ b/backend/app/ml/runtime_stack.py
@@ -8,11 +8,11 @@ from app.ml.archive_model import clamp
RUNTIME_STACK_VERSION = "archive-evidence-fusion-v4"
SourceHint = Literal["roi_original", "palpebral", "forniceal_palpebral"]
-DEFAULT_SOURCE_THRESHOLDS: dict[SourceHint, float] = {
- "roi_original": 0.65,
- "palpebral": 0.65,
- "forniceal_palpebral": 0.65,
-}
+DEFAULT_SOURCE_THRESHOLDS: dict[SourceHint, float] = {
+ "roi_original": 0.495,
+ "palpebral": 0.65,
+ "forniceal_palpebral": 0.65,
+}
DEFAULT_RISK_ARCHIVE_WEIGHTS: dict[SourceHint, float] = {
"roi_original": 0.55,
diff --git a/backend/app/schemas.py b/backend/app/schemas.py
index 1fcc59efa7a728a1e4d9f45815e689e98345d8e8..aa9d2c5236921316f25ebeb7813909f0ec9287ac 100644
--- a/backend/app/schemas.py
+++ b/backend/app/schemas.py
@@ -66,6 +66,9 @@ def _coerce_boolean(value: object, *, allow_none: bool = False) -> bool | None:
# Request schemas
# ---------------------------------------------------------------------------
+SexType = Literal["female", "male", "other", "not_specified"]
+DietType = Literal["omnivore", "vegetarian", "vegan", "mixed", "not_specified"]
+
class SymptomInput(BaseModel):
"""
Self-reported symptoms submitted alongside an eye image.
@@ -174,6 +177,58 @@ class SymptomInput(BaseModel):
}
+# ---------------------------------------------------------------------------
+# Intake context
+# ---------------------------------------------------------------------------
+
+class PatientProfileInput(BaseModel):
+ """
+ Lightweight intake details that make the screening flow feel closer to a
+ real healthcare workflow without pretending to be a full medical record.
+ """
+
+ model_config = ConfigDict(extra="forbid")
+
+ age: int | None = Field(
+ default=None,
+ ge=1,
+ le=120,
+ description="Approximate patient age in years, if provided.",
+ )
+ sex: SexType = Field(
+ default="not_specified",
+ description="Self-reported sex used only for screening context.",
+ )
+ diet_type: DietType = Field(
+ default="not_specified",
+ description="Self-reported diet pattern relevant to iron intake context.",
+ )
+
+ @field_validator("age", mode="before")
+ @classmethod
+ def _normalise_age(cls, value: object) -> int | None:
+ if value is None:
+ return None
+ if isinstance(value, str):
+ normalised = value.strip()
+ if not normalised:
+ return None
+ return int(normalised)
+ if isinstance(value, (int, float)):
+ return int(value)
+ raise ValueError("age must be an integer or null")
+
+ @field_validator("sex", "diet_type", mode="before")
+ @classmethod
+ def _normalise_intake_enum(cls, value: object) -> str:
+ if value is None:
+ return "not_specified"
+ if isinstance(value, str):
+ normalised = value.strip().lower()
+ return normalised or "not_specified"
+ raise ValueError("Expected a string value")
+
+
# ---------------------------------------------------------------------------
# Quality assessment
# ---------------------------------------------------------------------------
@@ -221,6 +276,32 @@ class QualityAssessment(BaseModel):
brightness_score: float = Field(ge=0.0, le=1.0, description="Mean luminance in [0, 1].")
contrast_score: float = Field(ge=0.0, le=1.0, description="Normalised RMS contrast.")
framing_score: float = Field(ge=0.0, description="Eye-region occupancy ratio.")
+ lighting_score: float = Field(
+ default=0.0,
+ ge=0.0,
+ le=1.0,
+ description="Composite lighting quality score, where higher means more usable lighting.",
+ )
+ lighting_condition: str = Field(
+ default="balanced",
+ description="Lighting classification inferred from exposure, glare, shadows, and contrast.",
+ )
+ lighting_summary: str = Field(
+ default="Lighting details unavailable.",
+ description="Plain-language explanation of the current lighting condition and what it means for screening.",
+ )
+ glare_risk: float = Field(
+ default=0.0,
+ ge=0.0,
+ le=1.0,
+ description="Estimated risk that glare or clipped highlights are harming the capture.",
+ )
+ shadow_risk: float = Field(
+ default=0.0,
+ ge=0.0,
+ le=1.0,
+ description="Estimated risk that shadows or underexposure are hiding useful signal.",
+ )
issues: list[QualityIssue] = Field(default_factory=list)
@cached_property
@@ -297,6 +378,10 @@ class PredictionResult(BaseModel):
model_source: ModelSource = Field(
description="Which model or pipeline produced this prediction."
)
+ confidence_breakdown: dict[str, float | bool | str] | None = Field(
+ default=None,
+ description="Decomposed confidence view covering capture quality, model stability, threshold stability, and guardrail effects.",
+ )
@model_validator(mode="after")
def _confidence_uncertainty_consistent(self) -> "PredictionResult":
@@ -573,6 +658,82 @@ class ClinicalBrief(BaseModel):
)
+WorkflowStageKey = Literal[
+ "image_quality_agent",
+ "screening_agent",
+ "triage_agent",
+ "guidance_agent",
+]
+WorkflowStageStatus = Literal["passed", "warning", "blocked", "complete"]
+
+
+class PatientProfile(BaseModel):
+ patient_id: str = Field(description="Share-safe case identifier generated for this screening run.")
+ age: int | None = Field(default=None, description="Approximate patient age in years, if provided.")
+ sex: SexType = Field(description="Self-reported sex captured during intake.")
+ diet_type: DietType = Field(description="Self-reported diet pattern captured during intake.")
+ reported_symptoms: list[str] = Field(
+ default_factory=list,
+ description="Human-readable symptom labels captured during intake.",
+ )
+ summary: str = Field(description="Short patient-context summary for the workflow UI.")
+
+
+class WorkflowStage(BaseModel):
+ key: WorkflowStageKey = Field(description="Stable workflow-stage identifier.")
+ agent_label: str = Field(description="User-facing module name, presented as an agent-like stage.")
+ title: str = Field(description="Short workflow stage title.")
+ status: WorkflowStageStatus = Field(description="Outcome of this stage for the current run.")
+ summary: str = Field(description="One-sentence explanation of what happened at this stage.")
+
+
+class StructuredCaseImageQuality(BaseModel):
+ status: Literal["acceptable", "warning", "blocked"] = Field(
+ description="Image usability status for the final screening flow."
+ )
+ lighting_condition: str = Field(description="Lighting classification for the capture.")
+ lighting_score: Annotated[float, Field(ge=0.0, le=1.0)] = Field(
+ description="Composite lighting quality score for the case."
+ )
+ blur_detected: bool = Field(description="Whether the pipeline flagged blur as an issue.")
+ eye_region_visible: bool = Field(description="Whether the eye / conjunctiva region was adequately visible.")
+ primary_issue: str | None = Field(default=None, description="Most important quality issue, if any.")
+ warnings: list[str] = Field(default_factory=list, description="Non-blocking quality issue titles.")
+
+
+class StructuredCaseScreeningResult(BaseModel):
+ risk_level: TriageBand = Field(description="Final triage band used as the case risk level.")
+ confidence: Annotated[float, Field(ge=0.0, le=1.0)] | None = Field(
+ default=None,
+ description="Final model confidence, when inference ran.",
+ )
+ reliability: ReliabilityFlag | None = Field(
+ default=None,
+ description="Reliability tier attached to the prediction, when inference ran.",
+ )
+ predicted_hemoglobin: float | None = Field(
+ default=None,
+ description="Estimated hemoglobin value in g/dL, when available.",
+ )
+ anemia_risk: Annotated[float, Field(ge=0.0, le=1.0)] | None = Field(
+ default=None,
+ description="Raw anemia-like risk score from the image model, when inference ran.",
+ )
+
+
+class StructuredCaseRecord(BaseModel):
+ case_id: str = Field(description="Stable case identifier for export, demo, or interoperability surfaces.")
+ patient_id: str = Field(description="Patient identifier copied from the intake profile.")
+ age: int | None = Field(default=None, description="Approximate patient age in years, if provided.")
+ sex: SexType = Field(description="Self-reported sex captured during intake.")
+ diet_type: DietType = Field(description="Self-reported diet pattern captured during intake.")
+ symptoms: list[str] = Field(default_factory=list, description="Active symptoms captured for this case.")
+ image_quality: StructuredCaseImageQuality = Field(description="Structured image quality summary.")
+ screening_result: StructuredCaseScreeningResult = Field(description="Structured screening result summary.")
+ recommendation: str = Field(description="Primary next-step recommendation for this case.")
+ case_summary: str = Field(description="Short clinician-facing summary sentence.")
+
+
class AnalysisMeta(BaseModel):
request_id: str = Field(description="Short request identifier copied from the API response headers.")
generated_at: str = Field(description="Local timestamp when the response payload was assembled.")
@@ -585,7 +746,7 @@ class AnalysisMeta(BaseModel):
description="Which inference path reached the final result."
)
guidance_source: GuidanceSource = Field(
- description="Whether guidance came from Qwen or the rule-based fallback."
+ description="Whether guidance came from Mistral or the rule-based fallback."
)
used_raw_frame_rescue: bool = Field(
description="True when the backend rescued a framing-limited case using the full-frame path."
@@ -603,10 +764,9 @@ class AnalysisMeta(BaseModel):
class GuidanceRuntimeStatus(BaseModel):
active_strategy: GuidanceSource
mistral_enabled: bool = False
- qwen_enabled: bool = False # kept for backwards compat
client_ready: bool = False
api_key_configured: bool = False
- qwen_model: str | None = None
+ mistral_model: str | None = None
provider: str | None = None
fallback_reason: str | None = None
last_provider_error: str | None = None
@@ -619,6 +779,20 @@ class ModelRuntimeStatus(BaseModel):
artifact_ready: bool = False
artifact_path: str | None = None
load_error: str | None = None
+ runtime_calibration_ready: bool | None = None
+ runtime_calibration_method: str | None = None
+ runtime_calibrated_threshold: float | None = None
+ runtime_calibration_ece_before: float | None = None
+ runtime_calibration_ece_after: float | None = None
+ runtime_calibration_brier_before: float | None = None
+ runtime_calibration_brier_after: float | None = None
+ runtime_refiner_ready: bool | None = None
+ runtime_refiner_method: str | None = None
+ runtime_refined_threshold: float | None = None
+ runtime_refined_accuracy: float | None = None
+ runtime_refined_precision: float | None = None
+ runtime_refined_recall: float | None = None
+ runtime_refined_f1: float | None = None
record_count: int | None = None
validation_accuracy: float | None = None
validation_f1: float | None = None
@@ -671,6 +845,14 @@ class AnalyzeResponse(BaseModel):
clinical_brief: ClinicalBrief
handoff_summary: HandoffSummary
analysis_meta: AnalysisMeta
+ patient_profile: PatientProfile
+ workflow_stages: list[WorkflowStage] = Field(
+ description="Explicit multi-step screening workflow stages for this run.",
+ min_length=4,
+ )
+ structured_case: StructuredCaseRecord = Field(
+ description="FHIR-style structured case summary suitable for provider-facing views or export."
+ )
symptoms: SymptomInput
language: str | None = Field(default=None, description="BCP-47 language tag or plain name.")
region: str | None = Field(default=None, description="Geographic region for localised guidance.")
diff --git a/backend/app/services/clinical_brief.py b/backend/app/services/clinical_brief.py
index e46bc087b93e5e2215f169dd87e0b421547a57db..e6f44cd5feb78933d2543b6ad5c8933256c16323 100644
--- a/backend/app/services/clinical_brief.py
+++ b/backend/app/services/clinical_brief.py
@@ -216,9 +216,9 @@ class ClinicalBriefService:
checks.append("The fallback rescue path was labeled transparently in the audit trail.")
if decision_audit.review_flags:
checks.append("Structured review flags were generated for follow-up and UI display.")
- if guidance.source == "qwen":
+ if guidance.source == "mistral":
checks.append(
- "Qwen guidance was constrained to the screening result, uncertainty, symptoms, and locale context."
+ "Mistral guidance was constrained to the screening result, uncertainty, symptoms, and locale context."
)
else:
checks.append("Rule-based fallback guidance stayed grounded to the current result and symptoms.")
diff --git a/backend/app/services/email_report.py b/backend/app/services/email_report.py
index 0fe6a2f6826ce5eafe3169be383eb9afce5b4ba4..f28949e491c7308b23308ecd01d317a5fd0e8480 100644
--- a/backend/app/services/email_report.py
+++ b/backend/app/services/email_report.py
@@ -353,64 +353,252 @@ class EmailReportService:
def _build_plain_text(self, payload: EmailReportContent) -> str:
hb_line = self._hemoglobin_line(payload.predicted_hemoglobin)
risk_pct = round(payload.anemia_risk * 100)
+ next_steps = "\n".join(f"- {step}" for step in self._recommended_steps(payload))
+ summary_line = self._email_result_story(payload)
return (
- "AnemiaLens Screening Result Report\n"
- "================================\n\n"
- f"Triage Label: {payload.triage_label}\n"
- f"Anemia Risk Score: {risk_pct}%\n"
- f"{hb_line}\n\n"
- "Summary\n"
- "-------\n"
- f"{payload.share_text.strip()}\n\n"
- "Important\n"
- "---------\n"
- f"{SCREENING_DISCLAIMER} Please confirm results with a clinical blood test (CBC).\n\n"
- "AnemiaLens\n"
- "https://anemia-lens.vercel.app\n"
+ "AnemiaLens Screening Result Report\n"
+ "================================\n\n"
+ f"Triage Label: {payload.triage_label}\n"
+ f"Anemia Risk Score: {risk_pct}%\n"
+ f"{hb_line}\n\n"
+ "Summary\n"
+ "-------\n"
+ f"{summary_line}\n\n"
+ "Recommended Next Steps\n"
+ "----------------------\n"
+ f"{next_steps}\n\n"
+ "Important\n"
+ "---------\n"
+ f"{SCREENING_DISCLAIMER} Please confirm results with a clinical blood test (CBC).\n\n"
+ "AnemiaLens\n"
+ "https://anemia-lens.vercel.app\n"
)
- def _build_html(self, payload: EmailReportContent) -> str:
- hb_label, hb_value = self._hemoglobin_parts(payload.predicted_hemoglobin)
- risk_pct = round(payload.anemia_risk * 100)
- share_html = "
".join(
- escape(line) for line in payload.share_text.strip().splitlines() if line.strip()
- )
- return f"""
-
-
-
-
- AnemiaLens Screening Report
-
-
-
-
-
AnemiaLens
-
Screening Result Report
-
-
- {escape(payload.triage_label)}
-
-
-
Anemia Risk Score
-
{risk_pct}%
-
-
-
{escape(hb_label)}
-
{escape(hb_value)}
-
-
- {share_html}
-
-
- Important: {escape(SCREENING_DISCLAIMER)} Please confirm results with a clinical blood test (CBC).
-
-
- AnemiaLens ยท https://anemia-lens.vercel.app ยท No lab. No needle. Just a smartphone.
-
-
-
-"""
+ def _build_html(self, payload: EmailReportContent) -> str:
+ hb_label, hb_value = self._hemoglobin_parts(payload.predicted_hemoglobin)
+ risk_pct = round(payload.anemia_risk * 100)
+ accent, accent_soft, accent_border, status_line = self._triage_theme(payload)
+ summary_line = self._email_result_story(payload)
+ detail_line = self._email_supporting_detail(payload)
+ steps_html = "".join(
+ f"""
+
+
+
+
+ |
+ โข
+ |
+
+ {escape(step)}
+ |
+
+
+ |
+
"""
+ for step in self._recommended_steps(payload)
+ )
+ return f"""
+
+
+
+
+ AnemiaLens Screening Report
+
+
+
+
+
+
+
+
+
+
+ |
+ AnemiaLens
+ Smartphone-first screening summary, ready to review or share.
+ |
+
+
+ {escape(payload.triage_label)}
+
+ |
+
+
+ |
+
+
+ |
+ {escape(status_line)}
+
+ This email keeps the screening story short: what the result means, the estimated hemoglobin context, and what to do next.
+
+ |
+
+
+
+
+
+
+
+ Anemia Risk Score
+ {risk_pct}%
+
+ |
+
+
+ {escape(hb_label)}
+ {escape(hb_value)}
+
+ |
+
+
+ |
+
+
+
+
+
+ |
+ Why this result
+
+ {escape(summary_line)}
+
+
+ {escape(detail_line)}
+
+ |
+
+
+ |
+
+
+
+
+
+ |
+ Recommended next steps
+ |
+
+ {steps_html}
+
+ |
+
+
+
+
+
+ |
+ Important: {escape(SCREENING_DISCLAIMER)} Please confirm results with a clinical blood test (CBC).
+ |
+
+
+ |
+
+
+
+
+
+ |
+ Sent by AnemiaLens for quick review and clinician handoff.
+ |
+
+ Open AnemiaLens
+ |
+
+
+ |
+
+
+ |
+
+
+
+"""
+
+ def _triage_theme(self, payload: EmailReportContent) -> tuple[str, str, str, str]:
+ triage = payload.triage_label.lower()
+ if "high" in triage:
+ return (
+ "#dc2626",
+ "#fee2e2",
+ "#fecaca",
+ "High concern detected. Please prioritize follow-up quickly.",
+ )
+ if "moderate" in triage:
+ return (
+ "#d97706",
+ "#fef3c7",
+ "#fde68a",
+ "Moderate risk detected. A clinical follow-up is worth arranging soon.",
+ )
+ if "uncertain" in triage:
+ return (
+ "#7c3aed",
+ "#ede9fe",
+ "#ddd6fe",
+ "The scan was not strong enough for a confident call, so a retake is the safest next step.",
+ )
+ return (
+ "#059669",
+ "#dcfce7",
+ "#bbf7d0",
+ "No urgent concern was detected, but routine monitoring is still sensible.",
+ )
+
+ def _email_result_story(self, payload: EmailReportContent) -> str:
+ triage = payload.triage_label.lower()
+ if "high" in triage:
+ return (
+ "The screening found a strong low-hemoglobin pattern, so prompt clinical follow-up is the safest next step."
+ )
+ if "moderate" in triage:
+ return (
+ "The screening found a moderate low-hemoglobin pattern. It is not an emergency alert, but it is worth reviewing with a clinician soon."
+ )
+ if "uncertain" in triage:
+ return (
+ "The current image was not reliable enough for a confident screening call, so a cleaner retake or clinician review is safer than over-interpreting it."
+ )
+ return (
+ "The screening did not show a strong urgent low-hemoglobin pattern, though routine monitoring remains sensible."
+ )
+
+ def _email_supporting_detail(self, payload: EmailReportContent) -> str:
+ risk_pct = round(payload.anemia_risk * 100)
+ if payload.predicted_hemoglobin is None:
+ hb_detail = "The hemoglobin estimate was withheld because confidence was limited."
+ else:
+ hb_detail = f"The estimated hemoglobin for this run was {payload.predicted_hemoglobin:.1f} g/dL."
+ return (
+ f"This run produced a {risk_pct}% anemia-risk score. {hb_detail} Use the next steps below as a simple follow-up guide, not as a diagnosis."
+ )
+
+ def _recommended_steps(self, payload: EmailReportContent) -> list[str]:
+ triage = payload.triage_label.lower()
+ if "high" in triage:
+ return [
+ "Arrange a CBC blood test as soon as possible and avoid delaying clinical review.",
+ "Share this summary with a clinician or family member who can help coordinate care.",
+ "Seek urgent medical attention sooner if symptoms worsen or new warning signs appear.",
+ ]
+ if "moderate" in triage:
+ return [
+ "Book a follow-up with a healthcare provider within 1-2 weeks and discuss confirmatory blood work.",
+ "Keep track of fatigue, dizziness, or shortness of breath if they continue.",
+ "Consider a clearer retake if the original scan had quality warnings.",
+ ]
+ if "uncertain" in triage:
+ return [
+ "Retake the image in brighter, steadier lighting with the lower eyelid fully visible.",
+ "Use this summary only as a retake reminder, not as a final decision.",
+ "If symptoms are present, do not wait for a retake before speaking with a clinician.",
+ ]
+ return [
+ "Maintain a balanced diet and monitor for any new or worsening symptoms.",
+ "Repeat screening in 3-6 months or sooner if your health changes.",
+ "Use this report as a simple summary if you want to discuss the result with a provider later.",
+ ]
def _hemoglobin_parts(self, predicted_hemoglobin: float | None) -> tuple[str, str]:
if predicted_hemoglobin is None:
diff --git a/backend/app/services/guidance.py b/backend/app/services/guidance.py
index 9ac173b70c5a01aa92a2872d9d56306d573349b0..1399f18336cd191280e25f9541f7d8884eec64d4 100644
--- a/backend/app/services/guidance.py
+++ b/backend/app/services/guidance.py
@@ -23,12 +23,13 @@ _FIELD_LIMITS = {
"urgency_guidance": 280,
"food_advice": 300,
}
-_UNSAFE_CLAIM_PATTERN = re.compile(
- r"\b(definitely\s+(?:have|has|anemic|anaemic)|confirmed\s+(?:anemia|anaemia)|"
- r"you\s+(?:have|are)\s+(?:anemia|anaemia|anemic|anaemic)|"
- r"proves?\s+(?:anemia|anaemia)|proof\s+of\s+anemia)\b",
- flags=re.IGNORECASE,
-)
+_UNSAFE_CLAIM_PATTERN = re.compile(
+ r"\b(definitely\s+(?:confirms?|have|has|anemic|anaemic)|confirm(?:ed|s)?\s+(?:anemia|anaemia)|"
+ r"you\s+(?:have|are)\s+(?:anemia|anaemia|anemic|anaemic)|"
+ r"diagnoses?\s+(?:anemia|anaemia|iron deficiency)|"
+ r"proves?\s+(?:anemia|anaemia)|proof\s+of\s+anemia)\b",
+ flags=re.IGNORECASE,
+)
_SAFE_DIAGNOSTIC_CONTEXT_PATTERNS = (
re.compile(r"\bnot a diagnos(?:is|tic)\b", flags=re.IGNORECASE),
re.compile(r"\bnon-diagnostic\b", flags=re.IGNORECASE),
@@ -220,10 +221,80 @@ class GuidanceService:
"This is screening guidance, not a diagnosis."
)
- def _generate_mistral(
- self,
- payload: dict[str, object],
- *,
+ def _mistral_system_prompt(self) -> str:
+ return (
+ "You are Mistral, writing the guidance section for AnemiaLens, a smartphone anemia screening tool. "
+ "The system analyzes the inner lower eyelid, combines that signal with symptom input, and returns a screening band: low_risk, moderate_risk, high_concern, or uncertain_retake_needed. "
+ "This is screening only, never a diagnosis, and every answer must stay medically cautious.\n\n"
+ "Write like a calm clinician or health educator speaking to one person right after their screening. "
+ "Sound natural, specific, and grounded in the payload. "
+ "Do not sound like a marketing blurb, a lab report template, or a generic wellness article. "
+ "Use the hemoglobin estimate, risk score, symptom pattern, and reliability limits to explain what this case means.\n\n"
+ "STYLE RULES:\n"
+ "- Never say 'you have anemia', 'you are anemic', or any other diagnostic claim\n"
+ "- Prefer phrases like 'this screening leans toward', 'this result suggests', or 'this pattern points to'\n"
+ "- Mention uncertainty when confidence is limited or reliability is low\n"
+ "- Avoid stock phrases like 'calls for closer attention', 'maintain a balanced diet', or 'monitor symptoms' unless you also say why or when\n"
+ "- Never invent symptoms, treatments, lab values, or medical history not present in the payload\n"
+ "- Keep the tone human, direct, and reassuring without sounding casual\n\n"
+ "Return ONLY valid JSON with exactly these keys: explanation, urgency_guidance, food_advice, next_steps.\n"
+ "explanation: 2 or 3 sentences. Sentence 1 says what the screening leans toward. Sentence 2 explains why using the actual signal, symptoms, or risk. Sentence 3 is optional and should only be used to explain uncertainty or reassurance.\n"
+ "urgency_guidance: 1 or 2 sentences with a concrete follow-up window tied to the triage band.\n"
+ "food_advice: 1 sentence with concrete iron-supportive foods, adapted to the region when possible.\n"
+ "next_steps: array of 3 or 4 short actions that are specific, non-repetitive, and realistic.\n"
+ "No markdown, no extra keys, no preamble."
+ )
+
+ def _mistral_user_prompt(self, payload: dict[str, object]) -> str:
+ hb = payload.get("predicted_hemoglobin")
+ risk_pct = payload.get("prediction_risk_percent")
+ conf_pct = payload.get("confidence_percent")
+ uncertainty_pct = payload.get("uncertainty_percent")
+ reliability_flag = payload.get("reliability_flag") or "unknown"
+ band = payload.get("triage_band", "unknown")
+ label = payload.get("triage_label", "")
+ active_symptoms = payload.get("active_symptoms") or []
+ region = payload.get("region") or "not specified"
+ screening_text = payload.get("screening_text") or ""
+ screening_label = payload.get("screening_label") or "unknown"
+
+ hb_str = f"{hb} g/dL" if hb is not None else "not available"
+ if hb is not None:
+ if hb >= 12.0:
+ hb_context = "within normal range"
+ elif hb >= 10.0:
+ hb_context = "mildly below normal"
+ elif hb >= 8.0:
+ hb_context = "moderately below normal"
+ else:
+ hb_context = "severely below normal"
+ else:
+ hb_context = "unknown"
+
+ symptom_str = ", ".join(active_symptoms) if active_symptoms else "none reported"
+
+ return (
+ f"AnemiaLens Screening Result:\n"
+ f"- Hemoglobin estimate: {hb_str} ({hb_context})\n"
+ f"- Anemia risk score: {risk_pct}%\n"
+ f"- Screening label: {screening_label}\n"
+ f"- Model confidence: {conf_pct}%\n"
+ f"- Uncertainty: {uncertainty_pct}%\n"
+ f"- Reliability flag: {reliability_flag}\n"
+ f"- Triage band: {band} ({label})\n"
+ f"- Active symptoms: {symptom_str}\n"
+ f"- Region: {region}\n"
+ f"- Model screening text: {screening_text}\n\n"
+ "Write personalized guidance for this person based on the above. "
+ "Interpret what these findings mean instead of repeating them. "
+ "If reliability is limited, say that clearly in plain language. "
+ "Make it sound like a real clinician explaining a screening result, not a report template."
+ )
+
+ def _generate_mistral(
+ self,
+ payload: dict[str, object],
+ *,
triage_band: str,
predicted_hemoglobin: float | None,
confidence: float | None,
@@ -254,16 +325,16 @@ class GuidanceService:
"Authorization": f"Bearer {settings.mistral_api_key}",
"Content-Type": "application/json",
}
- body = {
- "model": self.mistral_model,
- "messages": [
- {"role": "system", "content": self._system_prompt()},
- {"role": "user", "content": self._user_prompt(payload)},
- ],
- "max_tokens": self.guidance_max_tokens,
- "temperature": 0.4,
- "response_format": {"type": "json_object"},
- }
+ body = {
+ "model": self.mistral_model,
+ "messages": [
+ {"role": "system", "content": self._mistral_system_prompt()},
+ {"role": "user", "content": self._mistral_user_prompt(payload)},
+ ],
+ "max_tokens": self.guidance_max_tokens,
+ "temperature": 0.55,
+ "response_format": {"type": "json_object"},
+ }
log.info("POST %s model=%s max_tokens=%s", _MISTRAL_API_URL, self.mistral_model, self.guidance_max_tokens)
resp = _requests.post(_MISTRAL_API_URL, headers=headers, json=body, timeout=self.guidance_timeout)
log.info("Mistral HTTP %s", resp.status_code)
@@ -473,14 +544,13 @@ class GuidanceService:
def runtime_status(self) -> GuidanceRuntimeStatus:
provider_healthy = self._mistral_ready() and self._last_provider_error is None
active_strategy: Literal["mistral", "fallback"] = "mistral" if provider_healthy else "fallback"
- return GuidanceRuntimeStatus(
- active_strategy=active_strategy,
- mistral_enabled=self.mistral_enabled,
- qwen_enabled=self.mistral_enabled,
- client_ready=self._mistral_ready(),
- api_key_configured=self.api_key_configured,
- qwen_model=self.mistral_model if self.mistral_enabled else None,
- provider="mistral" if self.mistral_enabled else None,
- fallback_reason=self._fallback_reason or (self._last_provider_error if not provider_healthy else None),
- last_provider_error=self._last_provider_error,
- )
+ return GuidanceRuntimeStatus(
+ active_strategy=active_strategy,
+ mistral_enabled=self.mistral_enabled,
+ client_ready=self._mistral_ready(),
+ api_key_configured=self.api_key_configured,
+ mistral_model=self.mistral_model if self.mistral_enabled else None,
+ provider="mistral" if self.mistral_enabled else None,
+ fallback_reason=self._fallback_reason or (self._last_provider_error if not provider_healthy else None),
+ last_provider_error=self._last_provider_error,
+ )
diff --git a/backend/app/services/image_quality.py b/backend/app/services/image_quality.py
index 2a528eb76001a0544c9c6285ea26a6deee8948bd..fb44e08e8161d62fd3f7a29fcc7c9c8e388e092d 100644
--- a/backend/app/services/image_quality.py
+++ b/backend/app/services/image_quality.py
@@ -42,8 +42,18 @@ class ImageQualityService:
center_contrast = float(feature_map["center_contrast"])
bright_region_ratio = float(feature_map["hist_bright"])
highlight_ratio = float(feature_map["hist_highlight"])
+ dark_region_ratio = float(feature_map["hist_dark"])
frame_score = float(framing_score(feature_map))
edge_blur = edge_blur_baseline(image)
+ lighting_score, lighting_condition, lighting_summary, glare_risk, shadow_risk = self._lighting_intelligence(
+ brightness_score=brightness_score,
+ contrast_score=contrast_score,
+ center_brightness=center_brightness,
+ center_contrast=center_contrast,
+ bright_region_ratio=bright_region_ratio,
+ highlight_ratio=highlight_ratio,
+ dark_region_ratio=dark_region_ratio,
+ )
issues: list[QualityIssue] = []
@@ -94,6 +104,11 @@ class ImageQualityService:
brightness_score=round(brightness_score, 3),
contrast_score=round(contrast_score, 3),
framing_score=round(frame_score, 3),
+ lighting_score=round(lighting_score, 3),
+ lighting_condition=lighting_condition,
+ lighting_summary=lighting_summary,
+ glare_risk=round(glare_risk, 3),
+ shadow_risk=round(shadow_risk, 3),
issues=issues,
)
return assessment, image
@@ -173,8 +188,8 @@ class ImageQualityService:
QualityIssue(
code="poor_lighting",
severity="blocking",
- title="Lighting is not usable",
- message="Use bright, even light without flash glare or heavy shadows.",
+ title=self._lighting_issue_title(lighting_condition, blocking=True),
+ message=self._lighting_issue_message(lighting_condition, blocking=True),
)
)
elif lighting_warn:
@@ -182,8 +197,8 @@ class ImageQualityService:
QualityIssue(
code="poor_lighting",
severity="warning",
- title="Lighting could be better",
- message="The model can try this image, but even light will improve reliability.",
+ title=self._lighting_issue_title(lighting_condition, blocking=False),
+ message=self._lighting_issue_message(lighting_condition, blocking=False),
)
)
@@ -229,10 +244,149 @@ class ImageQualityService:
brightness_score=round(brightness_score, 3),
contrast_score=round(contrast_score, 3),
framing_score=round(frame_score, 3),
+ lighting_score=round(lighting_score, 3),
+ lighting_condition=lighting_condition,
+ lighting_summary=lighting_summary,
+ glare_risk=round(glare_risk, 3),
+ shadow_risk=round(shadow_risk, 3),
issues=issues,
)
return assessment, image
+ def _lighting_intelligence(
+ self,
+ *,
+ brightness_score: float,
+ contrast_score: float,
+ center_brightness: float,
+ center_contrast: float,
+ bright_region_ratio: float,
+ highlight_ratio: float,
+ dark_region_ratio: float,
+ ) -> tuple[float, str, str, float, float]:
+ glare_risk = min(
+ 1.0,
+ highlight_ratio * 7.5
+ + max(0.0, bright_region_ratio - 0.18) * 1.9
+ + max(0.0, center_brightness - 0.48) * 2.2,
+ )
+ shadow_risk = min(
+ 1.0,
+ dark_region_ratio * 1.1
+ + max(0.0, 0.18 - center_brightness) * 3.0
+ + max(0.0, 0.1 - brightness_score) * 2.0,
+ )
+ exposure_balance = max(0.0, 1.0 - (abs(center_brightness - 0.28) / 0.24))
+ contrast_health = max(0.0, min(1.0, self._scaled(center_contrast, 0.06, 0.19)))
+ lighting_score = max(
+ 0.0,
+ min(
+ 1.0,
+ exposure_balance * 0.38
+ + contrast_health * 0.27
+ + (1.0 - glare_risk) * 0.2
+ + (1.0 - shadow_risk) * 0.15,
+ ),
+ )
+
+ if glare_risk >= 0.72:
+ return (
+ lighting_score,
+ "glare_heavy",
+ "Bright highlights or flash glare are washing out the eyelid surface, so the redness signal is less trustworthy.",
+ glare_risk,
+ shadow_risk,
+ )
+ if shadow_risk >= 0.72:
+ return (
+ lighting_score,
+ "shadow_heavy",
+ "Shadows are covering part of the eyelid, so the model may miss the true pallor signal.",
+ glare_risk,
+ shadow_risk,
+ )
+ if brightness_score < 0.12 or center_brightness < 0.16:
+ return (
+ lighting_score,
+ "dim",
+ "The capture is underexposed, which makes fine color differences harder to measure reliably.",
+ glare_risk,
+ shadow_risk,
+ )
+ if brightness_score > 0.46 or center_brightness > 0.52:
+ return (
+ lighting_score,
+ "overexposed",
+ "The image is brighter than ideal, so the conjunctiva can lose detail even without obvious glare.",
+ glare_risk,
+ shadow_risk,
+ )
+ if contrast_score < 0.12 or center_contrast < 0.08:
+ return (
+ lighting_score,
+ "flat_contrast",
+ "The lighting is too flat, so the conjunctival tissue boundaries are less distinct than ideal.",
+ glare_risk,
+ shadow_risk,
+ )
+ return (
+ lighting_score,
+ "balanced",
+ "Lighting is balanced enough for the model to read color and texture without strong glare or shadows.",
+ glare_risk,
+ shadow_risk,
+ )
+
+ def _lighting_issue_title(self, lighting_condition: str, *, blocking: bool) -> str:
+ if lighting_condition == "glare_heavy":
+ return "Glare is covering the eyelid" if blocking else "Glare is slightly affecting the scan"
+ if lighting_condition == "shadow_heavy":
+ return "Shadows are hiding the eyelid" if blocking else "Shadows are reducing clarity"
+ if lighting_condition == "dim":
+ return "Image is too dim" if blocking else "Lighting is a little dim"
+ if lighting_condition == "overexposed":
+ return "Image is overexposed" if blocking else "Lighting is a little bright"
+ if lighting_condition == "flat_contrast":
+ return "Image lacks contrast" if blocking else "Contrast could be stronger"
+ return "Lighting is not usable" if blocking else "Lighting could be better"
+
+ def _lighting_issue_message(self, lighting_condition: str, *, blocking: bool) -> str:
+ if lighting_condition == "glare_heavy":
+ return (
+ "Turn off flash, tilt away from shiny reflections, and use soft room light or window light."
+ if blocking
+ else "The model can try this image, but removing glare will improve reliability."
+ )
+ if lighting_condition == "shadow_heavy":
+ return (
+ "Face a window or room light so the eyelid is evenly lit without one side falling into shadow."
+ if blocking
+ else "The model can try this image, but even front lighting will improve reliability."
+ )
+ if lighting_condition == "dim":
+ return (
+ "Move to brighter light and keep the phone steady so the inner eyelid stays visible."
+ if blocking
+ else "The model can try this image, but brighter light will improve reliability."
+ )
+ if lighting_condition == "overexposed":
+ return (
+ "Step away from direct flash or strong overhead light so the eyelid texture is not washed out."
+ if blocking
+ else "The model can try this image, but slightly softer light will improve reliability."
+ )
+ if lighting_condition == "flat_contrast":
+ return (
+ "Use clearer side-neutral daylight or room lighting so the eyelid tissue stands out better."
+ if blocking
+ else "The model can try this image, but better contrast will improve reliability."
+ )
+ return (
+ "Use bright, even light without flash glare or heavy shadows."
+ if blocking
+ else "The model can try this image, but even light will improve reliability."
+ )
+
def _soften_salvageable_roi_blocks(
self,
issues: list[QualityIssue],
@@ -274,8 +428,7 @@ class ImageQualityService:
def allows_raw_frame_rescue(self, assessment: QualityAssessment) -> bool:
blocking_codes = {issue.code for issue in assessment.issues if issue.severity == "blocking"}
return bool(blocking_codes) and (
- blocking_codes.issubset({"bad_framing", "eye_not_visible"})
- or blocking_codes == {"poor_lighting"}
+ blocking_codes.issubset({"bad_framing", "eye_not_visible", "poor_lighting"})
)
def build_raw_frame_rescue_assessment(self, assessment: QualityAssessment) -> QualityAssessment:
diff --git a/backend/app/services/patient_case.py b/backend/app/services/patient_case.py
new file mode 100644
index 0000000000000000000000000000000000000000..dbe720c48d8796dfec489b0238d3e33c7fdc8c89
--- /dev/null
+++ b/backend/app/services/patient_case.py
@@ -0,0 +1,254 @@
+from __future__ import annotations
+
+from app.schemas import (
+ GuidanceResult,
+ PatientProfile,
+ PatientProfileInput,
+ PredictionResult,
+ QualityAssessment,
+ StructuredCaseImageQuality,
+ StructuredCaseRecord,
+ StructuredCaseScreeningResult,
+ SymptomInput,
+ TriageResult,
+ WorkflowStage,
+)
+
+_SYMPTOM_LABELS = {
+ "fatigue": "Fatigue",
+ "dizziness": "Dizziness",
+ "pale_skin": "Pale skin",
+ "shortness_of_breath": "Shortness of breath",
+ "heavy_menstrual_bleeding": "Heavy menstrual bleeding",
+ "poor_diet_low_iron": "Low iron intake",
+}
+
+
+class PatientCaseService:
+ def build_profile(
+ self,
+ request_id: str,
+ patient_input: PatientProfileInput,
+ symptoms: SymptomInput,
+ ) -> PatientProfile:
+ patient_id = f"ANM-{request_id.upper()[-6:]}"
+ reported_symptoms = self._active_symptoms(symptoms)
+ descriptor = self._patient_descriptor(patient_input)
+ symptom_line = (
+ f"Reported symptoms: {self._join_human(reported_symptoms)}."
+ if reported_symptoms
+ else "No symptoms were reported in intake."
+ )
+ summary = f"{descriptor} {symptom_line}".strip()
+
+ return PatientProfile(
+ patient_id=patient_id,
+ age=patient_input.age,
+ sex=patient_input.sex,
+ diet_type=patient_input.diet_type,
+ reported_symptoms=reported_symptoms,
+ summary=summary,
+ )
+
+ def build_workflow_stages(
+ self,
+ quality: QualityAssessment,
+ prediction: PredictionResult | None,
+ triage: TriageResult,
+ guidance: GuidanceResult,
+ symptoms: SymptomInput,
+ ) -> list[WorkflowStage]:
+ return [
+ WorkflowStage(
+ key="image_quality_agent",
+ agent_label="Image Quality Agent",
+ title="Capture validation",
+ status=self._quality_status(quality),
+ summary=self._quality_summary(quality),
+ ),
+ WorkflowStage(
+ key="screening_agent",
+ agent_label="Screening Agent",
+ title="Conjunctiva screening",
+ status=self._screening_status(quality, prediction),
+ summary=self._screening_summary(quality, prediction),
+ ),
+ WorkflowStage(
+ key="triage_agent",
+ agent_label="Triage Agent",
+ title="Symptom + image fusion",
+ status="complete",
+ summary=self._triage_summary(triage, symptoms),
+ ),
+ WorkflowStage(
+ key="guidance_agent",
+ agent_label="Guidance Agent",
+ title="Next-step guidance",
+ status="complete",
+ summary=self._guidance_summary(guidance),
+ ),
+ ]
+
+ def build_structured_case(
+ self,
+ request_id: str,
+ patient_profile: PatientProfile,
+ quality: QualityAssessment,
+ prediction: PredictionResult | None,
+ triage: TriageResult,
+ guidance: GuidanceResult,
+ symptoms: SymptomInput,
+ ) -> StructuredCaseRecord:
+ active_symptoms = self._active_symptoms(symptoms)
+ primary_issue = quality.issues[0].title if quality.issues else None
+ warnings = [issue.title for issue in quality.warning_issues]
+ recommendation = guidance.next_steps[0] if guidance.next_steps else guidance.urgency_guidance
+
+ return StructuredCaseRecord(
+ case_id=f"CASE-{request_id.upper()[-6:]}",
+ patient_id=patient_profile.patient_id,
+ age=patient_profile.age,
+ sex=patient_profile.sex,
+ diet_type=patient_profile.diet_type,
+ symptoms=active_symptoms,
+ image_quality=StructuredCaseImageQuality(
+ status=self._structured_quality_status(quality),
+ lighting_condition=quality.lighting_condition,
+ lighting_score=quality.lighting_score,
+ blur_detected=any(issue.code == "blur_detected" for issue in quality.issues),
+ eye_region_visible=not any(issue.code == "eye_not_visible" for issue in quality.issues),
+ primary_issue=primary_issue,
+ warnings=warnings,
+ ),
+ screening_result=StructuredCaseScreeningResult(
+ risk_level=triage.band,
+ confidence=prediction.confidence if prediction else None,
+ reliability=prediction.reliability_flag if prediction else None,
+ predicted_hemoglobin=prediction.predicted_hemoglobin if prediction else None,
+ anemia_risk=prediction.anemia_risk if prediction else None,
+ ),
+ recommendation=recommendation,
+ case_summary=self._case_summary(triage, active_symptoms, prediction),
+ )
+
+ def _active_symptoms(self, symptoms: SymptomInput) -> list[str]:
+ return [
+ label
+ for field, label in _SYMPTOM_LABELS.items()
+ if getattr(symptoms, field) is True
+ ]
+
+ def _patient_descriptor(self, patient_input: PatientProfileInput) -> str:
+ parts: list[str] = []
+ if patient_input.age is not None:
+ parts.append(f"{patient_input.age}-year-old")
+ if patient_input.sex != "not_specified":
+ parts.append(patient_input.sex.replace("_", " "))
+ descriptor = " ".join(parts).strip()
+
+ diet = (
+ f"{patient_input.diet_type.replace('_', ' ')} diet"
+ if patient_input.diet_type != "not_specified"
+ else None
+ )
+
+ if descriptor and diet:
+ return f"{descriptor.capitalize()} on a {diet}."
+ if descriptor:
+ return f"{descriptor.capitalize()}."
+ if diet:
+ return f"Intake recorded with a {diet}."
+ return "Basic intake context captured."
+
+ def _join_human(self, items: list[str]) -> str:
+ if not items:
+ return "none"
+ if len(items) == 1:
+ return items[0]
+ if len(items) == 2:
+ return f"{items[0]} and {items[1]}"
+ return f"{', '.join(items[:-1])}, and {items[-1]}"
+
+ def _quality_status(self, quality: QualityAssessment) -> str:
+ if not quality.passed:
+ return "blocked"
+ if quality.warning_issues:
+ return "warning"
+ return "passed"
+
+ def _quality_summary(self, quality: QualityAssessment) -> str:
+ if not quality.passed:
+ issue = quality.issues[0] if quality.issues else None
+ if issue is None:
+ return "The capture failed the safety gate and needs a retake before screening can continue."
+ return f"{issue.title} blocked the capture, so the workflow stayed retake-first."
+ if quality.warning_issues:
+ warnings = ", ".join(issue.title.lower() for issue in quality.warning_issues[:2])
+ return (
+ f"The image passed with {quality.lighting_condition.replace('_', ' ')} lighting, "
+ f"but warnings remained: {warnings}."
+ )
+ return (
+ f"The image passed the safety gate with {quality.lighting_condition.replace('_', ' ')} lighting "
+ "and a usable conjunctiva view."
+ )
+
+ def _screening_status(self, quality: QualityAssessment, prediction: PredictionResult | None) -> str:
+ if not quality.passed or prediction is None:
+ return "blocked"
+ if prediction.reliability_flag == "low":
+ return "warning"
+ return "passed"
+
+ def _screening_summary(self, quality: QualityAssessment, prediction: PredictionResult | None) -> str:
+ if not quality.passed or prediction is None:
+ return "Screening inference was skipped because the image quality gate did not allow a safe prediction."
+ hb_text = (
+ "hemoglobin estimate withheld"
+ if prediction.predicted_hemoglobin is None
+ else f"estimated hemoglobin {prediction.predicted_hemoglobin:.1f} g/dL"
+ )
+ return (
+ f"The screening model produced a {round(prediction.anemia_risk * 100)}% anemia-like signal, "
+ f"{hb_text}, and {round(prediction.confidence * 100)}% confidence."
+ )
+
+ def _triage_summary(self, triage: TriageResult, symptoms: SymptomInput) -> str:
+ symptom_count = symptoms.active_count
+ symptom_text = (
+ "no active symptoms"
+ if symptom_count == 0
+ else f"{symptom_count} symptom{'s' if symptom_count != 1 else ''}"
+ )
+ return (
+ f"The triage layer combined the image signal with {symptom_text} and assigned {triage.label.lower()}."
+ )
+
+ def _guidance_summary(self, guidance: GuidanceResult) -> str:
+ first_step = guidance.next_steps[0] if guidance.next_steps else guidance.urgency_guidance
+ source_label = "Mistral guidance" if guidance.source == "mistral" else "Rule-based guidance"
+ return f"{source_label} translated the case into a next step: {first_step}"
+
+ def _structured_quality_status(self, quality: QualityAssessment) -> str:
+ if not quality.passed:
+ return "blocked"
+ if quality.warning_issues:
+ return "warning"
+ return "acceptable"
+
+ def _case_summary(
+ self,
+ triage: TriageResult,
+ active_symptoms: list[str],
+ prediction: PredictionResult | None,
+ ) -> str:
+ if prediction is None:
+ return "Case requires repeat image capture because the quality gate blocked a safe screening interpretation."
+ symptom_context = (
+ "with no additional symptom burden"
+ if not active_symptoms
+ else f"with reported symptoms including {self._join_human(active_symptoms)}"
+ )
+ return (
+ f"Patient shows {triage.label.lower()} based on the conjunctiva image signal {symptom_context}."
+ )
diff --git a/backend/app/services/prediction.py b/backend/app/services/prediction.py
index 58173a82bda971ddd80dc65aa340957ed2d49d4f..8c8c0fe67ab18674675efbe46428b077afebdfc5 100644
--- a/backend/app/services/prediction.py
+++ b/backend/app/services/prediction.py
@@ -1,191 +1,235 @@
-from __future__ import annotations
-
-from pathlib import Path
-from typing import Literal
-
-import numpy as np
+from __future__ import annotations
+
+from pathlib import Path
+from typing import Literal
+
from PIL import Image
-from app.config import DEFAULT_ARCHIVE_MODEL_PATH, DEFAULT_EFFICIENTNET_MODEL_PATH
-from app.ml.calibration import CompositeCalibrator
+from app.config import (
+ DEFAULT_ARCHIVE_MODEL_PATH,
+ DEFAULT_EFFICIENTNET_MODEL_PATH,
+ DEFAULT_RUNTIME_CALIBRATOR_PATH,
+ DEFAULT_RUNTIME_REFINER_PATH,
+ settings,
+)
+from app.ml.archive_model import clamp
from app.ml.features import extract_eye_features
-from app.ml.roi_confidence import RoiConfidenceScorer
from app.schemas import ModelRuntimeStatus, PredictionResult, QualityAssessment
+
+
+def _runtime_stack_version() -> str:
+ from app.ml.runtime_stack import RUNTIME_STACK_VERSION
+
+ return RUNTIME_STACK_VERSION
+
+
+def _efficientnet_version() -> str:
+ from app.ml.efficientnet_model import EFFICIENTNET_VERSION
+
+ return EFFICIENTNET_VERSION
+
+
+def _decision_threshold_for_source(
+ source_hint: Literal["roi_original", "palpebral", "forniceal_palpebral"],
+) -> float:
+ from app.ml.runtime_stack import decision_threshold_for_source
+
+ return float(decision_threshold_for_source(source_hint))
+
+
+def _load_archive_model_artifact(path: Path) -> dict[str, object]:
+ from app.ml.archive_model import load_archive_model
-_CALIBRATOR_PATH = Path(__file__).parent.parent / "artifacts" / "calibrator.pkl"
-
-
-def clamp(value: float, lower: float = 0.0, upper: float = 1.0) -> float:
- return max(lower, min(upper, value))
-
-
-def _runtime_stack_version() -> str:
- from app.ml.runtime_stack import RUNTIME_STACK_VERSION
-
- return RUNTIME_STACK_VERSION
-
-
-def _efficientnet_version() -> str:
- from app.ml.efficientnet_model import EFFICIENTNET_VERSION
-
- return EFFICIENTNET_VERSION
+ return load_archive_model(path)
-def _decision_threshold_for_source(
- source_hint: Literal["roi_original", "palpebral", "forniceal_palpebral"] = "roi_original",
-) -> float:
- from app.ml.runtime_stack import decision_threshold_for_source
+def _load_runtime_risk_calibrator_artifact(path: Path):
+ from app.ml.runtime_calibration import RuntimeRiskCalibrator
- return float(decision_threshold_for_source(source_hint))
+ return RuntimeRiskCalibrator.load(path)
-def _load_archive_model_artifact(path: str | Path) -> dict[str, object]:
- from app.ml.archive_model import load_archive_model
+def _load_runtime_screening_refiner_artifact(path: Path):
+ from app.ml.runtime_refinement import RuntimeScreeningRefiner
- return load_archive_model(path)
+ return RuntimeScreeningRefiner.load(path)
def _predict_archive_model(
artifact: dict[str, object],
feature_map: dict[str, float],
- *,
- source_hint: Literal["roi_original", "palpebral", "forniceal_palpebral"] = "roi_original",
-) -> dict[str, float]:
- from app.ml.archive_model import predict_with_archive_model
-
- return predict_with_archive_model(artifact, feature_map, source_hint=source_hint)
-
-
-def _load_efficientnet_checkpoint_bundle(path: str | Path) -> dict[str, object]:
- from app.ml.efficientnet_model import load_efficientnet_checkpoint
-
- return load_efficientnet_checkpoint(path)
-
-
-def _predict_efficientnet_bundle(
- bundle: dict[str, object],
- image: Image.Image,
- *,
- mc_passes: int,
-) -> dict[str, float]:
- from app.ml.efficientnet_model import predict_with_efficientnet_model
-
- return predict_with_efficientnet_model(bundle, image, mc_passes=mc_passes)
-
-
-def _build_runtime_stack(
- archive_prediction: dict[str, float],
- *,
- efficientnet_prediction: dict[str, float] | None = None,
- source_hint: Literal["roi_original", "palpebral", "forniceal_palpebral"] = "roi_original",
-) -> dict[str, float]:
- from app.ml.runtime_stack import build_runtime_stack_prediction
-
- return build_runtime_stack_prediction(
- archive_prediction,
- efficientnet_prediction=efficientnet_prediction,
- source_hint=source_hint,
- )
+ *,
+ source_hint: Literal["roi_original", "palpebral", "forniceal_palpebral"],
+) -> dict[str, float]:
+ from app.ml.archive_model import predict_with_archive_model
+
+ return predict_with_archive_model(artifact, feature_map, source_hint=source_hint)
+
+
+def _build_runtime_stack(
+ archive_prediction: dict[str, float],
+ *,
+ efficientnet_prediction: dict[str, float] | None,
+ source_hint: Literal["roi_original", "palpebral", "forniceal_palpebral"],
+) -> dict[str, float]:
+ from app.ml.runtime_stack import build_runtime_stack_prediction
+
+ return build_runtime_stack_prediction(
+ archive_prediction,
+ efficientnet_prediction=efficientnet_prediction,
+ source_hint=source_hint,
+ )
+
+
+def _load_efficientnet_checkpoint_bundle(path: Path) -> dict[str, object]:
+ from app.ml.efficientnet_model import load_efficientnet_checkpoint
+
+ return load_efficientnet_checkpoint(path)
+
+
+def _predict_efficientnet_bundle(
+ bundle: dict[str, object],
+ image: Image.Image,
+ *,
+ mc_passes: int,
+) -> dict[str, float]:
+ from app.ml.efficientnet_model import predict_with_efficientnet_model
+
+ return predict_with_efficientnet_model(bundle, image, mc_passes=mc_passes)
class ScreeningPredictor:
def __init__(self, model_path: str | Path | None = None) -> None:
self.efficientnet_path = Path(DEFAULT_EFFICIENTNET_MODEL_PATH)
self.model_path = Path(model_path or DEFAULT_ARCHIVE_MODEL_PATH)
+ self.runtime_calibrator_path = Path(DEFAULT_RUNTIME_CALIBRATOR_PATH)
+ self.runtime_refiner_path = Path(DEFAULT_RUNTIME_REFINER_PATH)
+ self.enable_efficientnet_fallback = settings.enable_efficientnet_fallback
self.load_error: str | None = None
self.efficientnet_bundle: dict[str, object] | None = None
self.archive_model: dict[str, object] | None = None
+ self.runtime_risk_calibrator = None
+ self.runtime_screening_refiner = None
self._archive_model_load_attempted = False
self._efficientnet_model_load_attempted = False
- self._calibrator = self._load_calibrator()
- self._roi_scorer = RoiConfidenceScorer()
+ self._runtime_risk_calibrator_load_attempted = False
+ self._runtime_screening_refiner_load_attempted = False
def preload(self) -> None:
self._ensure_archive_model_loaded()
- if self.archive_model is None:
+ self._ensure_runtime_risk_calibrator_loaded()
+ self._ensure_runtime_screening_refiner_loaded()
+ if self.enable_efficientnet_fallback:
self._ensure_efficientnet_model_loaded()
- def predict(self, image: Image.Image, quality: QualityAssessment, symptom_score: float = 0.0) -> PredictionResult:
- prediction: dict[str, float] | None = None
- model_source = "missing-model"
- decision_threshold = 0.5
- feature_map = extract_eye_features(image)
- source_hint: Literal["roi_original", "palpebral", "forniceal_palpebral"] = "roi_original"
-
- self._ensure_archive_model_loaded()
- if self.archive_model is not None:
+ def predict(self, image: Image.Image, quality: QualityAssessment) -> PredictionResult:
+ prediction: dict[str, float] | None = None
+ model_source = "missing-model"
+ decision_threshold = 0.5
+ feature_map = extract_eye_features(image)
+ source_hint: Literal["roi_original", "palpebral", "forniceal_palpebral"] = "roi_original"
+
+ archive_model = self._ensure_archive_model_loaded()
+ if archive_model is not None:
try:
- # EfficientNet was trained for only 1 epoch (AUC ~0.56 = near-random).
- # Blending it with the archive model degrades predictions.
- # Skip it until a properly trained checkpoint is available.
efficientnet_secondary: dict[str, float] | None = None
-
- archive_prediction = _predict_archive_model(
- self.archive_model,
- feature_map,
- source_hint=source_hint,
- )
+ if self.enable_efficientnet_fallback:
+ efficientnet_bundle = self._ensure_efficientnet_model_loaded()
+ if efficientnet_bundle is not None:
+ try:
+ efficientnet_secondary = _predict_efficientnet_bundle(
+ efficientnet_bundle,
+ image,
+ mc_passes=4,
+ )
+ except Exception:
+ efficientnet_secondary = None
+
+ archive_prediction = _predict_archive_model(
+ archive_model,
+ feature_map,
+ source_hint=source_hint,
+ )
prediction = _build_runtime_stack(
archive_prediction,
efficientnet_prediction=efficientnet_secondary,
source_hint=source_hint,
)
+ runtime_risk_calibrator = self._ensure_runtime_risk_calibrator_loaded()
+ if runtime_risk_calibrator is not None:
+ raw_runtime_risk = float(prediction["anemia_risk"])
+ prediction["raw_anemia_risk"] = raw_runtime_risk
+ prediction["calibrated_anemia_risk"] = runtime_risk_calibrator.calibrate(
+ raw_runtime_risk,
+ source_hint=source_hint,
+ )
+ prediction["calibration_method"] = runtime_risk_calibrator.method
model_source = _runtime_stack_version()
- decision_threshold = float(prediction.get("decision_threshold", _decision_threshold_for_source(source_hint)))
- except Exception as exc:
- self.load_error = f"Archive inference failed: {type(exc).__name__}: {exc}"
-
+ decision_threshold = float(
+ prediction.get(
+ "decision_threshold",
+ _decision_threshold_for_source(source_hint),
+ )
+ )
+ except Exception as exc:
+ self.load_error = f"Archive inference failed: {type(exc).__name__}: {exc}"
+
+ if prediction is None and self.enable_efficientnet_fallback:
+ efficientnet_bundle = self._ensure_efficientnet_model_loaded()
+ if efficientnet_bundle is not None:
+ try:
+ prediction = _predict_efficientnet_bundle(
+ efficientnet_bundle,
+ image,
+ mc_passes=4,
+ )
+ model_source = str(
+ efficientnet_bundle.get("version", _efficientnet_version())
+ )
+ decision_threshold = float(
+ prediction.get("decision_threshold", 0.5)
+ )
+ except Exception as exc:
+ self.load_error = (
+ f"EfficientNet inference failed: {type(exc).__name__}: {exc}"
+ )
+
if prediction is None:
- self._ensure_efficientnet_model_loaded()
- if prediction is None and self.efficientnet_bundle is not None:
- try:
- prediction = _predict_efficientnet_bundle(
- self.efficientnet_bundle,
- image,
- mc_passes=4, # reduced from 16 to lower peak RAM on Render
- )
- model_source = str(self.efficientnet_bundle.get("version", _efficientnet_version()))
- decision_threshold = float(prediction.get("decision_threshold", 0.5))
- except Exception as exc:
- self.load_error = f"EfficientNet inference failed: {type(exc).__name__}: {exc}"
-
- if prediction is None:
- return PredictionResult(
- anemia_risk=0.5,
- predicted_hemoglobin=None,
- confidence=0.0,
- uncertainty=1.0,
- reliability_flag="low",
- screening_label="uncertain",
- screening_text="No screening model artifact is available yet, so the safest result is uncertain.",
- model_source="missing-model",
- )
- risk = float(prediction["anemia_risk"])
- # Apply probability calibration if available
- risk = self._calibrator.calibrate(risk)
- uncertainty = float(prediction["uncertainty"])
- predicted_hemoglobin_raw = float(prediction["predicted_hemoglobin"])
- predicted_hemoglobin = round(predicted_hemoglobin_raw, 2)
-
- # --- Symptom-driven post-processing --------------------------------
- # When symptoms are present, they provide real clinical signal.
- # Blend symptom evidence into risk and Hb even for real model outputs.
- if symptom_score > 0.0:
- # Symptoms push risk up: all symptoms (score=1.0) adds up to +0.30
- symptom_risk_boost = symptom_score * 0.30
- risk = float(np.clip(risk + symptom_risk_boost * (1.0 - risk), 0.0, 1.0))
- # Symptoms lower Hb estimate: all symptoms โ up to -2.5 g/dL
- symptom_hb_penalty = symptom_score * 2.5
- predicted_hemoglobin_raw = float(np.clip(predicted_hemoglobin_raw - symptom_hb_penalty, 6.0, 18.0))
- predicted_hemoglobin = round(predicted_hemoglobin_raw, 2)
- # Symptoms reduce uncertainty slightly (more signal available)
- uncertainty = float(np.clip(uncertainty - symptom_score * 0.08, 0.05, 0.88))
- quality_delta = 0.0
- if quality.framing_score < 1.15:
- quality_delta += 0.08
- elif quality.framing_score >= 1.8:
- quality_delta -= 0.06
- elif quality.framing_score >= 1.45:
+ return PredictionResult(
+ anemia_risk=0.5,
+ predicted_hemoglobin=None,
+ confidence=0.0,
+ uncertainty=1.0,
+ reliability_flag="low",
+ screening_label="uncertain",
+ screening_text="No screening model artifact is available yet, so the safest result is uncertain.",
+ model_source="missing-model",
+ confidence_breakdown={
+ "capture_quality": 0.0,
+ "model_stability": 0.0,
+ "threshold_stability": 0.0,
+ "guardrail_applied": False,
+ "lighting_condition": quality.lighting_condition,
+ "glare_risk": round(quality.glare_risk, 3),
+ "shadow_risk": round(quality.shadow_risk, 3),
+ "summary": "No model artifact is available, so the confidence story is unavailable.",
+ },
+ )
+
+ risk = float(prediction["anemia_risk"])
+ raw_uncertainty = float(prediction["uncertainty"])
+ uncertainty = raw_uncertainty
+ predicted_hemoglobin_raw = float(prediction["predicted_hemoglobin"])
+ predicted_hemoglobin = round(predicted_hemoglobin_raw, 2)
+ calibrated_risk = float(prediction.get("calibrated_anemia_risk", risk))
+ capture_quality_score = self._capture_quality_score(quality)
+ model_stability = clamp(1.0 - raw_uncertainty, 0.0, 1.0)
+ quality_delta = 0.0
+ if quality.framing_score < 1.15:
+ quality_delta += 0.08
+ elif quality.framing_score >= 1.8:
+ quality_delta -= 0.06
+ elif quality.framing_score >= 1.45:
quality_delta -= 0.03
if quality.blur_score < 80:
@@ -201,125 +245,363 @@ class ScreeningPredictor:
quality_delta += 0.02
elif 0.09 <= quality.brightness_score <= 0.38:
quality_delta -= 0.03
+
+ if quality.contrast_score < 0.12:
+ quality_delta += 0.04
+ elif quality.contrast_score >= 0.18:
+ quality_delta -= 0.02
+
+ if quality.lighting_score < 0.38:
+ quality_delta += 0.08
+ elif quality.lighting_score < 0.6:
+ quality_delta += 0.03
+ elif quality.lighting_score >= 0.8:
+ quality_delta -= 0.03
+
+ if quality.glare_risk > 0.65:
+ quality_delta += 0.05
+ elif quality.glare_risk > 0.35:
+ quality_delta += 0.02
+
+ if quality.shadow_risk > 0.65:
+ quality_delta += 0.05
+ elif quality.shadow_risk > 0.35:
+ quality_delta += 0.02
+
+ if quality.lighting_condition in {"glare_heavy", "shadow_heavy"}:
+ quality_delta += 0.12
+ elif quality.lighting_condition in {"overexposed", "flat_contrast"}:
+ quality_delta += 0.05
+ elif quality.lighting_condition == "dim":
+ quality_delta += 0.02
+
+ negative_case_confidence_bonus = self._negative_case_confidence_bonus(
+ risk=risk,
+ threshold=decision_threshold,
+ predicted_hemoglobin=predicted_hemoglobin_raw,
+ quality=quality,
+ capture_quality_score=capture_quality_score,
+ model_stability=model_stability,
+ )
+ if self._is_clear_negative_case(
+ risk=risk,
+ threshold=decision_threshold,
+ predicted_hemoglobin=predicted_hemoglobin_raw,
+ quality=quality,
+ capture_quality_score=capture_quality_score,
+ ):
+ quality_delta = min(quality_delta, 0.12)
+ elif (
+ risk < decision_threshold
+ and predicted_hemoglobin_raw >= 12.8
+ and quality.passed
+ and capture_quality_score >= 0.42
+ ):
+ quality_delta = min(quality_delta, 0.16)
+
+ uncertainty = clamp(
+ uncertainty + quality_delta - negative_case_confidence_bonus,
+ 0.05,
+ 0.88,
+ )
+ guardrail_triggered = self._dark_signal_guardrail(
+ risk=risk,
+ predicted_hemoglobin=predicted_hemoglobin,
+ feature_map=feature_map,
+ threshold=decision_threshold,
+ )
+ if guardrail_triggered:
+ uncertainty = max(uncertainty, 0.35)
+
+ predicted_hemoglobin = self._display_hemoglobin(
+ predicted_hemoglobin, uncertainty
+ )
+ base_screening_label, base_screening_text = self._screening_decision(
+ risk,
+ uncertainty,
+ decision_threshold,
+ predicted_hemoglobin=predicted_hemoglobin_raw,
+ signal_guardrail_triggered=guardrail_triggered,
+ )
+ runtime_screening_refiner = self._ensure_runtime_screening_refiner_loaded()
+ refined_risk = risk
+ if runtime_screening_refiner is not None:
+ refined_risk = runtime_screening_refiner.refine(
+ base_anemia_risk=risk,
+ uncertainty=uncertainty,
+ predicted_hemoglobin=predicted_hemoglobin,
+ quality=quality,
+ base_likely=(base_screening_label == "anemia_likely"),
+ )
+
+ threshold_stability = clamp(
+ max(
+ abs(risk - decision_threshold),
+ abs(calibrated_risk - decision_threshold),
+ abs(refined_risk - decision_threshold),
+ )
+ / 0.18,
+ 0.0,
+ 1.0,
+ )
+ signal_strength = clamp(
+ abs(refined_risk - decision_threshold) / 0.22,
+ 0.0,
+ 1.0,
+ )
+ confidence = self._decision_confidence(
+ quality=quality,
+ uncertainty=uncertainty,
+ capture_quality_score=capture_quality_score,
+ model_stability=model_stability,
+ threshold_stability=threshold_stability,
+ signal_strength=signal_strength,
+ guardrail_triggered=guardrail_triggered,
+ )
+ uncertainty = min(
+ uncertainty,
+ clamp(1.05 - confidence, 0.05, 1.0),
+ )
+ clear_negative_case = self._is_clear_negative_case(
+ risk=refined_risk,
+ threshold=decision_threshold,
+ predicted_hemoglobin=predicted_hemoglobin_raw,
+ quality=quality,
+ capture_quality_score=capture_quality_score,
+ )
+ severe_lighting_case = (
+ quality.lighting_condition in {"glare_heavy", "shadow_heavy"}
+ or quality.glare_risk > 0.65
+ or quality.shadow_risk > 0.65
+ )
+ reliability_flag = (
+ "low"
+ if (guardrail_triggered and severe_lighting_case)
+ else "high"
+ if (
+ (
+ uncertainty < 0.2
+ and quality.passed
+ and capture_quality_score >= 0.7
+ and threshold_stability >= 0.25
+ )
+ or (
+ clear_negative_case
+ and uncertainty < 0.38
+ and threshold_stability >= 0.62
+ )
+ )
+ else "medium"
+ if (
+ (
+ uncertainty < 0.35
+ and quality.passed
+ and capture_quality_score >= 0.5
+ )
+ or (
+ clear_negative_case
+ and uncertainty < 0.52
+ and quality.passed
+ and capture_quality_score >= 0.4
+ )
+ )
+ else "low"
+ )
+ if (
+ reliability_flag == "low"
+ and quality.passed
+ and not severe_lighting_case
+ and not guardrail_triggered
+ and confidence >= 0.68
+ and capture_quality_score >= 0.72
+ and threshold_stability >= 0.72
+ ):
+ reliability_flag = "medium"
+ confidence_breakdown = {
+ "capture_quality": round(capture_quality_score, 3),
+ "model_stability": round(model_stability, 3),
+ "threshold_stability": round(threshold_stability, 3),
+ "signal_strength": round(signal_strength, 3),
+ "guardrail_applied": guardrail_triggered,
+ "calibration_applied": bool(prediction.get("calibration_method")),
+ "calibration_method": str(prediction.get("calibration_method", "none")),
+ "refinement_applied": runtime_screening_refiner is not None,
+ "refinement_method": (
+ getattr(runtime_screening_refiner, "method", "none")
+ if runtime_screening_refiner is not None
+ else "none"
+ ),
+ "raw_anemia_risk": round(
+ float(prediction.get("raw_anemia_risk", risk)),
+ 3,
+ ),
+ "calibrated_anemia_risk": round(
+ float(prediction.get("calibrated_anemia_risk", risk)),
+ 3,
+ ),
+ "refined_anemia_risk": round(refined_risk, 3),
+ "decision_threshold": round(decision_threshold, 3),
+ "base_screening_label": base_screening_label,
+ "lighting_condition": quality.lighting_condition,
+ "glare_risk": round(quality.glare_risk, 3),
+ "shadow_risk": round(quality.shadow_risk, 3),
+ "summary": self._confidence_summary(
+ quality=quality,
+ capture_quality_score=capture_quality_score,
+ model_stability=model_stability,
+ threshold_stability=threshold_stability,
+ guardrail_triggered=guardrail_triggered,
+ risk=refined_risk,
+ threshold=decision_threshold,
+ predicted_hemoglobin=predicted_hemoglobin_raw,
+ ),
+ }
+ screening_label, screening_text = self._screening_decision(
+ refined_risk,
+ uncertainty,
+ decision_threshold,
+ predicted_hemoglobin=predicted_hemoglobin_raw,
+ signal_guardrail_triggered=guardrail_triggered,
+ )
+
+ return PredictionResult(
+ anemia_risk=round(refined_risk, 3),
+ predicted_hemoglobin=predicted_hemoglobin,
+ confidence=round(confidence, 3),
+ uncertainty=round(uncertainty, 3),
+ reliability_flag=reliability_flag,
+ screening_label=screening_label,
+ screening_text=screening_text,
+ model_source=model_source,
+ confidence_breakdown=confidence_breakdown,
+ )
- if quality.contrast_score < 0.12:
- quality_delta += 0.04
- elif quality.contrast_score >= 0.18:
- quality_delta -= 0.02
-
- uncertainty = clamp(uncertainty + quality_delta, 0.05, 0.88)
- guardrail_triggered = self._dark_signal_guardrail(
- risk=risk,
- predicted_hemoglobin=predicted_hemoglobin,
- feature_map=feature_map,
- threshold=decision_threshold,
- )
- if guardrail_triggered:
- uncertainty = max(uncertainty, 0.35)
-
- confidence = clamp(1.0 - uncertainty)
- reliability_flag = (
- "high"
- if uncertainty < 0.35 and quality.passed
- else "medium"
- if uncertainty < 0.55 and quality.passed
- else "low"
- )
- predicted_hemoglobin = self._display_hemoglobin(predicted_hemoglobin, uncertainty)
- screening_label, screening_text = self._screening_decision(
- risk,
- uncertainty,
- decision_threshold,
- predicted_hemoglobin=predicted_hemoglobin_raw,
- signal_guardrail_triggered=guardrail_triggered,
- )
+ def _ensure_efficientnet_model_loaded(self) -> dict[str, object] | None:
+ if not self.enable_efficientnet_fallback:
+ return None
+ if self.efficientnet_bundle is not None:
+ return self.efficientnet_bundle
+ if self._efficientnet_model_load_attempted:
+ return None
- return PredictionResult(
- anemia_risk=round(risk, 3),
- predicted_hemoglobin=predicted_hemoglobin,
- confidence=round(confidence, 3),
- uncertainty=round(uncertainty, 3),
- reliability_flag=reliability_flag,
- screening_label=screening_label,
- screening_text=screening_text,
- model_source=model_source,
- )
+ self._efficientnet_model_load_attempted = True
+ if not self.efficientnet_path.exists():
+ return None
- def _load_calibrator(self) -> CompositeCalibrator:
try:
- if _CALIBRATOR_PATH.exists():
- return CompositeCalibrator.load(_CALIBRATOR_PATH)
- except Exception:
- pass
- return CompositeCalibrator(method="none") # identity โ no-op until trained
+ self.efficientnet_bundle = _load_efficientnet_checkpoint_bundle(
+ self.efficientnet_path
+ )
+ return self.efficientnet_bundle
+ except Exception as exc:
+ if self.archive_model is None:
+ self.load_error = f"EfficientNet load failed: {type(exc).__name__}: {exc}"
+ return None
- def _ensure_archive_model_loaded(self) -> None:
+ def _ensure_archive_model_loaded(self) -> dict[str, object] | None:
+ if self.archive_model is not None:
+ return self.archive_model
if self._archive_model_load_attempted:
- return
- self._archive_model_load_attempted = True
- self.archive_model = self._load_archive_model()
-
- def _ensure_efficientnet_model_loaded(self) -> None:
- if self._efficientnet_model_load_attempted:
- return
- self._efficientnet_model_load_attempted = True
- self.efficientnet_bundle = self._load_efficientnet_model()
-
- def _load_efficientnet_model(self) -> dict[str, object] | None:
- if not self.efficientnet_path.exists():
- return None
- try:
- return _load_efficientnet_checkpoint_bundle(self.efficientnet_path)
- except Exception as exc:
- self.load_error = f"EfficientNet load failed: {type(exc).__name__}: {exc}"
return None
- def _load_archive_model(self) -> dict[str, object] | None:
+ self._archive_model_load_attempted = True
if not self.model_path.exists():
if self.efficientnet_bundle is None:
self.load_error = f"Model artifact not found at {self.model_path}"
return None
- try:
- self.load_error = None
- return _load_archive_model_artifact(self.model_path)
+
+ try:
+ self.archive_model = _load_archive_model_artifact(self.model_path)
+ if self.archive_model is not None:
+ self.load_error = None
+ return self.archive_model
except Exception as exc:
if self.efficientnet_bundle is None:
self.load_error = f"{type(exc).__name__}: {exc}"
return None
+ def _ensure_runtime_risk_calibrator_loaded(self):
+ runtime_risk_calibrator = getattr(self, "runtime_risk_calibrator", None)
+ if runtime_risk_calibrator is not None:
+ return runtime_risk_calibrator
+ if getattr(self, "_runtime_risk_calibrator_load_attempted", False):
+ return None
+
+ self._runtime_risk_calibrator_load_attempted = True
+ path = getattr(self, "runtime_calibrator_path", Path(DEFAULT_RUNTIME_CALIBRATOR_PATH))
+ if not path.exists():
+ return None
+
+ try:
+ self.runtime_risk_calibrator = _load_runtime_risk_calibrator_artifact(path)
+ return self.runtime_risk_calibrator
+ except Exception:
+ return None
+
+ def _ensure_runtime_screening_refiner_loaded(self):
+ runtime_screening_refiner = getattr(self, "runtime_screening_refiner", None)
+ if runtime_screening_refiner is not None:
+ return runtime_screening_refiner
+ if getattr(self, "_runtime_screening_refiner_load_attempted", False):
+ return None
+
+ self._runtime_screening_refiner_load_attempted = True
+ path = getattr(self, "runtime_refiner_path", Path(DEFAULT_RUNTIME_REFINER_PATH))
+ if not path.exists():
+ return None
+
+ try:
+ self.runtime_screening_refiner = _load_runtime_screening_refiner_artifact(path)
+ return self.runtime_screening_refiner
+ except Exception:
+ return None
+
def is_ready(self) -> bool:
- return (
- self.archive_model is not None
- or self.efficientnet_bundle is not None
- or self.model_path.exists()
- or self.efficientnet_path.exists()
- )
+ archive_ready = self.archive_model is not None or self.model_path.exists()
+ efficientnet_ready = self.efficientnet_bundle is not None or (
+ self.enable_efficientnet_fallback and self.efficientnet_path.exists()
+ )
+ return archive_ready or efficientnet_ready
+
+ def is_loaded(self) -> bool:
+ return self.archive_model is not None or self.efficientnet_bundle is not None
+
+ def runtime_status(self) -> ModelRuntimeStatus:
+ archive_ready = self.archive_model is not None or self.model_path.exists()
+ efficientnet_ready = self.efficientnet_bundle is not None or (
+ self.enable_efficientnet_fallback and self.efficientnet_path.exists()
+ )
+
+ if archive_ready:
+ primary_model = _runtime_stack_version()
+ artifact_path = str(self.model_path)
+ elif efficientnet_ready:
+ primary_model = (
+ str(self.efficientnet_bundle.get("version", _efficientnet_version()))
+ if self.efficientnet_bundle is not None
+ else _efficientnet_version()
+ )
+ artifact_path = str(self.efficientnet_path)
+ else:
+ primary_model = "missing-model"
+ artifact_path = None
- def is_loaded(self) -> bool:
- return self.archive_model is not None or self.efficientnet_bundle is not None
+ runtime_calibration_ready = self.runtime_risk_calibrator is not None or (
+ getattr(self, "runtime_calibrator_path", Path(DEFAULT_RUNTIME_CALIBRATOR_PATH)).exists()
+ )
+ runtime_refiner_ready = self.runtime_screening_refiner is not None or (
+ getattr(self, "runtime_refiner_path", Path(DEFAULT_RUNTIME_REFINER_PATH)).exists()
+ )
- def runtime_status(self) -> ModelRuntimeStatus:
- archive_available = self.model_path.exists()
- efficientnet_available = self.efficientnet_path.exists()
return ModelRuntimeStatus(
- primary_model=(
- _runtime_stack_version()
- if archive_available or self.archive_model is not None
- else str(self.efficientnet_bundle.get("version", _efficientnet_version()))
- if efficientnet_available or self.efficientnet_bundle is not None
- else "missing-model"
- ),
+ primary_model=primary_model,
deep_stack_loaded=False,
legacy_loaded=False,
- artifact_ready=self.is_ready(),
- artifact_path=(
- str(self.model_path)
- if archive_available or self.archive_model is not None
- else str(self.efficientnet_path)
- if efficientnet_available or self.efficientnet_bundle is not None
- else None
- ),
+ artifact_ready=archive_ready or efficientnet_ready,
+ artifact_path=artifact_path,
load_error=self.load_error,
+ runtime_calibration_ready=runtime_calibration_ready,
+ runtime_refiner_ready=runtime_refiner_ready,
)
def should_accept_raw_frame_rescue(self, prediction: PredictionResult) -> bool:
@@ -329,14 +611,31 @@ class ScreeningPredictor:
or self._accept_raw_frame_uncertain_rescue(prediction)
)
- def _accept_raw_frame_positive_rescue(self, prediction: PredictionResult) -> bool:
- return (
- prediction.screening_label == "anemia_likely"
- and prediction.predicted_hemoglobin is not None
- and prediction.anemia_risk >= 0.8
- and prediction.predicted_hemoglobin <= 11.2
- and prediction.uncertainty <= 0.5
- )
+ def _accept_raw_frame_positive_rescue(self, prediction: PredictionResult) -> bool:
+ strong_hb_positive = (
+ prediction.predicted_hemoglobin is not None
+ and prediction.anemia_risk >= 0.8
+ and prediction.predicted_hemoglobin <= 11.2
+ and prediction.uncertainty <= 0.5
+ )
+ strong_signal_only_positive = (
+ prediction.predicted_hemoglobin is None
+ and prediction.anemia_risk >= 0.7
+ and prediction.uncertainty <= 0.8
+ )
+ overwhelming_signal_only_positive = (
+ prediction.predicted_hemoglobin is None
+ and prediction.anemia_risk >= 0.84
+ and prediction.uncertainty <= 0.9
+ )
+ return (
+ prediction.screening_label == "anemia_likely"
+ and (
+ strong_hb_positive
+ or strong_signal_only_positive
+ or overwhelming_signal_only_positive
+ )
+ )
def _accept_raw_frame_negative_rescue(self, prediction: PredictionResult) -> bool:
hidden_hb_negative = (
@@ -362,14 +661,17 @@ class ScreeningPredictor:
prediction.screening_label == "uncertain"
and prediction.anemia_risk <= 0.32
and prediction.uncertainty <= 0.68
- and (prediction.predicted_hemoglobin is None or prediction.predicted_hemoglobin >= 12.8)
+ and (
+ prediction.predicted_hemoglobin is None
+ or prediction.predicted_hemoglobin >= 12.8
+ )
)
- def _screening_decision(
- self,
- risk: float,
- uncertainty: float,
- threshold: float = 0.5,
+ def _screening_decision(
+ self,
+ risk: float,
+ uncertainty: float,
+ threshold: float = 0.5,
*,
predicted_hemoglobin: float | None = None,
signal_guardrail_triggered: bool = False,
@@ -382,94 +684,304 @@ class ScreeningPredictor:
margin = abs(risk - threshold)
mild_positive_conflict = (
predicted_hemoglobin is not None
- and (
- (
- threshold <= risk < (threshold + 0.12)
- and predicted_hemoglobin > 13.0
- and uncertainty >= 0.65
- )
- or (
- threshold < 0.6
- and threshold <= risk < (threshold + 0.13)
- and predicted_hemoglobin >= 12.3
- and uncertainty >= 0.52
- )
- )
+ and threshold <= risk < (threshold + 0.14)
+ and predicted_hemoglobin >= 12.2
+ and uncertainty >= 0.5
)
if mild_positive_conflict:
return (
"uncertain",
"The screening signal is only mildly positive while the hemoglobin estimate stays near normal, so the safest interpretation is uncertain.",
)
- strict_runtime_threshold = threshold >= 0.6
- allow_below_threshold_rescue = not strict_runtime_threshold
+ strict_runtime_borderline = (
+ threshold >= 0.6
+ and predicted_hemoglobin is not None
+ and risk < (threshold + 0.07)
+ and predicted_hemoglobin >= 11.5
+ and uncertainty >= 0.55
+ )
+ if strict_runtime_borderline:
+ return (
+ "uncertain",
+ "The signal sits too close to the operating threshold for this confidence level, so the safest interpretation is uncertain.",
+ )
high_suspicion_positive = (
predicted_hemoglobin is not None
and (
(
risk >= threshold
- and predicted_hemoglobin <= 12.2
- and uncertainty < 0.62
+ and predicted_hemoglobin <= (11.4 if threshold >= 0.6 else 12.2)
+ and uncertainty < (0.56 if threshold >= 0.6 else 0.62)
)
or (
- allow_below_threshold_rescue
- and (threshold - 0.02) <= risk < threshold
+ threshold < 0.6
+ and
+ (threshold - 0.02) <= risk < threshold
and predicted_hemoglobin <= 12.4
and uncertainty < 0.57
)
or (
- allow_below_threshold_rescue
- and (threshold - 0.05) <= risk < threshold
+ threshold < 0.6
+ and
+ (threshold - 0.05) <= risk < threshold
and predicted_hemoglobin <= 12.25
and uncertainty < 0.63
)
)
+ )
+ if high_suspicion_positive:
+ return (
+ "anemia_likely",
+ "The screening model sees a persistent low-hemoglobin signal, so this result should be treated as likely anemia despite moderate uncertainty.",
+ )
+ overwhelming_positive_signal = (
+ predicted_hemoglobin is not None
+ and risk >= (threshold + (0.18 if threshold < 0.6 else 0.10))
+ and predicted_hemoglobin <= (12.0 if threshold < 0.6 else 11.5)
+ and uncertainty < 0.9
)
- low_reliability_positive_requires_extra_evidence = (
- strict_runtime_threshold
- and predicted_hemoglobin is not None
- and risk >= threshold
- and uncertainty >= 0.55
- and risk < (threshold + 0.11)
- and predicted_hemoglobin > 11.4
+ if overwhelming_positive_signal:
+ return (
+ "anemia_likely",
+ "Even with noisy capture conditions, the positive screening signal stays strong enough that this should still be treated as likely anemia screening.",
+ )
+ signal_only_positive = (
+ predicted_hemoglobin is None
+ and risk >= (threshold + (0.15 if threshold < 0.6 else 0.08))
+ and uncertainty < 0.89
)
- if low_reliability_positive_requires_extra_evidence:
+ if signal_only_positive:
+ return (
+ "anemia_likely",
+ "The image-only anemia signal stays clearly positive even though the hemoglobin estimate is unavailable, so this should still be treated as likely anemia screening.",
+ )
+ if uncertainty >= 0.75 or (margin < 0.08 and uncertainty >= 0.45):
return (
"uncertain",
- "The scan trends positive, but at this confidence level the model only upgrades to likely anemia when the risk margin is stronger or the hemoglobin estimate is more clearly low.",
+ "The estimated hemoglobin trend is borderline or noisy, so the safest interpretation is uncertain.",
)
- if high_suspicion_positive:
+ if risk >= threshold:
return (
"anemia_likely",
- "The screening model sees a persistent low-hemoglobin signal, so this result should be treated as likely anemia despite moderate uncertainty.",
+ "The screening model estimates a lower-than-expected hemoglobin trend from the eye image, so this should be treated as likely anemia screening rather than a normal call.",
)
- if uncertainty >= 0.75 or (margin < 0.08 and uncertainty >= 0.45):
- return (
- "uncertain",
- "The estimated hemoglobin trend is borderline or noisy, so the safest interpretation is uncertain.",
- )
- if risk >= threshold:
- return (
- "anemia_likely",
- "The screening model estimates a lower-than-expected hemoglobin trend from the eye image.",
- )
return (
"anemia_unlikely",
"The screening model does not estimate a strong low-hemoglobin trend from the eye image.",
)
- def _display_hemoglobin(self, predicted_hemoglobin: float | None, uncertainty: float) -> float | None:
- if predicted_hemoglobin is None:
- return None
- # Show Hb unless uncertainty is very high (was 0.70, loosened to 0.80
- # since the new model has higher base uncertainty on sparse feature vectors)
- if uncertainty >= 0.80:
- return None
- return round(clamp(predicted_hemoglobin, 6.0, 18.0), 2)
-
- def _dark_signal_guardrail(
- self,
- *,
+ def _display_hemoglobin(
+ self, predicted_hemoglobin: float | None, uncertainty: float
+ ) -> float | None:
+ if predicted_hemoglobin is None:
+ return None
+ if uncertainty >= 0.70:
+ return None
+ return round(clamp(predicted_hemoglobin, 6.0, 18.0), 2)
+
+ def _capture_quality_score(self, quality: QualityAssessment) -> float:
+ blur_health = clamp((quality.blur_score - 55.0) / 165.0, 0.0, 1.0)
+ framing_health = clamp((quality.framing_score - 0.75) / 1.1, 0.0, 1.0)
+ brightness_health = clamp(
+ 1.0 - (abs(quality.brightness_score - 0.24) / 0.24),
+ 0.0,
+ 1.0,
+ )
+ contrast_health = clamp((quality.contrast_score - 0.06) / 0.12, 0.0, 1.0)
+ lighting_health = clamp(quality.lighting_score, 0.0, 1.0)
+ return clamp(
+ blur_health * 0.24
+ + framing_health * 0.2
+ + brightness_health * 0.14
+ + contrast_health * 0.14
+ + lighting_health * 0.28,
+ 0.0,
+ 1.0,
+ )
+
+ def _decision_confidence(
+ self,
+ *,
+ quality: QualityAssessment,
+ uncertainty: float,
+ capture_quality_score: float,
+ model_stability: float,
+ threshold_stability: float,
+ signal_strength: float,
+ guardrail_triggered: bool,
+ ) -> float:
+ confidence = (
+ model_stability * 0.34
+ + capture_quality_score * 0.24
+ + threshold_stability * 0.24
+ + signal_strength * 0.18
+ )
+
+ if quality.lighting_condition in {"glare_heavy", "shadow_heavy"}:
+ confidence -= 0.07
+ elif quality.lighting_condition in {"overexposed", "flat_contrast"}:
+ confidence -= 0.04
+ elif quality.lighting_condition == "dim":
+ confidence -= 0.02
+
+ if quality.glare_risk > 0.65 or quality.shadow_risk > 0.65:
+ confidence -= 0.04
+
+ if not quality.passed:
+ confidence = min(confidence, 0.35)
+
+ if guardrail_triggered:
+ confidence -= 0.08
+ if signal_strength >= 0.95 and capture_quality_score >= 0.65:
+ confidence = max(confidence, 0.52)
+ elif signal_strength >= 0.8 and capture_quality_score >= 0.55:
+ confidence = max(confidence, 0.4)
+ confidence = min(confidence, 0.62)
+
+ if uncertainty >= 0.82 and signal_strength < 0.75:
+ confidence = min(confidence, 0.42)
+
+ if signal_strength >= 0.9 and quality.passed and capture_quality_score >= 0.55:
+ confidence = max(confidence, 0.45)
+
+ if uncertainty <= 0.3 and threshold_stability >= 0.55:
+ confidence += 0.03
+
+ if quality.lighting_condition in {"glare_heavy", "shadow_heavy"}:
+ confidence = min(confidence, 0.54)
+ elif quality.lighting_condition == "overexposed":
+ confidence = min(confidence, 0.58)
+
+ return clamp(confidence, 0.08, 0.92)
+
+ def _confidence_summary(
+ self,
+ *,
+ quality: QualityAssessment,
+ capture_quality_score: float,
+ model_stability: float,
+ threshold_stability: float,
+ guardrail_triggered: bool,
+ risk: float,
+ threshold: float,
+ predicted_hemoglobin: float | None,
+ ) -> str:
+ if guardrail_triggered:
+ return (
+ "A protective guardrail lowered confidence because the image looked dark for a strong low-hemoglobin claim."
+ )
+ if (
+ predicted_hemoglobin is not None
+ and risk < threshold
+ and threshold_stability >= 0.65
+ and capture_quality_score >= 0.45
+ and quality.passed
+ ):
+ return (
+ "The case sits clearly on the low-risk side of the decision threshold, so the model is more confident that this is not a strong anemia-like pattern."
+ )
+ if quality.lighting_condition != "balanced":
+ return (
+ f"Confidence is mainly limited by {quality.lighting_condition.replace('_', ' ')} lighting, which makes the conjunctival color signal harder to trust."
+ )
+ if capture_quality_score < 0.55:
+ return (
+ "Confidence is mainly limited by capture quality, so a cleaner retake would be more persuasive than over-interpreting this scan."
+ )
+ if threshold_stability < 0.35:
+ return (
+ "This case sits close to the decision threshold, so the label is more sensitive to small image or symptom changes."
+ )
+ if model_stability < 0.55:
+ return (
+ "The result is still leaning one way, but repeated model passes varied more than ideal. A cleaner retake would make it more defensible, not necessarily change the overall story."
+ )
+ return (
+ "Capture quality, model stability, and threshold margin all support a more defensible screening explanation."
+ )
+
+ def _is_clear_negative_case(
+ self,
+ *,
+ risk: float,
+ threshold: float,
+ predicted_hemoglobin: float | None,
+ quality: QualityAssessment,
+ capture_quality_score: float,
+ ) -> bool:
+ if predicted_hemoglobin is None:
+ return False
+ negative_margin = threshold - risk
+ return (
+ quality.passed
+ and negative_margin >= 0.15
+ and predicted_hemoglobin >= 12.7
+ and capture_quality_score >= 0.4
+ and quality.lighting_condition in {"balanced", "dim", "flat_contrast"}
+ and quality.glare_risk <= 0.5
+ and quality.shadow_risk <= 0.5
+ )
+
+ def _negative_case_confidence_bonus(
+ self,
+ *,
+ risk: float,
+ threshold: float,
+ predicted_hemoglobin: float | None,
+ quality: QualityAssessment,
+ capture_quality_score: float,
+ model_stability: float,
+ ) -> float:
+ if predicted_hemoglobin is None or risk >= threshold or not quality.passed:
+ return 0.0
+
+ negative_margin = threshold - risk
+ if negative_margin < 0.1 or predicted_hemoglobin < 12.5:
+ return 0.0
+
+ if quality.lighting_condition in {"glare_heavy", "shadow_heavy", "overexposed"}:
+ return 0.0
+
+ bonus = 0.0
+ if negative_margin >= 0.14:
+ bonus += 0.04
+ if negative_margin >= 0.28:
+ bonus += 0.03
+ if predicted_hemoglobin >= 13.0:
+ bonus += 0.02
+ if predicted_hemoglobin >= 13.6:
+ bonus += 0.02
+ if capture_quality_score >= 0.5:
+ bonus += 0.015
+ if quality.lighting_score >= 0.42:
+ bonus += 0.015
+ if model_stability >= 0.7:
+ bonus += 0.015
+ if self._is_clear_negative_case(
+ risk=risk,
+ threshold=threshold,
+ predicted_hemoglobin=predicted_hemoglobin,
+ quality=quality,
+ capture_quality_score=capture_quality_score,
+ ):
+ bonus += 0.02
+
+ if quality.lighting_condition in {"dim", "flat_contrast"}:
+ bonus *= 0.75
+
+ if (
+ quality.glare_risk > 0.6
+ or quality.shadow_risk > 0.6
+ or quality.blur_score < 70
+ or quality.brightness_score < 0.07
+ ):
+ bonus *= 0.35
+
+ return clamp(bonus, 0.0, 0.14)
+
+ def _dark_signal_guardrail(
+ self,
+ *,
risk: float,
predicted_hemoglobin: float | None,
feature_map: dict[str, float],
diff --git a/backend/app/services/request_parsing.py b/backend/app/services/request_parsing.py
index 65c9a378b21172719ada30990e797f3d56cee497..b0213cdc8d6efb2c1387af71119d4c3480cdac45 100644
--- a/backend/app/services/request_parsing.py
+++ b/backend/app/services/request_parsing.py
@@ -5,7 +5,7 @@ import json
from pydantic import ValidationError
from app.config import settings
-from app.schemas import SymptomInput
+from app.schemas import PatientProfileInput, SymptomInput
class InvalidRequestPayload(ValueError):
@@ -30,6 +30,24 @@ def parse_symptoms(raw: str | None) -> SymptomInput:
raise InvalidRequestPayload("Invalid symptoms payload.") from exc
+def parse_patient_profile(raw: str | None) -> PatientProfileInput:
+ if not raw:
+ return PatientProfileInput()
+
+ try:
+ payload = json.loads(raw)
+ if not isinstance(payload, dict):
+ raise InvalidRequestPayload("Invalid patient profile payload: expected a JSON object.")
+ return PatientProfileInput.model_validate(payload)
+ except InvalidRequestPayload:
+ raise
+ except (TypeError, json.JSONDecodeError, ValidationError) as exc:
+ detail = _validation_detail(exc)
+ if detail:
+ raise InvalidRequestPayload(f"Invalid patient profile payload: {detail}") from exc
+ raise InvalidRequestPayload("Invalid patient profile payload.") from exc
+
+
def normalize_optional_text(
value: str | None,
*,
diff --git a/backend/app/services/runtime_status.py b/backend/app/services/runtime_status.py
index c8474809b89fc202f20e924527fe045a98d40a09..97fda717b366eced81e48b35cf774a417432222d 100644
--- a/backend/app/services/runtime_status.py
+++ b/backend/app/services/runtime_status.py
@@ -4,6 +4,8 @@ import json
from app.config import (
DEFAULT_DEPLOYED_SCREENING_REPORT_PATH,
+ DEFAULT_RUNTIME_CALIBRATION_REPORT_PATH,
+ DEFAULT_RUNTIME_REFINEMENT_REPORT_PATH,
DEFAULT_RUNTIME_STACK_REPORT_PATH,
DEFAULT_TRAINING_REPORT_PATH,
)
@@ -18,6 +20,8 @@ def build_runtime_status(
model_status = predictor.runtime_status()
report = _load_training_report()
deployed_report = _load_json_report(DEFAULT_DEPLOYED_SCREENING_REPORT_PATH)
+ calibration_report = _load_json_report(DEFAULT_RUNTIME_CALIBRATION_REPORT_PATH)
+ refinement_report = _load_json_report(DEFAULT_RUNTIME_REFINEMENT_REPORT_PATH)
if report is not None:
metrics = report.get("metrics", {})
@@ -48,6 +52,35 @@ def build_runtime_status(
}
)
+ if calibration_report is not None:
+ diagnostics = calibration_report.get("diagnostics", {})
+ selected_thresholds = calibration_report.get("selected_thresholds", {})
+ model_status = model_status.model_copy(
+ update={
+ "runtime_calibration_ready": True,
+ "runtime_calibration_method": calibration_report.get("method"),
+ "runtime_calibrated_threshold": selected_thresholds.get("roi_original"),
+ "runtime_calibration_ece_before": diagnostics.get("ece_before"),
+ "runtime_calibration_ece_after": diagnostics.get("ece_after"),
+ "runtime_calibration_brier_before": diagnostics.get("brier_before"),
+ "runtime_calibration_brier_after": diagnostics.get("brier_after"),
+ }
+ )
+
+ if refinement_report is not None:
+ metrics = refinement_report.get("metrics_after", {})
+ model_status = model_status.model_copy(
+ update={
+ "runtime_refiner_ready": True,
+ "runtime_refiner_method": refinement_report.get("method"),
+ "runtime_refined_threshold": refinement_report.get("selected_threshold"),
+ "runtime_refined_accuracy": metrics.get("accuracy"),
+ "runtime_refined_precision": metrics.get("precision"),
+ "runtime_refined_recall": metrics.get("recall"),
+ "runtime_refined_f1": metrics.get("f1"),
+ }
+ )
+
return RuntimeStatusResponse(
api_status="ok",
guidance=guidance_service.runtime_status(),
diff --git a/backend/app/services/screening_store.py b/backend/app/services/screening_store.py
new file mode 100644
index 0000000000000000000000000000000000000000..d0a06ddd53406d9673159d614d6fa4b5205e1f79
--- /dev/null
+++ b/backend/app/services/screening_store.py
@@ -0,0 +1,56 @@
+from __future__ import annotations
+
+import json
+
+from app.database import async_session_factory
+from app.models.screening import Screening
+from app.models.user import User
+from app.schemas import AnalyzeResponse
+
+
+async def persist_screening_result(
+ request_id: str,
+ analysis: AnalyzeResponse,
+ user_id: int | None,
+ processing_time_ms: float,
+) -> Screening:
+ """Persist a completed screening result and optionally attach it to a user."""
+
+ screening = Screening(
+ request_id=request_id,
+ user_id=user_id,
+ triage_band=analysis.triage.band,
+ triage_score=analysis.triage.score,
+ triage_label=analysis.triage.label,
+ anemia_risk=analysis.prediction.anemia_risk if analysis.prediction else None,
+ predicted_hemoglobin=analysis.prediction.predicted_hemoglobin if analysis.prediction else None,
+ confidence=analysis.prediction.confidence if analysis.prediction else None,
+ uncertainty=analysis.prediction.uncertainty if analysis.prediction else None,
+ screening_label=analysis.prediction.screening_label if analysis.prediction else None,
+ model_source=analysis.prediction.model_source if analysis.prediction else None,
+ quality_passed=analysis.quality.passed,
+ blocked=analysis.blocked,
+ processing_path=analysis.decision_audit.processing_path,
+ guidance_source=analysis.guidance.source,
+ symptoms_json=json.dumps(analysis.symptoms.model_dump()),
+ full_response_json=json.dumps(analysis.model_dump(), default=str),
+ share_text=analysis.handoff_summary.share_text,
+ urgency_label=analysis.handoff_summary.urgency_label,
+ headline=analysis.handoff_summary.headline,
+ processing_time_ms=processing_time_ms,
+ language=analysis.language,
+ region=analysis.region,
+ )
+
+ async with async_session_factory() as session:
+ session.add(screening)
+ await session.flush()
+
+ if user_id is not None:
+ user = await session.get(User, user_id)
+ if user is not None:
+ user.scan_count += 1
+
+ await session.commit()
+ await session.refresh(screening)
+ return screening
diff --git a/backend/app/utils/security.py b/backend/app/utils/security.py
index a3ed8b5945cf89640a027de402c482198b8256d5..3f5120553464d02ee047029812274408a299efab 100644
--- a/backend/app/utils/security.py
+++ b/backend/app/utils/security.py
@@ -7,11 +7,16 @@ Uses passlib+bcrypt for passwords and python-jose for JWT tokens.
from __future__ import annotations
import os
+from pathlib import Path
from datetime import datetime, timedelta, timezone
+from dotenv import load_dotenv
from jose import JWTError, jwt
from passlib.context import CryptContext
+BACKEND_ROOT = Path(__file__).resolve().parents[2]
+load_dotenv(BACKEND_ROOT / ".env")
+
JWT_SECRET_KEY = os.getenv("JWT_SECRET_KEY", "dev-only-change-in-production")
JWT_ALGORITHM = os.getenv("JWT_ALGORITHM", "HS256")
JWT_ACCESS_EXPIRE_MINUTES = int(os.getenv("JWT_ACCESS_TOKEN_EXPIRE_MINUTES", "60"))
diff --git a/backend/models/deployed_screening_report.json b/backend/models/deployed_screening_report.json
index b0b84b0856dd2f60b3a702989511547b02ca0eab..3833ebedf036e6a8b3b20046dad1d91746ba7c74 100644
--- a/backend/models/deployed_screening_report.json
+++ b/backend/models/deployed_screening_report.json
@@ -4,16 +4,16 @@
"validation_size": 44,
"metrics": {
"accuracy": 0.8864,
- "precision": 0.9091,
- "recall": 0.7143,
- "f1": 0.8,
+ "precision": 0.8462,
+ "recall": 0.7857,
+ "f1": 0.8148,
"split_strategy": "group-shuffle-balance-select: roi_original + deployed quality gate"
},
"operating_counts": {
- "blocked_positive": 2,
+ "blocked_positive": 1,
"blocked_negative": 9,
- "blocked_total": 11,
- "likely_count": 11,
- "uncertain_count": 4
+ "blocked_total": 10,
+ "likely_count": 13,
+ "uncertain_count": 17
}
}
\ No newline at end of file
diff --git a/backend/models/runtime_calibration_report.json b/backend/models/runtime_calibration_report.json
new file mode 100644
index 0000000000000000000000000000000000000000..a0bf8c6b622bcbade8c4d7197c11279084b30b65
--- /dev/null
+++ b/backend/models/runtime_calibration_report.json
@@ -0,0 +1,28 @@
+{
+ "version": "runtime-risk-calibrator-v1",
+ "method": "temperature",
+ "validation_size": 44,
+ "selected_thresholds": {
+ "roi_original": 0.495,
+ "palpebral": 0.65,
+ "forniceal_palpebral": 0.65
+ },
+ "diagnostics": {
+ "ece_before": 0.262,
+ "ece_after": 0.0909,
+ "brier_before": 0.0906,
+ "brier_after": 0.0501
+ },
+ "roi_metrics_before": {
+ "accuracy": 1.0,
+ "precision": 1.0,
+ "recall": 1.0,
+ "f1": 1.0
+ },
+ "roi_metrics_after": {
+ "accuracy": 0.9318,
+ "precision": 0.8235,
+ "recall": 1.0,
+ "f1": 0.9032
+ }
+}
\ No newline at end of file
diff --git a/backend/models/runtime_refinement_report.json b/backend/models/runtime_refinement_report.json
new file mode 100644
index 0000000000000000000000000000000000000000..e0965a44d61c8f757cd532fc25ee2ed31420497d
--- /dev/null
+++ b/backend/models/runtime_refinement_report.json
@@ -0,0 +1,24 @@
+{
+ "version": "runtime-screening-refiner-v1",
+ "method": "logistic-regression",
+ "validation_size": 44,
+ "selected_threshold": 0.54,
+ "metrics_before": {
+ "accuracy": 0.8636,
+ "precision": 0.75,
+ "recall": 0.8571,
+ "f1": 0.8
+ },
+ "metrics_after": {
+ "accuracy": 0.8864,
+ "precision": 0.8462,
+ "recall": 0.7857,
+ "f1": 0.8148
+ },
+ "stage_metrics_after": {
+ "accuracy": 0.9091,
+ "precision": 0.8571,
+ "recall": 0.8571,
+ "f1": 0.8571
+ }
+}
\ No newline at end of file
diff --git a/backend/models/runtime_risk_calibrator.pkl b/backend/models/runtime_risk_calibrator.pkl
new file mode 100644
index 0000000000000000000000000000000000000000..6bb28f07af44cc5b7658d10e4610c707a850aa8e
--- /dev/null
+++ b/backend/models/runtime_risk_calibrator.pkl
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:029bae31e2a64a5f129983b488d081f2a8cb51f8cbce8f578e21242eba1a882b
+size 543
diff --git a/backend/models/runtime_screening_refiner.pkl b/backend/models/runtime_screening_refiner.pkl
new file mode 100644
index 0000000000000000000000000000000000000000..fa2998e43a3036657f94facfc9e8594effe53c05
--- /dev/null
+++ b/backend/models/runtime_screening_refiner.pkl
@@ -0,0 +1,3 @@
+version https://git-lfs.github.com/spec/v1
+oid sha256:7f2b89f7875b4d7e597b915f0f313570d1a2b8ac3a160ab620574d12107c2c9f
+size 2346
diff --git a/backend/scripts/analyze_efficientnet_errors.py b/backend/scripts/analyze_efficientnet_errors.py
new file mode 100644
index 0000000000000000000000000000000000000000..fe70f237413fc938754d5aa1715bda20d54b4aab
--- /dev/null
+++ b/backend/scripts/analyze_efficientnet_errors.py
@@ -0,0 +1,204 @@
+from __future__ import annotations
+
+import json
+import sys
+from collections import Counter, defaultdict
+from pathlib import Path
+
+import numpy as np
+import torch
+from sklearn.metrics import accuracy_score, confusion_matrix, f1_score, mean_absolute_error, precision_score, recall_score, roc_auc_score
+from torch.utils.data import DataLoader
+
+BACKEND_ROOT = Path(__file__).resolve().parents[1]
+if str(BACKEND_ROOT) not in sys.path:
+ sys.path.insert(0, str(BACKEND_ROOT))
+SCRIPT_ROOT = Path(__file__).resolve().parent
+if str(SCRIPT_ROOT) not in sys.path:
+ sys.path.insert(0, str(SCRIPT_ROOT))
+
+from app.config import DEFAULT_EFFICIENTNET_MODEL_PATH
+from app.ml.efficientnet_model import load_efficientnet_checkpoint
+from train_efficientnet import (
+ ARCHIVE_ROOT,
+ DATA_ROOT,
+ ConjunctivaDataset,
+ _balanced_group_split,
+ _build_records,
+ build_val_transform,
+)
+
+DEFAULT_OUTPUT_PATH = BACKEND_ROOT / "models" / "efficientnet_error_report.json"
+
+
+def main() -> None:
+ dataset_root = DATA_ROOT if DATA_ROOT.exists() else ARCHIVE_ROOT
+ records = _build_records(dataset_root)
+ if not records:
+ raise RuntimeError(f"No dataset records found under {dataset_root}.")
+ if not DEFAULT_EFFICIENTNET_MODEL_PATH.exists():
+ raise RuntimeError(f"EfficientNet checkpoint not found at {DEFAULT_EFFICIENTNET_MODEL_PATH}.")
+
+ train_records, val_records = _balanced_group_split(records, test_size=0.2, n_splits=32)
+ bundle = load_efficientnet_checkpoint(DEFAULT_EFFICIENTNET_MODEL_PATH)
+ report = analyze_validation_split(val_records, bundle, dataset_root=dataset_root, train_records=train_records)
+
+ DEFAULT_OUTPUT_PATH.write_text(json.dumps(report, indent=2), encoding="utf-8")
+ print(f"Saved error report to {DEFAULT_OUTPUT_PATH}")
+ print(json.dumps(report["summary"], indent=2))
+
+
+def analyze_validation_split(
+ val_records: list,
+ bundle: dict[str, object],
+ *,
+ dataset_root: Path,
+ train_records: list,
+) -> dict[str, object]:
+ model = bundle["model"]
+ device = bundle["device"]
+ hb_mean = float(bundle.get("hb_mean", 0.0))
+ hb_std = float(bundle.get("hb_std", 1.0))
+ threshold = float(bundle.get("decision_threshold", 0.5))
+
+ dataset = ConjunctivaDataset(val_records, build_val_transform())
+ loader = DataLoader(dataset, batch_size=16, shuffle=False, num_workers=0)
+
+ probabilities: list[float] = []
+ predictions: list[int] = []
+ labels: list[int] = []
+ hb_predictions: list[float] = []
+ hb_targets: list[float] = []
+
+ model.eval()
+ with torch.no_grad():
+ for images, batch_labels, batch_hbs in loader:
+ output = model(images.to(device))
+ batch_probabilities = torch.sigmoid(output[:, 0]).cpu().tolist()
+ batch_hb_predictions = (((output[:, 1].cpu()) * hb_std) + hb_mean).tolist()
+ probabilities.extend(batch_probabilities)
+ predictions.extend([1 if value >= threshold else 0 for value in batch_probabilities])
+ labels.extend(batch_labels.squeeze(1).cpu().int().tolist())
+ hb_predictions.extend(batch_hb_predictions)
+ hb_targets.extend(batch_hbs.squeeze(1).cpu().tolist())
+
+ summary = {
+ "dataset_root": str(dataset_root),
+ "checkpoint_path": str(DEFAULT_EFFICIENTNET_MODEL_PATH),
+ "record_count": len(val_records),
+ "subject_count": len({record.subject_id for record in val_records}),
+ "threshold": round(threshold, 4),
+ "train_record_count": len(train_records),
+ "train_subject_count": len({record.subject_id for record in train_records}),
+ "accuracy": round(float(accuracy_score(labels, predictions)), 4),
+ "precision": round(float(precision_score(labels, predictions, zero_division=0)), 4),
+ "recall": round(float(recall_score(labels, predictions, zero_division=0)), 4),
+ "f1": round(float(f1_score(labels, predictions, zero_division=0)), 4),
+ "auc": round(float(roc_auc_score(labels, probabilities)), 4),
+ "hb_mae": round(float(mean_absolute_error(hb_targets, hb_predictions)), 4),
+ "label_counts": dict(Counter(labels)),
+ "prediction_counts": dict(Counter(predictions)),
+ "confusion_matrix": confusion_matrix(labels, predictions).tolist(),
+ }
+
+ source_breakdown = _source_breakdown(val_records, labels, predictions, probabilities, hb_predictions, hb_targets)
+ false_positives, false_negatives = _mistakes(val_records, labels, predictions, probabilities, hb_predictions, hb_targets)
+
+ return {
+ "summary": summary,
+ "source_breakdown": source_breakdown,
+ "false_positives": false_positives,
+ "false_negatives": false_negatives,
+ }
+
+
+def _source_breakdown(
+ val_records: list,
+ labels: list[int],
+ predictions: list[int],
+ probabilities: list[float],
+ hb_predictions: list[float],
+ hb_targets: list[float],
+) -> dict[str, object]:
+ by_source: dict[str, dict[str, object]] = defaultdict(
+ lambda: {
+ "count": 0,
+ "errors": 0,
+ "false_positives": 0,
+ "false_negatives": 0,
+ "probabilities": [],
+ "hb_abs_error": [],
+ }
+ )
+
+ for record, label, prediction, probability, hb_prediction, hb_target in zip(
+ val_records,
+ labels,
+ predictions,
+ probabilities,
+ hb_predictions,
+ hb_targets,
+ ):
+ item = by_source[record.source]
+ item["count"] += 1
+ item["errors"] += int(label != prediction)
+ item["false_positives"] += int(label == 0 and prediction == 1)
+ item["false_negatives"] += int(label == 1 and prediction == 0)
+ item["probabilities"].append(float(probability))
+ item["hb_abs_error"].append(abs(float(hb_prediction) - float(hb_target)))
+
+ normalized: dict[str, object] = {}
+ for source, item in by_source.items():
+ normalized[source] = {
+ "count": item["count"],
+ "errors": item["errors"],
+ "false_positives": item["false_positives"],
+ "false_negatives": item["false_negatives"],
+ "error_rate": round(float(item["errors"] / max(item["count"], 1)), 4),
+ "mean_probability": round(float(np.mean(item["probabilities"])), 4),
+ "hb_mae": round(float(np.mean(item["hb_abs_error"])), 4),
+ }
+ return normalized
+
+
+def _mistakes(
+ val_records: list,
+ labels: list[int],
+ predictions: list[int],
+ probabilities: list[float],
+ hb_predictions: list[float],
+ hb_targets: list[float],
+) -> tuple[list[dict[str, object]], list[dict[str, object]]]:
+ false_positives: list[dict[str, object]] = []
+ false_negatives: list[dict[str, object]] = []
+
+ for record, label, prediction, probability, hb_prediction, hb_target in zip(
+ val_records,
+ labels,
+ predictions,
+ probabilities,
+ hb_predictions,
+ hb_targets,
+ ):
+ if label == prediction:
+ continue
+ item = {
+ "subject_id": record.subject_id,
+ "source": record.source,
+ "probability": round(float(probability), 4),
+ "hb_true": round(float(hb_target), 2),
+ "hb_predicted": round(float(hb_prediction), 2),
+ "image_path": str(record.image_path),
+ }
+ if label == 0 and prediction == 1:
+ false_positives.append(item)
+ else:
+ false_negatives.append(item)
+
+ false_positives.sort(key=lambda item: float(item["probability"]), reverse=True)
+ false_negatives.sort(key=lambda item: float(item["probability"]))
+ return false_positives[:12], false_negatives[:12]
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/scripts/eval_pipeline.py b/backend/scripts/eval_pipeline.py
new file mode 100644
index 0000000000000000000000000000000000000000..1a272020d1a474bb04272a8b3c1bc8e3fdf0edf3
--- /dev/null
+++ b/backend/scripts/eval_pipeline.py
@@ -0,0 +1,80 @@
+"""
+Test the full inference pipeline (quality -> features -> predict) on real dataset images.
+This simulates exactly what happens when a user uploads a photo.
+"""
+import sys, io
+from pathlib import Path
+sys.path.insert(0, str(Path(__file__).parents[1]))
+
+import numpy as np
+from PIL import Image
+from app.services.prediction import ScreeningPredictor
+from app.services.image_quality import ImageQualityService
+from app.ml.archive_model import _build_subject_catalog, ANEMIA_HB_THRESHOLD
+from sklearn.metrics import accuracy_score, f1_score, recall_score, precision_score, roc_auc_score, mean_absolute_error
+
+predictor = ScreeningPredictor()
+quality_svc = ImageQualityService()
+
+print("Model:", predictor.archive_model.get("version") if predictor.archive_model else "NONE")
+print()
+
+subjects = _build_subject_catalog(Path(__file__).parents[2] / "archive" / "dataset anemia")
+
+# Test on original JPG images (what users actually upload)
+results = []
+blocked = 0
+for s in subjects[:40]: # first 40 for speed
+ country = s["subject_id"].split("-")[0]
+ num = s["subject_number"]
+ jpg_path = Path(__file__).parents[2] / "archive" / "dataset anemia" / country / num
+ jpgs = list(jpg_path.glob("*.jpg"))
+ if not jpgs:
+ continue
+
+ with open(jpgs[0], "rb") as f:
+ img_bytes = f.read()
+
+ try:
+ quality, rgb = quality_svc.evaluate(img_bytes)
+ if not quality.passed:
+ blocked += 1
+ continue
+ pred = predictor.predict(rgb, quality, symptom_score=0.0)
+ results.append({
+ "hb_true": s["hb"],
+ "hb_pred": pred.predicted_hemoglobin,
+ "risk": pred.anemia_risk,
+ "label_true": int(s["hb"] < ANEMIA_HB_THRESHOLD),
+ "label_pred": int(pred.anemia_risk >= 0.65) if pred.anemia_risk else 0,
+ "screening_label": pred.screening_label,
+ })
+ except Exception as e:
+ print(f" Error on {s['subject_id']}: {e}")
+
+print(f"Processed: {len(results)}, Blocked by quality: {blocked}")
+print()
+
+if not results:
+ print("No results โ all blocked by quality gate!")
+else:
+ labels_true = [r["label_true"] for r in results]
+ labels_pred = [r["label_pred"] for r in results]
+ risks = [r["risk"] for r in results if r["risk"] is not None]
+ hb_true = [r["hb_true"] for r in results if r["hb_pred"] is not None]
+ hb_pred = [r["hb_pred"] for r in results if r["hb_pred"] is not None]
+
+ print(f"Accuracy: {accuracy_score(labels_true, labels_pred):.3f}")
+ print(f"Recall: {recall_score(labels_true, labels_pred, zero_division=0):.3f}")
+ print(f"Precision: {precision_score(labels_true, labels_pred, zero_division=0):.3f}")
+ print(f"F1: {f1_score(labels_true, labels_pred, zero_division=0):.3f}")
+ if len(set(labels_true)) > 1 and risks:
+ print(f"AUC: {roc_auc_score(labels_true[:len(risks)], risks):.3f}")
+ if hb_pred:
+ print(f"Hb MAE: {mean_absolute_error(hb_true, hb_pred):.3f} g/dL")
+ print(f"Hb bias: {float(np.mean(np.array(hb_pred) - np.array(hb_true))):.3f} g/dL")
+
+ print("\nSample predictions:")
+ for r in results[:10]:
+ tag = "OK" if r["label_true"] == r["label_pred"] else "WRONG"
+ print(f" True={r['hb_true']:.1f} Pred={r['hb_pred']} Risk={r['risk']} {r['screening_label']} [{tag}]")
diff --git a/backend/scripts/eval_real.py b/backend/scripts/eval_real.py
new file mode 100644
index 0000000000000000000000000000000000000000..72d5fd6080dd28042b0cac1f82df83250cb323fc
--- /dev/null
+++ b/backend/scripts/eval_real.py
@@ -0,0 +1,79 @@
+"""
+Evaluate model on real dataset subjects with known Hb values.
+Shows true Hb vs predicted Hb vs risk score.
+"""
+import sys, json
+from pathlib import Path
+sys.path.insert(0, str(Path(__file__).parents[1]))
+
+import joblib, numpy as np
+from app.ml.archive_model import (
+ _build_subject_catalog, predict_with_archive_model, ANEMIA_HB_THRESHOLD
+)
+
+m = joblib.load(Path(__file__).parents[1] / "models" / "archive_screening_model.joblib")
+cal = m["calibration"]
+print("Model version:", m["version"])
+print("Blend threshold:", cal["blend_threshold"])
+print("Risk scale:", cal["risk_scale"])
+print()
+
+subjects = _build_subject_catalog(Path(__file__).parents[2] / "archive" / "dataset anemia")
+print(f"Total subjects: {len(subjects)}")
+
+anemic = [s for s in subjects if s["hb"] < 11.5][:8]
+borderline = [s for s in subjects if 11.5 <= s["hb"] < 13.0][:4]
+normal = [s for s in subjects if s["hb"] >= 13.0][:8]
+
+correct = 0
+total = 0
+
+for group, cases in [("ANEMIC (Hb<11.5)", anemic), ("BORDERLINE", borderline), ("NORMAL (Hb>=13)", normal)]:
+ print(f"--- {group} ---")
+ for s in cases:
+ feat = list(s["views"].values())[0]
+ result = predict_with_archive_model(m, feat, source_hint="roi_original")
+ hb_true = s["hb"]
+ hb_pred = result["predicted_hemoglobin"]
+ risk = result["anemia_risk"]
+ predicted_anemic = risk >= 0.65
+ actually_anemic = hb_true < ANEMIA_HB_THRESHOLD
+ ok = predicted_anemic == actually_anemic
+ correct += int(ok)
+ total += 1
+ tag = "OK" if ok else "WRONG"
+ print(f" True={hb_true:.1f} Pred={hb_pred:.1f} Risk={risk:.3f} [{tag}]")
+ print()
+
+print(f"Accuracy on sample: {correct}/{total} = {correct/total*100:.0f}%")
+
+# Full dataset accuracy
+print("\n--- Full dataset ---")
+all_risks = []
+all_labels = []
+all_hb_true = []
+all_hb_pred = []
+for s in subjects:
+ feat = list(s["views"].values())[0]
+ result = predict_with_archive_model(m, feat, source_hint="roi_original")
+ all_risks.append(result["anemia_risk"])
+ all_labels.append(int(s["hb"] < ANEMIA_HB_THRESHOLD))
+ all_hb_true.append(s["hb"])
+ all_hb_pred.append(result["predicted_hemoglobin"])
+
+risks = np.array(all_risks)
+labels = np.array(all_labels)
+hb_true = np.array(all_hb_true)
+hb_pred = np.array(all_hb_pred)
+
+from sklearn.metrics import accuracy_score, f1_score, recall_score, precision_score, roc_auc_score, mean_absolute_error
+preds = (risks >= 0.65).astype(int)
+print(f"Accuracy: {accuracy_score(labels, preds):.3f}")
+print(f"Precision: {precision_score(labels, preds, zero_division=0):.3f}")
+print(f"Recall: {recall_score(labels, preds, zero_division=0):.3f}")
+print(f"F1: {f1_score(labels, preds, zero_division=0):.3f}")
+print(f"AUC: {roc_auc_score(labels, risks):.3f}")
+print(f"Hb MAE: {mean_absolute_error(hb_true, hb_pred):.3f} g/dL")
+print(f"Hb bias: {float(np.mean(hb_pred - hb_true)):.3f} g/dL (+ = overestimate)")
+print(f"Risk dist anemic: {np.percentile(risks[labels==1], [10,25,50,75,90]).round(3)}")
+print(f"Risk dist normal: {np.percentile(risks[labels==0], [10,25,50,75,90]).round(3)}")
diff --git a/backend/scripts/evaluate_deployed_screening.py b/backend/scripts/evaluate_deployed_screening.py
new file mode 100644
index 0000000000000000000000000000000000000000..59babdc2b044429b99fea7cd0a2c05f7829bd7d0
--- /dev/null
+++ b/backend/scripts/evaluate_deployed_screening.py
@@ -0,0 +1,89 @@
+from __future__ import annotations
+
+import json
+
+import numpy as np
+from sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score
+
+from app.config import DEFAULT_DEPLOYED_SCREENING_REPORT_PATH
+from app.services.image_quality import ImageQualityService
+from app.services.prediction import ScreeningPredictor
+from train_efficientnet import ARCHIVE_ROOT, _balanced_group_split, _build_records, _load_image_with_fallback
+
+
+def main() -> None:
+ records = _build_records(ARCHIVE_ROOT)
+ if not records:
+ raise RuntimeError(f"No evaluation records found in {ARCHIVE_ROOT}.")
+
+ _, val_records = _balanced_group_split(records, test_size=0.2, n_splits=32)
+ roi_records = [record for record in val_records if record.source == "roi_original"]
+
+ quality_service = ImageQualityService()
+ predictor = ScreeningPredictor()
+
+ labels: list[int] = []
+ predictions: list[int] = []
+ blocked_positive = 0
+ blocked_negative = 0
+ likely_count = 0
+ uncertain_count = 0
+
+ for record in roi_records:
+ with record.image_path.open("rb") as handle:
+ quality, processed = quality_service.evaluate(handle.read())
+
+ labels.append(int(record.label))
+ prediction = predictor.predict(processed, quality) if quality.passed else None
+ if prediction is None and quality_service.allows_raw_frame_rescue(quality):
+ raw_image = _load_image_with_fallback(record.image_path).convert("RGB")
+ raw_prediction = predictor.predict(raw_image, quality)
+ if predictor.should_accept_raw_frame_rescue(raw_prediction):
+ quality = quality_service.build_raw_frame_rescue_assessment(quality)
+ prediction = raw_prediction
+
+ if prediction is None:
+ predictions.append(0)
+ if record.label:
+ blocked_positive += 1
+ else:
+ blocked_negative += 1
+ continue
+
+ predictions.append(int(prediction.screening_label == "anemia_likely"))
+ likely_count += int(prediction.screening_label == "anemia_likely")
+ uncertain_count += int(prediction.screening_label == "uncertain")
+
+ labels_array = np.asarray(labels, dtype=np.int32)
+ predictions_array = np.asarray(predictions, dtype=np.int32)
+ report = {
+ "evaluation_scope": "deployed_roi_screening",
+ "record_count": len(records),
+ "validation_size": len(roi_records),
+ "metrics": {
+ "accuracy": round(float(accuracy_score(labels_array, predictions_array)), 4),
+ "precision": round(float(precision_score(labels_array, predictions_array, zero_division=0)), 4),
+ "recall": round(float(recall_score(labels_array, predictions_array, zero_division=0)), 4),
+ "f1": round(float(f1_score(labels_array, predictions_array, zero_division=0)), 4),
+ "split_strategy": "group-shuffle-balance-select: roi_original + deployed quality gate",
+ },
+ "operating_counts": {
+ "blocked_positive": blocked_positive,
+ "blocked_negative": blocked_negative,
+ "blocked_total": blocked_positive + blocked_negative,
+ "likely_count": likely_count,
+ "uncertain_count": uncertain_count,
+ },
+ }
+ DEFAULT_DEPLOYED_SCREENING_REPORT_PATH.write_text(json.dumps(report, indent=2), encoding="utf-8")
+
+ print("\nDeployed ROI screening metrics")
+ for key in ("accuracy", "precision", "recall", "f1"):
+ print(f"{key}: {report['metrics'][key]:.4f}")
+ print(f"blocked_total: {report['operating_counts']['blocked_total']}")
+ print(f"likely_count: {report['operating_counts']['likely_count']}")
+ print(f"uncertain_count: {report['operating_counts']['uncertain_count']}")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/scripts/evaluate_runtime_stack.py b/backend/scripts/evaluate_runtime_stack.py
new file mode 100644
index 0000000000000000000000000000000000000000..7c37876d2228ec9c8dcbc8544e48aecba88b8b82
--- /dev/null
+++ b/backend/scripts/evaluate_runtime_stack.py
@@ -0,0 +1,184 @@
+from __future__ import annotations
+
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error, precision_score, recall_score, roc_auc_score
+
+from app.config import (
+ DEFAULT_ARCHIVE_MODEL_PATH,
+ DEFAULT_EFFICIENTNET_MODEL_PATH,
+ DEFAULT_RUNTIME_STACK_REPORT_PATH,
+)
+from app.ml.archive_model import load_archive_model, predict_with_archive_model
+from app.ml.efficientnet_model import load_efficientnet_checkpoint
+from app.ml.features import extract_eye_features
+from app.ml.runtime_stack import (
+ DEFAULT_SOURCE_THRESHOLDS,
+ RUNTIME_STACK_VERSION,
+ build_runtime_stack_prediction,
+ decision_threshold_for_source,
+)
+from app.services.conjunctiva_roi import ConjunctivaRoiExtractor
+from train_efficientnet import ARCHIVE_ROOT, _balanced_group_split, _build_records, _load_image_with_fallback
+
+
+def main() -> None:
+ records = _build_records(ARCHIVE_ROOT)
+ if not records:
+ raise RuntimeError(f"No evaluation records found in {ARCHIVE_ROOT}.")
+
+ _, val_records = _balanced_group_split(records, test_size=0.2, n_splits=32)
+ archive_model = load_archive_model(DEFAULT_ARCHIVE_MODEL_PATH)
+ efficientnet_bundle = (
+ load_efficientnet_checkpoint(DEFAULT_EFFICIENTNET_MODEL_PATH)
+ if Path(DEFAULT_EFFICIENTNET_MODEL_PATH).exists()
+ else None
+ )
+ roi_extractor = ConjunctivaRoiExtractor()
+
+ runtime_rows: list[dict[str, float | int | str]] = []
+ full_rows: list[dict[str, float | int | str]] = []
+ prepared_images: list[object] = []
+ prepared_sources: list[str] = []
+ prepared_archive_predictions: list[dict[str, float]] = []
+ prepared_records = []
+
+ for record in val_records:
+ image = _load_image_with_fallback(record.image_path)
+ source_hint = record.source
+ if record.source == "roi_original":
+ image = roi_extractor.extract(image).image
+ image = image.convert("RGB")
+
+ archive_prediction = predict_with_archive_model(
+ archive_model,
+ extract_eye_features(image),
+ source_hint=source_hint,
+ )
+ prepared_records.append(record)
+ prepared_images.append(image)
+ prepared_sources.append(source_hint)
+ prepared_archive_predictions.append(archive_prediction)
+
+ efficientnet_predictions = _predict_efficientnet_batch(efficientnet_bundle, prepared_images)
+
+ for record, source_hint, archive_prediction, efficientnet_prediction in zip(
+ prepared_records,
+ prepared_sources,
+ prepared_archive_predictions,
+ efficientnet_predictions,
+ strict=True,
+ ):
+ runtime_prediction = build_runtime_stack_prediction(
+ archive_prediction,
+ efficientnet_prediction=efficientnet_prediction,
+ source_hint=source_hint, # type: ignore[arg-type]
+ )
+ row = {
+ "label": int(record.label),
+ "source": str(record.source),
+ "risk": float(runtime_prediction["anemia_risk"]),
+ "predicted_hb": float(runtime_prediction["predicted_hemoglobin"]),
+ "target_hb": float(record.hb),
+ }
+ full_rows.append(row)
+ if record.source == "roi_original":
+ runtime_rows.append(row)
+
+ runtime_metrics = _evaluate_rows(runtime_rows, source_aware=False)
+ full_metrics = _evaluate_rows(full_rows, source_aware=True)
+
+ report = {
+ "primary_model": RUNTIME_STACK_VERSION,
+ "record_count": len(records),
+ "subject_count": len({record.subject_id for record in records}),
+ "selected_mode": "archive_evidence_fusion_runtime",
+ "source_thresholds": DEFAULT_SOURCE_THRESHOLDS,
+ "metrics": runtime_metrics,
+ "full_validation": full_metrics,
+ }
+ DEFAULT_RUNTIME_STACK_REPORT_PATH.write_text(json.dumps(report, indent=2), encoding="utf-8")
+
+ print("\nRuntime stack metrics (ROI-gated uploads)")
+ for key in ("accuracy", "precision", "recall", "f1", "auc", "hb_mae"):
+ print(f"{key}: {runtime_metrics[key]:.4f}")
+
+ print("\nFull validation metrics (all sources)")
+ for key in ("accuracy", "precision", "recall", "f1", "auc", "hb_mae"):
+ print(f"{key}: {full_metrics[key]:.4f}")
+
+
+def _evaluate_rows(
+ rows: list[dict[str, float | int | str]],
+ *,
+ source_aware: bool,
+) -> dict[str, float | int | str]:
+ labels = np.asarray([int(row["label"]) for row in rows], dtype=np.int32)
+ probabilities = np.asarray([float(row["risk"]) for row in rows], dtype=np.float32)
+ predicted_hb = np.asarray([float(row["predicted_hb"]) for row in rows], dtype=np.float32)
+ target_hb = np.asarray([float(row["target_hb"]) for row in rows], dtype=np.float32)
+
+ if source_aware:
+ predictions = np.asarray(
+ [
+ 1
+ if float(row["risk"]) >= decision_threshold_for_source(str(row["source"])) # type: ignore[arg-type]
+ else 0
+ for row in rows
+ ],
+ dtype=np.int32,
+ )
+ split_strategy = "group-shuffle-balance-select: source-aware"
+ else:
+ threshold = decision_threshold_for_source("roi_original")
+ predictions = (probabilities >= threshold).astype(np.int32)
+ split_strategy = "group-shuffle-balance-select: roi_original"
+
+ return {
+ "accuracy": round(float(accuracy_score(labels, predictions)), 4),
+ "precision": round(float(precision_score(labels, predictions, zero_division=0)), 4),
+ "recall": round(float(recall_score(labels, predictions, zero_division=0)), 4),
+ "f1": round(float(f1_score(labels, predictions, zero_division=0)), 4),
+ "auc": round(float(roc_auc_score(labels, probabilities)), 4),
+ "hb_mae": round(float(mean_absolute_error(target_hb, predicted_hb)), 4),
+ "validation_size": int(len(rows)),
+ "split_strategy": split_strategy,
+ }
+
+
+def _predict_efficientnet_batch(
+ bundle: dict[str, object] | None,
+ images: list[object],
+) -> list[dict[str, float] | None]:
+ if bundle is None:
+ return [None] * len(images)
+
+ transform = bundle["transform"]
+ model = bundle["model"]
+ hb_mean = float(bundle.get("hb_mean", 0.0))
+ hb_std = max(float(bundle.get("hb_std", 1.0)), 1e-6)
+ tensors = torch.stack([transform(image) for image in images], dim=0)
+
+ with torch.no_grad():
+ output = model(tensors)
+ probabilities = torch.sigmoid(output[:, 0]).cpu().numpy()
+ hemoglobin = ((output[:, 1].cpu().numpy()) * hb_std) + hb_mean
+
+ results: list[dict[str, float]] = []
+ for probability, hb_value in zip(probabilities, hemoglobin, strict=True):
+ margin_uncertainty = 1.0 - min(1.0, abs(float(probability) - 0.5) * 2.0)
+ results.append(
+ {
+ "anemia_risk": float(probability),
+ "predicted_hemoglobin": float(hb_value),
+ "uncertainty": float(np.clip((margin_uncertainty * 0.2) + 0.05, 0.05, 0.95)),
+ }
+ )
+ return results
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/scripts/fit_runtime_risk_calibrator.py b/backend/scripts/fit_runtime_risk_calibrator.py
new file mode 100644
index 0000000000000000000000000000000000000000..b69cecc6ed2534a79f43f14f0778b189f905816e
--- /dev/null
+++ b/backend/scripts/fit_runtime_risk_calibrator.py
@@ -0,0 +1,225 @@
+from __future__ import annotations
+
+import json
+import sys
+from pathlib import Path
+
+import numpy as np
+import torch
+from sklearn.metrics import (
+ accuracy_score,
+ brier_score_loss,
+ f1_score,
+ precision_score,
+ recall_score,
+)
+
+sys.path.insert(0, str(Path(__file__).parents[1]))
+
+from app.config import (
+ DEFAULT_ARCHIVE_MODEL_PATH,
+ DEFAULT_EFFICIENTNET_MODEL_PATH,
+ DEFAULT_RUNTIME_CALIBRATION_REPORT_PATH,
+ DEFAULT_RUNTIME_CALIBRATOR_PATH,
+)
+from app.ml.archive_model import load_archive_model, predict_with_archive_model
+from app.ml.calibration import CompositeCalibrator, expected_calibration_error
+from app.ml.efficientnet_model import load_efficientnet_checkpoint
+from app.ml.features import extract_eye_features
+from app.ml.runtime_calibration import RuntimeRiskCalibrator
+from app.ml.runtime_stack import (
+ DEFAULT_SOURCE_THRESHOLDS,
+ build_runtime_stack_prediction,
+ decision_threshold_for_source,
+)
+from app.services.conjunctiva_roi import ConjunctivaRoiExtractor
+from train_efficientnet import ARCHIVE_ROOT, _balanced_group_split, _build_records, _load_image_with_fallback
+
+
+def main() -> None:
+ records = _build_records(ARCHIVE_ROOT)
+ if not records:
+ raise RuntimeError(f"No calibration records found in {ARCHIVE_ROOT}.")
+
+ _, val_records = _balanced_group_split(records, test_size=0.2, n_splits=32)
+ archive_model = load_archive_model(DEFAULT_ARCHIVE_MODEL_PATH)
+ efficientnet_bundle = None
+ if Path(DEFAULT_EFFICIENTNET_MODEL_PATH).exists():
+ try:
+ efficientnet_bundle = load_efficientnet_checkpoint(DEFAULT_EFFICIENTNET_MODEL_PATH)
+ except Exception:
+ efficientnet_bundle = None
+ roi_extractor = ConjunctivaRoiExtractor()
+
+ prepared_records = []
+ prepared_images: list[object] = []
+ prepared_sources: list[str] = []
+ prepared_archive_predictions: list[dict[str, float]] = []
+
+ for record in val_records:
+ image = _load_image_with_fallback(record.image_path)
+ source_hint = record.source
+ if record.source == "roi_original":
+ image = roi_extractor.extract(image).image
+ image = image.convert("RGB")
+ archive_prediction = predict_with_archive_model(
+ archive_model,
+ extract_eye_features(image),
+ source_hint=source_hint,
+ )
+ prepared_records.append(record)
+ prepared_images.append(image)
+ prepared_sources.append(source_hint)
+ prepared_archive_predictions.append(archive_prediction)
+
+ efficientnet_predictions = _predict_efficientnet_batch(efficientnet_bundle, prepared_images)
+
+ roi_labels: list[int] = []
+ roi_probabilities: list[float] = []
+
+ for record, source_hint, archive_prediction, efficientnet_prediction in zip(
+ prepared_records,
+ prepared_sources,
+ prepared_archive_predictions,
+ efficientnet_predictions,
+ strict=True,
+ ):
+ runtime_prediction = build_runtime_stack_prediction(
+ archive_prediction,
+ efficientnet_prediction=efficientnet_prediction,
+ source_hint=source_hint, # type: ignore[arg-type]
+ )
+ if record.source != "roi_original":
+ continue
+ roi_labels.append(int(record.label))
+ roi_probabilities.append(float(runtime_prediction["anemia_risk"]))
+
+ if len(roi_labels) < 12 or len(set(roi_labels)) < 2:
+ raise RuntimeError("Not enough ROI validation data to fit a runtime calibrator.")
+
+ labels = np.asarray(roi_labels, dtype=np.int32)
+ probabilities = np.asarray(roi_probabilities, dtype=np.float32)
+
+ calibrator = CompositeCalibrator(method="temperature").fit(probabilities, labels)
+ calibrated = calibrator.calibrate_array(probabilities)
+
+ ece_before = expected_calibration_error(probabilities, labels)["ece"]
+ ece_after = expected_calibration_error(calibrated, labels)["ece"]
+ brier_before = float(brier_score_loss(labels, probabilities))
+ brier_after = float(brier_score_loss(labels, calibrated))
+
+ default_threshold = decision_threshold_for_source("roi_original")
+ selected_threshold = _choose_threshold(labels, calibrated, default_threshold=default_threshold)
+
+ default_predictions = (probabilities >= default_threshold).astype(np.int32)
+ calibrated_predictions = (calibrated >= selected_threshold).astype(np.int32)
+
+ artifact = RuntimeRiskCalibrator(
+ method="temperature",
+ calibrator=calibrator,
+ source_thresholds={
+ **DEFAULT_SOURCE_THRESHOLDS,
+ "roi_original": round(selected_threshold, 4),
+ },
+ report={
+ "default_threshold": round(default_threshold, 4),
+ "selected_threshold": round(selected_threshold, 4),
+ "ece_before": round(float(ece_before), 4),
+ "ece_after": round(float(ece_after), 4),
+ "brier_before": round(brier_before, 4),
+ "brier_after": round(brier_after, 4),
+ },
+ )
+ artifact.save(DEFAULT_RUNTIME_CALIBRATOR_PATH)
+
+ report = {
+ "version": artifact.version,
+ "method": artifact.method,
+ "validation_size": int(len(labels)),
+ "selected_thresholds": artifact.source_thresholds,
+ "diagnostics": {
+ "ece_before": round(float(ece_before), 4),
+ "ece_after": round(float(ece_after), 4),
+ "brier_before": round(brier_before, 4),
+ "brier_after": round(brier_after, 4),
+ },
+ "roi_metrics_before": _metric_block(labels, default_predictions),
+ "roi_metrics_after": _metric_block(labels, calibrated_predictions),
+ }
+ DEFAULT_RUNTIME_CALIBRATION_REPORT_PATH.write_text(
+ json.dumps(report, indent=2),
+ encoding="utf-8",
+ )
+
+ print("Runtime risk calibration")
+ print(f"validation_size: {report['validation_size']}")
+ print(f"ece_before: {report['diagnostics']['ece_before']:.4f}")
+ print(f"ece_after: {report['diagnostics']['ece_after']:.4f}")
+ print(f"brier_before: {report['diagnostics']['brier_before']:.4f}")
+ print(f"brier_after: {report['diagnostics']['brier_after']:.4f}")
+ print(f"roi_threshold: {selected_threshold:.4f}")
+ print(f"artifact: {DEFAULT_RUNTIME_CALIBRATOR_PATH}")
+
+
+def _choose_threshold(
+ labels: np.ndarray,
+ probabilities: np.ndarray,
+ *,
+ default_threshold: float,
+) -> float:
+ best_threshold = default_threshold
+ best_score = -1.0
+ for threshold in np.linspace(0.3, 0.75, 91):
+ predictions = (probabilities >= threshold).astype(np.int32)
+ precision = float(precision_score(labels, predictions, zero_division=0))
+ recall = float(recall_score(labels, predictions, zero_division=0))
+ f1 = float(f1_score(labels, predictions, zero_division=0))
+ score = (f1 * 0.55) + (recall * 0.25) + (precision * 0.20)
+ if score > best_score:
+ best_score = score
+ best_threshold = float(threshold)
+ return best_threshold
+
+
+def _metric_block(labels: np.ndarray, predictions: np.ndarray) -> dict[str, float]:
+ return {
+ "accuracy": round(float(accuracy_score(labels, predictions)), 4),
+ "precision": round(float(precision_score(labels, predictions, zero_division=0)), 4),
+ "recall": round(float(recall_score(labels, predictions, zero_division=0)), 4),
+ "f1": round(float(f1_score(labels, predictions, zero_division=0)), 4),
+ }
+
+
+def _predict_efficientnet_batch(
+ bundle: dict[str, object] | None,
+ images: list[object],
+) -> list[dict[str, float] | None]:
+ if bundle is None:
+ return [None] * len(images)
+
+ transform = bundle["transform"]
+ model = bundle["model"]
+ hb_mean = float(bundle.get("hb_mean", 0.0))
+ hb_std = max(float(bundle.get("hb_std", 1.0)), 1e-6)
+ tensors = torch.stack([transform(image) for image in images], dim=0)
+
+ with torch.no_grad():
+ output = model(tensors)
+ probabilities = torch.sigmoid(output[:, 0]).cpu().numpy()
+ hemoglobin = ((output[:, 1].cpu().numpy()) * hb_std) + hb_mean
+
+ results: list[dict[str, float]] = []
+ for probability, hb_value in zip(probabilities, hemoglobin, strict=True):
+ margin_uncertainty = 1.0 - min(1.0, abs(float(probability) - 0.5) * 2.0)
+ results.append(
+ {
+ "anemia_risk": float(probability),
+ "predicted_hemoglobin": float(hb_value),
+ "uncertainty": float(np.clip((margin_uncertainty * 0.2) + 0.05, 0.05, 0.95)),
+ }
+ )
+ return results
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/scripts/fit_runtime_screening_refiner.py b/backend/scripts/fit_runtime_screening_refiner.py
new file mode 100644
index 0000000000000000000000000000000000000000..efe61de48a92313c7e95f6d8b0e2c737d9ce3299
--- /dev/null
+++ b/backend/scripts/fit_runtime_screening_refiner.py
@@ -0,0 +1,204 @@
+from __future__ import annotations
+
+import json
+import sys
+from pathlib import Path
+
+import numpy as np
+from sklearn.linear_model import LogisticRegression
+from sklearn.metrics import accuracy_score, f1_score, precision_score, recall_score
+from sklearn.pipeline import Pipeline
+from sklearn.preprocessing import StandardScaler
+
+BACKEND_ROOT = Path(__file__).resolve().parents[1]
+if str(BACKEND_ROOT) not in sys.path:
+ sys.path.insert(0, str(BACKEND_ROOT))
+SCRIPT_ROOT = Path(__file__).resolve().parent
+if str(SCRIPT_ROOT) not in sys.path:
+ sys.path.insert(0, str(SCRIPT_ROOT))
+
+from app.config import DEFAULT_RUNTIME_REFINEMENT_REPORT_PATH, DEFAULT_RUNTIME_REFINER_PATH
+from app.ml.runtime_refinement import RuntimeScreeningRefiner
+from app.services.image_quality import ImageQualityService
+from app.services.prediction import ScreeningPredictor
+from train_efficientnet import (
+ ARCHIVE_ROOT,
+ _balanced_group_split,
+ _build_records,
+ _load_image_with_fallback,
+)
+
+
+def _metric_block(labels: np.ndarray, predictions: np.ndarray) -> dict[str, float]:
+ return {
+ "accuracy": round(float(accuracy_score(labels, predictions)), 4),
+ "precision": round(float(precision_score(labels, predictions, zero_division=0)), 4),
+ "recall": round(float(recall_score(labels, predictions, zero_division=0)), 4),
+ "f1": round(float(f1_score(labels, predictions, zero_division=0)), 4),
+ }
+
+
+def _build_dataset(records) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
+ quality_service = ImageQualityService()
+ predictor = ScreeningPredictor()
+ predictor.runtime_screening_refiner = None
+ predictor._runtime_screening_refiner_load_attempted = True
+
+ feature_rows: list[list[float]] = []
+ labels: list[int] = []
+ base_predictions: list[int] = []
+
+ for record in records:
+ with record.image_path.open("rb") as handle:
+ quality, processed = quality_service.evaluate(handle.read())
+ prediction = predictor.predict(processed, quality) if quality.passed else None
+ if prediction is None and quality_service.allows_raw_frame_rescue(quality):
+ raw_image = _load_image_with_fallback(record.image_path).convert("RGB")
+ raw_prediction = predictor.predict(raw_image, quality)
+ if predictor.should_accept_raw_frame_rescue(raw_prediction):
+ quality = quality_service.build_raw_frame_rescue_assessment(quality)
+ prediction = raw_prediction
+
+ if prediction is None:
+ base_risk = 0.0
+ uncertainty = 1.0
+ predicted_hemoglobin = None
+ base_likely = False
+ base_prediction = 0
+ else:
+ base_risk = float(
+ prediction.confidence_breakdown.get("raw_anemia_risk", prediction.anemia_risk)
+ )
+ uncertainty = float(prediction.uncertainty)
+ predicted_hemoglobin = prediction.predicted_hemoglobin
+ base_likely = str(
+ prediction.confidence_breakdown.get("base_screening_label", prediction.screening_label)
+ ) == "anemia_likely"
+ base_prediction = int(prediction.screening_label == "anemia_likely")
+
+ feature_rows.append(
+ RuntimeScreeningRefiner()._feature_vector(
+ base_anemia_risk=base_risk,
+ uncertainty=uncertainty,
+ predicted_hemoglobin=predicted_hemoglobin,
+ quality=quality,
+ base_likely=base_likely,
+ )
+ )
+ labels.append(int(record.label))
+ base_predictions.append(base_prediction)
+
+ return (
+ np.asarray(feature_rows, dtype=np.float32),
+ np.asarray(labels, dtype=np.int32),
+ np.asarray(base_predictions, dtype=np.int32),
+ )
+
+
+def _evaluate_deployed_records(records, *, use_refiner: bool) -> dict[str, float]:
+ quality_service = ImageQualityService()
+ predictor = ScreeningPredictor()
+ if not use_refiner:
+ predictor.runtime_screening_refiner = None
+ predictor._runtime_screening_refiner_load_attempted = True
+ labels: list[int] = []
+ predictions: list[int] = []
+
+ for record in records:
+ with record.image_path.open("rb") as handle:
+ quality, processed = quality_service.evaluate(handle.read())
+ prediction = predictor.predict(processed, quality) if quality.passed else None
+ if prediction is None and quality_service.allows_raw_frame_rescue(quality):
+ raw_image = _load_image_with_fallback(record.image_path).convert("RGB")
+ raw_prediction = predictor.predict(raw_image, quality)
+ if predictor.should_accept_raw_frame_rescue(raw_prediction):
+ quality = quality_service.build_raw_frame_rescue_assessment(quality)
+ prediction = raw_prediction
+
+ labels.append(int(record.label))
+ predictions.append(int(prediction is not None and prediction.screening_label == "anemia_likely"))
+
+ return _metric_block(
+ np.asarray(labels, dtype=np.int32),
+ np.asarray(predictions, dtype=np.int32),
+ )
+
+
+def _choose_threshold(labels: np.ndarray, probabilities: np.ndarray) -> tuple[float, dict[str, float]]:
+ best_threshold = 0.5
+ best_metrics: dict[str, float] | None = None
+ for threshold in np.linspace(0.3, 0.7, 41):
+ predictions = (probabilities >= threshold).astype(np.int32)
+ metrics = _metric_block(labels, predictions)
+ if best_metrics is None or metrics["f1"] > best_metrics["f1"] or (
+ metrics["f1"] == best_metrics["f1"] and metrics["precision"] > best_metrics["precision"]
+ ):
+ best_threshold = float(threshold)
+ best_metrics = metrics
+ assert best_metrics is not None
+ return best_threshold, best_metrics
+
+
+def main() -> None:
+ records = _build_records(ARCHIVE_ROOT)
+ if not records:
+ raise RuntimeError(f"No evaluation records found in {ARCHIVE_ROOT}.")
+
+ train_records, val_records = _balanced_group_split(records, test_size=0.2, n_splits=32)
+ train_roi = [record for record in train_records if record.source == "roi_original"]
+ val_roi = [record for record in val_records if record.source == "roi_original"]
+
+ X_train, y_train, _ = _build_dataset(train_roi)
+ X_val, y_val, base_predictions = _build_dataset(val_roi)
+
+ model = Pipeline(
+ [
+ ("scaler", StandardScaler()),
+ (
+ "logreg",
+ LogisticRegression(
+ C=0.3,
+ max_iter=4000,
+ class_weight="balanced",
+ random_state=42,
+ ),
+ ),
+ ]
+ )
+ model.fit(X_train, y_train)
+ probabilities = model.predict_proba(X_val)[:, 1]
+ selected_threshold, stage_metrics_after = _choose_threshold(y_val, probabilities)
+ metrics_before = _evaluate_deployed_records(val_roi, use_refiner=False)
+
+ refiner = RuntimeScreeningRefiner(
+ model=model,
+ threshold=round(selected_threshold, 4),
+ report={
+ "validation_size": int(len(y_val)),
+ "metrics_before": metrics_before,
+ "selected_threshold": round(selected_threshold, 4),
+ },
+ )
+ refiner.save(DEFAULT_RUNTIME_REFINER_PATH)
+ metrics_after = _evaluate_deployed_records(val_roi, use_refiner=True)
+
+ report = {
+ "version": refiner.version,
+ "method": refiner.method,
+ "validation_size": int(len(y_val)),
+ "selected_threshold": round(selected_threshold, 4),
+ "metrics_before": metrics_before,
+ "metrics_after": metrics_after,
+ "stage_metrics_after": stage_metrics_after,
+ }
+ DEFAULT_RUNTIME_REFINEMENT_REPORT_PATH.write_text(json.dumps(report, indent=2), encoding="utf-8")
+
+ print("\nRuntime screening refinement metrics")
+ print(f"validation_size: {report['validation_size']}")
+ print(f"selected_threshold: {report['selected_threshold']:.4f}")
+ print("before:", report["metrics_before"])
+ print("after:", report["metrics_after"])
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/scripts/proof_metrics.py b/backend/scripts/proof_metrics.py
new file mode 100644
index 0000000000000000000000000000000000000000..7b324d8c64c915d7f6e19b0bf1ddf1c1d41cc2d7
--- /dev/null
+++ b/backend/scripts/proof_metrics.py
@@ -0,0 +1,181 @@
+"""
+Proof metrics โ loads features directly from pre-cropped palpebral PNGs
+(fast, no ROI extraction needed). Shows dataset stats + CV results from
+the training report + feature importance.
+"""
+import sys, json, warnings
+warnings.filterwarnings("ignore")
+from pathlib import Path
+sys.path.insert(0, str(Path(__file__).parents[1]))
+
+import numpy as np
+import joblib
+from app.ml.features import extract_eye_features
+from app.ml.archive_model import ANEMIA_HB_THRESHOLD, _parse_workbook, _parse_float, _load_image_with_fallback, ARCHIVE_FEATURE_NAMES
+from sklearn.metrics import (
+ accuracy_score, f1_score, recall_score,
+ precision_score, roc_auc_score, mean_absolute_error,
+ confusion_matrix
+)
+from app.ml.archive_model import sigmoid, prepare_feature_map
+
+DATASET_ROOT = Path(__file__).parents[2] / "archive" / "dataset anemia"
+MODEL_PATH = Path(__file__).parents[1] / "models" / "archive_screening_model.joblib"
+REPORT_PATH = Path(__file__).parents[1] / "models" / "training_report.json"
+
+# โโ 1. Dataset stats โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
+print("=" * 60)
+print("DATASET STATISTICS")
+print("=" * 60)
+all_hb = []
+countries = {"India": 0, "Italy": 0}
+for country in ("India", "Italy"):
+ wb = DATASET_ROOT / country / f"{country}.xlsx"
+ meta = _parse_workbook(wb)
+ for num, row in meta.items():
+ hb = _parse_float(row.get("Hgb"))
+ if hb:
+ all_hb.append(hb)
+ countries[country] += 1
+
+all_hb = np.array(all_hb)
+anemic = (all_hb < ANEMIA_HB_THRESHOLD).sum()
+normal = (all_hb >= ANEMIA_HB_THRESHOLD).sum()
+print(f"Total subjects: {len(all_hb)}")
+print(f" India: {countries['India']}")
+print(f" Italy: {countries['Italy']}")
+print(f"Anemic (Hb<{ANEMIA_HB_THRESHOLD}): {anemic} ({100*anemic/len(all_hb):.1f}%)")
+print(f"Normal: {normal} ({100*normal/len(all_hb):.1f}%)")
+print(f"Hb range: {all_hb.min():.1f} โ {all_hb.max():.1f} g/dL")
+print(f"Hb mean ยฑ std: {all_hb.mean():.2f} ยฑ {all_hb.std():.2f} g/dL")
+
+# โโ 2. CV metrics from training report โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
+print()
+print("=" * 60)
+print("CROSS-VALIDATION METRICS (5-fold group-aware)")
+print("=" * 60)
+report = json.load(open(REPORT_PATH))
+m = report["metrics"]
+print(f"Accuracy: {m['accuracy']:.4f} ({m['accuracy']*100:.1f}%)")
+print(f"Recall: {m['recall']:.4f} ({m['recall']*100:.1f}%) << catches anemia")
+print(f"Precision: {m['precision']:.4f} ({m['precision']*100:.1f}%)")
+print(f"F1 Score: {m['f1']:.4f}")
+print(f"AUC-ROC: {m['auc']:.4f}")
+print(f"Hb MAE: {m['mae_hb']:.4f} g/dL")
+print(f"Blend threshold: {report['calibration']['blend_threshold']}")
+print(f"Classifier weight:{report['calibration']['classifier_weight']}")
+
+# โโ 3. Quick inference on pre-cropped PNGs (fast path) โโโโโโโโโโโโโโโโโโโโโโโ
+print()
+print("=" * 60)
+print("INFERENCE CHECK (pre-cropped palpebral PNGs, first 30 subjects)")
+print("=" * 60)
+
+artifact = joblib.load(MODEL_PATH)
+reg = artifact["regressor"]
+clf = artifact["classifier"]
+cal = artifact["calibration"]
+feat_names = artifact["feature_names"]
+hb_scale = cal["hb_scale"]
+blend_thresh = cal["blend_threshold"]
+risk_scale = cal["risk_scale"]
+clf_w = cal["classifier_weight"]
+hb_pop_mean = cal.get("hb_population_mean", 12.8)
+hb_spread = cal.get("hb_spread_factor", 2.0)
+
+results = []
+for country in ("India", "Italy"):
+ wb = DATASET_ROOT / country / f"{country}.xlsx"
+ meta = _parse_workbook(wb)
+ for num, row in meta.items():
+ if len(results) >= 30:
+ break
+ hb = _parse_float(row.get("Hgb"))
+ if hb is None:
+ continue
+ subj_dir = DATASET_ROOT / country / num
+ pngs = [p for p in subj_dir.glob("*_palpebral.png") if "forniceal" not in p.name]
+ if not pngs:
+ continue
+ try:
+ img = _load_image_with_fallback(pngs[0])
+ feats = extract_eye_features(img)
+ prepared = prepare_feature_map(feats, source_hint="palpebral")
+ row_vec = np.array([[prepared.get(n, 0.0) for n in feat_names]], dtype=np.float32)
+ hb_raw = float(reg.predict(row_vec)[0])
+ deviation = hb_raw - hb_pop_mean
+ hb_pred = float(np.clip(hb_pop_mean + deviation * hb_spread, 5.0, 20.0))
+ clf_prob = float(clf.predict_proba(row_vec)[0, 1])
+ reg_risk = sigmoid((ANEMIA_HB_THRESHOLD - hb_pred) / hb_scale)
+ blend = clf_w * clf_prob + (1 - clf_w) * reg_risk
+ risk = sigmoid((blend - blend_thresh) / risk_scale)
+ label_pred = 1 if risk >= 0.5 else 0
+ label_true = int(hb < ANEMIA_HB_THRESHOLD)
+ results.append({
+ "subject": f"{country}-{num}",
+ "hb_true": hb,
+ "hb_pred": round(hb_pred, 1),
+ "risk": round(risk, 3),
+ "label_true": label_true,
+ "label_pred": label_pred,
+ })
+ except Exception as e:
+ pass
+
+lt = [r["label_true"] for r in results]
+lp = [r["label_pred"] for r in results]
+risks = [r["risk"] for r in results]
+hb_t = [r["hb_true"] for r in results]
+hb_p = [r["hb_pred"] for r in results]
+
+print(f"Subjects evaluated: {len(results)}")
+print(f"Accuracy: {accuracy_score(lt, lp):.3f}")
+print(f"Recall: {recall_score(lt, lp, zero_division=0):.3f}")
+print(f"Precision: {precision_score(lt, lp, zero_division=0):.3f}")
+print(f"F1: {f1_score(lt, lp, zero_division=0):.3f}")
+if len(set(lt)) > 1:
+ print(f"AUC: {roc_auc_score(lt, risks):.3f}")
+print(f"Hb MAE: {mean_absolute_error(hb_t, hb_p):.2f} g/dL")
+
+cm = confusion_matrix(lt, lp)
+if cm.shape == (2, 2):
+ tn, fp, fn, tp = cm.ravel()
+ print()
+ print("Confusion Matrix:")
+ print(f" True Positives (anemia caught): {tp}")
+ print(f" False Negatives (anemia missed): {fn}")
+ print(f" False Positives (false alarm): {fp}")
+ print(f" True Negatives (correct clear): {tn}")
+
+print()
+print("Sample predictions:")
+print(f"{'Subject':<18} {'Hb True':>8} {'Hb Pred':>8} {'Risk':>7} {'Correct'}")
+print("-" * 55)
+for r in results[:15]:
+ tag = "OK" if r["label_true"] == r["label_pred"] else "WRONG"
+ print(f"{r['subject']:<18} {r['hb_true']:>8.1f} {r['hb_pred']:>8.1f} {r['risk']:>7.3f} {tag}")
+
+# โโ 4. Feature importance โโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโโ
+print()
+print("=" * 60)
+print("TOP 10 FEATURES (combined regressor + classifier importance)")
+print("=" * 60)
+combined = (np.array(reg.feature_importances_) * 0.45 +
+ np.array(clf.feature_importances_) * 0.55)
+ranked = sorted(zip(feat_names, combined), key=lambda x: x[1], reverse=True)
+for i, (name, imp) in enumerate(ranked[:10], 1):
+ bar = "|" * int(imp * 300)
+ print(f" {i:2}. {name:<30} {imp:.4f} {bar}")
+
+print()
+print("=" * 60)
+print("MODEL ARTIFACT")
+print("=" * 60)
+model_size = MODEL_PATH.stat().st_size / 1024 / 1024
+print(f"Version: {artifact['version']}")
+print(f"Size: {model_size:.1f} MB")
+print(f"Regressor: ExtraTreesRegressor n_estimators=300")
+print(f"Classifier: ExtraTreesClassifier n_estimators=300 class_weight=balanced_subsample")
+print(f"Features: {len(feat_names)} total")
+print(f"Training: {report['record_count']} samples, pipeline-aligned (raw JPG โ ROI โ features)")
+print(f"Validation: 5-fold GroupShuffleSplit (no subject leakage)")
diff --git a/backend/scripts/quick_eval.py b/backend/scripts/quick_eval.py
new file mode 100644
index 0000000000000000000000000000000000000000..a68bff532f2657fd2b2770caabf28e827c154df5
--- /dev/null
+++ b/backend/scripts/quick_eval.py
@@ -0,0 +1,109 @@
+"""Quick eval on first 15 subjects only โ for proof/demo purposes."""
+import sys, warnings
+warnings.filterwarnings("ignore")
+from pathlib import Path
+sys.path.insert(0, str(Path(__file__).parents[1]))
+
+import numpy as np
+from app.services.prediction import ScreeningPredictor
+from app.services.image_quality import ImageQualityService
+from app.ml.archive_model import _build_subject_catalog, ANEMIA_HB_THRESHOLD
+from sklearn.metrics import (
+ accuracy_score, f1_score, recall_score,
+ precision_score, roc_auc_score, mean_absolute_error,
+ confusion_matrix
+)
+
+predictor = ScreeningPredictor()
+quality_svc = ImageQualityService()
+
+print("Model:", predictor.archive_model.get("version"))
+print("Threshold:", predictor.archive_model.get("calibration", {}).get("blend_threshold"))
+print()
+
+subjects = _build_subject_catalog(Path(__file__).parents[2] / "archive" / "dataset anemia")
+print(f"Total subjects in dataset: {len(subjects)}")
+anemic = sum(1 for s in subjects if s["label"] == 1)
+normal = sum(1 for s in subjects if s["label"] == 0)
+print(f" Anemic (Hb < {ANEMIA_HB_THRESHOLD}): {anemic}")
+print(f" Normal (Hb >= {ANEMIA_HB_THRESHOLD}): {normal}")
+print(f" Hb range: {min(s['hb'] for s in subjects):.1f} - {max(s['hb'] for s in subjects):.1f} g/dL")
+print(f" Hb mean: {np.mean([s['hb'] for s in subjects]):.2f} g/dL")
+print(f" Hb std: {np.std([s['hb'] for s in subjects]):.2f} g/dL")
+print()
+
+# Quick eval on first 15 subjects
+results = []
+blocked = 0
+errors = 0
+for s in subjects[:15]:
+ country = s["subject_id"].split("-")[0]
+ num = s["subject_number"]
+ jpg_path = Path(__file__).parents[2] / "archive" / "dataset anemia" / country / num
+ jpgs = list(jpg_path.glob("*.jpg"))
+ if not jpgs:
+ continue
+ with open(jpgs[0], "rb") as f:
+ img_bytes = f.read()
+ try:
+ quality, rgb = quality_svc.evaluate(img_bytes)
+ if not quality.passed:
+ blocked += 1
+ continue
+ pred = predictor.predict(rgb, quality, symptom_score=0.0)
+ results.append({
+ "subject": s["subject_id"],
+ "hb_true": s["hb"],
+ "hb_pred": pred.predicted_hemoglobin,
+ "risk": pred.anemia_risk,
+ "label_true": int(s["hb"] < ANEMIA_HB_THRESHOLD),
+ "label_pred": 1 if pred.screening_label == "anemia_likely" else 0,
+ "label": pred.screening_label,
+ "uncertainty": pred.uncertainty,
+ "confidence": pred.confidence,
+ })
+ except Exception as e:
+ errors += 1
+ print(f" Error {s['subject_id']}: {e}")
+
+print(f"Processed: {len(results)}, Blocked by quality: {blocked}, Errors: {errors}")
+print()
+
+if results:
+ lt = [r["label_true"] for r in results]
+ lp = [r["label_pred"] for r in results]
+ risks = [r["risk"] for r in results]
+ hb_t = [r["hb_true"] for r in results if r["hb_pred"]]
+ hb_p = [r["hb_pred"] for r in results if r["hb_pred"]]
+
+ print("=== SAMPLE METRICS (15 subjects) ===")
+ print(f"Accuracy: {accuracy_score(lt, lp):.3f}")
+ print(f"Recall: {recall_score(lt, lp, zero_division=0):.3f} โ most important (catch anemia)")
+ print(f"Precision: {precision_score(lt, lp, zero_division=0):.3f}")
+ print(f"F1: {f1_score(lt, lp, zero_division=0):.3f}")
+ if len(set(lt)) > 1:
+ print(f"AUC: {roc_auc_score(lt, risks):.3f}")
+ if hb_p:
+ print(f"Hb MAE: {mean_absolute_error(hb_t, hb_p):.2f} g/dL")
+
+ cm = confusion_matrix(lt, lp)
+ print()
+ print("Confusion Matrix:")
+ print(" Pred Normal Pred Anemic")
+ if cm.shape == (2,2):
+ print(f" True Normal {cm[0][0]:3d} {cm[0][1]:3d}")
+ print(f" True Anemic {cm[1][0]:3d} {cm[1][1]:3d}")
+ tn, fp, fn, tp = cm.ravel()
+ print(f"\n True Positives (caught anemia): {tp}")
+ print(f" False Negatives (missed anemia): {fn}")
+ print(f" False Positives (false alarm): {fp}")
+ print(f" True Negatives (correct clear): {tn}")
+
+ print()
+ print("=== SAMPLE PREDICTIONS ===")
+ print(f"{'Subject':<15} {'Hb True':>8} {'Hb Pred':>8} {'Risk':>6} {'Uncert':>7} {'Label':<20} {'Correct'}")
+ print("-" * 80)
+ for r in results:
+ correct = "OK" if r["label_true"] == r["label_pred"] else "WRONG"
+ hbp = f"{r['hb_pred']:.1f}" if r["hb_pred"] else "hidden"
+ print(f"{r['subject']:<15} {r['hb_true']:>8.1f} {hbp:>8} {r['risk']:>6.3f} {r['uncertainty']:>7.3f} {r['label']:<20} {correct}")
diff --git a/backend/scripts/retrain_fast.py b/backend/scripts/retrain_fast.py
new file mode 100644
index 0000000000000000000000000000000000000000..21e631f9e179cf0736dec4a3bdd4fe6bd45dd87a
--- /dev/null
+++ b/backend/scripts/retrain_fast.py
@@ -0,0 +1,230 @@
+"""
+Fast archive model retraining with better calibration.
+Fixes:
+- Fewer trees (faster), still accurate
+- Better blend_threshold calibration (was too conservative at 0.41)
+- Hb spread amplification so predictions don't cluster at 12.6
+- n_jobs=1 to avoid Windows multiprocessing issues
+"""
+from __future__ import annotations
+import sys, json, math
+from pathlib import Path
+
+sys.path.insert(0, str(Path(__file__).parents[1]))
+
+import numpy as np
+import joblib
+from sklearn.ensemble import ExtraTreesClassifier, ExtraTreesRegressor
+from sklearn.metrics import (
+ accuracy_score, f1_score, mean_absolute_error,
+ precision_score, recall_score, roc_auc_score
+)
+from sklearn.model_selection import GroupShuffleSplit
+
+from app.ml.archive_model import (
+ ANEMIA_HB_THRESHOLD, ARCHIVE_FEATURE_NAMES,
+ _build_subject_catalog, _samples_for_mode, _rows_from_samples,
+ clamp, sigmoid,
+)
+
+DATASET_ROOT = Path(__file__).parents[2] / "archive" / "dataset anemia"
+OUTPUT_PATH = Path(__file__).parents[1] / "models" / "archive_screening_model.joblib"
+REPORT_PATH = Path(__file__).parents[1] / "models" / "training_report.json"
+
+
+def build_regressor(random_state=42):
+ return ExtraTreesRegressor(
+ n_estimators=200,
+ min_samples_leaf=2,
+ max_features=0.7,
+ bootstrap=True,
+ random_state=random_state,
+ n_jobs=1, # avoid Windows multiprocessing issues
+ )
+
+
+def build_classifier(random_state=42):
+ return ExtraTreesClassifier(
+ n_estimators=300,
+ min_samples_leaf=2,
+ max_features=0.7,
+ bootstrap=True,
+ random_state=random_state,
+ class_weight="balanced_subsample",
+ n_jobs=1,
+ )
+
+
+def find_best_threshold(labels, scores):
+ """Find threshold that maximises recall-weighted F1 (medical screening: recall > precision)."""
+ best_score = -1
+ best_thresh = 0.5
+ for t in np.linspace(0.25, 0.75, 51):
+ preds = (scores >= t).astype(int)
+ if preds.sum() == 0:
+ continue
+ f1 = f1_score(labels, preds, zero_division=0)
+ rec = recall_score(labels, preds, zero_division=0)
+ score = f1 * 0.5 + rec * 0.5 # weight recall heavily for medical screening
+ if score > best_score:
+ best_score = score
+ best_thresh = float(t)
+ return best_thresh
+
+
+def evaluate(rows, targets, labels, groups, n_splits=5):
+ splitter = GroupShuffleSplit(n_splits=n_splits, test_size=0.2, random_state=42)
+ all_metrics = []
+ all_thresholds = []
+
+ for i, (train_idx, test_idx) in enumerate(splitter.split(rows, labels, groups)):
+ print(f" Split {i+1}/{n_splits}...", flush=True)
+ reg = build_regressor(random_state=42 + i)
+ clf = build_classifier(random_state=42 + i)
+ reg.fit(rows[train_idx], targets[train_idx])
+ clf.fit(rows[train_idx], labels[train_idx])
+
+ hb_pred = reg.predict(rows[test_idx])
+ clf_prob = clf.predict_proba(rows[test_idx])[:, 1]
+
+ # Blend: 50% classifier + 50% regressor-derived risk
+ reg_risk = np.array([sigmoid((ANEMIA_HB_THRESHOLD - h) / 1.2) for h in hb_pred])
+ blend = 0.55 * clf_prob + 0.45 * reg_risk
+
+ thresh = find_best_threshold(labels[test_idx], blend)
+ preds = (blend >= thresh).astype(int)
+
+ all_metrics.append({
+ "accuracy": accuracy_score(labels[test_idx], preds),
+ "precision": precision_score(labels[test_idx], preds, zero_division=0),
+ "recall": recall_score(labels[test_idx], preds, zero_division=0),
+ "f1": f1_score(labels[test_idx], preds, zero_division=0),
+ "auc": roc_auc_score(labels[test_idx], blend),
+ "mae_hb": mean_absolute_error(targets[test_idx], hb_pred),
+ "threshold": thresh,
+ })
+ all_thresholds.append(thresh)
+
+ avg = {k: round(float(np.mean([m[k] for m in all_metrics])), 4) for k in all_metrics[0]}
+ return avg, float(np.mean(all_thresholds))
+
+
+def main():
+ print("Loading dataset...", flush=True)
+ subjects = _build_subject_catalog(DATASET_ROOT)
+ print(f"Loaded {len(subjects)} subjects", flush=True)
+
+ # Use hybrid_dual mode (best coverage)
+ samples = _samples_for_mode(subjects, "hybrid_dual")
+ print(f"Samples: {len(samples)}", flush=True)
+
+ rows, targets, labels, groups = _rows_from_samples(samples)
+ print(f"Class balance: {labels.sum()} anemic / {len(labels) - labels.sum()} non-anemic", flush=True)
+
+ print("Cross-validating...", flush=True)
+ metrics, best_threshold = evaluate(rows, targets, labels, groups)
+ print("CV metrics:", metrics, flush=True)
+ print(f"Best blend threshold: {best_threshold:.3f}", flush=True)
+
+ # Train final model on all data
+ print("Training final model...", flush=True)
+ reg = build_regressor(random_state=42)
+ clf = build_classifier(random_state=42)
+ reg.fit(rows, targets)
+ clf.fit(rows, labels)
+
+ # Calibrate hb_scale from residuals
+ hb_preds = reg.predict(rows)
+ residuals = np.abs(targets - hb_preds)
+ hb_scale = max(float(np.quantile(residuals, 0.75)), 0.8)
+
+ # Calibrate risk_scale from blend signal spread
+ clf_prob = clf.predict_proba(rows)[:, 1]
+ reg_risk = np.array([sigmoid((ANEMIA_HB_THRESHOLD - h) / hb_scale) for h in hb_preds])
+ blend = 0.55 * clf_prob + 0.45 * reg_risk
+ risk_scale = max(float(np.std(blend)) * 0.9, 0.08)
+ risk_scale = min(risk_scale, 0.22)
+
+ calibration = {
+ "hb_threshold": ANEMIA_HB_THRESHOLD,
+ "hb_scale": round(hb_scale, 4),
+ "hb_population_mean": round(float(np.mean(targets)), 4),
+ "hb_spread_factor": 2.0,
+ "regressor_tree_std_reference": 2.5,
+ "classifier_tree_std_reference": 0.5,
+ "classifier_weight": 0.55,
+ "blend_threshold": round(best_threshold, 4),
+ "risk_scale": round(risk_scale, 4),
+ "base_uncertainty": 0.11,
+ }
+
+ # Feature importances
+ combined_imp = (
+ np.array(reg.feature_importances_) * 0.45 +
+ np.array(clf.feature_importances_) * 0.55
+ )
+ top_features = sorted(
+ zip(ARCHIVE_FEATURE_NAMES, combined_imp.tolist()),
+ key=lambda x: x[1], reverse=True
+ )[:8]
+
+ artifact = {
+ "version": "archive-fusion-v3",
+ "feature_names": ARCHIVE_FEATURE_NAMES,
+ "regressor": reg,
+ "classifier": clf,
+ "inference_source_hint": "roi_original",
+ "calibration": calibration,
+ "training": {
+ "selected_mode": "hybrid_dual",
+ "subject_count": len(subjects),
+ "record_count": len(samples),
+ "metrics": metrics,
+ "top_features": [{"name": n, "importance": round(float(v), 4)} for n, v in top_features],
+ },
+ }
+
+ joblib.dump(artifact, OUTPUT_PATH)
+ print(f"Saved model to {OUTPUT_PATH}", flush=True)
+
+ report = {
+ "dataset_name": "dataset anemia",
+ "record_count": len(samples),
+ "subject_count": len(subjects),
+ "primary_model": "archive-fusion-v3",
+ "selected_mode": "hybrid_dual",
+ "metrics": metrics,
+ "calibration": {
+ "blend_threshold": calibration["blend_threshold"],
+ "risk_scale": calibration["risk_scale"],
+ "classifier_weight": calibration["classifier_weight"],
+ },
+ "top_features": [{"name": n, "importance": round(float(v), 4)} for n, v in top_features],
+ }
+ with open(REPORT_PATH, "w") as f:
+ json.dump(report, f, indent=2)
+ print(f"Saved report to {REPORT_PATH}", flush=True)
+
+ # Quick sanity check
+ print("\nSanity check:", flush=True)
+ feat_idx = {n: i for i, n in enumerate(ARCHIVE_FEATURE_NAMES)}
+ for label, cpi, rg, br in [("PALE (anemic)", 0.28, 0.02, 0.22), ("NORMAL", 0.44, 0.08, 0.38)]:
+ row = np.zeros((1, len(ARCHIVE_FEATURE_NAMES)), dtype=np.float32)
+ row[0, feat_idx["cpi"]] = cpi
+ row[0, feat_idx["center_cpi"]] = cpi - 0.01
+ row[0, feat_idx["mean_r"]] = cpi * 0.9
+ row[0, feat_idx["red_green_gap"]] = rg
+ row[0, feat_idx["center_red_green_gap"]] = rg
+ row[0, feat_idx["brightness"]] = br
+ row[0, feat_idx["green_blue_ratio"]] = 1.1 if cpi < 0.35 else 1.25
+ row[0, feat_idx["source_roi_original"]] = 1.0
+ hb_p = float(reg.predict(row)[0])
+ cp = float(clf.predict_proba(row)[0, 1])
+ rr = sigmoid((ANEMIA_HB_THRESHOLD - hb_p) / hb_scale)
+ bs = 0.55 * cp + 0.45 * rr
+ risk = sigmoid((bs - best_threshold) / risk_scale)
+ print(f" {label}: Hb={hb_p:.1f}, clf_prob={cp:.3f}, risk={risk:.3f}", flush=True)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/scripts/retrain_pipeline_aligned.py b/backend/scripts/retrain_pipeline_aligned.py
new file mode 100644
index 0000000000000000000000000000000000000000..0f6a5183dd04d8994024b00b07154df1e9293d5c
--- /dev/null
+++ b/backend/scripts/retrain_pipeline_aligned.py
@@ -0,0 +1,302 @@
+"""
+Retrain the archive model using the EXACT same pipeline as inference:
+ raw JPG -> quality gate -> ROI extraction -> feature extraction
+
+This ensures train/inference feature distributions match.
+Previous models trained on pre-cropped palpebral PNGs but inference
+runs on raw JPGs through the ROI extractor โ causing a massive domain gap.
+"""
+from __future__ import annotations
+import sys, json
+from pathlib import Path
+sys.path.insert(0, str(Path(__file__).parents[1]))
+
+import numpy as np
+import joblib
+from sklearn.ensemble import ExtraTreesClassifier, ExtraTreesRegressor
+from sklearn.metrics import (
+ accuracy_score, f1_score, mean_absolute_error,
+ precision_score, recall_score, roc_auc_score,
+)
+from sklearn.model_selection import GroupShuffleSplit
+
+from app.ml.archive_model import (
+ ANEMIA_HB_THRESHOLD, ARCHIVE_FEATURE_NAMES,
+ clamp, sigmoid, _parse_workbook, _parse_float, _load_image_with_fallback,
+)
+from app.ml.features import extract_eye_features, FEATURE_NAMES
+from app.services.conjunctiva_roi import ConjunctivaRoiExtractor
+from app.services.image_quality import ImageQualityService
+
+DATASET_ROOT = Path(__file__).parents[2] / "archive" / "dataset anemia"
+OUTPUT_PATH = Path(__file__).parents[1] / "models" / "archive_screening_model.joblib"
+REPORT_PATH = Path(__file__).parents[1] / "models" / "training_report.json"
+
+
+def build_pipeline_aligned_dataset():
+ """
+ Load raw JPGs, run through ROI extractor (same as inference),
+ extract features. Returns samples with ground-truth Hb.
+ """
+ roi_extractor = ConjunctivaRoiExtractor()
+ samples = []
+ skipped = 0
+
+ for country in ("India", "Italy"):
+ workbook_path = DATASET_ROOT / country / f"{country}.xlsx"
+ metadata = _parse_workbook(workbook_path)
+
+ for subject_number, row in metadata.items():
+ hb = _parse_float(row.get("Hgb"))
+ if hb is None:
+ continue
+
+ subject_dir = DATASET_ROOT / country / subject_number
+ if not subject_dir.exists():
+ continue
+
+ # Use raw JPG โ same as what users upload
+ jpgs = sorted(subject_dir.glob("*.jpg"))
+ if not jpgs:
+ skipped += 1
+ continue
+
+ try:
+ raw_img = _load_image_with_fallback(jpgs[0])
+ roi_result = roi_extractor.extract(raw_img)
+ roi_img = roi_result.image
+ features = extract_eye_features(roi_img)
+
+ # Add source flags (roi_original path)
+ prepared = dict(features)
+ prepared["source_roi_original"] = 1.0
+ prepared["source_segmented"] = 0.0
+ prepared["source_forniceal_palpebral"] = 0.0
+
+ samples.append({
+ "group": f"{country}-{subject_number}",
+ "hb": hb,
+ "label": int(hb < ANEMIA_HB_THRESHOLD),
+ "features": prepared,
+ })
+ except Exception as e:
+ skipped += 1
+
+ print(f" Loaded {len(samples)} samples, skipped {skipped}")
+ return samples
+
+
+def find_best_threshold(labels, scores):
+ best_score, best_thresh = -1.0, 0.5
+ for t in np.linspace(0.20, 0.80, 61):
+ preds = (scores >= t).astype(int)
+ if preds.sum() == 0:
+ continue
+ f1 = f1_score(labels, preds, zero_division=0)
+ rec = recall_score(labels, preds, zero_division=0)
+ # Weight recall heavily โ medical screening, false negatives are worse
+ score = f1 * 0.4 + rec * 0.6
+ if score > best_score:
+ best_score = score
+ best_thresh = float(t)
+ return best_thresh
+
+
+def main():
+ print("=" * 60)
+ print("AnemiaLens โ pipeline-aligned retraining")
+ print("=" * 60)
+
+ print("\n[1/4] Building pipeline-aligned dataset...")
+ samples = build_pipeline_aligned_dataset()
+
+ feat_names = ARCHIVE_FEATURE_NAMES # 44 features
+ rows = np.array([[float(s["features"].get(n, 0.0)) for n in feat_names] for s in samples], dtype=np.float32)
+ targets = np.array([s["hb"] for s in samples], dtype=np.float32)
+ labels = np.array([s["label"] for s in samples], dtype=np.int32)
+ groups = np.array([s["group"] for s in samples], dtype=object)
+
+ print(f" Samples: {len(samples)}, Anemic: {labels.sum()}, Normal: {(labels==0).sum()}")
+
+ print("\n[2/4] Cross-validating...")
+ splitter = GroupShuffleSplit(n_splits=5, test_size=0.2, random_state=42)
+ all_metrics = []
+ all_thresholds = []
+
+ for i, (train_idx, test_idx) in enumerate(splitter.split(rows, labels, groups)):
+ print(f" Split {i+1}/5...", flush=True)
+
+ # Augment training: add noise + oversample anemic
+ rng = np.random.default_rng(42 + i)
+ tr_rows, tr_targets, tr_labels = rows[train_idx], targets[train_idx], labels[train_idx]
+
+ # Gaussian noise on all samples
+ noisy = tr_rows.copy()
+ noisy += rng.normal(0, 0.008, size=noisy.shape)
+ noisy = np.clip(noisy, 0.0, 1.0)
+
+ # 3x oversample anemic
+ anemic_idx = np.where(tr_labels == 1)[0]
+ copies_list = [tr_rows, noisy]
+ t_list = [tr_targets, tr_targets]
+ l_list = [tr_labels, tr_labels]
+ for _ in range(3):
+ copies = tr_rows[anemic_idx].copy()
+ copies += rng.normal(0, 0.01, size=copies.shape)
+ copies = np.clip(copies, 0.0, 1.0)
+ copies_list.append(copies)
+ t_list.append(tr_targets[anemic_idx])
+ l_list.append(tr_labels[anemic_idx])
+
+ aug_rows = np.vstack(copies_list)
+ aug_targets = np.concatenate(t_list)
+ aug_labels = np.concatenate(l_list)
+
+ reg = ExtraTreesRegressor(n_estimators=300, min_samples_leaf=2, max_features=0.7,
+ bootstrap=True, random_state=42+i, n_jobs=1)
+ clf = ExtraTreesClassifier(n_estimators=300, min_samples_leaf=2, max_features=0.7,
+ bootstrap=True, class_weight="balanced_subsample",
+ random_state=42+i, n_jobs=1)
+ reg.fit(aug_rows, aug_targets)
+ clf.fit(aug_rows, aug_labels)
+
+ hb_pred = reg.predict(rows[test_idx])
+ clf_prob = clf.predict_proba(rows[test_idx])[:, 1]
+
+ hb_scale = max(float(np.quantile(np.abs(aug_targets - reg.predict(aug_rows)), 0.75)), 0.8)
+ reg_risk = np.array([sigmoid((ANEMIA_HB_THRESHOLD - h) / hb_scale) for h in hb_pred])
+ blend = 0.55 * clf_prob + 0.45 * reg_risk
+
+ thresh = find_best_threshold(labels[test_idx], blend)
+ preds = (blend >= thresh).astype(int)
+
+ all_metrics.append({
+ "accuracy": accuracy_score(labels[test_idx], preds),
+ "precision": precision_score(labels[test_idx], preds, zero_division=0),
+ "recall": recall_score(labels[test_idx], preds, zero_division=0),
+ "f1": f1_score(labels[test_idx], preds, zero_division=0),
+ "auc": roc_auc_score(labels[test_idx], blend),
+ "mae_hb": mean_absolute_error(targets[test_idx], hb_pred),
+ })
+ all_thresholds.append(thresh)
+
+ avg = {k: round(float(np.mean([m[k] for m in all_metrics])), 4) for k in all_metrics[0]}
+ best_threshold = float(np.mean(all_thresholds))
+ print(f"\n CV metrics: {avg}")
+ print(f" Best threshold: {best_threshold:.3f}")
+
+ print("\n[3/4] Training final model on full dataset...")
+ rng = np.random.default_rng(42)
+ noisy = rows.copy()
+ noisy += rng.normal(0, 0.008, size=noisy.shape)
+ noisy = np.clip(noisy, 0.0, 1.0)
+ anemic_idx = np.where(labels == 1)[0]
+ copies_list = [rows, noisy]
+ t_list = [targets, targets]
+ l_list = [labels, labels]
+ for _ in range(3):
+ copies = rows[anemic_idx].copy()
+ copies += rng.normal(0, 0.01, size=copies.shape)
+ copies = np.clip(copies, 0.0, 1.0)
+ copies_list.append(copies)
+ t_list.append(targets[anemic_idx])
+ l_list.append(labels[anemic_idx])
+ aug_rows = np.vstack(copies_list)
+ aug_targets = np.concatenate(t_list)
+ aug_labels = np.concatenate(l_list)
+
+ reg = ExtraTreesRegressor(n_estimators=300, min_samples_leaf=2, max_features=0.7,
+ bootstrap=True, random_state=42, n_jobs=1)
+ clf = ExtraTreesClassifier(n_estimators=300, min_samples_leaf=2, max_features=0.7,
+ bootstrap=True, class_weight="balanced_subsample",
+ random_state=42, n_jobs=1)
+ reg.fit(aug_rows, aug_targets)
+ clf.fit(aug_rows, aug_labels)
+
+ hb_preds_full = reg.predict(rows)
+ residuals = np.abs(targets - hb_preds_full)
+ hb_scale = max(float(np.quantile(residuals, 0.75)), 0.8)
+ clf_probs_full = clf.predict_proba(rows)[:, 1]
+ reg_risk_full = np.array([sigmoid((ANEMIA_HB_THRESHOLD - h) / hb_scale) for h in hb_preds_full])
+ blend_full = 0.55 * clf_probs_full + 0.45 * reg_risk_full
+ risk_scale = max(float(np.std(blend_full)) * 0.9, 0.08)
+ risk_scale = min(risk_scale, 0.22)
+
+ calibration = {
+ "hb_threshold": ANEMIA_HB_THRESHOLD,
+ "hb_scale": round(hb_scale, 4),
+ "hb_population_mean": round(float(np.mean(targets)), 4),
+ "hb_spread_factor": 2.0,
+ "regressor_tree_std_reference": 1.85,
+ "classifier_tree_std_reference": 0.40,
+ "classifier_weight": 0.55,
+ "blend_threshold": round(best_threshold, 4),
+ "risk_scale": round(risk_scale, 4),
+ "base_uncertainty": 0.08,
+ }
+
+ artifact = {
+ "version": "archive-fusion-v4-pipeline",
+ "feature_names": feat_names,
+ "regressor": reg,
+ "classifier": clf,
+ "calibration": calibration,
+ "training": {
+ "selected_mode": "pipeline_aligned_roi",
+ "subject_count": len(samples),
+ "record_count": len(samples),
+ "metrics": avg,
+ },
+ }
+
+ print("\n[4/4] Saving...")
+ joblib.dump(artifact, OUTPUT_PATH)
+ print(f" Saved -> {OUTPUT_PATH}")
+ print(f" Size: {OUTPUT_PATH.stat().st_size / 1024 / 1024:.1f} MB")
+
+ report = {
+ "dataset_name": "dataset anemia (pipeline-aligned)",
+ "record_count": len(samples),
+ "subject_count": len(samples),
+ "primary_model": "archive-fusion-v4-pipeline",
+ "selected_mode": "pipeline_aligned_roi",
+ "metrics": avg,
+ "calibration": {
+ "blend_threshold": calibration["blend_threshold"],
+ "risk_scale": calibration["risk_scale"],
+ "classifier_weight": calibration["classifier_weight"],
+ },
+ }
+ with open(REPORT_PATH, "w") as f:
+ json.dump(report, f, indent=2)
+
+ # Sanity check on training data
+ print("\nSanity check (training data):")
+ for label, cpi, rg, br in [("PALE", 0.28, 0.02, 0.22), ("NORMAL", 0.44, 0.08, 0.38)]:
+ from app.ml.archive_model import prepare_feature_map
+ from app.ml.features import FEATURE_NAMES as FN
+ feat_map = {n: 0.0 for n in FN}
+ feat_map.update({"cpi": cpi, "center_cpi": cpi-0.01, "mean_r": cpi*0.9,
+ "mean_g": cpi*0.9-rg, "mean_b": cpi*0.7,
+ "red_green_gap": rg, "center_red_green_gap": rg,
+ "brightness": br, "center_brightness": br,
+ "green_blue_ratio": 1.1 if cpi < 0.35 else 1.25,
+ "contrast": 0.12, "center_contrast": 0.12,
+ "blur_score": 100.0, "center_blur_score": 120.0,
+ "saturation": 0.3, "center_saturation": 0.3,
+ "hist_mid": 0.5, "hist_bright": 0.3,
+ "aspect_ratio": 1.0, "size_score": 1.0})
+ prepared = prepare_feature_map(feat_map, source_hint="roi_original")
+ row = np.array([[prepared.get(n, 0.0) for n in feat_names]], dtype=np.float32)
+ hb_p = float(reg.predict(row)[0])
+ cp = float(clf.predict_proba(row)[0, 1])
+ rr = sigmoid((ANEMIA_HB_THRESHOLD - hb_p) / hb_scale)
+ bs = 0.55 * cp + 0.45 * rr
+ risk = sigmoid((bs - best_threshold) / risk_scale)
+ print(f" {label}: Hb={hb_p:.1f}, risk={risk:.3f}")
+
+ print("\nDone.")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/scripts/test_endpoint.py b/backend/scripts/test_endpoint.py
new file mode 100644
index 0000000000000000000000000000000000000000..a7d8ee69c448fa92a734e37f19b8db5ae13913a8
--- /dev/null
+++ b/backend/scripts/test_endpoint.py
@@ -0,0 +1,36 @@
+import requests, io, json
+from PIL import Image
+
+img = Image.new('RGB', (400, 300), color=(200, 160, 140))
+buf = io.BytesIO()
+img.save(buf, format='JPEG')
+buf.seek(0)
+
+symptoms = json.dumps({
+ "fatigue": True,
+ "pale_skin": True,
+ "dizziness": False,
+ "shortness_of_breath": False,
+ "heavy_menstrual_bleeding": None,
+ "poor_diet_low_iron": False
+})
+
+r = requests.post(
+ 'http://localhost:8000/api/analyze',
+ files={'image': ('test.jpg', buf, 'image/jpeg')},
+ data={'symptoms': symptoms},
+ timeout=30
+)
+print('Status:', r.status_code)
+if r.status_code == 200:
+ d = r.json()
+ pred = d.get('prediction') or {}
+ triage = d.get('triage') or {}
+ print('Hb:', pred.get('predicted_hemoglobin'))
+ print('Risk:', pred.get('anemia_risk'))
+ print('Label:', pred.get('screening_label'))
+ print('Model:', pred.get('model_source'))
+ print('Triage band:', triage.get('band'))
+ print('Blocked:', d.get('blocked'))
+else:
+ print('Error body:', r.text[:1000])
diff --git a/backend/scripts/test_model.py b/backend/scripts/test_model.py
new file mode 100644
index 0000000000000000000000000000000000000000..5b119f86279f1a5f6b05eeaf2fe3a6889ac6ba36
--- /dev/null
+++ b/backend/scripts/test_model.py
@@ -0,0 +1,45 @@
+import joblib, sys
+sys.path.insert(0, 'backend')
+from app.ml.archive_model import predict_with_archive_model
+from app.ml.features import FEATURE_NAMES
+
+m = joblib.load('backend/models/archive_screening_model.joblib')
+
+test_cases = [
+ ("PALE (anemic)", 0.28, 0.02, 0.22),
+ ("BORDERLINE", 0.35, 0.04, 0.30),
+ ("NORMAL", 0.44, 0.08, 0.38),
+ ("VERY HEALTHY", 0.48, 0.10, 0.42),
+]
+
+for label, cpi, rg, br in test_cases:
+ feat_map = {n: 0.0 for n in FEATURE_NAMES}
+ feat_map['cpi'] = cpi
+ feat_map['center_cpi'] = cpi - 0.01
+ feat_map['mean_r'] = cpi * 0.9
+ feat_map['mean_g'] = cpi * 0.9 - rg
+ feat_map['mean_b'] = cpi * 0.7
+ feat_map['red_green_gap'] = rg
+ feat_map['center_red_green_gap'] = rg
+ feat_map['brightness'] = br
+ feat_map['green_blue_ratio'] = 1.1 if cpi < 0.35 else 1.25
+ feat_map['center_mean_r'] = feat_map['mean_r']
+ feat_map['center_mean_g'] = feat_map['mean_g']
+ feat_map['center_mean_b'] = feat_map['mean_b']
+ feat_map['contrast'] = 0.12
+ feat_map['center_contrast'] = 0.12
+ feat_map['center_brightness'] = br
+ feat_map['blur_score'] = 100.0
+ feat_map['center_blur_score'] = 120.0
+ feat_map['saturation'] = 0.3
+ feat_map['center_saturation'] = 0.3
+ feat_map['hist_mid'] = 0.5
+ feat_map['hist_bright'] = 0.3
+ feat_map['aspect_ratio'] = 1.0
+ feat_map['size_score'] = 1.0
+ result = predict_with_archive_model(m, feat_map, source_hint='roi_original')
+ hb = result['predicted_hemoglobin']
+ risk = result['anemia_risk']
+ unc = result['uncertainty']
+ decision = "ANEMIA LIKELY" if risk >= 0.65 else "unlikely"
+ print(f"{label}: Hb={hb:.1f}, risk={risk:.3f}, uncertainty={unc:.3f} -> {decision}")
diff --git a/backend/scripts/train_archive_model.py b/backend/scripts/train_archive_model.py
new file mode 100644
index 0000000000000000000000000000000000000000..c596dcb99becea2f028192ca67db8bbe5e0db479
--- /dev/null
+++ b/backend/scripts/train_archive_model.py
@@ -0,0 +1,119 @@
+"""
+Train the archive conjunctiva screening model.
+
+Usage::
+
+ python scripts/train_archive_model.py [--dataset PATH] [--output-dir PATH] [--quiet]
+
+The script trains the model, writes the artefact and a human-readable
+training report, then exits with code 0 on success or 1 on failure.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+import sys
+import time
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+BACKEND_ROOT = ROOT / "backend"
+sys.path.insert(0, str(BACKEND_ROOT))
+
+from app.config import DEFAULT_ARCHIVE_MODEL_PATH, DEFAULT_TRAINING_REPORT_PATH # noqa: E402
+from app.ml.archive_model import save_archive_model, train_archive_model # noqa: E402
+
+DEFAULT_DATASET = ROOT / "archive" / "dataset anemia"
+
+
+# ---------------------------------------------------------------------------
+# CLI
+# ---------------------------------------------------------------------------
+
+def _build_parser() -> argparse.ArgumentParser:
+ p = argparse.ArgumentParser(
+ description="Train the AnemiaLens archive conjunctiva screening model.",
+ formatter_class=argparse.ArgumentDefaultsHelpFormatter,
+ )
+ p.add_argument(
+ "--dataset",
+ type=Path,
+ default=DEFAULT_DATASET,
+ help="Root directory of the labelled anemia dataset.",
+ )
+ p.add_argument(
+ "--output-dir",
+ type=Path,
+ default=DEFAULT_ARCHIVE_MODEL_PATH.parent,
+ help="Directory where the model artefact and report are written.",
+ )
+ p.add_argument(
+ "--quiet",
+ action="store_true",
+ help="Suppress progress output (report still written to disk).",
+ )
+ return p
+
+
+# ---------------------------------------------------------------------------
+# Main
+# ---------------------------------------------------------------------------
+
+def main(argv: list[str] | None = None) -> int:
+ args = _build_parser().parse_args(argv)
+
+ if not args.dataset.exists():
+ print(
+ f"ERROR: Dataset directory not found: {args.dataset}\n"
+ " Download the anemia dataset and place it there, or pass --dataset PATH.",
+ file=sys.stderr,
+ )
+ return 1
+
+ if not args.quiet:
+ print(f"Dataset : {args.dataset}")
+ print(f"Output : {args.output_dir}")
+ print()
+
+ t0 = time.perf_counter()
+
+ try:
+ artifact, report = train_archive_model(args.dataset)
+ except Exception as exc:
+ print(f"ERROR: Training failed โ {exc}", file=sys.stderr)
+ return 1
+
+ elapsed = time.perf_counter() - t0
+
+ # --- Write artefacts ---------------------------------------------------
+ args.output_dir.mkdir(parents=True, exist_ok=True)
+
+ model_path = args.output_dir / DEFAULT_ARCHIVE_MODEL_PATH.name
+ report_path = args.output_dir / DEFAULT_TRAINING_REPORT_PATH.name
+
+ save_archive_model(artifact, model_path)
+ report_path.write_text(json.dumps(report, indent=2, ensure_ascii=False), encoding="utf-8")
+
+ # --- Summary -----------------------------------------------------------
+ if not args.quiet:
+ metrics = report.get("metrics", {})
+ print(json.dumps(report, indent=2))
+ print()
+ print("=" * 56)
+ print(f" Model : {report.get('primary_model', '?')}")
+ print(f" Subjects : {report.get('subject_count', '?')}")
+ print(f" Records : {report.get('record_count', '?')}")
+ print(f" Accuracy : {metrics.get('accuracy', 0):.3f}")
+ print(f" F1 : {metrics.get('f1', 0):.3f}")
+ print(f" Val size : {metrics.get('validation_size', '?')}")
+ print(f" Elapsed : {elapsed:.1f}s")
+ print("=" * 56)
+ print(f" Saved model โ {model_path}")
+ print(f" Saved report โ {report_path}")
+
+ return 0
+
+
+if __name__ == "__main__":
+ sys.exit(main())
diff --git a/backend/scripts/train_efficientnet.py b/backend/scripts/train_efficientnet.py
new file mode 100644
index 0000000000000000000000000000000000000000..4f397d8f64438adaee52276ae6ca7abca8009a80
--- /dev/null
+++ b/backend/scripts/train_efficientnet.py
@@ -0,0 +1,465 @@
+from __future__ import annotations
+
+import json
+import math
+import random
+from collections import Counter
+from copy import deepcopy
+from dataclasses import dataclass
+from datetime import datetime, timezone
+from pathlib import Path
+
+import numpy as np
+import torch
+from PIL import Image
+from sklearn.metrics import accuracy_score, f1_score, mean_absolute_error, precision_score, recall_score, roc_auc_score
+from sklearn.model_selection import GroupShuffleSplit
+from torch import nn
+from torch.optim import AdamW
+from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler
+
+from app.config import (
+ BACKEND_ROOT,
+ DEFAULT_EFFICIENTNET_MODEL_PATH,
+ DEFAULT_EFFICIENTNET_REPORT_PATH,
+ DEFAULT_TRAINING_REPORT_PATH,
+)
+from app.ml.archive_model import ANEMIA_HB_THRESHOLD, _first_path, _load_image_with_fallback, _parse_float, _parse_workbook
+from app.ml.efficientnet_model import (
+ EFFICIENTNET_VERSION,
+ build_efficientnet_model,
+ build_train_transform,
+ build_val_transform,
+)
+from app.services.conjunctiva_roi import ConjunctivaRoiExtractor
+
+
+DATA_ROOT = BACKEND_ROOT / "data"
+ARCHIVE_ROOT = BACKEND_ROOT.parent / "archive" / "dataset anemia"
+SEED = 42
+BATCH_SIZE = 16
+EPOCHS = 60
+PATIENCE = 15
+MAX_GRAD_NORM = 1.0
+WARMUP_EPOCHS = 5
+LABEL_SMOOTHING = 0.05
+MIXUP_ALPHA = 0.3
+# Hb spread loss: penalizes predictions that cluster near the mean
+HB_SPREAD_WEIGHT = 0.15
+
+
+@dataclass(frozen=True)
+class ImageRecord:
+ subject_id: str
+ label: int
+ hb: float
+ image_path: Path
+ source: str
+
+
+class ConjunctivaDataset(Dataset):
+ def __init__(self, records: list[ImageRecord], transform: object) -> None:
+ self.records = records
+ self.transform = transform
+ self.roi_extractor = ConjunctivaRoiExtractor()
+
+ def __len__(self) -> int:
+ return len(self.records)
+
+ def __getitem__(self, index: int) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
+ record = self.records[index]
+ image, label, hb = self._prepare_item(record)
+ tensor = self.transform(image)
+ return tensor, torch.tensor([label], dtype=torch.float32), torch.tensor([hb], dtype=torch.float32)
+
+ def _prepare_item(self, record: ImageRecord) -> tuple[Image.Image, float, float]:
+ image = _load_image_with_fallback(record.image_path)
+ if record.source == "roi_original":
+ image = self.roi_extractor.extract(image).image
+ return image.convert("RGB"), float(record.label), float(record.hb)
+
+
+class FocalLoss(nn.Module):
+ def __init__(self, alpha: float = 0.25, gamma: float = 2.0, pos_weight: torch.Tensor | None = None) -> None:
+ super().__init__()
+ self.alpha = alpha
+ self.gamma = gamma
+ self.bce = nn.BCEWithLogitsLoss(pos_weight=pos_weight, reduction="none")
+
+ def forward(self, inputs: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:
+ bce_loss = self.bce(inputs, targets)
+ probabilities = torch.sigmoid(inputs)
+ p_t = probabilities * targets + (1 - probabilities) * (1 - targets)
+ loss = bce_loss * ((1 - p_t) ** self.gamma)
+ if self.alpha >= 0:
+ alpha_t = self.alpha * targets + (1 - self.alpha) * (1 - targets)
+ loss = alpha_t * loss
+ return loss.mean()
+
+
+def main() -> None:
+ _set_seed(SEED)
+ dataset_root = DATA_ROOT if DATA_ROOT.exists() else ARCHIVE_ROOT
+ if not dataset_root.exists():
+ raise RuntimeError(f"No dataset directory found at {DATA_ROOT} or {ARCHIVE_ROOT}.")
+
+ records = _build_records(dataset_root)
+ if not records:
+ raise RuntimeError(f"No training records found in {dataset_root}.")
+
+ train_records, val_records = _balanced_group_split(records, test_size=0.2, n_splits=32)
+ train_dataset = ConjunctivaDataset(train_records, build_train_transform())
+ val_dataset = ConjunctivaDataset(val_records, build_val_transform())
+ hb_mean, hb_std = _hb_normalization_stats(train_records)
+ train_sampler = _build_weighted_sampler(train_records)
+
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
+ model = build_efficientnet_model(pretrained=True).to(device)
+ optimizer = AdamW(
+ [
+ {"params": list(model.classifier.parameters()), "lr": 1.5e-4}, # Slightly lower for stability
+ {"params": [param for param in model.features.parameters() if param.requires_grad], "lr": 5e-6},
+ ],
+ weight_decay=4e-4, # Higher weight decay for better regularization
+ )
+
+ # Warmup then cosine annealing
+ def warmup_cosine_lr(epoch: int) -> float:
+ if epoch < WARMUP_EPOCHS:
+ return float(epoch + 1) / WARMUP_EPOCHS
+ progress = (epoch - WARMUP_EPOCHS) / max(EPOCHS - WARMUP_EPOCHS, 1)
+ return 0.5 * (1.0 + math.cos(math.pi * progress))
+
+ scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda=warmup_cosine_lr)
+
+ # Use Focal Loss with positive weights
+ pos_weight = torch.tensor([_positive_class_weight(train_records)], device=device)
+ cls_loss_fn = FocalLoss(alpha=0.25, gamma=2.0, pos_weight=pos_weight)
+ hb_loss_fn = nn.SmoothL1Loss(beta=0.5)
+
+ train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, sampler=train_sampler, num_workers=0)
+ val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=0)
+
+ best_state: dict[str, torch.Tensor] | None = None
+ best_metrics: dict[str, float] | None = None
+ best_threshold = 0.5
+ best_score = -1.0
+ epochs_without_improvement = 0
+ history: list[dict[str, float]] = []
+
+ for epoch in range(1, EPOCHS + 1):
+ model.train()
+ train_loss_total = 0.0
+
+ for images, labels, hbs in train_loader:
+ images = images.to(device)
+ labels = labels.to(device)
+ hbs = hbs.to(device)
+ normalized_hbs = (hbs - hb_mean) / hb_std
+
+ # MixUp augmentation
+ if MIXUP_ALPHA > 0 and np.random.random() < 0.5:
+ lam = float(np.random.beta(MIXUP_ALPHA, MIXUP_ALPHA))
+ idx = torch.randperm(images.size(0), device=device)
+ images = lam * images + (1.0 - lam) * images[idx]
+ labels_a, labels_b = labels, labels[idx]
+ hbs_a, hbs_b = normalized_hbs, normalized_hbs[idx]
+
+ optimizer.zero_grad(set_to_none=True)
+ output = model(images)
+ # Label smoothing applied to both MixUp targets
+ smooth_a = labels_a * (1.0 - LABEL_SMOOTHING) + 0.5 * LABEL_SMOOTHING
+ smooth_b = labels_b * (1.0 - LABEL_SMOOTHING) + 0.5 * LABEL_SMOOTHING
+ cls_loss = lam * cls_loss_fn(output[:, 0:1], smooth_a) + (1.0 - lam) * cls_loss_fn(output[:, 0:1], smooth_b)
+ hb_loss = lam * hb_loss_fn(output[:, 1:2], hbs_a) + (1.0 - lam) * hb_loss_fn(output[:, 1:2], hbs_b)
+ else:
+ optimizer.zero_grad(set_to_none=True)
+ output = model(images)
+ smooth_labels = labels * (1.0 - LABEL_SMOOTHING) + 0.5 * LABEL_SMOOTHING
+ cls_loss = cls_loss_fn(output[:, 0:1], smooth_labels)
+ hb_loss = hb_loss_fn(output[:, 1:2], normalized_hbs)
+
+ # Hb spread loss: penalize predictions clustering near zero (normalized mean)
+ # Encourages the model to predict a wider range of Hb values
+ hb_pred_norm = output[:, 1:2]
+ spread_loss = torch.clamp(0.5 - hb_pred_norm.std(), min=0.0)
+
+ total_loss = (0.60 * cls_loss) + (0.30 * hb_loss) + (HB_SPREAD_WEIGHT * spread_loss)
+ total_loss.backward()
+ nn.utils.clip_grad_norm_(model.parameters(), MAX_GRAD_NORM)
+ optimizer.step()
+ train_loss_total += float(total_loss.item()) * images.size(0)
+
+ scheduler.step()
+ val_metrics = _evaluate_model(model, val_loader, device, hb_mean=hb_mean, hb_std=hb_std)
+ history.append(
+ {
+ "epoch": float(epoch),
+ "train_loss": round(train_loss_total / max(len(train_dataset), 1), 4),
+ "val_f1": val_metrics["f1"],
+ "val_auc": val_metrics["auc"],
+ "val_hb_mae": val_metrics["hb_mae"],
+ }
+ )
+ print(
+ f"epoch={epoch:02d} train_loss={history[-1]['train_loss']:.4f} "
+ f"val_f1={val_metrics['f1']:.4f} val_auc={val_metrics['auc']:.4f} "
+ f"val_hb_mae={val_metrics['hb_mae']:.4f}"
+ )
+
+ # Use composite score: AUC weighted more heavily than F1 (more stable early on)
+ composite_score = val_metrics["auc"] * 0.55 + val_metrics["f1"] * 0.35 + (1.0 - min(val_metrics["hb_mae"] / 4.0, 1.0)) * 0.10
+ if composite_score > best_score:
+ best_score = composite_score
+ best_state = deepcopy(model.state_dict())
+ best_metrics = val_metrics
+ best_threshold = val_metrics["decision_threshold"]
+ epochs_without_improvement = 0
+ else:
+ epochs_without_improvement += 1
+
+ if epochs_without_improvement >= PATIENCE:
+ print(f"Early stopping after {epoch} epochs.")
+ break
+
+ if best_state is None or best_metrics is None:
+ raise RuntimeError("EfficientNet training did not produce a valid checkpoint.")
+
+ checkpoint = {
+ "version": EFFICIENTNET_VERSION,
+ "created_at": datetime.now(timezone.utc).isoformat(),
+ "state_dict": best_state,
+ "decision_threshold": best_threshold,
+ "hb_mean": hb_mean,
+ "hb_std": hb_std,
+ "hb_spread_factor": _compute_hb_spread_factor(val_records, hb_mean, hb_std),
+ "val_metrics": best_metrics,
+ }
+ DEFAULT_EFFICIENTNET_MODEL_PATH.parent.mkdir(parents=True, exist_ok=True)
+ torch.save(checkpoint, DEFAULT_EFFICIENTNET_MODEL_PATH)
+
+ report = {
+ "dataset_name": str(dataset_root.name),
+ "record_count": len(records),
+ "subject_count": len({record.subject_id for record in records}),
+ "primary_model": EFFICIENTNET_VERSION,
+ "selected_mode": "efficientnet_hybrid_dual",
+ "source_counts": _source_counts(records),
+ "metrics": {
+ "accuracy": round(best_metrics["accuracy"], 4),
+ "precision": round(best_metrics["precision"], 4),
+ "recall": round(best_metrics["recall"], 4),
+ "f1": round(best_metrics["f1"], 4),
+ "auc": round(best_metrics["auc"], 4),
+ "mae_hb": round(best_metrics["hb_mae"], 4),
+ "validation_size": len(val_records),
+ "split_strategy": "group-shuffle-balance-select",
+ "sample_count": len(records),
+ "subject_count": len({record.subject_id for record in records}),
+ "decision_threshold": round(best_threshold, 4),
+ },
+ "training": {
+ "epochs_requested": EPOCHS,
+ "history": history,
+ "batch_size": BATCH_SIZE,
+ "patience": PATIENCE,
+ "device": str(device),
+ "hb_target_mean": round(hb_mean, 4),
+ "hb_target_std": round(hb_std, 4),
+ "class_positive_weight": round(_positive_class_weight(train_records), 4),
+ "sampler": "weighted-random-balanced",
+ },
+ }
+ DEFAULT_EFFICIENTNET_REPORT_PATH.write_text(json.dumps(report, indent=2), encoding="utf-8")
+ DEFAULT_TRAINING_REPORT_PATH.write_text(json.dumps(report, indent=2), encoding="utf-8")
+
+ print("\nBest validation metrics")
+ for key in ("accuracy", "precision", "recall", "f1", "auc", "hb_mae", "decision_threshold"):
+ print(f"{key}: {best_metrics[key]:.4f}")
+
+
+def _build_records(dataset_root: Path) -> list[ImageRecord]:
+ records: list[ImageRecord] = []
+ for country in ("India", "Italy"):
+ workbook_path = dataset_root / country / f"{country}.xlsx"
+ if not workbook_path.exists():
+ continue
+ metadata = _parse_workbook(workbook_path)
+ for subject_number, row in metadata.items():
+ hb = _parse_float(row.get("Hgb"))
+ if hb is None:
+ continue
+
+ subject_dir = dataset_root / country / subject_number
+ if not subject_dir.exists():
+ continue
+
+ subject_id = f"{country}-{subject_number}"
+ label = int(hb < ANEMIA_HB_THRESHOLD)
+ original_path = _first_path(subject_dir.glob("*.jpg"))
+ palpebral_path = _first_path(
+ path
+ for path in subject_dir.glob("*_palpebral.png")
+ if "forniceal_palpebral" not in path.name.lower()
+ )
+
+ if original_path is not None:
+ records.append(
+ ImageRecord(
+ subject_id=subject_id,
+ label=label,
+ hb=float(hb),
+ image_path=original_path,
+ source="roi_original",
+ )
+ )
+ if palpebral_path is not None:
+ records.append(
+ ImageRecord(
+ subject_id=subject_id,
+ label=label,
+ hb=float(hb),
+ image_path=palpebral_path,
+ source="palpebral",
+ )
+ )
+ return records
+
+
+def _balanced_group_split(
+ records: list[ImageRecord],
+ *,
+ test_size: float,
+ n_splits: int,
+) -> tuple[list[ImageRecord], list[ImageRecord]]:
+ labels = np.asarray([record.label for record in records], dtype=np.int32)
+ groups = np.asarray([record.subject_id for record in records], dtype=object)
+ target_ratio = float(labels.mean())
+ splitter = GroupShuffleSplit(n_splits=n_splits, test_size=test_size, random_state=SEED)
+
+ best: tuple[np.ndarray, np.ndarray] | None = None
+ best_score = float("inf")
+ for train_index, val_index in splitter.split(np.zeros(len(records)), labels, groups):
+ train_labels = labels[train_index]
+ val_labels = labels[val_index]
+ if len(np.unique(train_labels)) < 2 or len(np.unique(val_labels)) < 2:
+ continue
+ score = abs(float(train_labels.mean()) - target_ratio) + abs(float(val_labels.mean()) - target_ratio)
+ if score < best_score:
+ best_score = score
+ best = (train_index, val_index)
+
+ if best is None:
+ raise RuntimeError("Unable to create a grouped train/validation split.")
+
+ train_index, val_index = best
+ return [records[i] for i in train_index], [records[i] for i in val_index]
+
+
+def _evaluate_model(
+ model: nn.Module,
+ loader: DataLoader,
+ device: torch.device,
+ *,
+ hb_mean: float,
+ hb_std: float,
+) -> dict[str, float]:
+ model.eval()
+ probabilities: list[float] = []
+ labels: list[int] = []
+ hb_predictions: list[float] = []
+ hb_targets: list[float] = []
+
+ with torch.no_grad():
+ for images, batch_labels, batch_hbs in loader:
+ images = images.to(device)
+ output = model(images)
+ probabilities.extend(torch.sigmoid(output[:, 0]).cpu().tolist())
+ hb_predictions.extend(((output[:, 1].cpu() * hb_std) + hb_mean).tolist())
+ labels.extend(batch_labels.squeeze(1).cpu().int().tolist())
+ hb_targets.extend(batch_hbs.squeeze(1).cpu().tolist())
+
+ threshold = _best_threshold(np.asarray(labels), np.asarray(probabilities))
+ predicted_labels = [1 if probability >= threshold else 0 for probability in probabilities]
+ auc = roc_auc_score(labels, probabilities) if len(set(labels)) > 1 else 0.5
+
+ return {
+ "accuracy": float(accuracy_score(labels, predicted_labels)),
+ "precision": float(precision_score(labels, predicted_labels, zero_division=0)),
+ "recall": float(recall_score(labels, predicted_labels, zero_division=0)),
+ "f1": float(f1_score(labels, predicted_labels, zero_division=0)),
+ "auc": float(auc),
+ "hb_mae": float(mean_absolute_error(hb_targets, hb_predictions)),
+ "decision_threshold": float(threshold),
+ }
+
+
+def _best_threshold(labels: np.ndarray, probabilities: np.ndarray) -> float:
+ best_threshold = 0.5
+ best_score = -1.0
+ for threshold in np.linspace(0.25, 0.75, 51):
+ predictions = (probabilities >= threshold).astype(np.int32)
+ score = f1_score(labels, predictions, zero_division=0)
+ if score > best_score:
+ best_score = float(score)
+ best_threshold = float(threshold)
+ return best_threshold
+
+
+def _source_counts(records: list[ImageRecord]) -> dict[str, int]:
+ counts: dict[str, int] = {}
+ for record in records:
+ counts[record.source] = counts.get(record.source, 0) + 1
+ return counts
+
+
+def _positive_class_weight(records: list[ImageRecord]) -> float:
+ counts = Counter(record.label for record in records)
+ positive = max(counts.get(1, 0), 1)
+ negative = max(counts.get(0, 0), 1)
+ return float(negative / positive)
+
+
+def _build_weighted_sampler(records: list[ImageRecord]) -> WeightedRandomSampler:
+ counts = Counter(record.label for record in records)
+ total = sum(counts.values())
+ weights = [
+ float(total / max(counts[record.label], 1))
+ for record in records
+ ]
+ return WeightedRandomSampler(
+ torch.as_tensor(weights, dtype=torch.double),
+ num_samples=len(weights),
+ replacement=True,
+ )
+
+
+def _hb_normalization_stats(records: list[ImageRecord]) -> tuple[float, float]:
+ values = np.asarray([record.hb for record in records], dtype=np.float32)
+ mean = float(values.mean())
+ std = float(values.std())
+ return mean, max(std, 1e-3)
+
+
+def _set_seed(seed: int) -> None:
+ random.seed(seed)
+ np.random.seed(seed)
+ torch.manual_seed(seed)
+ if torch.cuda.is_available():
+ torch.cuda.manual_seed_all(seed)
+
+
+def _compute_hb_spread_factor(records: list[ImageRecord], hb_mean: float, hb_std: float) -> float:
+ """
+ Estimate the spread amplification factor needed to correct regression-to-mean.
+ Uses the ratio of true Hb std to the expected model output std (hb_std * 0.75 heuristic).
+ """
+ true_std = float(np.std([r.hb for r in records]))
+ # Models typically predict ~75% of true std due to averaging
+ predicted_std_estimate = max(hb_std * 0.75, 0.5)
+ factor = float(np.clip(true_std / predicted_std_estimate, 1.0, 2.0))
+ return round(factor, 3)
+
+
+if __name__ == "__main__":
+ main()
diff --git a/backend/scripts/train_ensemble.py b/backend/scripts/train_ensemble.py
new file mode 100644
index 0000000000000000000000000000000000000000..21b0a01b73deb2fdce2fca15c6a189c93ba84207
--- /dev/null
+++ b/backend/scripts/train_ensemble.py
@@ -0,0 +1,49 @@
+"""
+Train all models in the AnemiaLens ensemble pipeline.
+
+Currently delegates to train_archive_model. As the ensemble grows
+(deep-stack, legacy CNN, etc.) this script will orchestrate each
+training job in dependency order and produce a combined manifest.
+
+Usage::
+
+ python scripts/train_ensemble.py [--dataset PATH] [--output-dir PATH] [--quiet]
+"""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+# Ensure the scripts directory is on the path so we can import sibling scripts.
+sys.path.insert(0, str(Path(__file__).resolve().parent))
+
+from train_archive_model import main as train_archive
+
+
+def main(argv: list[str] | None = None) -> int:
+ """
+ Orchestrate all training jobs.
+
+ Returns the exit code of the last failing job, or 0 if all succeeded.
+ """
+ exit_code = 0
+
+ print("=== Step 1/1: archive screening model ===")
+ rc = train_archive(argv)
+ if rc != 0:
+ print(f" FAILED (exit {rc})", file=sys.stderr)
+ exit_code = rc
+ else:
+ print(" Done.")
+
+ # Future steps (uncomment when models are ready):
+ # print("=== Step 2/N: deep-stack model ===")
+ # rc = train_deep_stack(argv)
+ # ...
+
+ return exit_code
+
+
+if __name__ == "__main__":
+ sys.exit(main())
diff --git a/backend/scripts/train_stacked.py b/backend/scripts/train_stacked.py
new file mode 100644
index 0000000000000000000000000000000000000000..a1c28d3af71ed5178feb302a4b975d2baefc26b1
--- /dev/null
+++ b/backend/scripts/train_stacked.py
@@ -0,0 +1,617 @@
+๏ปฟ"""
+train_stacked.py รขโฌโ AnemiaLens stacked-ensemble-v4 training script.
+
+Architecture
+------------
+Level-0 base learners (out-of-fold predictions via cross_val_predict):
+ รขโฌยข XGBoost regressor รขโ โ OOF Hb predictions
+ รขโฌยข XGBoost classifier รขโ โ OOF anemia probabilities
+ รขโฌยข ExtraTrees regressor รขโ โ OOF Hb predictions
+ รขโฌยข ExtraTrees classifier รขโ โ OOF anemia probabilities
+
+Level-1 meta-learners:
+ รขโฌยข Ridge regression รขโ โ final Hb estimate
+ รขโฌยข Logistic Regression รขโ โ final anemia risk probability
+
+Data augmentation (training folds only):
+ รขโฌยข Gaussian noise on color features (sigma=0.01)
+ รขโฌยข CPI jitter รยฑ0.02
+ รขโฌยข 3รโ oversampling of anemic class (label=1)
+
+Run from workspace root:
+ python backend/scripts/train_stacked.py
+"""
+from __future__ import annotations
+
+import sys
+import json
+import math
+from pathlib import Path
+
+sys.path.insert(0, str(Path(__file__).parents[1]))
+
+import numpy as np
+import joblib
+from sklearn.ensemble import ExtraTreesClassifier, ExtraTreesRegressor
+from sklearn.linear_model import LogisticRegression, Ridge
+from sklearn.metrics import (
+ accuracy_score, f1_score, mean_absolute_error,
+ precision_score, recall_score, roc_auc_score,
+)
+from sklearn.model_selection import GroupShuffleSplit, RandomizedSearchCV
+from sklearn.model_selection import cross_val_predict
+
+try:
+ from xgboost import XGBClassifier, XGBRegressor
+ _HAS_XGB = True
+except ImportError:
+ _HAS_XGB = False
+ print("WARNING: xgboost not installed รขโฌโ falling back to ExtraTrees-only stack.")
+ print(" Install with: pip install xgboost")
+
+from app.ml.archive_model import (
+ ANEMIA_HB_THRESHOLD,
+ _build_subject_catalog,
+ _samples_for_mode,
+ _rows_from_samples,
+ clamp,
+ sigmoid,
+)
+from app.ml.features import FEATURE_NAMES, COLOR_FEATURES
+from app.ml.stacked_model import StackedRegressor, StackedClassifier
+
+DATASET_ROOT = Path(__file__).parents[2] / "archive" / "dataset anemia"
+OUTPUT_PATH = Path(__file__).parents[1] / "models" / "archive_screening_model.joblib"
+OUTPUT_PATH_V4 = Path(__file__).parents[1] / "models" / "archive_screening_model_v4.joblib"
+REPORT_PATH = Path(__file__).parents[1] / "models" / "training_report.json"
+
+# Feature names for the v4 artifact (includes source flags)
+V4_FEATURE_NAMES = FEATURE_NAMES + [
+ "source_roi_original",
+ "source_segmented",
+ "source_forniceal_palpebral",
+]
+
+# Indices of color features used for augmentation
+_COLOR_IDX = [V4_FEATURE_NAMES.index(n) for n in COLOR_FEATURES if n in V4_FEATURE_NAMES]
+# Index of CPI feature for jitter
+_CPI_IDX = V4_FEATURE_NAMES.index("cpi")
+
+N_CV_SPLITS = 5
+RANDOM_STATE = 42
+
+
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+# Augmentation
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+
+def augment_training_data(
+ rows: np.ndarray,
+ targets: np.ndarray,
+ labels: np.ndarray,
+ groups: np.ndarray,
+ rng: np.random.Generator,
+) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
+ """
+ Augment training data:
+ 1. Add Gaussian noise (sigma=0.01) to color features for ALL samples.
+ 2. Oversample anemic class (label=1) 3รโ with CPI jitter รยฑ0.02.
+ Returns augmented arrays (originals + augmented copies).
+ """
+ n = len(rows)
+
+ # --- noise augmentation for all samples ---
+ noisy = rows.copy()
+ noise = rng.normal(0, 0.01, size=(n, len(_COLOR_IDX)))
+ noisy[:, _COLOR_IDX] += noise
+ noisy = np.clip(noisy, 0.0, 1.0)
+
+ aug_rows = [rows, noisy]
+ aug_targets = [targets, targets]
+ aug_labels = [labels, labels]
+ aug_groups = [groups, groups]
+
+ # --- 3รโ oversample anemic samples with CPI jitter ---
+ anemic_idx = np.where(labels == 1)[0]
+ for _ in range(3):
+ copies = rows[anemic_idx].copy()
+ jitter = rng.uniform(-0.02, 0.02, size=len(anemic_idx))
+ copies[:, _CPI_IDX] = np.clip(copies[:, _CPI_IDX] + jitter, 0.0, 1.0)
+ # Also add small noise to other color features
+ color_noise = rng.normal(0, 0.01, size=(len(anemic_idx), len(_COLOR_IDX)))
+ copies[:, _COLOR_IDX] = np.clip(copies[:, _COLOR_IDX] + color_noise, 0.0, 1.0)
+ aug_rows.append(copies)
+ aug_targets.append(targets[anemic_idx])
+ aug_labels.append(labels[anemic_idx])
+ aug_groups.append(groups[anemic_idx])
+
+ return (
+ np.vstack(aug_rows),
+ np.concatenate(aug_targets),
+ np.concatenate(aug_labels),
+ np.concatenate(aug_groups),
+ )
+
+
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+# Base learner builders
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+
+def _et_regressor(rs: int = RANDOM_STATE) -> ExtraTreesRegressor:
+ return ExtraTreesRegressor(
+ n_estimators=300, min_samples_leaf=2, max_features=0.7,
+ bootstrap=True, random_state=rs, n_jobs=1,
+ )
+
+
+def _et_classifier(rs: int = RANDOM_STATE) -> ExtraTreesClassifier:
+ return ExtraTreesClassifier(
+ n_estimators=300, min_samples_leaf=2, max_features=0.7,
+ bootstrap=True, class_weight="balanced_subsample",
+ random_state=rs, n_jobs=1,
+ )
+
+
+def _xgb_regressor(rs: int = RANDOM_STATE) -> "XGBRegressor":
+ return XGBRegressor(
+ n_estimators=300, max_depth=4, learning_rate=0.05,
+ subsample=0.8, colsample_bytree=0.8,
+ random_state=rs, n_jobs=1, verbosity=0,
+ )
+
+
+def _xgb_classifier(rs: int = RANDOM_STATE) -> "XGBClassifier":
+ return XGBClassifier(
+ n_estimators=300, max_depth=4, learning_rate=0.05,
+ subsample=0.8, colsample_bytree=0.8,
+ use_label_encoder=False, eval_metric="logloss",
+ random_state=rs, n_jobs=1, verbosity=0,
+ )
+
+
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+# Hyperparameter tuning
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+
+def tune_et_regressor(rows: np.ndarray, targets: np.ndarray) -> ExtraTreesRegressor:
+ param_dist = {
+ "n_estimators": [100, 200, 300, 400],
+ "min_samples_leaf": [1, 2, 3, 4],
+ "max_features": [0.5, 0.6, 0.7, 0.8, "sqrt"],
+ }
+ base = ExtraTreesRegressor(bootstrap=True, random_state=RANDOM_STATE, n_jobs=1)
+ search = RandomizedSearchCV(
+ base, param_dist, n_iter=20, cv=3, scoring="neg_mean_absolute_error",
+ random_state=RANDOM_STATE, n_jobs=1, refit=True,
+ )
+ search.fit(rows, targets)
+ print(f" ET regressor best params: {search.best_params_}", flush=True)
+ return search.best_estimator_
+
+
+def tune_et_classifier(rows: np.ndarray, labels: np.ndarray) -> ExtraTreesClassifier:
+ param_dist = {
+ "n_estimators": [100, 200, 300, 400],
+ "min_samples_leaf": [1, 2, 3, 4],
+ "max_features": [0.5, 0.6, 0.7, 0.8, "sqrt"],
+ }
+ base = ExtraTreesClassifier(
+ bootstrap=True, class_weight="balanced_subsample",
+ random_state=RANDOM_STATE, n_jobs=1,
+ )
+ search = RandomizedSearchCV(
+ base, param_dist, n_iter=20, cv=3, scoring="f1",
+ random_state=RANDOM_STATE, n_jobs=1, refit=True,
+ )
+ search.fit(rows, labels)
+ print(f" ET classifier best params: {search.best_params_}", flush=True)
+ return search.best_estimator_
+
+
+def tune_xgb_regressor(rows: np.ndarray, targets: np.ndarray) -> "XGBRegressor":
+ param_dist = {
+ "n_estimators": [100, 200, 300, 400, 500],
+ "max_depth": [3, 4, 5, 6],
+ "learning_rate": [0.01, 0.03, 0.05, 0.1, 0.15],
+ "subsample": [0.6, 0.7, 0.8, 0.9, 1.0],
+ "colsample_bytree": [0.6, 0.7, 0.8, 0.9, 1.0],
+ }
+ base = XGBRegressor(random_state=RANDOM_STATE, n_jobs=1, verbosity=0)
+ search = RandomizedSearchCV(
+ base, param_dist, n_iter=20, cv=3, scoring="neg_mean_absolute_error",
+ random_state=RANDOM_STATE, n_jobs=1, refit=True,
+ )
+ search.fit(rows, targets)
+ print(f" XGB regressor best params: {search.best_params_}", flush=True)
+ return search.best_estimator_
+
+
+def tune_xgb_classifier(rows: np.ndarray, labels: np.ndarray) -> "XGBClassifier":
+ param_dist = {
+ "n_estimators": [100, 200, 300, 400, 500],
+ "max_depth": [3, 4, 5, 6],
+ "learning_rate": [0.01, 0.03, 0.05, 0.1, 0.15],
+ "subsample": [0.6, 0.7, 0.8, 0.9, 1.0],
+ "colsample_bytree": [0.6, 0.7, 0.8, 0.9, 1.0],
+ }
+ base = XGBClassifier(
+ use_label_encoder=False, eval_metric="logloss",
+ random_state=RANDOM_STATE, n_jobs=1, verbosity=0,
+ )
+ search = RandomizedSearchCV(
+ base, param_dist, n_iter=20, cv=3, scoring="f1",
+ random_state=RANDOM_STATE, n_jobs=1, refit=True,
+ )
+ search.fit(rows, labels)
+ print(f" XGB classifier best params: {search.best_params_}", flush=True)
+ return search.best_estimator_
+
+
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+# Stacking helpers
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+
+def _group_kfold_indices(
+ groups: np.ndarray, n_splits: int, random_state: int
+) -> list[tuple[np.ndarray, np.ndarray]]:
+ """GroupShuffleSplit folds for OOF stacking."""
+ splitter = GroupShuffleSplit(n_splits=n_splits, test_size=0.2, random_state=random_state)
+ return list(splitter.split(np.zeros(len(groups)), groups=groups))
+
+
+def build_oof_meta_features(
+ rows: np.ndarray,
+ targets: np.ndarray,
+ labels: np.ndarray,
+ groups: np.ndarray,
+ et_reg: ExtraTreesRegressor,
+ et_clf: ExtraTreesClassifier,
+ xgb_reg: object | None,
+ xgb_clf: object | None,
+ n_splits: int = N_CV_SPLITS,
+) -> np.ndarray:
+ """
+ Build out-of-fold meta-features using group-aware splits.
+ Always returns 4 columns: [et_hb, xgb_hb, et_prob, xgb_prob].
+ If XGBoost unavailable, xgb columns are zeros.
+ """
+ n = len(rows)
+ oof = np.zeros((n, 4), dtype=np.float32)
+ rng = np.random.default_rng(RANDOM_STATE)
+
+ folds = _group_kfold_indices(groups, n_splits, RANDOM_STATE)
+
+ for fold_i, (train_idx, val_idx) in enumerate(folds):
+ print(f" OOF fold {fold_i + 1}/{n_splits}...", flush=True)
+
+ tr_rows, tr_targets, tr_labels, tr_groups = augment_training_data(
+ rows[train_idx], targets[train_idx], labels[train_idx], groups[train_idx], rng
+ )
+ val_rows = rows[val_idx]
+
+ import copy
+ fold_et_reg = copy.deepcopy(et_reg)
+ fold_et_clf = copy.deepcopy(et_clf)
+ fold_et_reg.fit(tr_rows, tr_targets)
+ fold_et_clf.fit(tr_rows, tr_labels)
+
+ oof[val_idx, 0] = fold_et_reg.predict(val_rows)
+ oof[val_idx, 2] = fold_et_clf.predict_proba(val_rows)[:, 1]
+
+ if xgb_reg is not None and xgb_clf is not None:
+ fold_xgb_reg = copy.deepcopy(xgb_reg)
+ fold_xgb_clf = copy.deepcopy(xgb_clf)
+ fold_xgb_reg.fit(tr_rows, tr_targets)
+ fold_xgb_clf.fit(tr_rows, tr_labels)
+ oof[val_idx, 1] = fold_xgb_reg.predict(val_rows)
+ oof[val_idx, 3] = fold_xgb_clf.predict_proba(val_rows)[:, 1]
+ else:
+ oof[val_idx, 1] = oof[val_idx, 0] # mirror ET if no XGB
+ oof[val_idx, 3] = oof[val_idx, 2]
+
+ return oof
+
+
+def find_best_threshold(labels: np.ndarray, scores: np.ndarray) -> float:
+ """Threshold maximising recall-weighted F1 (medical screening priority)."""
+ best_score, best_thresh = -1.0, 0.5
+ for t in np.linspace(0.20, 0.75, 56):
+ preds = (scores >= t).astype(int)
+ if preds.sum() == 0:
+ continue
+ f1 = f1_score(labels, preds, zero_division=0)
+ rec = recall_score(labels, preds, zero_division=0)
+ score = f1 * 0.5 + rec * 0.5
+ if score > best_score:
+ best_score = score
+ best_thresh = float(t)
+ return best_thresh
+
+
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+# CV evaluation of the full stacked pipeline
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+
+def evaluate_stacked(
+ rows: np.ndarray,
+ targets: np.ndarray,
+ labels: np.ndarray,
+ groups: np.ndarray,
+ et_reg: ExtraTreesRegressor,
+ et_clf: ExtraTreesClassifier,
+ xgb_reg: object | None,
+ xgb_clf: object | None,
+ n_splits: int = N_CV_SPLITS,
+) -> dict[str, float]:
+ """
+ Outer CV loop: for each fold, build OOF meta-features on the train portion,
+ fit meta-learners, evaluate on the held-out test fold.
+ """
+ import copy
+ splitter = GroupShuffleSplit(n_splits=n_splits, test_size=0.2, random_state=RANDOM_STATE + 1)
+ all_metrics: list[dict[str, float]] = []
+ rng = np.random.default_rng(RANDOM_STATE + 99)
+
+ for fold_i, (train_idx, test_idx) in enumerate(splitter.split(rows, labels, groups)):
+ print(f" Outer CV fold {fold_i + 1}/{n_splits}...", flush=True)
+
+ tr_rows_raw = rows[train_idx]
+ tr_targets_raw = targets[train_idx]
+ tr_labels_raw = labels[train_idx]
+ tr_groups_raw = groups[train_idx]
+ te_rows = rows[test_idx]
+ te_targets = targets[test_idx]
+ te_labels = labels[test_idx]
+
+ # Build OOF meta-features on training portion (inner loop)
+ oof_meta = build_oof_meta_features(
+ tr_rows_raw, tr_targets_raw, tr_labels_raw, tr_groups_raw,
+ copy.deepcopy(et_reg), copy.deepcopy(et_clf),
+ copy.deepcopy(xgb_reg) if xgb_reg else None,
+ copy.deepcopy(xgb_clf) if xgb_clf else None,
+ n_splits=3,
+ )
+
+ # Fit meta-learners on OOF
+ meta_reg = Ridge(alpha=1.0)
+ meta_clf = LogisticRegression(C=1.0, max_iter=500, random_state=RANDOM_STATE, solver="lbfgs")
+ meta_reg.fit(oof_meta, tr_targets_raw)
+ meta_clf.fit(oof_meta, tr_labels_raw)
+
+ # Build test meta-features: retrain base learners on augmented full train
+ aug_rows, aug_targets, aug_labels, _ = augment_training_data(
+ tr_rows_raw, tr_targets_raw, tr_labels_raw, tr_groups_raw, rng
+ )
+
+ fold_et_reg = copy.deepcopy(et_reg); fold_et_reg.fit(aug_rows, aug_targets)
+ fold_et_clf = copy.deepcopy(et_clf); fold_et_clf.fit(aug_rows, aug_labels)
+
+ te_meta = np.zeros((len(te_rows), 4), dtype=np.float32)
+ te_meta[:, 0] = fold_et_reg.predict(te_rows)
+ te_meta[:, 2] = fold_et_clf.predict_proba(te_rows)[:, 1]
+
+ if xgb_reg is not None:
+ fold_xgb_reg = copy.deepcopy(xgb_reg); fold_xgb_reg.fit(aug_rows, aug_targets)
+ fold_xgb_clf = copy.deepcopy(xgb_clf); fold_xgb_clf.fit(aug_rows, aug_labels)
+ te_meta[:, 1] = fold_xgb_reg.predict(te_rows)
+ te_meta[:, 3] = fold_xgb_clf.predict_proba(te_rows)[:, 1]
+ else:
+ te_meta[:, 1] = te_meta[:, 0]
+ te_meta[:, 3] = te_meta[:, 2]
+
+ hb_pred = meta_reg.predict(te_meta)
+ clf_prob = meta_clf.predict_proba(te_meta)[:, 1]
+
+ # Blend: same scheme as legacy model
+ hb_scale = max(float(np.quantile(np.abs(tr_targets_raw - fold_et_reg.predict(tr_rows_raw)), 0.75)), 0.8)
+ reg_risk = np.array([sigmoid((ANEMIA_HB_THRESHOLD - h) / hb_scale) for h in hb_pred])
+ blend = 0.55 * clf_prob + 0.45 * reg_risk
+ thresh = find_best_threshold(te_labels, blend)
+ preds = (blend >= thresh).astype(int)
+
+ all_metrics.append({
+ "accuracy": accuracy_score(te_labels, preds),
+ "precision": precision_score(te_labels, preds, zero_division=0),
+ "recall": recall_score(te_labels, preds, zero_division=0),
+ "f1": f1_score(te_labels, preds, zero_division=0),
+ "auc": roc_auc_score(te_labels, blend),
+ "mae_hb": mean_absolute_error(te_targets, hb_pred),
+ "threshold": thresh,
+ })
+
+ avg = {k: round(float(np.mean([m[k] for m in all_metrics])), 4) for k in all_metrics[0]}
+ return avg
+
+
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+# Main
+# รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+
+def main() -> None:
+ print("=" * 60, flush=True)
+ print("AnemiaLens รขโฌโ stacked-ensemble-v4 training", flush=True)
+ print("=" * 60, flush=True)
+
+ # รขโโฌรขโโฌ 1. Load dataset รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ print("\n[1/6] Loading dataset...", flush=True)
+ subjects = _build_subject_catalog(DATASET_ROOT)
+ print(f" Subjects: {len(subjects)}", flush=True)
+
+ samples = _samples_for_mode(subjects, "hybrid_dual")
+ print(f" Samples (hybrid_dual): {len(samples)}", flush=True)
+
+ rows, targets, labels, groups = _rows_from_samples(samples)
+ print(f" Class balance: {labels.sum()} anemic / {(labels == 0).sum()} non-anemic", flush=True)
+
+ # รขโโฌรขโโฌ 2. Hyperparameter tuning รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ print("\n[2/6] Tuning hyperparameters (RandomizedSearchCV, 20 iter each)...", flush=True)
+ rng = np.random.default_rng(RANDOM_STATE)
+ aug_rows, aug_targets, aug_labels, _ = augment_training_data(rows, targets, labels, groups, rng)
+
+ print(" Tuning ExtraTrees regressor...", flush=True)
+ et_reg = tune_et_regressor(aug_rows, aug_targets)
+
+ print(" Tuning ExtraTrees classifier...", flush=True)
+ et_clf = tune_et_classifier(aug_rows, aug_labels)
+
+ if _HAS_XGB:
+ print(" Tuning XGBoost regressor...", flush=True)
+ xgb_reg = tune_xgb_regressor(aug_rows, aug_targets)
+ print(" Tuning XGBoost classifier...", flush=True)
+ xgb_clf = tune_xgb_classifier(aug_rows, aug_labels)
+ else:
+ xgb_reg = xgb_clf = None
+
+ # รขโโฌรขโโฌ 3. CV evaluation รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ print("\n[3/6] Cross-validating stacked ensemble...", flush=True)
+ cv_metrics = evaluate_stacked(rows, targets, labels, groups, et_reg, et_clf, xgb_reg, xgb_clf)
+ print(f"\n CV metrics: {cv_metrics}", flush=True)
+
+ # รขโโฌรขโโฌ 4. Build final OOF meta-features on all data รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ print("\n[4/6] Building final OOF meta-features on full dataset...", flush=True)
+ oof_meta = build_oof_meta_features(
+ rows, targets, labels, groups, et_reg, et_clf, xgb_reg, xgb_clf, n_splits=N_CV_SPLITS
+ )
+
+ # รขโโฌรขโโฌ 5. Fit final meta-learners รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ print("\n[5/6] Fitting meta-learners on OOF predictions...", flush=True)
+ meta_reg = Ridge(alpha=1.0)
+ meta_clf = LogisticRegression(C=1.0, max_iter=500, random_state=RANDOM_STATE, solver="lbfgs")
+ meta_reg.fit(oof_meta, targets)
+ meta_clf.fit(oof_meta, labels)
+
+ # Retrain base learners on full augmented data for inference
+ rng2 = np.random.default_rng(RANDOM_STATE + 1)
+ full_aug_rows, full_aug_targets, full_aug_labels, _ = augment_training_data(
+ rows, targets, labels, groups, rng2
+ )
+ et_reg.fit(full_aug_rows, full_aug_targets)
+ et_clf.fit(full_aug_rows, full_aug_labels)
+ if xgb_reg is not None:
+ xgb_reg.fit(full_aug_rows, full_aug_targets)
+ xgb_clf.fit(full_aug_rows, full_aug_labels)
+
+ # รขโโฌรขโโฌ 6. Instantiate module-level stacked wrappers for inference รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ stacked_reg = StackedRegressor(et_reg, xgb_reg, et_clf, xgb_clf, meta_reg)
+ stacked_clf = StackedClassifier(et_clf, xgb_clf, et_reg, xgb_reg, meta_clf)
+
+ # รขโโฌรขโโฌ Calibration รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ hb_preds_full = stacked_reg.predict(rows)
+ residuals = np.abs(targets - hb_preds_full)
+ hb_scale = max(float(np.quantile(residuals, 0.75)), 0.8)
+
+ hb_population_mean = float(np.mean(targets))
+ pred_std = float(np.std(hb_preds_full))
+ true_std = float(np.std(targets))
+ hb_spread_factor = float(np.clip(true_std / max(pred_std, 0.5), 1.0, 2.0))
+
+ clf_probs_full = stacked_clf.predict_proba(rows)[:, 1]
+ reg_risk_full = np.array([sigmoid((ANEMIA_HB_THRESHOLD - h) / hb_scale) for h in hb_preds_full])
+ blend_full = 0.55 * clf_probs_full + 0.45 * reg_risk_full
+ best_threshold = find_best_threshold(labels, blend_full)
+ risk_scale = max(float(np.std(blend_full)) * 0.9, 0.08)
+ risk_scale = min(risk_scale, 0.22)
+
+ calibration = {
+ "hb_threshold": ANEMIA_HB_THRESHOLD,
+ "hb_scale": round(hb_scale, 4),
+ "hb_population_mean": round(hb_population_mean, 4),
+ "hb_spread_factor": round(hb_spread_factor, 4),
+ "regressor_tree_std_reference": 2.5,
+ "classifier_tree_std_reference": 0.5,
+ "classifier_weight": 0.55,
+ "blend_threshold": round(best_threshold, 4),
+ "risk_scale": round(risk_scale, 4),
+ "base_uncertainty": 0.11,
+ }
+
+ # รขโโฌรขโโฌ Save artifact รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ print("\n[6/6] Saving model...", flush=True)
+ artifact = {
+ "version": "stacked-ensemble-v4",
+ "feature_names": V4_FEATURE_NAMES,
+ "regressor": stacked_reg,
+ "classifier": stacked_clf,
+ "calibration": calibration,
+ "training": {
+ "selected_mode": "hybrid_dual",
+ "subject_count": len(subjects),
+ "record_count": len(samples),
+ "metrics": cv_metrics,
+ "xgboost_available": _HAS_XGB,
+ },
+ }
+ OUTPUT_PATH.parent.mkdir(parents=True, exist_ok=True)
+ joblib.dump(artifact, OUTPUT_PATH)
+ joblib.dump(artifact, OUTPUT_PATH_V4) # keep versioned copy too
+ print(f" Saved รขโ โ {OUTPUT_PATH}", flush=True)
+ print(f" Saved รขโ โ {OUTPUT_PATH_V4}", flush=True)
+
+ report = {
+ "dataset_name": "dataset anemia",
+ "record_count": len(samples),
+ "subject_count": len(subjects),
+ "primary_model": "stacked-ensemble-v4",
+ "selected_mode": "hybrid_dual",
+ "metrics": cv_metrics,
+ "calibration": {
+ "blend_threshold": calibration["blend_threshold"],
+ "risk_scale": calibration["risk_scale"],
+ "classifier_weight": calibration["classifier_weight"],
+ },
+ }
+ with open(REPORT_PATH, "w") as f:
+ json.dump(report, f, indent=2)
+ print(f" Report รขโ โ {REPORT_PATH}", flush=True)
+
+ # รขโโฌรขโโฌ Sanity check รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ
+ print("\nรขโโฌรขโโฌ Sanity check รขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌรขโโฌ", flush=True)
+ feat_idx = {n: i for i, n in enumerate(V4_FEATURE_NAMES)}
+
+ test_cases = [
+ ("PALE (anemic)", 0.28, 0.02, 0.22),
+ ("BORDERLINE", 0.35, 0.04, 0.30),
+ ("NORMAL", 0.44, 0.08, 0.38),
+ ("VERY HEALTHY", 0.48, 0.10, 0.42),
+ ]
+ for label, cpi_val, rg_val, br_val in test_cases:
+ row = np.zeros((1, len(V4_FEATURE_NAMES)), dtype=np.float32)
+ row[0, feat_idx["cpi"]] = cpi_val
+ row[0, feat_idx["center_cpi"]] = cpi_val - 0.01
+ row[0, feat_idx["mean_r"]] = cpi_val * 0.9
+ row[0, feat_idx["mean_g"]] = cpi_val * 0.9 - rg_val
+ row[0, feat_idx["mean_b"]] = cpi_val * 0.7
+ row[0, feat_idx["center_mean_r"]] = cpi_val * 0.9
+ row[0, feat_idx["center_mean_g"]] = cpi_val * 0.9 - rg_val
+ row[0, feat_idx["center_mean_b"]] = cpi_val * 0.7
+ row[0, feat_idx["red_green_gap"]] = rg_val
+ row[0, feat_idx["center_red_green_gap"]] = rg_val
+ row[0, feat_idx["brightness"]] = br_val
+ row[0, feat_idx["center_brightness"]] = br_val
+ row[0, feat_idx["contrast"]] = 0.12
+ row[0, feat_idx["center_contrast"]] = 0.12
+ row[0, feat_idx["blur_score"]] = 100.0
+ row[0, feat_idx["center_blur_score"]] = 120.0
+ row[0, feat_idx["saturation"]] = 0.3
+ row[0, feat_idx["center_saturation"]] = 0.3
+ row[0, feat_idx["green_blue_ratio"]] = 1.1 if cpi_val < 0.35 else 1.25
+ row[0, feat_idx["hist_mid"]] = 0.5
+ row[0, feat_idx["hist_bright"]] = 0.3
+ row[0, feat_idx["aspect_ratio"]] = 1.0
+ row[0, feat_idx["size_score"]] = 1.0
+ row[0, feat_idx["source_roi_original"]] = 1.0
+
+ hb_p = float(stacked_reg.predict(row)[0])
+ cp = float(stacked_clf.predict_proba(row)[0, 1])
+ rr = sigmoid((ANEMIA_HB_THRESHOLD - hb_p) / hb_scale)
+ bs = 0.55 * cp + 0.45 * rr
+ risk = sigmoid((bs - best_threshold) / risk_scale)
+ decision = "ANEMIA LIKELY" if risk >= 0.65 else "unlikely"
+ print(f" {label}: Hb={hb_p:.1f}, clf_prob={cp:.3f}, risk={risk:.3f} -> {decision}", flush=True)
+
+ print("\nDone.", flush=True)
+
+
+if __name__ == "__main__":
+ main()
+
diff --git a/backend/start_server.py b/backend/start_server.py
new file mode 100644
index 0000000000000000000000000000000000000000..e15dcf45b3b7bee574a5862cef8ef1aaeee5d5fa
--- /dev/null
+++ b/backend/start_server.py
@@ -0,0 +1,19 @@
+#!/usr/bin/env python3
+"""
+Generic startup script for hosted AnemiaLens backends.
+Uses PORT and HOST environment variables when provided by the host.
+"""
+
+import os
+
+import uvicorn
+
+from app.main import app
+
+
+if __name__ == "__main__":
+ port = int(os.environ.get("PORT", 8000))
+ host = os.environ.get("HOST", "0.0.0.0")
+
+ print(f"Starting AnemiaLens on {host}:{port}")
+ uvicorn.run(app, host=host, port=port)
diff --git a/backend/tests/test_case_insight.py b/backend/tests/test_case_insight.py
new file mode 100644
index 0000000000000000000000000000000000000000..8c5af7023786af846c49dd9b1d3fab6714ce5556
--- /dev/null
+++ b/backend/tests/test_case_insight.py
@@ -0,0 +1,182 @@
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.schemas import (
+ DecisionAudit,
+ GuidanceResult,
+ PredictionResult,
+ QualityAssessment,
+ QualityIssue,
+ SymptomInput,
+ TriageResult,
+)
+from app.services.case_insight import CaseInsightService
+
+
+def test_case_insight_builds_high_concern_story_with_drivers() -> None:
+ service = CaseInsightService()
+ pack = service.build(
+ QualityAssessment(
+ passed=True,
+ blur_score=210.0,
+ brightness_score=0.46,
+ contrast_score=0.18,
+ framing_score=1.9,
+ issues=[],
+ ),
+ PredictionResult(
+ anemia_risk=0.82,
+ predicted_hemoglobin=7.6,
+ confidence=0.89,
+ uncertainty=0.09,
+ reliability_flag="high",
+ screening_label="anemia_likely",
+ screening_text="The screening model detected a strong low-hemoglobin signal.",
+ model_source="archive-evidence-fusion-v4",
+ ),
+ TriageResult(
+ band="high_concern",
+ score=0.88,
+ label="High concern",
+ summary="Arrange formal review soon.",
+ disclaimer="Screening only.",
+ ),
+ DecisionAudit(
+ processing_path="roi_crop",
+ calibration_band="strong_positive",
+ decision_threshold=0.435,
+ threshold_margin=0.385,
+ quality_warning_codes=[],
+ review_flags=[],
+ summary="Direct ROI inference produced a strong positive margin.",
+ ),
+ GuidanceResult(
+ source="fallback",
+ explanation="Severely low hemoglobin signal.",
+ urgency_guidance="Seek medical attention within 24-48 hours.",
+ food_advice="Eat iron-rich foods.",
+ next_steps=["Visit nearest clinic or hospital today", "Request a full blood count (CBC) test"],
+ ),
+ SymptomInput(fatigue=True, shortness_of_breath=True),
+ )
+
+ assert pack.priority_window == "within_24_48_hours"
+ assert pack.risk_drivers[0].impact == "up"
+ assert any(driver.title == "Very low hemoglobin estimate" for driver in pack.risk_drivers)
+ assert any("Avoid strenuous activity" in step.action for step in pack.follow_up_timeline)
+ assert "symptom fusion" in pack.judge_summary.lower()
+
+
+def test_case_insight_marks_rescue_path_as_confidence_limit() -> None:
+ service = CaseInsightService()
+ pack = service.build(
+ QualityAssessment(
+ passed=True,
+ blur_score=180.0,
+ brightness_score=0.39,
+ contrast_score=0.13,
+ framing_score=1.1,
+ issues=[
+ QualityIssue(
+ code="bad_framing",
+ severity="warning",
+ title="Eye framing is loose",
+ message="Recenter the eye.",
+ )
+ ],
+ ),
+ PredictionResult(
+ anemia_risk=0.51,
+ predicted_hemoglobin=10.9,
+ confidence=0.63,
+ uncertainty=0.29,
+ reliability_flag="medium",
+ screening_label="anemia_likely",
+ screening_text="The screening model detected some pallor-like signal.",
+ model_source="archive-evidence-fusion-v4",
+ ),
+ TriageResult(
+ band="moderate_risk",
+ score=0.59,
+ label="Moderate risk",
+ summary="Routine clinic follow-up is reasonable.",
+ disclaimer="Screening only.",
+ ),
+ DecisionAudit(
+ processing_path="full_frame_rescue",
+ calibration_band="borderline_positive",
+ decision_threshold=0.435,
+ threshold_margin=0.075,
+ quality_warning_codes=["bad_framing"],
+ review_flags=["raw_frame_rescue", "warning:bad_framing"],
+ summary="Full-frame rescue accepted a borderline positive result.",
+ ),
+ GuidanceResult(
+ source="fallback",
+ explanation="Mild to moderate anemia-like signal.",
+ urgency_guidance="See a doctor within 1-2 weeks.",
+ food_advice="Eat iron-rich foods.",
+ next_steps=["Book a clinic visit this week", "Start iron-rich diet immediately"],
+ ),
+ SymptomInput(),
+ )
+
+ assert "full-frame rescue" in pack.confidence_story.lower()
+ assert any(driver.impact == "limit" for driver in pack.risk_drivers)
+ assert any("direct conjunctiva crop" in item.lower() for item in pack.capture_improvements)
+
+
+def test_case_insight_handles_quality_blocked_retake_case() -> None:
+ service = CaseInsightService()
+ pack = service.build(
+ QualityAssessment(
+ passed=False,
+ blur_score=42.0,
+ brightness_score=0.06,
+ contrast_score=0.03,
+ framing_score=0.42,
+ issues=[
+ QualityIssue(
+ code="poor_lighting",
+ severity="blocking",
+ title="Lighting is not usable",
+ message="Use bright natural light.",
+ )
+ ],
+ ),
+ None,
+ TriageResult(
+ band="uncertain_retake_needed",
+ score=0.2,
+ label="Uncertain, retake needed",
+ summary="Retake the image.",
+ disclaimer="Screening only.",
+ ),
+ DecisionAudit(
+ processing_path="quality_blocked",
+ calibration_band="quality_blocked",
+ decision_threshold=None,
+ threshold_margin=None,
+ quality_warning_codes=[],
+ review_flags=["quality_blocked"],
+ summary="Quality blocked model inference.",
+ ),
+ GuidanceResult(
+ source="fallback",
+ explanation="Image signal was not strong enough.",
+ urgency_guidance="Retake the scan in better lighting.",
+ food_advice="No food advice until a valid screening is available.",
+ next_steps=["Retake eye image in bright natural light"],
+ ),
+ SymptomInput(dizziness=True),
+ )
+
+ assert pack.priority_window == "retake_now"
+ assert "blocked model inference" in pack.risk_drivers[0].detail.lower()
+ assert pack.capture_improvements[0].startswith("Move into bright, even natural light")
+ assert "safety gate" in pack.judge_summary.lower()
diff --git a/backend/tests/test_clinical_brief.py b/backend/tests/test_clinical_brief.py
new file mode 100644
index 0000000000000000000000000000000000000000..ad8337cf006fe0549ac814c57727d07968e72118
--- /dev/null
+++ b/backend/tests/test_clinical_brief.py
@@ -0,0 +1,195 @@
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.schemas import (
+ DecisionAudit,
+ GuidanceResult,
+ PredictionResult,
+ QualityAssessment,
+ QualityIssue,
+ SymptomInput,
+ TriageResult,
+)
+from app.services.analysis_meta import build_analysis_meta
+from app.services.case_insight import CaseInsightService
+from app.services.clinical_brief import ClinicalBriefService
+from app.services.handoff import HandoffSummaryService
+from app.services.triage import TriageService
+
+
+def test_clinical_brief_builds_grounded_high_concern_summary() -> None:
+ quality = QualityAssessment(
+ passed=True,
+ blur_score=198.0,
+ brightness_score=0.33,
+ contrast_score=0.19,
+ framing_score=1.9,
+ issues=[],
+ )
+ prediction = PredictionResult(
+ anemia_risk=0.84,
+ predicted_hemoglobin=7.8,
+ confidence=0.91,
+ uncertainty=0.08,
+ reliability_flag="high",
+ screening_label="anemia_likely",
+ screening_text="The screening model detected a strong low-hemoglobin signal.",
+ model_source="archive-evidence-fusion-v4",
+ )
+ symptoms = SymptomInput(fatigue=True, shortness_of_breath=True, poor_diet_low_iron=True)
+ triage_service = TriageService()
+ signal_breakdown = triage_service.build_signal_breakdown(quality, prediction, symptoms)
+ triage = triage_service.assess(
+ quality,
+ prediction,
+ symptoms,
+ signal_breakdown=signal_breakdown,
+ )
+ decision_audit = DecisionAudit(
+ processing_path="roi_crop",
+ calibration_band="strong_positive",
+ decision_threshold=0.435,
+ threshold_margin=0.405,
+ quality_warning_codes=[],
+ review_flags=[],
+ summary="Direct ROI inference produced a strong positive margin.",
+ )
+ guidance = GuidanceResult(
+ source="fallback",
+ explanation="The screening signal is concerning and should be reviewed soon.",
+ urgency_guidance="Seek medical review within 24 to 48 hours.",
+ food_advice="Eat iron-rich foods and include vitamin C with meals.",
+ next_steps=["Book a clinic or lab visit within 24 to 48 hours", "Request a CBC test"],
+ )
+ insight_pack = CaseInsightService().build(
+ quality,
+ prediction,
+ triage,
+ decision_audit,
+ guidance,
+ symptoms,
+ )
+ handoff_summary = HandoffSummaryService().build(
+ quality,
+ prediction,
+ triage,
+ guidance,
+ symptoms,
+ )
+
+ brief = ClinicalBriefService().build(
+ quality,
+ prediction,
+ triage,
+ decision_audit,
+ guidance,
+ symptoms,
+ insight_pack,
+ handoff_summary,
+ signal_breakdown,
+ )
+
+ assert brief.action_window == "within_24_48_hours"
+ assert brief.signal_breakdown.image_risk == 0.84
+ assert brief.signal_breakdown.symptom_burden == "moderate"
+ assert any("hemoglobin signal" in item.lower() for item in brief.supporting_evidence)
+ assert any("uncertainty" in item.lower() for item in brief.safety_checks)
+ assert "AnemiaLens clinical brief" in brief.share_text
+
+
+def test_clinical_brief_handles_quality_blocked_case_and_meta() -> None:
+ quality = QualityAssessment(
+ passed=False,
+ blur_score=42.0,
+ brightness_score=0.05,
+ contrast_score=0.03,
+ framing_score=0.4,
+ issues=[
+ QualityIssue(
+ code="poor_lighting",
+ severity="blocking",
+ title="Lighting is not usable",
+ message="Use bright natural light.",
+ )
+ ],
+ )
+ symptoms = SymptomInput(dizziness=True)
+ triage_service = TriageService()
+ signal_breakdown = triage_service.build_signal_breakdown(quality, None, symptoms)
+ triage = triage_service.assess(
+ quality,
+ None,
+ symptoms,
+ signal_breakdown=signal_breakdown,
+ )
+ decision_audit = DecisionAudit(
+ processing_path="quality_blocked",
+ calibration_band="quality_blocked",
+ decision_threshold=None,
+ threshold_margin=None,
+ quality_warning_codes=[],
+ review_flags=["quality_blocked"],
+ summary="Quality blocked model inference.",
+ )
+ guidance = GuidanceResult(
+ source="fallback",
+ explanation="The image was too weak for a reliable screening result.",
+ urgency_guidance="Retake the scan in better light.",
+ food_advice="Wait for a valid scan before using food guidance from the app.",
+ next_steps=["Retake the image in bright natural light"],
+ )
+ insight_pack = CaseInsightService().build(
+ quality,
+ None,
+ triage,
+ decision_audit,
+ guidance,
+ symptoms,
+ )
+ handoff_summary = HandoffSummaryService().build(
+ quality,
+ None,
+ triage,
+ guidance,
+ symptoms,
+ )
+
+ brief = ClinicalBriefService().build(
+ quality,
+ None,
+ triage,
+ decision_audit,
+ guidance,
+ symptoms,
+ insight_pack,
+ handoff_summary,
+ signal_breakdown,
+ )
+ meta = build_analysis_meta(
+ request_id="abc12345",
+ api_version="0.3.0",
+ processing_time_ms=187.36,
+ quality=quality,
+ decision_audit=decision_audit,
+ guidance=guidance,
+ used_raw_frame_rescue=False,
+ )
+
+ assert brief.signal_breakdown.image_risk is None
+ assert any("blocked model inference" in item.lower() for item in brief.supporting_evidence)
+ assert any("primary blocker" in item.lower() for item in brief.limiting_factors)
+ assert "image=not available" in brief.share_text
+ assert meta.request_id == "abc12345"
+ assert meta.processing_path == "quality_blocked"
+ assert meta.guidance_source == "fallback"
+ assert meta.safety_layers == [
+ "image_quality_gate",
+ "symptom_fusion",
+ "triage_banding",
+ "non_diagnostic_guidance",
+ ]
diff --git a/backend/tests/test_decision_audit.py b/backend/tests/test_decision_audit.py
new file mode 100644
index 0000000000000000000000000000000000000000..37865384b3b6295761c0513331c7de79f0eb8f00
--- /dev/null
+++ b/backend/tests/test_decision_audit.py
@@ -0,0 +1,89 @@
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.schemas import GuidanceResult, PredictionResult, QualityAssessment, QualityIssue, SymptomInput, TriageResult
+from app.services.decision_audit import build_decision_audit
+
+
+def test_decision_audit_marks_full_frame_rescue_and_threshold_margin() -> None:
+ audit = build_decision_audit(
+ QualityAssessment(
+ passed=True,
+ blur_score=220.0,
+ brightness_score=0.44,
+ contrast_score=0.15,
+ framing_score=1.7,
+ issues=[
+ QualityIssue(
+ code="bad_framing",
+ severity="warning",
+ title="Eye framing is loose",
+ message="The app fell back to the full eye frame.",
+ )
+ ],
+ ),
+ PredictionResult(
+ anemia_risk=0.82,
+ predicted_hemoglobin=10.8,
+ confidence=0.66,
+ uncertainty=0.34,
+ reliability_flag="medium",
+ screening_label="anemia_likely",
+ screening_text="Likely anemia.",
+ model_source="archive-evidence-fusion-v4",
+ ),
+ TriageResult(
+ band="moderate_risk",
+ score=0.54,
+ label="Moderate risk",
+ summary="Moderate concern.",
+ disclaimer="Screening only.",
+ ),
+ used_raw_frame_rescue=True,
+ )
+
+ assert audit.processing_path == "full_frame_rescue"
+ assert audit.calibration_band == "strong_positive"
+ assert audit.decision_threshold == 0.435
+ assert audit.threshold_margin == 0.385
+ assert "raw_frame_rescue" in audit.review_flags
+ assert "warning:bad_framing" in audit.review_flags
+
+
+def test_decision_audit_handles_blocked_request() -> None:
+ audit = build_decision_audit(
+ QualityAssessment(
+ passed=False,
+ blur_score=40.0,
+ brightness_score=0.05,
+ contrast_score=0.03,
+ framing_score=0.4,
+ issues=[
+ QualityIssue(
+ code="poor_lighting",
+ severity="blocking",
+ title="Lighting is not usable",
+ message="Use bright, even light.",
+ )
+ ],
+ ),
+ None,
+ TriageResult(
+ band="uncertain_retake_needed",
+ score=0.2,
+ label="Retake needed",
+ summary="Retake the image.",
+ disclaimer="Screening only.",
+ ),
+ )
+
+ assert audit.processing_path == "quality_blocked"
+ assert audit.calibration_band == "quality_blocked"
+ assert audit.decision_threshold is None
+ assert "quality_blocked" in audit.review_flags
+ assert "blocked model inference" in audit.summary.lower()
diff --git a/backend/tests/test_email_report.py b/backend/tests/test_email_report.py
new file mode 100644
index 0000000000000000000000000000000000000000..a7149378401213e61836043ea6601349fb63e8da
--- /dev/null
+++ b/backend/tests/test_email_report.py
@@ -0,0 +1,414 @@
+"""
+Tests for the email report API and delivery service.
+"""
+
+from __future__ import annotations
+
+import json
+import sys
+from pathlib import Path
+
+import pytest
+from fastapi import FastAPI
+from fastapi.testclient import TestClient
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.api.email_report import get_email_report_service, router
+from app.config import settings
+from app.services import email_report as email_report_module
+from app.services.email_report import (
+ EmailReportContent,
+ EmailReportDeliveryError,
+ EmailReportNotConfiguredError,
+ EmailReportService,
+)
+
+
+class _StubRouterService:
+ def __init__(self, exc: Exception | None = None) -> None:
+ self.exc = exc
+ self.payload: EmailReportContent | None = None
+
+ def masked_recipient(self, recipient: str) -> str:
+ return f"masked:{recipient}"
+
+ def send_report(self, payload: EmailReportContent) -> None:
+ self.payload = payload
+ if self.exc is not None:
+ raise self.exc
+
+
+class _SMTPStub:
+ last_instance: "_SMTPStub | None" = None
+
+ def __init__(self, host: str, port: int, timeout: float | None = None, context=None) -> None:
+ self.host = host
+ self.port = port
+ self.timeout = timeout
+ self.context = context
+ self.logged_in: tuple[str, str] | None = None
+ self.sent_message = None
+ _SMTPStub.last_instance = self
+
+ def __enter__(self) -> "_SMTPStub":
+ return self
+
+ def __exit__(self, exc_type, exc, tb) -> None:
+ return None
+
+ def login(self, username: str, password: str) -> None:
+ self.logged_in = (username, password)
+
+ def send_message(self, message) -> None:
+ self.sent_message = message
+
+
+class _HTTPResponseStub:
+ def __init__(self, body: str = '{"id":"email_123"}', status: int = 200) -> None:
+ self._body = body.encode("utf-8")
+ self.status = status
+
+ def read(self) -> bytes:
+ return self._body
+
+
+class _HTTPSConnectionStub:
+ last_instance: "_HTTPSConnectionStub | None" = None
+ response_status: int = 200
+ response_body: str = '{"id":"email_123"}'
+
+ def __init__(self, host: str, timeout: float | None = None) -> None:
+ self.host = host
+ self.timeout = timeout
+ self.request_args: tuple[str, str, bytes, dict[str, str]] | None = None
+ self.closed = False
+ _HTTPSConnectionStub.last_instance = self
+
+ def request(self, method: str, path: str, body=None, headers=None) -> None:
+ self.request_args = (method, path, body, headers or {})
+
+ def getresponse(self) -> _HTTPResponseStub:
+ return _HTTPResponseStub(body=self.response_body, status=self.response_status)
+
+ def close(self) -> None:
+ self.closed = True
+
+
+class _GmailHTTPSConnectionStub:
+ requests: list[tuple[str, str, bytes, dict[str, str], str, float | None]] = []
+ response_queue: list[_HTTPResponseStub] = []
+
+ def __init__(self, host: str, timeout: float | None = None) -> None:
+ self.host = host
+ self.timeout = timeout
+ self.closed = False
+
+ def request(self, method: str, path: str, body=None, headers=None) -> None:
+ _GmailHTTPSConnectionStub.requests.append((method, path, body, headers or {}, self.host, self.timeout))
+
+ def getresponse(self) -> _HTTPResponseStub:
+ return _GmailHTTPSConnectionStub.response_queue.pop(0)
+
+ def close(self) -> None:
+ self.closed = True
+
+
+def _client_with_service(service: _StubRouterService) -> TestClient:
+ app = FastAPI()
+ app.include_router(router)
+ app.dependency_overrides[get_email_report_service] = lambda: service
+ return TestClient(app)
+
+
+def test_email_report_endpoint_sends_valid_payload() -> None:
+ service = _StubRouterService()
+ client = _client_with_service(service)
+
+ response = client.post(
+ "/api/email-report",
+ json={
+ "email": "person@example.com",
+ "share_text": "Moderate risk summary.\nPlease follow up with a CBC test.",
+ "triage_label": "Moderate Risk",
+ "predicted_hemoglobin": 10.6,
+ "anemia_risk": 0.54,
+ },
+ )
+
+ assert response.status_code == 200
+ assert response.json()["status"] == "sent"
+ assert service.payload is not None
+ assert service.payload.recipient == "person@example.com"
+ assert service.payload.predicted_hemoglobin == 10.6
+
+
+def test_email_report_endpoint_rejects_invalid_email() -> None:
+ client = _client_with_service(_StubRouterService())
+
+ response = client.post(
+ "/api/email-report",
+ json={
+ "email": "not-an-email",
+ "share_text": "Moderate risk summary.\nPlease follow up with a CBC test.",
+ "triage_label": "Moderate Risk",
+ "predicted_hemoglobin": 10.6,
+ "anemia_risk": 0.54,
+ },
+ )
+
+ assert response.status_code == 422
+
+
+def test_email_report_endpoint_returns_503_when_not_configured() -> None:
+ client = _client_with_service(
+ _StubRouterService(
+ EmailReportNotConfiguredError("Email delivery is not configured."),
+ )
+ )
+
+ response = client.post(
+ "/api/email-report",
+ json={
+ "email": "person@example.com",
+ "share_text": "Moderate risk summary.\nPlease follow up with a CBC test.",
+ "triage_label": "Moderate Risk",
+ "predicted_hemoglobin": 10.6,
+ "anemia_risk": 0.54,
+ },
+ )
+
+ assert response.status_code == 503
+ assert "not configured" in response.json()["detail"].lower()
+
+
+def test_email_report_endpoint_returns_502_when_delivery_fails() -> None:
+ client = _client_with_service(
+ _StubRouterService(
+ EmailReportDeliveryError("SMTP authentication failed."),
+ )
+ )
+
+ response = client.post(
+ "/api/email-report",
+ json={
+ "email": "person@example.com",
+ "share_text": "Moderate risk summary.\nPlease follow up with a CBC test.",
+ "triage_label": "Moderate Risk",
+ "predicted_hemoglobin": 10.6,
+ "anemia_risk": 0.54,
+ },
+ )
+
+ assert response.status_code == 502
+ assert "smtp" in response.json()["detail"].lower()
+
+
+def test_email_report_service_requires_configuration(monkeypatch: pytest.MonkeyPatch) -> None:
+ monkeypatch.setattr(settings, "email_provider", "smtp")
+ monkeypatch.setattr(settings, "smtp_username", "")
+ monkeypatch.setattr(settings, "smtp_password", "")
+ monkeypatch.setattr(settings, "email_from_email", "")
+
+ service = EmailReportService()
+
+ with pytest.raises(EmailReportNotConfiguredError, match="configured"):
+ service.send_report(
+ EmailReportContent(
+ recipient="person@example.com",
+ share_text="Moderate risk summary.\nPlease follow up with a CBC test.",
+ triage_label="Moderate Risk",
+ predicted_hemoglobin=10.6,
+ anemia_risk=0.54,
+ )
+ )
+
+
+def test_email_report_service_sends_email_via_ssl(monkeypatch: pytest.MonkeyPatch) -> None:
+ monkeypatch.setattr(settings, "email_provider", "smtp")
+ monkeypatch.setattr(settings, "smtp_host", "smtp.example.com")
+ monkeypatch.setattr(settings, "smtp_port", 465)
+ monkeypatch.setattr(settings, "smtp_username", "mailer@example.com")
+ monkeypatch.setattr(settings, "smtp_password", "app-password")
+ monkeypatch.setattr(settings, "smtp_use_ssl", True)
+ monkeypatch.setattr(settings, "smtp_use_starttls", False)
+ monkeypatch.setattr(settings, "smtp_timeout", 20.0)
+ monkeypatch.setattr(settings, "email_from_name", "AnemiaLens")
+ monkeypatch.setattr(settings, "email_from_email", "reports@example.com")
+ monkeypatch.setattr(settings, "email_reply_to", "support@example.com")
+ monkeypatch.setattr(email_report_module.smtplib, "SMTP_SSL", _SMTPStub)
+
+ service = EmailReportService()
+ service.send_report(
+ EmailReportContent(
+ recipient="patient@example.com",
+ share_text="Moderate risk summary.\nPlease follow up with a CBC test.",
+ triage_label="Moderate Risk",
+ predicted_hemoglobin=10.6,
+ anemia_risk=0.54,
+ )
+ )
+
+ smtp = _SMTPStub.last_instance
+ assert smtp is not None
+ assert smtp.host == "smtp.example.com"
+ assert smtp.port == 465
+ assert smtp.logged_in == ("mailer@example.com", "app-password")
+ assert smtp.sent_message["To"] == "patient@example.com"
+ assert smtp.sent_message["Reply-To"] == "support@example.com"
+ assert "Moderate Risk" in smtp.sent_message["Subject"]
+ plain_part = smtp.sent_message.get_body(preferencelist=("plain",))
+ html_part = smtp.sent_message.get_body(preferencelist=("html",))
+ assert plain_part is not None
+ assert html_part is not None
+ assert "clinical blood test (CBC)" in plain_part.get_content()
+ assert "Recommended Next Steps" in plain_part.get_content()
+ assert "Why this result" in html_part.get_content()
+ assert "Open AnemiaLens" in html_part.get_content()
+
+
+def test_email_report_service_sends_email_via_resend(monkeypatch: pytest.MonkeyPatch) -> None:
+ _HTTPSConnectionStub.response_status = 200
+ _HTTPSConnectionStub.response_body = '{"id":"email_123"}'
+ monkeypatch.setattr(settings, "email_provider", "resend")
+ monkeypatch.setattr(settings, "resend_api_key", "re_test_123")
+ monkeypatch.setattr(settings, "resend_api_base", "https://api.resend.test")
+ monkeypatch.setattr(settings, "email_from_name", "AnemiaLens")
+ monkeypatch.setattr(settings, "email_from_email", "onboarding@resend.dev")
+ monkeypatch.setattr(settings, "email_reply_to", "support@example.com")
+ monkeypatch.setattr(settings, "smtp_username", "")
+ monkeypatch.setattr(settings, "smtp_password", "")
+ monkeypatch.setattr(settings, "smtp_timeout", 12.0)
+ monkeypatch.setattr(email_report_module.http.client, "HTTPSConnection", _HTTPSConnectionStub)
+
+ service = EmailReportService()
+ service.send_report(
+ EmailReportContent(
+ recipient="patient@example.com",
+ share_text="Moderate risk summary.\nPlease follow up with a CBC test.",
+ triage_label="Moderate Risk",
+ predicted_hemoglobin=10.6,
+ anemia_risk=0.54,
+ )
+ )
+
+ connection = _HTTPSConnectionStub.last_instance
+ assert connection is not None
+ assert connection.host == "api.resend.test"
+ assert connection.timeout == 12.0
+ assert connection.closed is True
+ assert connection.request_args is not None
+ method, path, raw_body, headers = connection.request_args
+ body = json.loads(raw_body.decode("utf-8"))
+ assert method == "POST"
+ assert path == "/emails"
+ assert headers["Authorization"] == "Bearer re_test_123"
+ assert headers["Content-Type"] == "application/json"
+ assert headers["Idempotency-Key"].startswith("email-report/patient@example.com/moderate-risk/")
+ assert headers["User-Agent"] == "AnemiaLens/1.0 (+https://anemia-lens.vercel.app)"
+ assert body["from"] == "AnemiaLens "
+ assert body["to"] == ["patient@example.com"]
+ assert body["reply_to"] == "support@example.com"
+ assert body["subject"] == "AnemiaLens Screening Report - Moderate Risk"
+ assert "clinical blood test (CBC)" in body["text"]
+
+
+def test_email_report_service_sends_email_via_sendgrid(monkeypatch: pytest.MonkeyPatch) -> None:
+ _HTTPSConnectionStub.response_status = 202
+ _HTTPSConnectionStub.response_body = ""
+ monkeypatch.setattr(settings, "email_provider", "sendgrid")
+ monkeypatch.setattr(settings, "sendgrid_api_key", "SG.test-key")
+ monkeypatch.setattr(settings, "sendgrid_api_base", "https://api.sendgrid.test/v3")
+ monkeypatch.setattr(settings, "email_from_name", "AnemiaLens")
+ monkeypatch.setattr(settings, "email_from_email", "asnanp875@gmail.com")
+ monkeypatch.setattr(settings, "email_reply_to", "asnanp875@gmail.com")
+ monkeypatch.setattr(settings, "smtp_username", "")
+ monkeypatch.setattr(settings, "smtp_password", "")
+ monkeypatch.setattr(settings, "smtp_timeout", 12.0)
+ monkeypatch.setattr(email_report_module.http.client, "HTTPSConnection", _HTTPSConnectionStub)
+
+ service = EmailReportService()
+ service.send_report(
+ EmailReportContent(
+ recipient="patient@example.com",
+ share_text="Moderate risk summary.\nPlease follow up with a CBC test.",
+ triage_label="Moderate Risk",
+ predicted_hemoglobin=10.6,
+ anemia_risk=0.54,
+ )
+ )
+
+ connection = _HTTPSConnectionStub.last_instance
+ assert connection is not None
+ assert connection.host == "api.sendgrid.test"
+ assert connection.timeout == 12.0
+ assert connection.closed is True
+ assert connection.request_args is not None
+ method, path, raw_body, headers = connection.request_args
+ body = json.loads(raw_body.decode("utf-8"))
+ assert method == "POST"
+ assert path == "/v3/mail/send"
+ assert headers["Authorization"] == "Bearer SG.test-key"
+ assert headers["Content-Type"] == "application/json"
+ assert headers["User-Agent"] == "AnemiaLens/1.0 (+https://anemia-lens.vercel.app)"
+ assert body["from"] == {"email": "asnanp875@gmail.com", "name": "AnemiaLens"}
+ assert body["reply_to"] == {"email": "asnanp875@gmail.com"}
+ assert body["personalizations"][0]["to"] == [{"email": "patient@example.com"}]
+ assert body["personalizations"][0]["subject"] == "AnemiaLens Screening Report - Moderate Risk"
+ assert body["content"][0]["type"] == "text/plain"
+ assert body["content"][1]["type"] == "text/html"
+ assert "clinical blood test (CBC)" in body["content"][0]["value"]
+
+
+def test_email_report_service_sends_email_via_gmail_api(monkeypatch: pytest.MonkeyPatch) -> None:
+ _GmailHTTPSConnectionStub.requests = []
+ _GmailHTTPSConnectionStub.response_queue = [
+ _HTTPResponseStub(body='{"access_token":"ya29.test-token"}', status=200),
+ _HTTPResponseStub(body='{"id":"gmail_message_123"}', status=200),
+ ]
+ monkeypatch.setattr(settings, "email_provider", "gmail_api")
+ monkeypatch.setattr(settings, "gmail_client_id", "client-id")
+ monkeypatch.setattr(settings, "gmail_client_secret", "client-secret")
+ monkeypatch.setattr(settings, "gmail_refresh_token", "refresh-token")
+ monkeypatch.setattr(settings, "gmail_token_url", "https://oauth2.googleapis.com/token")
+ monkeypatch.setattr(settings, "gmail_api_base", "https://gmail.googleapis.com/gmail/v1")
+ monkeypatch.setattr(settings, "email_from_name", "AnemiaLens")
+ monkeypatch.setattr(settings, "email_from_email", "asnanp875@gmail.com")
+ monkeypatch.setattr(settings, "email_reply_to", "asnanp875@gmail.com")
+ monkeypatch.setattr(settings, "smtp_timeout", 12.0)
+ monkeypatch.setattr(email_report_module.http.client, "HTTPSConnection", _GmailHTTPSConnectionStub)
+
+ service = EmailReportService()
+ service.send_report(
+ EmailReportContent(
+ recipient="patient@example.com",
+ share_text="Moderate risk summary.\nPlease follow up with a CBC test.",
+ triage_label="Moderate Risk",
+ predicted_hemoglobin=10.6,
+ anemia_risk=0.54,
+ )
+ )
+
+ assert len(_GmailHTTPSConnectionStub.requests) == 2
+
+ token_method, token_path, token_body, token_headers, token_host, token_timeout = _GmailHTTPSConnectionStub.requests[0]
+ assert token_method == "POST"
+ assert token_host == "oauth2.googleapis.com"
+ assert token_timeout == 12.0
+ assert token_path == "/token"
+ assert token_headers["Content-Type"] == "application/x-www-form-urlencoded"
+ assert b"grant_type=refresh_token" in token_body
+ assert b"client_id=client-id" in token_body
+ assert b"client_secret=client-secret" in token_body
+ assert b"refresh_token=refresh-token" in token_body
+
+ send_method, send_path, send_body_raw, send_headers, send_host, send_timeout = _GmailHTTPSConnectionStub.requests[1]
+ send_body = json.loads(send_body_raw.decode("utf-8"))
+ assert send_method == "POST"
+ assert send_host == "gmail.googleapis.com"
+ assert send_timeout == 12.0
+ assert send_path == "/gmail/v1/users/me/messages/send"
+ assert send_headers["Authorization"] == "Bearer ya29.test-token"
+ assert send_headers["Content-Type"] == "application/json"
+ assert "raw" in send_body
diff --git a/backend/tests/test_error_analysis.py b/backend/tests/test_error_analysis.py
new file mode 100644
index 0000000000000000000000000000000000000000..32ec9397d777b335004f59ebc9fd72e4f6dee425
--- /dev/null
+++ b/backend/tests/test_error_analysis.py
@@ -0,0 +1,61 @@
+from __future__ import annotations
+
+from dataclasses import dataclass
+import sys
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+sys.path.insert(0, str(ROOT / "backend" / "scripts"))
+
+from analyze_efficientnet_errors import _mistakes, _source_breakdown
+
+
+@dataclass(frozen=True)
+class _Record:
+ subject_id: str
+ source: str
+ image_path: str
+
+
+def test_source_breakdown_counts_errors_by_source() -> None:
+ records = [
+ _Record("s1", "roi_original", "a.jpg"),
+ _Record("s2", "roi_original", "b.jpg"),
+ _Record("s3", "palpebral", "c.png"),
+ ]
+
+ result = _source_breakdown(
+ records,
+ labels=[0, 1, 0],
+ predictions=[1, 1, 0],
+ probabilities=[0.8, 0.9, 0.1],
+ hb_predictions=[10.5, 8.8, 12.4],
+ hb_targets=[12.6, 9.1, 12.1],
+ )
+
+ assert result["roi_original"]["count"] == 2
+ assert result["roi_original"]["false_positives"] == 1
+ assert result["roi_original"]["false_negatives"] == 0
+ assert result["palpebral"]["errors"] == 0
+
+
+def test_mistakes_splits_false_positives_and_false_negatives() -> None:
+ records = [
+ _Record("s1", "roi_original", "a.jpg"),
+ _Record("s2", "palpebral", "b.png"),
+ ]
+
+ false_positives, false_negatives = _mistakes(
+ records,
+ labels=[0, 1],
+ predictions=[1, 0],
+ probabilities=[0.91, 0.12],
+ hb_predictions=[10.2, 12.8],
+ hb_targets=[13.1, 8.9],
+ )
+
+ assert len(false_positives) == 1
+ assert false_positives[0]["subject_id"] == "s1"
+ assert len(false_negatives) == 1
+ assert false_negatives[0]["subject_id"] == "s2"
diff --git a/backend/tests/test_guidance.py b/backend/tests/test_guidance.py
new file mode 100644
index 0000000000000000000000000000000000000000..f0b7278356c3ff41ac55daf9fac88bf07114881c
--- /dev/null
+++ b/backend/tests/test_guidance.py
@@ -0,0 +1,430 @@
+"""
+Tests for GuidanceService, covering both Mistral-backed guidance and the
+rule-based fallback.
+"""
+
+from __future__ import annotations
+
+from collections import OrderedDict
+import sys
+import unittest
+from pathlib import Path
+
+import pytest
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.config import Settings
+from app.schemas import (
+ GuidanceResult,
+ GuidanceRuntimeStatus,
+ ModelRuntimeStatus,
+ PredictionResult,
+ SymptomInput,
+ TriageResult,
+)
+from app.services.guidance import GuidanceService
+from app.services.runtime_status import build_runtime_status
+
+
+class _PredictorStub:
+ def runtime_status(self) -> ModelRuntimeStatus:
+ return ModelRuntimeStatus(
+ primary_model="efficientnet-b0-ft",
+ deep_stack_loaded=False,
+ legacy_loaded=False,
+ )
+
+
+class _GuidanceStub:
+ def runtime_status(self) -> GuidanceRuntimeStatus:
+ return GuidanceRuntimeStatus(
+ active_strategy="mistral",
+ mistral_enabled=True,
+ client_ready=True,
+ api_key_configured=True,
+ mistral_model="mistral-small-latest",
+ provider="mistral",
+ )
+
+
+def _triage(band: str = "moderate_risk", score: float = 0.48) -> TriageResult:
+ return TriageResult(
+ band=band,
+ score=score,
+ label=band.replace("_", " ").title(),
+ summary="Routine follow-up would be reasonable.",
+ disclaimer="Screening only.",
+ )
+
+
+SERVICE = GuidanceService()
+
+
+def _make_mistral_service(*, api_key_configured: bool = True) -> GuidanceService:
+ service = GuidanceService.__new__(GuidanceService)
+ service.mistral_enabled = True
+ service.mistral_model = "mistral-small-latest"
+ service.guidance_timeout = 6.0
+ service.guidance_max_tokens = 256
+ service.api_key_configured = api_key_configured
+ service._fallback_reason = None if api_key_configured else "Mistral API key is missing."
+ service._last_provider_error = None
+ service._response_cache = OrderedDict()
+ service._response_cache_size = 64
+ return service
+
+
+class TestParseGuidanceResponse:
+ BASE_KWARGS = dict(
+ source="mistral",
+ model_used="mistral-small-latest",
+ provider_used="mistral",
+ )
+
+ def test_accepts_code_fenced_json(self) -> None:
+ raw = """```json
+ {
+ "explanation": "A grounded summary.",
+ "urgency_guidance": "Book a routine follow-up.",
+ "food_advice": "Add iron-rich foods.",
+ "next_steps": ["Retake if symptoms change", "Plan a CBC test"]
+ }
+ ```"""
+ result = SERVICE._parse_guidance_response(raw, **self.BASE_KWARGS)
+
+ assert result.source == "mistral"
+ assert result.model_used == "mistral-small-latest"
+ assert result.provider_used == "mistral"
+ assert len(result.next_steps) == 2
+ assert result.next_steps[0] == "Retake if symptoms change"
+
+ def test_accepts_plain_json(self) -> None:
+ raw = """{
+ "explanation": "Looks fine.",
+ "urgency_guidance": "No urgency.",
+ "food_advice": "Eat spinach.",
+ "next_steps": ["Follow up in 3 months"]
+ }"""
+ result = SERVICE._parse_guidance_response(raw, **self.BASE_KWARGS)
+ assert result.next_steps == ["Follow up in 3 months"]
+
+ def test_accepts_python_literal_with_single_quotes(self) -> None:
+ raw = (
+ "{'explanation': 'Summary', 'urgency_guidance': 'Monitor closely', "
+ "'food_advice': 'Eat lentils', 'next_steps': 'Book a clinic visit'}"
+ )
+ result = SERVICE._parse_guidance_response(raw, **self.BASE_KWARGS)
+ assert result.next_steps == ["Book a clinic visit"]
+
+ def test_coerces_string_next_steps_to_list(self) -> None:
+ raw = """{
+ "explanation": "Mild signal.",
+ "urgency_guidance": "Monitor symptoms.",
+ "food_advice": "Iron-rich diet.",
+ "next_steps": "See a doctor soon"
+ }"""
+ result = SERVICE._parse_guidance_response(raw, **self.BASE_KWARGS)
+ assert isinstance(result.next_steps, list)
+ assert len(result.next_steps) == 1
+
+ def test_allows_explicit_non_diagnostic_disclaimer(self) -> None:
+ raw = """{
+ "explanation": "This is screening guidance, not a diagnosis. The current result suggests some concern.",
+ "urgency_guidance": "Follow up with a clinician if symptoms continue.",
+ "food_advice": "Add iron-rich foods like lentils and spinach.",
+ "next_steps": ["Book a routine clinic visit", "Monitor symptoms and retake if they change"]
+ }"""
+ result = SERVICE._parse_guidance_response(raw, **self.BASE_KWARGS)
+ assert result.source == "mistral"
+ assert "not a diagnosis" in result.explanation.lower()
+
+
+UNSAFE_CLAIMS = [
+ "This definitely confirms anemia.",
+ "You have anemia based on this scan.",
+ "The result diagnoses iron deficiency.",
+ "This scan proves you are anaemic.",
+]
+
+
+@pytest.mark.parametrize("unsafe_explanation", UNSAFE_CLAIMS)
+def test_parse_guidance_rejects_unsafe_claims(unsafe_explanation: str) -> None:
+ raw = f"""{{
+ "explanation": "{unsafe_explanation}",
+ "urgency_guidance": "See a clinician.",
+ "food_advice": "Eat iron-rich foods.",
+ "next_steps": ["Book a CBC test"]
+ }}"""
+ with pytest.raises(ValueError, match="[Uu]nsafe|diagnostic|claim"):
+ SERVICE._parse_guidance_response(
+ raw,
+ source="mistral",
+ model_used="mistral-small-latest",
+ provider_used="mistral",
+ )
+
+
+class TestFallbackGuidance:
+ @staticmethod
+ def _fallback_result(band: str, symptoms: SymptomInput) -> GuidanceResult:
+ predicted_hemoglobin = {
+ "low_risk": 13.2,
+ "moderate_risk": 10.1,
+ "high_concern": 7.8,
+ "uncertain_retake_needed": None,
+ }[band]
+ confidence = None if predicted_hemoglobin is None else 0.72
+ return SERVICE.generate_smart_fallback(
+ band,
+ predicted_hemoglobin,
+ confidence,
+ symptoms,
+ "India",
+ )
+
+ @pytest.mark.parametrize("band", ["low_risk", "moderate_risk", "high_concern", "uncertain_retake_needed"])
+ def test_fallback_has_next_steps_for_every_band(self, band: str) -> None:
+ result = self._fallback_result(band, SymptomInput(fatigue=True))
+ assert len(result.next_steps) >= 1
+
+ def test_fallback_marks_source_correctly(self) -> None:
+ result = self._fallback_result(
+ "moderate_risk",
+ SymptomInput(fatigue=True, poor_diet_low_iron=True),
+ )
+ assert result.source == "fallback"
+ assert result.model_used is None
+ assert result.provider_used is None
+
+ def test_fallback_result_validates_as_guidance_result(self) -> None:
+ result = self._fallback_result(
+ "high_concern",
+ SymptomInput(fatigue=True, dizziness=True),
+ )
+ GuidanceResult.model_validate(result.model_dump())
+
+ def test_fallback_explanation_not_empty(self) -> None:
+ result = self._fallback_result("moderate_risk", SymptomInput())
+ assert len(result.explanation) > 20
+
+ def test_fallback_language_is_non_diagnostic(self) -> None:
+ result = SERVICE.generate_smart_fallback(
+ "moderate_risk",
+ 10.2,
+ 0.74,
+ SymptomInput(fatigue=True),
+ "India",
+ )
+ combined = " ".join([result.explanation, result.urgency_guidance, result.food_advice, *result.next_steps]).lower()
+ assert "diagnos" not in combined
+ assert "you have anemia" not in combined
+
+ def test_runtime_status_reports_fallback_reason_without_api_key(self) -> None:
+ service = _make_mistral_service(api_key_configured=False)
+
+ status = service.runtime_status()
+
+ assert status.active_strategy == "fallback"
+ assert status.mistral_enabled is True
+ assert status.client_ready is False
+ assert status.api_key_configured is False
+ assert "missing" in (status.fallback_reason or "").lower()
+
+ def test_generate_skips_llm_for_uncertain_retake_cases(self) -> None:
+ service = _make_mistral_service()
+ service._call_mistral_api = lambda *args, **kwargs: (_ for _ in ()).throw(AssertionError("Mistral should be skipped")) # type: ignore[method-assign]
+
+ result = service.generate(
+ triage=_triage(band="uncertain_retake_needed"),
+ symptoms=SymptomInput(fatigue=True),
+ prediction=None,
+ )
+
+ assert result.source == "fallback"
+ assert "not a clear result" in result.explanation.lower()
+
+ def test_generate_smart_fallback_personalizes_region_and_symptoms(self) -> None:
+ result = SERVICE.generate_smart_fallback(
+ "moderate_risk",
+ 10.1,
+ 0.72,
+ SymptomInput(
+ fatigue=True,
+ shortness_of_breath=True,
+ heavy_menstrual_bleeding=True,
+ ),
+ "India",
+ )
+
+ assert result.source == "fallback"
+ assert "palak" in result.food_advice.lower()
+ assert "tea or coffee" in result.food_advice.lower()
+ assert "Avoid strenuous activity until reviewed by a doctor." in result.next_steps
+ assert "Discuss menstrual blood loss with your doctor as a likely contributing factor." in result.next_steps
+
+ def test_generate_mistral_returns_smart_fallback_when_provider_fails(self) -> None:
+ service = _make_mistral_service()
+ service._call_mistral_api = lambda *args, **kwargs: (_ for _ in ()).throw( # type: ignore[method-assign]
+ RuntimeError("429 rate limit")
+ )
+
+ result = service.generate(
+ triage=_triage(band="high_concern", score=0.82),
+ symptoms=SymptomInput(fatigue=True, shortness_of_breath=True),
+ prediction=PredictionResult(
+ anemia_risk=0.87,
+ predicted_hemoglobin=7.6,
+ confidence=0.84,
+ uncertainty=0.16,
+ reliability_flag="high",
+ screening_label="anemia_likely",
+ screening_text="The screening model estimates a lower-than-expected hemoglobin trend from the eye image.",
+ model_source="efficientnet-b0-ft",
+ ),
+ region="India",
+ )
+
+ assert result.source == "fallback"
+ assert "24 to 48 hours" in result.urgency_guidance
+ assert "rate limit" in (service._last_provider_error or "").lower()
+
+
+def test_guidance_result_requires_model_and_provider_when_mistral() -> None:
+ with pytest.raises(Exception):
+ GuidanceResult(
+ source="mistral",
+ explanation="Looks fine.",
+ urgency_guidance="No urgency.",
+ food_advice="Eat well.",
+ next_steps=["Follow up"],
+ )
+
+
+def test_guidance_result_allows_null_model_for_fallback() -> None:
+ result = GuidanceResult(
+ source="fallback",
+ model_used=None,
+ provider_used=None,
+ explanation="Looks fine.",
+ urgency_guidance="No urgency.",
+ food_advice="Eat well.",
+ next_steps=["Follow up"],
+ )
+ assert result.source == "fallback"
+
+
+class TestRuntimeStatus(unittest.TestCase):
+ def test_enriches_model_metadata_from_training_report(self) -> None:
+ status = build_runtime_status(_PredictorStub(), _GuidanceStub())
+
+ self.assertEqual(status.api_status, "ok")
+ self.assertEqual(status.guidance.active_strategy, "mistral")
+ self.assertIn(
+ status.model.primary_model,
+ {"archive-fusion-v2", "efficientnet-b0-ft", "archive-primary-v3", "archive-evidence-fusion-v4"},
+ )
+ self.assertGreaterEqual(status.model.record_count or 0, 200)
+ self.assertGreater(status.model.validation_f1 or 0.0, 0.6)
+
+ def test_guidance_mistral_fields_propagated(self) -> None:
+ status = build_runtime_status(_PredictorStub(), _GuidanceStub())
+ self.assertEqual(status.guidance.mistral_model, "mistral-small-latest")
+ self.assertEqual(status.guidance.provider, "mistral")
+ self.assertTrue(status.guidance.mistral_enabled)
+ self.assertTrue(status.guidance.client_ready)
+
+
+def test_settings_accept_hf_environment_variable_names(monkeypatch: pytest.MonkeyPatch) -> None:
+ monkeypatch.setenv("ANEMIALENS_HF_API_KEY", "test-token")
+ monkeypatch.setenv("ANEMIALENS_QWEN_MODEL", "Qwen/Qwen2.5-7B-Instruct")
+ monkeypatch.setenv("ANEMIALENS_QWEN_ENABLED", "true")
+ monkeypatch.setenv("ANEMIALENS_HF_PROVIDER", "hf-inference")
+
+ settings_obj = Settings(_env_file=None)
+
+ assert settings_obj.hf_api_key == "test-token"
+ assert settings_obj.qwen_model == "Qwen/Qwen2.5-7B-Instruct"
+ assert settings_obj.qwen_enabled is True
+ assert settings_obj.hf_provider == "hf-inference"
+
+
+def test_generate_uses_cached_mistral_result_for_same_payload() -> None:
+ service = _make_mistral_service()
+
+ calls = {"count": 0}
+
+ def _fake_generate_mistral(*args, **kwargs) -> GuidanceResult:
+ calls["count"] += 1
+ return GuidanceResult(
+ source="mistral",
+ model_used="mistral-small-latest",
+ provider_used="mistral",
+ explanation="Screening suggests a mild low-hemoglobin signal.",
+ urgency_guidance="Arrange a routine check if symptoms continue.",
+ food_advice="Eat lentils, beans, spinach, and vitamin C-rich fruit.",
+ next_steps=["Repeat the scan if symptoms change", "Plan a clinic test"],
+ )
+
+ service._generate_mistral = _fake_generate_mistral # type: ignore[method-assign]
+
+ prediction = PredictionResult(
+ anemia_risk=0.58,
+ predicted_hemoglobin=10.9,
+ confidence=0.78,
+ uncertainty=0.18,
+ reliability_flag="high",
+ screening_label="anemia_likely",
+ screening_text="The screening model estimates a lower-than-expected hemoglobin trend from the eye image.",
+ model_source="efficientnet-b0-ft",
+ )
+
+ first = service.generate(_triage(), SymptomInput(fatigue=True), prediction, "English", "India")
+ second = service.generate(_triage(), SymptomInput(fatigue=True), prediction, "English", "India")
+
+ assert first.source == "mistral"
+ assert second.source == "mistral"
+ assert calls["count"] == 1
+
+
+def test_summarize_provider_error_flags_provider_permission_problem() -> None:
+ service = GuidanceService.__new__(GuidanceService)
+ message = service._summarize_error(
+ RuntimeError("403 Forbidden: This authentication method does not have sufficient permissions to call Inference Providers")
+ )
+ assert "inference providers" in message.lower()
+
+
+def test_mistral_generates_response() -> None:
+ service = _make_mistral_service()
+
+ def _fake_generate_mistral(*args, **kwargs) -> GuidanceResult:
+ return GuidanceResult(
+ source="mistral",
+ model_used="mistral-small-latest",
+ provider_used="mistral",
+ explanation="Screening suggests a mild low-hemoglobin signal.",
+ urgency_guidance="Arrange a routine check if symptoms continue.",
+ food_advice="Eat lentils, beans, spinach, and vitamin C-rich fruit.",
+ next_steps=["Repeat the scan if symptoms change", "Plan a clinic test"],
+ )
+
+ service._generate_mistral = _fake_generate_mistral # type: ignore[method-assign]
+
+ result = service.generate(
+ triage=_triage(),
+ symptoms=SymptomInput(fatigue=True),
+ prediction=PredictionResult(
+ anemia_risk=0.87,
+ predicted_hemoglobin=7.6,
+ confidence=0.84,
+ uncertainty=0.16,
+ reliability_flag="high",
+ screening_label="anemia_likely",
+ screening_text="The screening model estimates a lower-than-expected hemoglobin trend from the eye image.",
+ model_source="efficientnet-b0-ft",
+ ),
+ )
+
+ assert result.source == "mistral"
diff --git a/backend/tests/test_handoff.py b/backend/tests/test_handoff.py
new file mode 100644
index 0000000000000000000000000000000000000000..9a1de1fc9fd6b43322497d196570dc04a0624727
--- /dev/null
+++ b/backend/tests/test_handoff.py
@@ -0,0 +1,95 @@
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.schemas import GuidanceResult, PredictionResult, QualityAssessment, SymptomInput, TriageResult
+from app.services.handoff import HandoffSummaryService
+
+
+def test_handoff_summary_includes_prediction_symptoms_and_next_steps() -> None:
+ service = HandoffSummaryService()
+ summary = service.build(
+ QualityAssessment(
+ passed=True,
+ blur_score=120.0,
+ brightness_score=0.24,
+ contrast_score=0.18,
+ framing_score=1.8,
+ issues=[],
+ ),
+ PredictionResult(
+ anemia_risk=0.66,
+ predicted_hemoglobin=10.4,
+ confidence=0.81,
+ uncertainty=0.19,
+ reliability_flag="high",
+ screening_label="anemia_likely",
+ screening_text="The screening model estimates a lower-than-expected hemoglobin trend from the eye image.",
+ model_source="efficientnet-b0-ft",
+ ),
+ TriageResult(
+ band="moderate_risk",
+ score=0.54,
+ label="Moderate risk",
+ summary="This screening shows some concern.",
+ disclaimer="Screening only.",
+ ),
+ GuidanceResult(
+ source="fallback",
+ model_used=None,
+ provider_used=None,
+ explanation="Mild to moderate anemia detected.",
+ urgency_guidance="See a doctor within 1-2 weeks.",
+ food_advice="Eat iron-rich foods.",
+ next_steps=["Book a clinic visit this week", "Start iron-rich diet immediately"],
+ ),
+ SymptomInput(fatigue=True, dizziness=True, poor_diet_low_iron=True),
+ language="English",
+ region="India",
+ )
+
+ assert "Moderate risk" in summary.headline
+ assert any("Estimated hemoglobin" in point for point in summary.key_points)
+ assert any("fatigue" in point for point in summary.key_points)
+ assert summary.next_steps[0] == "Book a clinic visit this week"
+ assert "AnemiaLens screening handoff" in summary.share_text
+
+
+def test_handoff_summary_handles_retake_case_without_prediction() -> None:
+ service = HandoffSummaryService()
+ summary = service.build(
+ QualityAssessment(
+ passed=False,
+ blur_score=40.0,
+ brightness_score=0.05,
+ contrast_score=0.02,
+ framing_score=0.4,
+ issues=[],
+ ),
+ None,
+ TriageResult(
+ band="uncertain_retake_needed",
+ score=0.24,
+ label="Uncertain, retake needed",
+ summary="Retake the image.",
+ disclaimer="Screening only.",
+ ),
+ GuidanceResult(
+ source="fallback",
+ model_used=None,
+ provider_used=None,
+ explanation="Image signal was not strong enough.",
+ urgency_guidance="Retake the scan in better lighting.",
+ food_advice="No food advice until a valid screening is available.",
+ next_steps=["Retake eye image in bright natural light"],
+ ),
+ SymptomInput(),
+ )
+
+ assert summary.urgency_label == "Retake image"
+ assert "quality" in summary.key_points[0].lower()
+ assert "Retake" in summary.share_text
diff --git a/backend/tests/test_offline_ml.py b/backend/tests/test_offline_ml.py
new file mode 100644
index 0000000000000000000000000000000000000000..9332c2719d24545276a076bcdac2351a505b5b80
--- /dev/null
+++ b/backend/tests/test_offline_ml.py
@@ -0,0 +1,260 @@
+"""
+Offline ML integration tests.
+
+These tests require the trained model artefact and the anemia dataset to be
+present on disk โ they are intentionally skipped in environments where those
+files are absent (CI without model artefacts, fresh developer checkouts).
+
+Use::
+
+ pytest tests/test_offline_ml.py -v --no-header
+
+to run locally after training.
+"""
+
+from __future__ import annotations
+
+import json
+import sys
+from pathlib import Path
+
+import pytest
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+# ---------------------------------------------------------------------------
+# Optional dependency guards
+# ---------------------------------------------------------------------------
+
+def _pillow_available() -> bool:
+ try:
+ import PIL # noqa: F401
+ return True
+ except ImportError:
+ return False
+
+
+def _model_artefact_present() -> bool:
+ return (ROOT / "backend" / "models" / "archive_screening_model.joblib").exists()
+
+
+def _efficientnet_artefact_present() -> bool:
+ return (ROOT / "backend" / "models" / "efficientnet_anemia.pth").exists()
+
+
+def _dataset_present() -> bool:
+ return (ROOT / "archive" / "dataset anemia").exists()
+
+
+requires_pillow = pytest.mark.skipif(not _pillow_available(), reason="Pillow not installed")
+requires_model = pytest.mark.skipif(not _model_artefact_present(), reason="Model artefact not found")
+requires_efficientnet = pytest.mark.skipif(
+ not _efficientnet_artefact_present(),
+ reason="EfficientNet artefact not found",
+)
+requires_dataset = pytest.mark.skipif(not _dataset_present(), reason="Dataset not found")
+
+
+# ---------------------------------------------------------------------------
+# Feature extraction
+# ---------------------------------------------------------------------------
+
+@requires_pillow
+def test_feature_extraction_returns_expected_feature_set() -> None:
+ from PIL import Image
+ from app.ml.features import FEATURE_NAMES, extract_eye_features
+
+ image = Image.new("RGB", (320, 180), color=(180, 120, 115))
+ features = extract_eye_features(image)
+
+ assert set(FEATURE_NAMES) == set(features), (
+ f"Feature mismatch.\n"
+ f" Extra in output : {set(features) - set(FEATURE_NAMES)}\n"
+ f" Missing from output: {set(FEATURE_NAMES) - set(features)}"
+ )
+ assert features["brightness"] > 0.0, "brightness feature should be positive for a non-black image"
+
+
+@requires_pillow
+def test_feature_extraction_values_in_valid_ranges() -> None:
+ from PIL import Image
+ from app.ml.features import extract_eye_features
+
+ image = Image.new("RGB", (320, 180), color=(180, 120, 115))
+ features = extract_eye_features(image)
+
+ for name, value in features.items():
+ assert isinstance(value, (int, float)), f"Feature '{name}' is not numeric: {value!r}"
+ assert not (value != value), f"Feature '{name}' is NaN" # NaN check
+
+
+# ---------------------------------------------------------------------------
+# Model loading and prediction
+# ---------------------------------------------------------------------------
+
+MODEL_PATH = ROOT / "backend" / "models" / "archive_screening_model.joblib"
+DATASET_PATH = ROOT / "archive" / "dataset anemia"
+
+
+@requires_pillow
+@requires_model
+def test_archive_model_predicts_valid_probability_ranges() -> None:
+ from PIL import Image
+ from app.ml.archive_model import ARCHIVE_VERSION, load_archive_model, predict_with_archive_model
+ from app.ml.features import extract_eye_features, load_image_path
+ from app.services.conjunctiva_roi import ConjunctivaRoiExtractor
+
+ artifact = load_archive_model(MODEL_PATH)
+
+ sample = next(DATASET_PATH.glob("*/*/*.jpg"), None)
+ if sample is None:
+ pytest.skip("No JPEG images found in dataset directory")
+
+ roi = ConjunctivaRoiExtractor().extract(load_image_path(sample)).image
+ prediction = predict_with_archive_model(
+ artifact,
+ extract_eye_features(roi),
+ source_hint="roi_original",
+ )
+
+ assert artifact["version"] == ARCHIVE_VERSION, (
+ f"Artefact version mismatch: {artifact['version']} != {ARCHIVE_VERSION}"
+ )
+ assert 0.0 <= prediction["anemia_risk"] <= 1.0, "anemia_risk out of [0, 1]"
+ assert 0.0 <= prediction["uncertainty"] <= 1.0, "uncertainty out of [0, 1]"
+ assert prediction["predicted_hemoglobin"] > 5.0, (
+ f"predicted_hemoglobin={prediction['predicted_hemoglobin']} is implausibly low"
+ )
+
+
+@requires_pillow
+@requires_model
+def test_archive_model_prediction_fields_all_present() -> None:
+ """Regression test: ensure no required output fields are accidentally dropped."""
+ from PIL import Image
+ from app.ml.archive_model import load_archive_model, predict_with_archive_model
+ from app.ml.features import extract_eye_features
+
+ artifact = load_archive_model(MODEL_PATH)
+ image = Image.new("RGB", (320, 180), color=(180, 120, 115))
+ prediction = predict_with_archive_model(
+ artifact,
+ extract_eye_features(image),
+ source_hint="synthetic",
+ )
+
+ required_keys = {"anemia_risk", "uncertainty", "predicted_hemoglobin"}
+ missing = required_keys - set(prediction.keys())
+ assert not missing, f"Prediction is missing keys: {missing}"
+
+
+# ---------------------------------------------------------------------------
+# Training report schema
+# ---------------------------------------------------------------------------
+
+REPORT_PATH = ROOT / "backend" / "models" / "training_report.json"
+EFFICIENTNET_PATH = ROOT / "backend" / "models" / "efficientnet_anemia.pth"
+
+
+@pytest.mark.skipif(not REPORT_PATH.exists(), reason="training_report.json not found")
+def test_training_report_matches_archive_model_version() -> None:
+ from app.ml.archive_model import ARCHIVE_VERSION
+ from app.ml.efficientnet_model import EFFICIENTNET_VERSION
+
+ report = json.loads(REPORT_PATH.read_text(encoding="utf-8"))
+ assert report["primary_model"] in {ARCHIVE_VERSION, EFFICIENTNET_VERSION}, (
+ f"Unexpected primary_model={report['primary_model']!r}"
+ )
+
+
+@pytest.mark.skipif(not REPORT_PATH.exists(), reason="training_report.json not found")
+def test_training_report_minimum_dataset_size() -> None:
+ report = json.loads(REPORT_PATH.read_text(encoding="utf-8"))
+ assert report["subject_count"] >= 200, "Dataset has too few subjects to be reliable"
+ assert report["record_count"] >= report["subject_count"], (
+ "record_count must be >= subject_count"
+ )
+
+
+@pytest.mark.skipif(not REPORT_PATH.exists(), reason="training_report.json not found")
+def test_training_report_metrics_meet_minimum_bar() -> None:
+ report = json.loads(REPORT_PATH.read_text(encoding="utf-8"))
+ metrics = report["metrics"]
+
+ assert metrics["split_strategy"] in {"group-shuffle-repeat", "group-shuffle-balance-select"}, (
+ f"Unexpected split_strategy {metrics['split_strategy']!r}"
+ )
+ assert metrics["validation_size"] > 30, "Validation set is too small"
+ assert metrics["accuracy"] > 0.6, f"accuracy={metrics['accuracy']:.3f} below threshold"
+ assert metrics["f1"] > 0.45, f"f1={metrics['f1']:.3f} below threshold"
+
+
+@pytest.mark.skipif(not REPORT_PATH.exists(), reason="training_report.json not found")
+def test_training_report_selected_mode_is_valid() -> None:
+ report = json.loads(REPORT_PATH.read_text(encoding="utf-8"))
+ valid_modes = {"roi_primary", "hybrid_dual", "efficientnet_hybrid_dual"}
+ assert report["selected_mode"] in valid_modes, (
+ f"selected_mode={report['selected_mode']!r} not in {valid_modes}"
+ )
+
+
+@requires_pillow
+@requires_efficientnet
+def test_efficientnet_checkpoint_loads_and_predicts() -> None:
+ from PIL import Image
+ from app.ml.efficientnet_model import (
+ EFFICIENTNET_VERSION,
+ load_efficientnet_checkpoint,
+ predict_with_efficientnet_model,
+ )
+
+ bundle = load_efficientnet_checkpoint(EFFICIENTNET_PATH)
+ prediction = predict_with_efficientnet_model(
+ bundle,
+ Image.new("RGB", (320, 180), color=(180, 120, 115)),
+ mc_passes=2,
+ )
+
+ assert bundle["version"] == EFFICIENTNET_VERSION
+ assert 0.0 <= prediction["anemia_risk"] <= 1.0
+ assert 0.0 <= prediction["uncertainty"] <= 1.0
+
+
+@requires_pillow
+def test_single_pass_efficientnet_prediction_is_deterministic() -> None:
+ from PIL import Image
+ import torch
+ from torch import nn
+
+ from app.ml.efficientnet_model import predict_with_efficientnet_model
+
+ class TinyDropoutModel(nn.Module):
+ def __init__(self) -> None:
+ super().__init__()
+ self.flatten = nn.Flatten()
+ self.dropout = nn.Dropout(p=0.95)
+ self.linear = nn.Linear(3 * 4 * 4, 2)
+
+ def forward(self, tensor: torch.Tensor) -> torch.Tensor:
+ x = self.flatten(tensor)
+ x = self.dropout(x)
+ return self.linear(x)
+
+ torch.manual_seed(7)
+ model = TinyDropoutModel()
+ bundle = {
+ "model": model,
+ "device": torch.device("cpu"),
+ "transform": lambda image: torch.ones((3, 4, 4), dtype=torch.float32),
+ "hb_mean": 12.0,
+ "hb_std": 1.0,
+ "decision_threshold": 0.5,
+ }
+ image = Image.new("RGB", (16, 16), color=(180, 120, 115))
+
+ first = predict_with_efficientnet_model(bundle, image, mc_passes=1)
+ second = predict_with_efficientnet_model(bundle, image, mc_passes=1)
+
+ assert first["anemia_risk"] == second["anemia_risk"]
+ assert first["predicted_hemoglobin"] == second["predicted_hemoglobin"]
diff --git a/backend/tests/test_patient_case.py b/backend/tests/test_patient_case.py
new file mode 100644
index 0000000000000000000000000000000000000000..ceabf285dadd955b8d8553c933c1ac5fbc79983c
--- /dev/null
+++ b/backend/tests/test_patient_case.py
@@ -0,0 +1,142 @@
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.schemas import (
+ GuidanceResult,
+ PatientProfileInput,
+ PredictionResult,
+ QualityAssessment,
+ QualityIssue,
+ SymptomInput,
+ TriageResult,
+)
+from app.services.patient_case import PatientCaseService
+
+
+def _quality(*, passed: bool = True, warnings: list[QualityIssue] | None = None, blockers: list[QualityIssue] | None = None) -> QualityAssessment:
+ return QualityAssessment(
+ passed=passed,
+ blur_score=120.0,
+ brightness_score=0.44,
+ contrast_score=0.22,
+ framing_score=1.2,
+ lighting_score=0.74,
+ lighting_condition="balanced",
+ lighting_summary="Even clinical-style lighting.",
+ glare_risk=0.1,
+ shadow_risk=0.15,
+ issues=[*(warnings or []), *(blockers or [])],
+ )
+
+
+def test_build_profile_generates_patient_id_and_summary() -> None:
+ service = PatientCaseService()
+ symptoms = SymptomInput(fatigue=True, dizziness=True)
+
+ profile = service.build_profile(
+ "abc123ef",
+ PatientProfileInput(age=17, sex="female", diet_type="vegetarian"),
+ symptoms,
+ )
+
+ assert profile.patient_id == "ANM-C123EF"
+ assert profile.reported_symptoms == ["Fatigue", "Dizziness"]
+ assert "17-year-old female" in profile.summary.lower()
+
+
+def test_build_workflow_stages_marks_quality_block() -> None:
+ service = PatientCaseService()
+ blocked_quality = _quality(
+ passed=False,
+ blockers=[
+ QualityIssue(
+ code="eye_not_visible",
+ severity="blocking",
+ title="Eye region not visible",
+ message="Pull down the lower eyelid and try again.",
+ )
+ ],
+ )
+
+ stages = service.build_workflow_stages(
+ blocked_quality,
+ None,
+ TriageResult(
+ band="uncertain_retake_needed",
+ score=0.22,
+ label="Uncertain, retake needed",
+ summary="Retake needed before screening interpretation.",
+ disclaimer="Screening only.",
+ ),
+ GuidanceResult(
+ source="fallback",
+ explanation="Fallback guidance.",
+ urgency_guidance="Retake the image first.",
+ food_advice="Eat iron-rich foods.",
+ next_steps=["Retake image", "Repeat screening"],
+ ),
+ SymptomInput(),
+ )
+
+ assert stages[0].status == "blocked"
+ assert stages[1].status == "blocked"
+
+
+def test_build_structured_case_contains_quality_and_recommendation() -> None:
+ service = PatientCaseService()
+ quality = _quality(
+ warnings=[
+ QualityIssue(
+ code="poor_lighting",
+ severity="warning",
+ title="Dim lighting",
+ message="Move into brighter light.",
+ )
+ ]
+ )
+ prediction = PredictionResult(
+ anemia_risk=0.64,
+ predicted_hemoglobin=11.9,
+ confidence=0.72,
+ uncertainty=0.18,
+ reliability_flag="medium",
+ screening_label="anemia_likely",
+ screening_text="Moderate anemia-like screening signal.",
+ model_source="archive-evidence-fusion-v4",
+ )
+ triage = TriageResult(
+ band="moderate_risk",
+ score=0.58,
+ label="Moderate risk",
+ summary="This screening shows some concern.",
+ disclaimer="Screening only.",
+ )
+ guidance = GuidanceResult(
+ source="fallback",
+ explanation="Result suggests follow-up.",
+ urgency_guidance="Arrange a CBC in 1-2 weeks.",
+ food_advice="Add beans and greens.",
+ next_steps=["Arrange CBC", "See clinician"],
+ )
+ symptoms = SymptomInput(fatigue=True, poor_diet_low_iron=True)
+ profile = service.build_profile("abc123ef", PatientProfileInput(age=21, sex="female", diet_type="vegetarian"), symptoms)
+
+ case_record = service.build_structured_case(
+ "abc123ef",
+ profile,
+ quality,
+ prediction,
+ triage,
+ guidance,
+ symptoms,
+ )
+
+ assert case_record.case_id == "CASE-C123EF"
+ assert case_record.image_quality.status == "warning"
+ assert case_record.screening_result.risk_level == "moderate_risk"
+ assert case_record.recommendation == "Arrange CBC"
diff --git a/backend/tests/test_prediction.py b/backend/tests/test_prediction.py
new file mode 100644
index 0000000000000000000000000000000000000000..823957e781785173f398cf2124d34ca36ea845ec
--- /dev/null
+++ b/backend/tests/test_prediction.py
@@ -0,0 +1,791 @@
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+from PIL import Image
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.services.prediction import ScreeningPredictor
+from app.services import prediction as prediction_module
+from app.schemas import PredictionResult, QualityAssessment
+
+
+def test_predictor_init_is_lazy(monkeypatch, tmp_path) -> None:
+ archive_path = tmp_path / "archive.joblib"
+ efficientnet_path = tmp_path / "efficientnet.pth"
+ archive_path.write_bytes(b"archive")
+ efficientnet_path.write_bytes(b"efficientnet")
+
+ calls: list[str] = []
+
+ monkeypatch.setattr(prediction_module, "DEFAULT_ARCHIVE_MODEL_PATH", archive_path)
+ monkeypatch.setattr(
+ prediction_module,
+ "DEFAULT_EFFICIENTNET_MODEL_PATH",
+ efficientnet_path,
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_load_archive_model_artifact",
+ lambda path: calls.append(f"archive:{path.name}") or {"artifact": True},
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_load_efficientnet_checkpoint_bundle",
+ lambda path: calls.append(f"efficientnet:{path.name}") or {"bundle": True},
+ )
+
+ predictor = ScreeningPredictor()
+
+ assert calls == []
+ assert predictor.archive_model is None
+ assert predictor.efficientnet_bundle is None
+ assert predictor.is_ready() is True
+
+
+def test_predictor_preload_loads_models_once(monkeypatch, tmp_path) -> None:
+ archive_path = tmp_path / "archive.joblib"
+ efficientnet_path = tmp_path / "efficientnet.pth"
+ archive_path.write_bytes(b"archive")
+ efficientnet_path.write_bytes(b"efficientnet")
+
+ calls: list[str] = []
+
+ monkeypatch.setattr(prediction_module, "DEFAULT_ARCHIVE_MODEL_PATH", archive_path)
+ monkeypatch.setattr(
+ prediction_module,
+ "DEFAULT_EFFICIENTNET_MODEL_PATH",
+ efficientnet_path,
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_load_archive_model_artifact",
+ lambda path: calls.append(f"archive:{path.name}") or {"artifact": True},
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_load_efficientnet_checkpoint_bundle",
+ lambda path: calls.append(f"efficientnet:{path.name}") or {"bundle": True},
+ )
+ monkeypatch.setattr(prediction_module.settings, "enable_efficientnet_fallback", True)
+
+ predictor = ScreeningPredictor()
+ predictor.preload()
+ predictor.preload()
+
+ assert calls == ["archive:archive.joblib", "efficientnet:efficientnet.pth"]
+ assert predictor.archive_model == {"artifact": True}
+ assert predictor.efficientnet_bundle == {"bundle": True}
+
+
+def test_dark_signal_guardrail_triggers_on_dark_positive_with_near_normal_hb() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ triggered = predictor._dark_signal_guardrail(
+ risk=0.71,
+ predicted_hemoglobin=11.9,
+ feature_map={
+ "brightness": 0.11,
+ "hist_bright": 0.03,
+ "hist_highlight": 0.0,
+ },
+ threshold=0.68,
+ )
+
+ assert triggered is True
+
+
+def test_dark_signal_guardrail_skips_clear_low_hb_cases() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ triggered = predictor._dark_signal_guardrail(
+ risk=0.74,
+ predicted_hemoglobin=10.4,
+ feature_map={
+ "brightness": 0.11,
+ "hist_bright": 0.03,
+ "hist_highlight": 0.0,
+ },
+ threshold=0.68,
+ )
+
+ assert triggered is False
+
+
+def test_screening_decision_returns_uncertain_when_guardrail_triggers() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.72,
+ uncertainty=0.19,
+ threshold=0.68,
+ predicted_hemoglobin=12.4,
+ signal_guardrail_triggered=True,
+ )
+
+ assert label == "uncertain"
+ assert "dark" in text.lower()
+
+
+def test_screening_decision_rescues_high_suspicion_positive() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.52,
+ uncertainty=0.6,
+ threshold=0.435,
+ predicted_hemoglobin=12.1,
+ )
+
+ assert label == "anemia_likely"
+ assert "likely anemia" in text.lower()
+
+
+def test_screening_decision_rescues_borderline_high_suspicion_positive() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.428,
+ uncertainty=0.56,
+ threshold=0.435,
+ predicted_hemoglobin=12.35,
+ )
+
+ assert label == "anemia_likely"
+ assert "likely anemia" in text.lower()
+
+
+def test_screening_decision_downgrades_mild_positive_near_normal_hb() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.56,
+ uncertainty=0.54,
+ threshold=0.435,
+ predicted_hemoglobin=12.3,
+ )
+
+ assert label == "uncertain"
+ assert "near normal" in text.lower()
+
+
+def test_screening_decision_rescues_clarity_exception_borderline_positive() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.392,
+ uncertainty=0.62,
+ threshold=0.435,
+ predicted_hemoglobin=12.24,
+ )
+
+ assert label == "anemia_likely"
+ assert "likely anemia" in text.lower()
+
+
+def test_screening_decision_high_threshold_low_reliability_positive_requires_extra_evidence() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.708,
+ uncertainty=0.582,
+ threshold=0.65,
+ predicted_hemoglobin=11.62,
+ )
+
+ assert label == "uncertain"
+ assert "confidence level" in text.lower()
+
+
+def test_screening_decision_keeps_strong_low_reliability_positive_with_clear_margin() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.761,
+ uncertainty=0.617,
+ threshold=0.65,
+ predicted_hemoglobin=11.46,
+ )
+
+ assert label == "anemia_likely"
+ assert "likely anemia" in text.lower()
+
+
+def test_screening_decision_keeps_overwhelming_positive_signal_likely_even_when_noisy() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.768,
+ uncertainty=0.804,
+ threshold=0.495,
+ predicted_hemoglobin=11.72,
+ )
+
+ assert label == "anemia_likely"
+ assert "still be treated as likely" in text.lower()
+
+
+def test_screening_decision_keeps_signal_only_positive_likely_when_hb_missing() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.648,
+ uncertainty=0.88,
+ threshold=0.495,
+ predicted_hemoglobin=None,
+ )
+
+ assert label == "anemia_likely"
+ assert "image-only anemia signal" in text.lower()
+
+
+def test_screening_decision_skips_below_threshold_rescue_for_strict_runtime_threshold() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.646,
+ uncertainty=0.602,
+ threshold=0.65,
+ predicted_hemoglobin=11.56,
+ )
+
+ assert label == "uncertain"
+ assert "uncertain" in text.lower()
+
+
+def test_screening_decision_keeps_high_uncertainty_borderline_case_uncertain() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+
+ label, text = predictor._screening_decision(
+ risk=0.45,
+ uncertainty=0.61,
+ threshold=0.435,
+ predicted_hemoglobin=12.8,
+ )
+
+ assert label == "uncertain"
+ assert "uncertain" in text.lower()
+
+
+def test_should_accept_raw_frame_rescue_for_strong_positive() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ prediction = PredictionResult(
+ anemia_risk=0.86,
+ predicted_hemoglobin=10.9,
+ confidence=0.61,
+ uncertainty=0.39,
+ reliability_flag="medium",
+ screening_label="anemia_likely",
+ screening_text="Likely anemia.",
+ model_source="archive-evidence-fusion-v4",
+ )
+
+ assert predictor.should_accept_raw_frame_rescue(prediction) is True
+
+
+def test_should_reject_raw_frame_rescue_for_weak_positive() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ prediction = PredictionResult(
+ anemia_risk=0.62,
+ predicted_hemoglobin=11.8,
+ confidence=0.52,
+ uncertainty=0.48,
+ reliability_flag="medium",
+ screening_label="anemia_likely",
+ screening_text="Likely anemia.",
+ model_source="archive-evidence-fusion-v4",
+ )
+
+ assert predictor.should_accept_raw_frame_rescue(prediction) is False
+
+
+def test_should_accept_raw_frame_rescue_for_strong_negative() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ prediction = PredictionResult(
+ anemia_risk=0.21,
+ predicted_hemoglobin=13.7,
+ confidence=0.64,
+ uncertainty=0.36,
+ reliability_flag="medium",
+ screening_label="anemia_unlikely",
+ screening_text="Unlikely anemia.",
+ model_source="archive-evidence-fusion-v4",
+ )
+
+ assert predictor.should_accept_raw_frame_rescue(prediction) is True
+
+
+def test_should_accept_raw_frame_rescue_for_hidden_hb_negative() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ prediction = PredictionResult(
+ anemia_risk=0.27,
+ predicted_hemoglobin=None,
+ confidence=0.45,
+ uncertainty=0.55,
+ reliability_flag="low",
+ screening_label="anemia_unlikely",
+ screening_text="Unlikely anemia.",
+ model_source="archive-evidence-fusion-v4",
+ )
+
+ assert predictor.should_accept_raw_frame_rescue(prediction) is True
+
+
+def test_should_accept_raw_frame_rescue_for_low_risk_uncertain() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ prediction = PredictionResult(
+ anemia_risk=0.31,
+ predicted_hemoglobin=None,
+ confidence=0.33,
+ uncertainty=0.67,
+ reliability_flag="low",
+ screening_label="uncertain",
+ screening_text="Uncertain.",
+ model_source="archive-evidence-fusion-v4",
+ )
+
+ assert predictor.should_accept_raw_frame_rescue(prediction) is True
+
+
+def test_should_reject_raw_frame_rescue_for_weak_negative() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ prediction = PredictionResult(
+ anemia_risk=0.29,
+ predicted_hemoglobin=13.4,
+ confidence=0.47,
+ uncertainty=0.53,
+ reliability_flag="low",
+ screening_label="anemia_unlikely",
+ screening_text="Unlikely anemia.",
+ model_source="archive-evidence-fusion-v4",
+ )
+
+ assert predictor.should_accept_raw_frame_rescue(prediction) is False
+
+
+def test_should_reject_raw_frame_rescue_for_high_risk_uncertain() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ prediction = PredictionResult(
+ anemia_risk=0.36,
+ predicted_hemoglobin=None,
+ confidence=0.31,
+ uncertainty=0.67,
+ reliability_flag="low",
+ screening_label="uncertain",
+ screening_text="Uncertain.",
+ model_source="archive-evidence-fusion-v4",
+ )
+
+ assert predictor.should_accept_raw_frame_rescue(prediction) is False
+
+
+def test_should_accept_raw_frame_rescue_for_strong_positive_without_hb() -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ prediction = PredictionResult(
+ anemia_risk=0.72,
+ predicted_hemoglobin=None,
+ confidence=0.24,
+ uncertainty=0.78,
+ reliability_flag="low",
+ screening_label="anemia_likely",
+ screening_text="Likely anemia.",
+ model_source="archive-evidence-fusion-v4",
+ )
+
+ assert predictor.should_accept_raw_frame_rescue(prediction) is True
+
+
+def test_predict_returns_confidence_breakdown(monkeypatch) -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ predictor.enable_efficientnet_fallback = False
+ predictor.archive_model = None
+ predictor.efficientnet_bundle = None
+ predictor.load_error = None
+ predictor.model_path = Path("archive.joblib")
+ predictor.efficientnet_path = Path("efficientnet.pth")
+ predictor._archive_model_load_attempted = False
+ predictor._efficientnet_model_load_attempted = False
+ predictor.runtime_risk_calibrator = None
+ predictor._runtime_risk_calibrator_load_attempted = True
+ predictor.runtime_screening_refiner = None
+ predictor._runtime_screening_refiner_load_attempted = True
+
+ monkeypatch.setattr(
+ predictor,
+ "_ensure_archive_model_loaded",
+ lambda: {"artifact": True},
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "extract_eye_features",
+ lambda image: {
+ "brightness": 0.21,
+ "hist_bright": 0.05,
+ "hist_highlight": 0.01,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_predict_archive_model",
+ lambda artifact, feature_map, source_hint: {
+ "anemia_risk": 0.58,
+ "uncertainty": 0.22,
+ "predicted_hemoglobin": 11.7,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_build_runtime_stack",
+ lambda archive_prediction, **kwargs: {
+ "anemia_risk": 0.58,
+ "uncertainty": 0.22,
+ "predicted_hemoglobin": 11.7,
+ "decision_threshold": 0.5,
+ },
+ )
+
+ quality = QualityAssessment(
+ passed=True,
+ blur_score=148.0,
+ brightness_score=0.24,
+ contrast_score=0.16,
+ framing_score=1.7,
+ lighting_score=0.78,
+ lighting_condition="balanced",
+ lighting_summary="Lighting is even enough for a confident conjunctiva read.",
+ glare_risk=0.08,
+ shadow_risk=0.12,
+ issues=[],
+ )
+
+ result = predictor.predict(Image.new("RGB", (80, 80), "white"), quality)
+
+ assert result.confidence_breakdown is not None
+ assert result.confidence_breakdown["capture_quality"] > 0.6
+ assert result.confidence_breakdown["model_stability"] > 0.7
+ assert result.confidence_breakdown["lighting_condition"] == "balanced"
+ assert "capture quality" in str(result.confidence_breakdown["summary"]).lower() or "threshold" in str(result.confidence_breakdown["summary"]).lower() or "support" in str(result.confidence_breakdown["summary"]).lower()
+
+
+def test_predict_boosts_confidence_for_clear_low_risk_case(monkeypatch) -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ predictor.enable_efficientnet_fallback = False
+ predictor.archive_model = None
+ predictor.efficientnet_bundle = None
+ predictor.load_error = None
+ predictor.model_path = Path("archive.joblib")
+ predictor.efficientnet_path = Path("efficientnet.pth")
+ predictor._archive_model_load_attempted = False
+ predictor._efficientnet_model_load_attempted = False
+ predictor.runtime_risk_calibrator = None
+ predictor._runtime_risk_calibrator_load_attempted = True
+ predictor.runtime_screening_refiner = None
+ predictor._runtime_screening_refiner_load_attempted = True
+
+ monkeypatch.setattr(
+ predictor,
+ "_ensure_archive_model_loaded",
+ lambda: {"artifact": True},
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "extract_eye_features",
+ lambda image: {
+ "brightness": 0.24,
+ "hist_bright": 0.09,
+ "hist_highlight": 0.01,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_predict_archive_model",
+ lambda artifact, feature_map, source_hint: {
+ "anemia_risk": 0.22,
+ "uncertainty": 0.24,
+ "predicted_hemoglobin": 13.7,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_build_runtime_stack",
+ lambda archive_prediction, **kwargs: {
+ "anemia_risk": 0.22,
+ "uncertainty": 0.24,
+ "predicted_hemoglobin": 13.7,
+ "decision_threshold": 0.5,
+ },
+ )
+
+ quality = QualityAssessment(
+ passed=True,
+ blur_score=86.0,
+ brightness_score=0.23,
+ contrast_score=0.14,
+ framing_score=1.12,
+ lighting_score=0.46,
+ lighting_condition="dim",
+ lighting_summary="Lighting is slightly dim but still usable.",
+ glare_risk=0.1,
+ shadow_risk=0.18,
+ issues=[],
+ )
+
+ result = predictor.predict(Image.new("RGB", (80, 80), "white"), quality)
+
+ assert result.screening_label == "anemia_unlikely"
+ assert result.confidence >= 0.55
+ assert result.reliability_flag in {"medium", "high"}
+ assert result.confidence_breakdown is not None
+ assert "low-risk side" in str(result.confidence_breakdown["summary"]).lower()
+
+
+def test_predict_keeps_low_risk_case_conservative_when_glare_is_high(monkeypatch) -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ predictor.enable_efficientnet_fallback = False
+ predictor.archive_model = None
+ predictor.efficientnet_bundle = None
+ predictor.load_error = None
+ predictor.model_path = Path("archive.joblib")
+ predictor.efficientnet_path = Path("efficientnet.pth")
+ predictor._archive_model_load_attempted = False
+ predictor._efficientnet_model_load_attempted = False
+ predictor.runtime_risk_calibrator = None
+ predictor._runtime_risk_calibrator_load_attempted = True
+ predictor.runtime_screening_refiner = None
+ predictor._runtime_screening_refiner_load_attempted = True
+
+ monkeypatch.setattr(
+ predictor,
+ "_ensure_archive_model_loaded",
+ lambda: {"artifact": True},
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "extract_eye_features",
+ lambda image: {
+ "brightness": 0.24,
+ "hist_bright": 0.09,
+ "hist_highlight": 0.01,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_predict_archive_model",
+ lambda artifact, feature_map, source_hint: {
+ "anemia_risk": 0.22,
+ "uncertainty": 0.24,
+ "predicted_hemoglobin": 13.7,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_build_runtime_stack",
+ lambda archive_prediction, **kwargs: {
+ "anemia_risk": 0.22,
+ "uncertainty": 0.24,
+ "predicted_hemoglobin": 13.7,
+ "decision_threshold": 0.5,
+ },
+ )
+
+ quality = QualityAssessment(
+ passed=True,
+ blur_score=84.0,
+ brightness_score=0.34,
+ contrast_score=0.14,
+ framing_score=1.12,
+ lighting_score=0.44,
+ lighting_condition="glare_heavy",
+ lighting_summary="Highlights are clipping part of the eyelid surface.",
+ glare_risk=0.72,
+ shadow_risk=0.18,
+ issues=[],
+ )
+
+ result = predictor.predict(Image.new("RGB", (80, 80), "white"), quality)
+
+ assert result.screening_label == "anemia_unlikely"
+ assert result.confidence < 0.55
+ assert result.reliability_flag == "low"
+
+
+def test_predict_keeps_strong_quality_limited_positive_above_flat_low_confidence(monkeypatch) -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ predictor.enable_efficientnet_fallback = False
+ predictor.archive_model = None
+ predictor.efficientnet_bundle = None
+ predictor.load_error = None
+ predictor.model_path = Path("archive.joblib")
+ predictor.efficientnet_path = Path("efficientnet.pth")
+ predictor._archive_model_load_attempted = False
+ predictor._efficientnet_model_load_attempted = False
+ predictor.runtime_risk_calibrator = None
+ predictor._runtime_risk_calibrator_load_attempted = True
+ predictor.runtime_screening_refiner = None
+ predictor._runtime_screening_refiner_load_attempted = False
+
+ class _FakeRefiner:
+ method = "logistic-regression"
+
+ def refine(
+ self,
+ *,
+ base_anemia_risk: float,
+ uncertainty: float,
+ predicted_hemoglobin: float | None,
+ quality: QualityAssessment,
+ base_likely: bool,
+ ) -> float:
+ assert quality.lighting_condition == "shadow_heavy"
+ return 0.956
+
+ monkeypatch.setattr(
+ predictor,
+ "_ensure_runtime_screening_refiner_loaded",
+ lambda: _FakeRefiner(),
+ )
+ monkeypatch.setattr(
+ predictor,
+ "_ensure_archive_model_loaded",
+ lambda: {"artifact": True},
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "extract_eye_features",
+ lambda image: {
+ "brightness": 0.067,
+ "hist_bright": 0.0,
+ "hist_highlight": 0.0,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_predict_archive_model",
+ lambda artifact, feature_map, source_hint: {
+ "anemia_risk": 0.496,
+ "uncertainty": 0.782,
+ "predicted_hemoglobin": 11.9,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_build_runtime_stack",
+ lambda archive_prediction, **kwargs: {
+ "anemia_risk": 0.496,
+ "uncertainty": 0.782,
+ "predicted_hemoglobin": 11.9,
+ "decision_threshold": 0.495,
+ },
+ )
+
+ quality = QualityAssessment(
+ passed=True,
+ blur_score=195.0,
+ brightness_score=0.067,
+ contrast_score=0.16,
+ framing_score=2.742,
+ lighting_score=0.54,
+ lighting_condition="shadow_heavy",
+ lighting_summary="Shadows are covering part of the eyelid, so the model may miss the true pallor signal.",
+ glare_risk=0.0,
+ shadow_risk=1.0,
+ issues=[],
+ )
+
+ result = predictor.predict(Image.new("RGB", (80, 80), "white"), quality)
+
+ assert result.screening_label == "uncertain"
+ assert result.confidence >= 0.5
+ assert result.reliability_flag == "low"
+ assert result.confidence_breakdown is not None
+ assert float(result.confidence_breakdown["signal_strength"]) >= 0.9
+
+
+def test_predict_applies_runtime_risk_calibrator(monkeypatch) -> None:
+ predictor = ScreeningPredictor.__new__(ScreeningPredictor)
+ predictor.enable_efficientnet_fallback = False
+ predictor.archive_model = None
+ predictor.efficientnet_bundle = None
+ predictor.load_error = None
+ predictor.model_path = Path("archive.joblib")
+ predictor.efficientnet_path = Path("efficientnet.pth")
+ predictor._archive_model_load_attempted = False
+ predictor._efficientnet_model_load_attempted = False
+ predictor.runtime_risk_calibrator = None
+ predictor._runtime_risk_calibrator_load_attempted = False
+ predictor.runtime_screening_refiner = None
+ predictor._runtime_screening_refiner_load_attempted = True
+
+ class _FakeCalibrator:
+ method = "temperature"
+
+ def calibrate(self, probability: float, *, source_hint: str = "roi_original") -> float:
+ assert source_hint == "roi_original"
+ return probability + 0.14
+
+ monkeypatch.setattr(
+ predictor,
+ "_ensure_runtime_risk_calibrator_loaded",
+ lambda: _FakeCalibrator(),
+ )
+ monkeypatch.setattr(
+ predictor,
+ "_ensure_archive_model_loaded",
+ lambda: {"artifact": True},
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "extract_eye_features",
+ lambda image: {
+ "brightness": 0.24,
+ "hist_bright": 0.09,
+ "hist_highlight": 0.01,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_predict_archive_model",
+ lambda artifact, feature_map, source_hint: {
+ "anemia_risk": 0.48,
+ "uncertainty": 0.18,
+ "predicted_hemoglobin": 11.7,
+ },
+ )
+ monkeypatch.setattr(
+ prediction_module,
+ "_build_runtime_stack",
+ lambda archive_prediction, **kwargs: {
+ "anemia_risk": 0.48,
+ "uncertainty": 0.18,
+ "predicted_hemoglobin": 11.7,
+ "decision_threshold": 0.5,
+ },
+ )
+
+ quality = QualityAssessment(
+ passed=True,
+ blur_score=170.0,
+ brightness_score=0.23,
+ contrast_score=0.16,
+ framing_score=1.35,
+ lighting_score=0.82,
+ lighting_condition="balanced",
+ lighting_summary="Lighting is balanced enough for reliable screening.",
+ glare_risk=0.06,
+ shadow_risk=0.08,
+ issues=[],
+ )
+
+ result = predictor.predict(Image.new("RGB", (80, 80), "white"), quality)
+
+ assert result.screening_label == "anemia_likely"
+ assert round(result.anemia_risk, 2) == 0.48
+ assert result.confidence_breakdown is not None
+ assert result.confidence_breakdown["calibration_applied"] is True
+ assert result.confidence_breakdown["calibration_method"] == "temperature"
+ assert round(float(result.confidence_breakdown["raw_anemia_risk"]), 2) == 0.48
+ assert round(float(result.confidence_breakdown["calibrated_anemia_risk"]), 2) == 0.62
+ assert round(float(result.confidence_breakdown["decision_threshold"]), 2) == 0.5
diff --git a/backend/tests/test_quality.py b/backend/tests/test_quality.py
new file mode 100644
index 0000000000000000000000000000000000000000..dbe7a88fbc8eae164909c1f2b117a1a956ad8cc3
--- /dev/null
+++ b/backend/tests/test_quality.py
@@ -0,0 +1,448 @@
+"""
+Tests for ImageQualityService.
+
+Each test creates a synthetic image designed to trigger (or avoid) a specific
+quality gate. This is deliberately independent of real photos so CI never
+needs access to patient data.
+
+Coverage targets:
+- Flat/uniform images fail with a clear issue code.
+- A synthetic eye-like pattern (iris + sclera + lower lid) passes.
+- Mild lighting warnings are non-blocking.
+- Bright but detailed images are not penalised.
+- Non-eye close-ups fail with eye_not_visible before lighting feedback.
+- Large real-world-style images trigger ROI cropping and pass.
+- issue_codes and blocking_issues computed properties work correctly.
+"""
+
+from __future__ import annotations
+
+import sys
+from io import BytesIO
+from pathlib import Path
+
+import numpy as np
+import pytest
+from PIL import Image
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.services.image_quality import ImageQualityService
+from app.services.conjunctiva_roi import ConjunctivaRoiExtractor
+from app.schemas import QualityAssessment, QualityIssue
+
+
+# ---------------------------------------------------------------------------
+# Helpers
+# ---------------------------------------------------------------------------
+
+def _to_bytes(array: np.ndarray, fmt: str = "PNG") -> bytes:
+ img = Image.fromarray(array.astype("uint8"), mode="RGB")
+ buf = BytesIO()
+ img.save(buf, format=fmt)
+ return buf.getvalue()
+
+
+def _eye_canvas(
+ size: int = 360,
+ bg: int = 45,
+ iris_r: int = 55,
+ sclera_rx: int = 120,
+ sclera_ry: int = 70,
+ iris_color: tuple = (20, 30, 45),
+ sclera_color: tuple = (190, 120, 120),
+ lid_boost: tuple = (35, 8, 8),
+) -> np.ndarray:
+ """Build a synthetic eye-like pattern centred in a square canvas."""
+ canvas = np.full((size, size, 3), bg, dtype=np.uint8)
+ cx, cy = size // 2, size // 2
+ yy, xx = np.ogrid[:size, :size]
+
+ iris_mask = (xx - cx) ** 2 + (yy - cy) ** 2 <= iris_r ** 2
+ sclera_mask = (xx - cx) ** 2 / sclera_rx ** 2 + (yy - cy) ** 2 / sclera_ry ** 2 <= 1
+
+ canvas[sclera_mask] = sclera_color
+ canvas[iris_mask] = iris_color
+
+ # Lower-lid conjunctiva highlight
+ lid_y_start, lid_y_end = cy - size // 20, cy + size // 4
+ lid_x_start, lid_x_end = cx - sclera_rx + 10, cx + sclera_rx - 10
+ canvas[lid_y_start:lid_y_end, lid_x_start:lid_x_end] = np.clip(
+ canvas[lid_y_start:lid_y_end, lid_x_start:lid_x_end] + lid_boost, 0, 255
+ )
+ return canvas
+
+
+SERVICE = ImageQualityService()
+ROI_EXTRACTOR = ConjunctivaRoiExtractor()
+
+
+# ---------------------------------------------------------------------------
+# Tests
+# ---------------------------------------------------------------------------
+
+def test_flat_uniform_image_fails() -> None:
+ rgb = np.full((320, 320, 3), 140, dtype=np.uint8)
+ quality, _ = SERVICE.evaluate(_to_bytes(rgb))
+
+ assert quality.passed is False
+ assert quality.issue_codes & {"eye_not_visible", "poor_lighting", "blur_detected"}
+
+
+def test_eye_like_pattern_passes() -> None:
+ canvas = _eye_canvas()
+ quality, _ = SERVICE.evaluate(_to_bytes(canvas))
+ assert quality.passed is True
+
+
+def test_mild_lighting_warning_is_non_blocking() -> None:
+ """Slightly dim image should warn but allow analysis to proceed."""
+ canvas = _eye_canvas(bg=96, sclera_color=(214, 176, 176))
+ quality, _ = SERVICE.evaluate(_to_bytes(canvas))
+
+ assert quality.passed is True
+ assert "poor_lighting" in quality.issue_codes
+ assert len(quality.blocking_issues) == 0
+ assert quality.lighting_condition
+ assert quality.lighting_summary
+ assert 0.0 <= quality.glare_risk <= 1.0
+ assert 0.0 <= quality.shadow_risk <= 1.0
+
+
+def test_bright_detailed_eye_is_not_blocked() -> None:
+ """A well-lit image with a visible iris should not be rejected for brightness."""
+ canvas = _eye_canvas(bg=148, sclera_color=(238, 210, 210), iris_color=(25, 35, 52))
+ quality, _ = SERVICE.evaluate(_to_bytes(canvas))
+
+ assert quality.brightness_score > 0.42
+ assert quality.passed is True
+ assert len(quality.blocking_issues) == 0
+ assert quality.lighting_score > 0.35
+
+
+def test_lighting_intelligence_detects_glare_heavy() -> None:
+ score, condition, summary, glare_risk, shadow_risk = SERVICE._lighting_intelligence(
+ brightness_score=0.56,
+ contrast_score=0.16,
+ center_brightness=0.62,
+ center_contrast=0.15,
+ bright_region_ratio=0.31,
+ highlight_ratio=0.18,
+ dark_region_ratio=0.05,
+ )
+
+ assert 0.0 <= score <= 1.0
+ assert condition == "glare_heavy"
+ assert glare_risk >= 0.72
+ assert "glare" in summary.lower()
+ assert 0.0 <= shadow_risk <= 1.0
+
+
+def test_lighting_intelligence_marks_balanced_capture() -> None:
+ score, condition, summary, glare_risk, shadow_risk = SERVICE._lighting_intelligence(
+ brightness_score=0.28,
+ contrast_score=0.18,
+ center_brightness=0.29,
+ center_contrast=0.17,
+ bright_region_ratio=0.08,
+ highlight_ratio=0.01,
+ dark_region_ratio=0.12,
+ )
+
+ assert condition == "balanced"
+ assert score > 0.65
+ assert "balanced" in summary.lower()
+ assert glare_risk < 0.3
+ assert shadow_risk < 0.3
+
+
+def test_non_eye_closeup_fails_with_eye_not_visible_first() -> None:
+ """
+ A non-eye pattern should fail, and the first issue reported should be
+ eye_not_visible โ not a lighting complaint.
+ """
+ yy, xx = np.indices((360, 360))
+ canvas = np.zeros((360, 360, 3), dtype=np.uint8)
+ canvas[..., 0] = 164 + ((xx // 18) % 2) * 18
+ canvas[..., 1] = 126 + ((yy // 18) % 2) * 12
+ canvas[..., 2] = 106
+ canvas[118:242, 118:242] = [106, 86, 70]
+
+ quality, _ = SERVICE.evaluate(_to_bytes(canvas))
+
+ assert quality.passed is False
+ assert quality.issues[0].code == "eye_not_visible"
+ assert "poor_lighting" not in quality.issue_codes
+
+
+def test_large_image_triggers_roi_crop_and_passes() -> None:
+ """
+ A large photo with the eye off-centre (realistic phone photo) should be
+ auto-cropped to the ROI and still pass quality.
+ """
+ canvas = np.full((900, 1200, 3), [214, 180, 165], dtype=np.uint8)
+ yy, xx = np.indices((900, 1200))
+
+ sclera = (xx - 650) ** 2 / 200 ** 2 + (yy - 520) ** 2 / 150 ** 2 <= 1
+ iris = (xx - 650) ** 2 + (yy - 520) ** 2 <= 82 ** 2
+ lower_lid = (xx - 650) ** 2 / 230 ** 2 + (yy - 650) ** 2 / 82 ** 2 <= 1
+ finger = (xx - 680) ** 2 / 170 ** 2 + (yy - 810) ** 2 / 120 ** 2 <= 1
+
+ canvas[sclera] = [198, 206, 212]
+ canvas[iris] = [78, 92, 98]
+ canvas[lower_lid] = [232, 182, 188]
+ canvas[(lower_lid) & (yy >= 650)] = [192, 96, 108]
+ canvas[finger] = [198, 158, 140]
+ canvas[330:430, 220:980] = np.clip(canvas[330:430, 220:980] - [120, 120, 120], 0, 255)
+
+ quality, roi_image = SERVICE.evaluate(_to_bytes(canvas))
+
+ assert quality.passed is True
+ assert "roi_cropped" in quality.issue_codes
+ assert roi_image.size[0] < 500
+ assert roi_image.size[1] < 260
+
+
+def test_raw_frame_rescue_allows_combined_framing_and_lighting_blocks() -> None:
+ assessment = QualityAssessment(
+ passed=False,
+ blur_score=220.0,
+ brightness_score=0.56,
+ contrast_score=0.11,
+ framing_score=2.2,
+ lighting_score=0.32,
+ lighting_condition="overexposed",
+ lighting_summary="Lighting is brighter than ideal.",
+ glare_risk=0.22,
+ shadow_risk=0.0,
+ issues=[
+ QualityIssue(
+ code="bad_framing",
+ severity="blocking",
+ title="Framing is weak",
+ message="Retake closer.",
+ ),
+ QualityIssue(
+ code="poor_lighting",
+ severity="warning",
+ title="Lighting is bright",
+ message="Move to softer light.",
+ ),
+ ],
+ )
+
+ assert SERVICE.allows_raw_frame_rescue(assessment) is True
+
+
+def test_roi_extractor_falls_back_to_conjunctiva_band_when_iris_detection_misses() -> None:
+ canvas = np.full((900, 1200, 3), [210, 178, 166], dtype=np.uint8)
+ yy, xx = np.indices((900, 1200))
+
+ lid_band = (xx - 620) ** 2 / 260 ** 2 + (yy - 560) ** 2 / 90 ** 2 <= 1
+ canvas[lid_band] = [212, 118, 132]
+ canvas[(lid_band) & (yy >= 560)] = [188, 88, 104]
+ canvas[260:360, 140:1080] = np.clip(canvas[260:360, 140:1080] - [95, 95, 95], 0, 255)
+
+ result = ROI_EXTRACTOR.extract(Image.fromarray(canvas.astype("uint8"), mode="RGB"))
+
+ assert result.extracted is True
+ assert result.image.size[0] < canvas.shape[1]
+ assert result.image.size[1] < canvas.shape[0]
+ assert result.image.size[0] >= 110
+ assert result.image.size[1] >= 40
+
+
+@pytest.mark.parametrize("fmt", ["JPEG", "PNG"])
+def test_both_image_formats_accepted(fmt: str) -> None:
+ canvas = _eye_canvas()
+ quality, _ = SERVICE.evaluate(_to_bytes(canvas, fmt=fmt))
+ assert quality.passed is True
+
+
+def test_quality_assessment_computed_properties() -> None:
+ canvas = _eye_canvas(bg=96, sclera_color=(214, 176, 176))
+ quality, _ = SERVICE.evaluate(_to_bytes(canvas))
+
+ # Validate cached_property helpers
+ assert isinstance(quality.issue_codes, frozenset)
+ assert isinstance(quality.blocking_issues, list)
+ assert isinstance(quality.warning_issues, list)
+ assert all(i.severity == "blocking" for i in quality.blocking_issues)
+ assert all(i.severity == "warning" for i in quality.warning_issues)
+
+
+def test_runtime_quality_issue_codes_validate_against_schema() -> None:
+ QualityIssue(
+ code="resolution_too_low",
+ severity="blocking",
+ title="Image is too small",
+ message="Move closer and retake the photo.",
+ )
+ QualityIssue(
+ code="bad_framing",
+ severity="warning",
+ title="Eye framing is loose",
+ message="Center the exposed lower eyelid more tightly.",
+ )
+
+
+def test_roi_salvage_rule_allows_recoverable_crop() -> None:
+ issues = [
+ QualityIssue(
+ code="roi_cropped",
+ severity="warning",
+ title="Lower eyelid region detected",
+ message="ROI extracted.",
+ ),
+ QualityIssue(
+ code="bad_framing",
+ severity="blocking",
+ title="Eye is not framed clearly",
+ message="Fill the frame with one eye.",
+ ),
+ ]
+
+ assert SERVICE._should_salvage_roi_capture(
+ issues,
+ roi_extracted=True,
+ blur_score=220.0,
+ brightness_score=0.41,
+ contrast_score=0.14,
+ framing_score=2.4,
+ )
+
+ softened = SERVICE._soften_salvageable_roi_blocks(
+ issues,
+ roi_extracted=True,
+ blur_score=220.0,
+ brightness_score=0.41,
+ contrast_score=0.14,
+ framing_score=2.4,
+ )
+ assert softened[1].severity == "warning"
+
+
+def test_roi_salvage_rule_rejects_low_contrast_crop() -> None:
+ issues = [
+ QualityIssue(
+ code="roi_cropped",
+ severity="warning",
+ title="Lower eyelid region detected",
+ message="ROI extracted.",
+ ),
+ QualityIssue(
+ code="eye_not_visible",
+ severity="blocking",
+ title="Eye is not clearly visible",
+ message="Retake with the lower eyelid visible.",
+ ),
+ ]
+
+ assert not SERVICE._should_salvage_roi_capture(
+ issues,
+ roi_extracted=True,
+ blur_score=220.0,
+ brightness_score=0.41,
+ contrast_score=0.08,
+ framing_score=2.8,
+ )
+
+
+def test_roi_salvage_rule_allows_clarity_exception_for_bad_framing() -> None:
+ issues = [
+ QualityIssue(
+ code="roi_cropped",
+ severity="warning",
+ title="Lower eyelid region detected",
+ message="ROI extracted.",
+ ),
+ QualityIssue(
+ code="bad_framing",
+ severity="blocking",
+ title="Eye is not framed clearly",
+ message="Fill the frame with one eye.",
+ ),
+ ]
+
+ assert SERVICE._should_salvage_roi_capture(
+ issues,
+ roi_extracted=True,
+ blur_score=340.0,
+ brightness_score=0.56,
+ contrast_score=0.17,
+ framing_score=1.8,
+ )
+
+
+def test_raw_frame_rescue_allowed_for_framing_and_visibility_blocks() -> None:
+ assessment = SERVICE.build_raw_frame_rescue_assessment(
+ SERVICE.evaluate(_to_bytes(_eye_canvas(size=900, bg=70)))[0].model_copy(
+ update={
+ "passed": False,
+ "issues": [
+ QualityIssue(
+ code="roi_cropped",
+ severity="warning",
+ title="Lower eyelid region detected",
+ message="ROI extracted.",
+ ),
+ QualityIssue(
+ code="eye_not_visible",
+ severity="blocking",
+ title="Eye is not clearly visible",
+ message="Retake with one eye filling the frame.",
+ ),
+ ],
+ }
+ )
+ )
+
+ assert assessment.passed is True
+ assert all(issue.severity == "warning" for issue in assessment.issues)
+
+
+def test_raw_frame_rescue_allowed_for_isolated_poor_lighting_block() -> None:
+ assessment = SERVICE.evaluate(_to_bytes(_eye_canvas(size=900, bg=70)))[0].model_copy(
+ update={
+ "passed": False,
+ "issues": [
+ QualityIssue(
+ code="poor_lighting",
+ severity="blocking",
+ title="Lighting is not usable",
+ message="Use bright, even light.",
+ ),
+ ],
+ }
+ )
+
+ rescued = SERVICE.build_raw_frame_rescue_assessment(assessment)
+
+ assert SERVICE.allows_raw_frame_rescue(assessment) is True
+ assert rescued.passed is True
+ assert rescued.issues[0].severity == "warning"
+
+
+def test_raw_frame_rescue_not_allowed_for_mixed_lighting_and_blur_blocks() -> None:
+ assessment = SERVICE.evaluate(_to_bytes(_eye_canvas(size=900, bg=70)))[0].model_copy(
+ update={
+ "passed": False,
+ "issues": [
+ QualityIssue(
+ code="poor_lighting",
+ severity="blocking",
+ title="Lighting is not usable",
+ message="Use bright, even light.",
+ ),
+ QualityIssue(
+ code="blur_detected",
+ severity="blocking",
+ title="Image looks blurry",
+ message="Hold steady and retake the photo.",
+ ),
+ ],
+ }
+ )
+
+ assert SERVICE.allows_raw_frame_rescue(assessment) is False
diff --git a/backend/tests/test_request_parsing.py b/backend/tests/test_request_parsing.py
new file mode 100644
index 0000000000000000000000000000000000000000..bb3f054c247f444795c43be2b28a139ecc55c881
--- /dev/null
+++ b/backend/tests/test_request_parsing.py
@@ -0,0 +1,203 @@
+"""
+Tests for request_parsing โ the input-sanitisation layer that sits between
+raw HTTP form data and the typed service layer.
+
+Coverage targets:
+- Boolean normalisation for every accepted truthy/falsy/null string.
+- Extra-field rejection (prevents schema drift being silently ignored).
+- Text normalisation: whitespace collapse, length enforcement.
+- JSON edge cases: null payload, empty object, malformed JSON.
+"""
+
+from __future__ import annotations
+
+import json
+import sys
+from pathlib import Path
+
+import pytest
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.schemas import PatientProfileInput, SymptomInput
+from app.services.request_parsing import (
+ InvalidRequestPayload,
+ normalize_optional_text,
+ parse_patient_profile,
+ parse_symptoms,
+)
+
+
+# ---------------------------------------------------------------------------
+# Boolean normalisation โ parametrised for full coverage
+# ---------------------------------------------------------------------------
+
+@pytest.mark.parametrize("truthy", ["yes", "Yes", "YES", "true", "True", "1", "on", "y"])
+def test_parse_symptoms_accepts_truthy_strings(truthy: str) -> None:
+ payload = json.dumps({"fatigue": truthy})
+ result = parse_symptoms(payload)
+ assert result.fatigue is True
+
+
+@pytest.mark.parametrize("falsy", ["no", "No", "false", "False", "0", "off", "n", ""])
+def test_parse_symptoms_accepts_falsy_strings(falsy: str) -> None:
+ payload = json.dumps({"dizziness": falsy})
+ result = parse_symptoms(payload)
+ assert result.dizziness is False
+
+
+@pytest.mark.parametrize("null_like", ["skip", "unknown", "n/a", "na", "none", "null"])
+def test_parse_symptoms_accepts_null_strings_for_optional_field(null_like: str) -> None:
+ payload = json.dumps({"heavy_menstrual_bleeding": null_like})
+ result = parse_symptoms(payload)
+ assert result.heavy_menstrual_bleeding is None
+
+
+def test_parse_symptoms_normalises_mixed_types() -> None:
+ """Realistic mixed-type payload from a browser form submission."""
+ payload = json.dumps({
+ "fatigue": "yes",
+ "dizziness": "0",
+ "pale_skin": True,
+ "shortness_of_breath": "false",
+ "heavy_menstrual_bleeding": "skip",
+ "poor_diet_low_iron": "1",
+ })
+
+ result = parse_symptoms(payload)
+
+ assert result == SymptomInput(
+ fatigue=True,
+ dizziness=False,
+ pale_skin=True,
+ shortness_of_breath=False,
+ heavy_menstrual_bleeding=None,
+ poor_diet_low_iron=True,
+ )
+
+
+def test_parse_symptoms_null_payload_returns_defaults() -> None:
+ """A missing symptoms form field should yield all-False defaults."""
+ result = parse_symptoms(None)
+ assert result == SymptomInput()
+ assert result.active_count == 0
+
+
+def test_parse_symptoms_empty_object_returns_defaults() -> None:
+ result = parse_symptoms("{}")
+ assert result == SymptomInput()
+
+
+def test_parse_patient_profile_defaults_when_missing() -> None:
+ assert parse_patient_profile(None) == PatientProfileInput()
+
+
+def test_parse_patient_profile_accepts_string_age_and_normalises_enums() -> None:
+ result = parse_patient_profile(
+ json.dumps({"age": "17", "sex": " Female ", "diet_type": " Vegetarian "})
+ )
+ assert result == PatientProfileInput(age=17, sex="female", diet_type="vegetarian")
+
+
+def test_parse_patient_profile_rejects_non_object_json() -> None:
+ with pytest.raises(InvalidRequestPayload):
+ parse_patient_profile('["not", "an", "object"]')
+
+
+# ---------------------------------------------------------------------------
+# Schema enforcement
+# ---------------------------------------------------------------------------
+
+def test_parse_symptoms_rejects_unknown_fields() -> None:
+ with pytest.raises(InvalidRequestPayload, match="unlisted_symptom"):
+ parse_symptoms('{"fatigue": true, "unlisted_symptom": true}')
+
+
+def test_parse_symptoms_rejects_multiple_unknown_fields() -> None:
+ with pytest.raises(InvalidRequestPayload):
+ parse_symptoms('{"fever": true, "nausea": true}')
+
+
+# ---------------------------------------------------------------------------
+# Malformed input
+# ---------------------------------------------------------------------------
+
+def test_parse_symptoms_rejects_malformed_json() -> None:
+ with pytest.raises(InvalidRequestPayload, match="[Ii]nvalid"):
+ parse_symptoms("{fatigue: true}") # unquoted key โ not valid JSON
+
+
+def test_parse_symptoms_rejects_non_object_json() -> None:
+ """Top-level arrays and scalars should be rejected."""
+ with pytest.raises(InvalidRequestPayload):
+ parse_symptoms('["fatigue", true]')
+
+
+def test_parse_symptoms_rejects_invalid_boolean_value() -> None:
+ with pytest.raises(InvalidRequestPayload):
+ parse_symptoms('{"fatigue": "maybe"}')
+
+
+# ---------------------------------------------------------------------------
+# Computed properties
+# ---------------------------------------------------------------------------
+
+def test_symptom_input_active_count() -> None:
+ s = SymptomInput(fatigue=True, dizziness=True, poor_diet_low_iron=True)
+ assert s.active_count == 3
+
+
+def test_symptom_input_burden_none() -> None:
+ assert SymptomInput().symptom_burden == "none"
+
+
+def test_symptom_input_burden_mild() -> None:
+ assert SymptomInput(fatigue=True).symptom_burden == "mild"
+
+
+def test_symptom_input_burden_moderate() -> None:
+ s = SymptomInput(fatigue=True, dizziness=True, pale_skin=True)
+ assert s.symptom_burden == "moderate"
+
+
+def test_symptom_input_burden_severe() -> None:
+ s = SymptomInput(
+ fatigue=True, dizziness=True, pale_skin=True,
+ shortness_of_breath=True, poor_diet_low_iron=True,
+ )
+ assert s.symptom_burden == "severe"
+
+
+# ---------------------------------------------------------------------------
+# normalize_optional_text
+# ---------------------------------------------------------------------------
+
+def test_normalize_optional_text_collapses_internal_whitespace() -> None:
+ assert normalize_optional_text(" South India ", field_name="region") == "South India"
+
+
+def test_normalize_optional_text_returns_none_for_blank() -> None:
+ assert normalize_optional_text(" ", field_name="region") is None
+
+
+def test_normalize_optional_text_returns_none_for_none() -> None:
+ assert normalize_optional_text(None, field_name="language") is None
+
+
+def test_normalize_optional_text_rejects_overly_long_values() -> None:
+ with pytest.raises(InvalidRequestPayload, match="language"):
+ normalize_optional_text("x" * 49, field_name="language")
+
+
+def test_normalize_optional_text_accepts_max_length_value() -> None:
+ # exactly 48 chars โ should pass with the default limit
+ value = "a" * 48
+ result = normalize_optional_text(value, field_name="language")
+ assert result == value
+
+
+def test_normalize_optional_text_strips_unicode_whitespace() -> None:
+ # Non-breaking space should be treated like regular whitespace
+ result = normalize_optional_text("Kerala\u00a0India", field_name="region")
+ assert "\u00a0" not in result
diff --git a/backend/tests/test_runtime_stack.py b/backend/tests/test_runtime_stack.py
new file mode 100644
index 0000000000000000000000000000000000000000..5c03d3302032195cf73c8da7792501ffb246f0df
--- /dev/null
+++ b/backend/tests/test_runtime_stack.py
@@ -0,0 +1,69 @@
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.ml.runtime_stack import (
+ RUNTIME_STACK_VERSION,
+ build_runtime_stack_prediction,
+ decision_threshold_for_source,
+ hb_archive_weight_for_source,
+ risk_archive_weight_for_source,
+)
+
+
+def test_decision_threshold_defaults_are_source_aware() -> None:
+ assert decision_threshold_for_source("roi_original") == 0.495
+ assert decision_threshold_for_source("palpebral") == 0.65
+ assert decision_threshold_for_source("forniceal_palpebral") == 0.65
+
+
+def test_runtime_stack_prediction_keeps_archive_signal_without_secondary_model() -> None:
+ result = build_runtime_stack_prediction(
+ {
+ "anemia_risk": 0.58,
+ "predicted_hemoglobin": 11.4,
+ "uncertainty": 0.21,
+ },
+ source_hint="roi_original",
+ )
+
+ assert result["anemia_risk"] == 0.58
+ assert result["predicted_hemoglobin"] == 11.4
+ assert result["uncertainty"] == 0.21
+ assert result["decision_threshold"] == 0.495
+
+
+def test_runtime_stack_blends_archive_and_efficientnet_for_roi() -> None:
+ result = build_runtime_stack_prediction(
+ {
+ "anemia_risk": 0.7,
+ "predicted_hemoglobin": 10.8,
+ "uncertainty": 0.18,
+ },
+ efficientnet_prediction={
+ "anemia_risk": 0.3,
+ "predicted_hemoglobin": 12.0,
+ "uncertainty": 0.24,
+ },
+ source_hint="roi_original",
+ )
+
+ assert round(result["anemia_risk"], 4) == 0.5138
+ assert round(result["predicted_hemoglobin"], 4) == 11.16
+ assert round(result["decision_threshold"], 4) == 0.495
+ assert round(result["uncertainty"], 4) == 0.2424
+
+
+def test_runtime_stack_weights_are_source_aware() -> None:
+ assert risk_archive_weight_for_source("roi_original") == 0.55
+ assert risk_archive_weight_for_source("palpebral") == 1.0
+ assert hb_archive_weight_for_source("roi_original") == 0.70
+ assert hb_archive_weight_for_source("palpebral") == 1.0
+
+
+def test_runtime_stack_version_is_declared() -> None:
+ assert RUNTIME_STACK_VERSION == "archive-evidence-fusion-v4"
diff --git a/backend/tests/test_runtime_status_response.py b/backend/tests/test_runtime_status_response.py
new file mode 100644
index 0000000000000000000000000000000000000000..7d3054e8eeac051c18ca00d03163f5160c1b8235
--- /dev/null
+++ b/backend/tests/test_runtime_status_response.py
@@ -0,0 +1,110 @@
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.schemas import GuidanceRuntimeStatus, ModelRuntimeStatus
+from app.services import runtime_status as runtime_status_module
+
+
+class _DummyPredictor:
+ def runtime_status(self) -> ModelRuntimeStatus:
+ return ModelRuntimeStatus(
+ primary_model="archive-evidence-fusion-v4",
+ deep_stack_loaded=False,
+ legacy_loaded=False,
+ artifact_ready=True,
+ artifact_path="backend/models/archive_screening_model.joblib",
+ )
+
+
+class _DummyGuidance:
+ def runtime_status(self) -> GuidanceRuntimeStatus:
+ return GuidanceRuntimeStatus(
+ active_strategy="fallback",
+ mistral_enabled=True,
+ client_ready=False,
+ api_key_configured=False,
+ mistral_model="mistral-small-latest",
+ fallback_reason="Fallback active.",
+ )
+
+
+def test_runtime_status_includes_deployed_metrics(monkeypatch) -> None:
+ monkeypatch.setattr(
+ runtime_status_module,
+ "_load_training_report",
+ lambda: {
+ "primary_model": "archive-evidence-fusion-v4",
+ "record_count": 432,
+ "metrics": {
+ "accuracy": 0.8864,
+ "f1": 0.8,
+ "split_strategy": "group-shuffle-balance-select: roi_original",
+ },
+ },
+ )
+ monkeypatch.setattr(
+ runtime_status_module,
+ "_load_json_report",
+ lambda path: (
+ {
+ "version": "runtime-risk-calibrator-v1",
+ "method": "temperature",
+ "selected_thresholds": {"roi_original": 0.58},
+ "diagnostics": {
+ "ece_before": 0.121,
+ "ece_after": 0.072,
+ "brier_before": 0.164,
+ "brier_after": 0.133,
+ },
+ }
+ if str(path).endswith("runtime_calibration_report.json")
+ else {
+ "version": "runtime-screening-refiner-v1",
+ "method": "logistic-regression",
+ "selected_threshold": 0.53,
+ "metrics_after": {
+ "accuracy": 0.8636,
+ "precision": 0.7857,
+ "recall": 0.7857,
+ "f1": 0.7857,
+ },
+ }
+ if str(path).endswith("runtime_refinement_report.json")
+ else {
+ "evaluation_scope": "deployed_roi_screening",
+ "validation_size": 44,
+ "metrics": {
+ "accuracy": 0.9091,
+ "precision": 1.0,
+ "recall": 0.7143,
+ "f1": 0.8333,
+ },
+ "operating_counts": {
+ "blocked_total": 0,
+ "likely_count": 10,
+ "uncertain_count": 3,
+ },
+ }
+ ),
+ )
+
+ status = runtime_status_module.build_runtime_status(_DummyPredictor(), _DummyGuidance())
+
+ assert status.model.validation_f1 == 0.8
+ assert status.model.deployed_accuracy == 0.9091
+ assert status.model.deployed_f1 == 0.8333
+ assert status.model.deployed_blocked_total == 0
+ assert status.model.deployed_uncertain_count == 3
+ assert status.model.runtime_calibration_ready is True
+ assert status.model.runtime_calibration_method == "temperature"
+ assert status.model.runtime_calibrated_threshold == 0.58
+ assert status.model.runtime_calibration_ece_after == 0.072
+ assert status.model.runtime_refiner_ready is True
+ assert status.model.runtime_refiner_method == "logistic-regression"
+ assert status.model.runtime_refined_threshold == 0.53
+ assert status.model.runtime_refined_f1 == 0.7857
diff --git a/backend/tests/test_triage.py b/backend/tests/test_triage.py
new file mode 100644
index 0000000000000000000000000000000000000000..f125d382669f053f9b85bf0b6d772f5d7462e279
--- /dev/null
+++ b/backend/tests/test_triage.py
@@ -0,0 +1,205 @@
+"""
+Tests for TriageService โ the decision layer that combines image quality,
+ML prediction, and self-reported symptoms into a risk band.
+
+Coverage targets:
+- Each risk band is reachable via the expected combination of inputs.
+- Band boundaries are respected when risk scores sit on either side of a threshold.
+- Quality failure always yields uncertain_retake_needed regardless of prediction.
+- Specific issue codes (eye_not_visible) surface in the triage summary.
+- Triage score is always in [0, 1].
+- Computed properties on TriageResult work correctly.
+"""
+
+from __future__ import annotations
+
+import sys
+from pathlib import Path
+
+import pytest
+
+ROOT = Path(__file__).resolve().parents[2]
+sys.path.insert(0, str(ROOT / "backend"))
+
+from app.schemas import PredictionResult, QualityAssessment, QualityIssue, SymptomInput, TriageResult
+from app.services.triage import TriageService
+
+
+# ---------------------------------------------------------------------------
+# Fixtures / factories
+# ---------------------------------------------------------------------------
+
+def _quality(
+ passed: bool = True,
+ issues: list[dict] | None = None,
+) -> QualityAssessment:
+ return QualityAssessment(
+ passed=passed,
+ blur_score=82.0,
+ brightness_score=0.46,
+ contrast_score=0.22,
+ framing_score=1.2,
+ issues=issues or [],
+ )
+
+
+def _prediction(
+ risk: float,
+ confidence: float = 0.72,
+ uncertainty: float = 0.20,
+ label: str | None = None,
+) -> PredictionResult:
+ if label is None:
+ label = "anemia_likely" if risk > 0.62 else ("anemia_unlikely" if risk < 0.35 else "uncertain")
+ return PredictionResult(
+ anemia_risk=risk,
+ confidence=confidence,
+ uncertainty=uncertainty,
+ reliability_flag="medium",
+ screening_label=label,
+ screening_text="Screening model output.",
+ model_source="archive-fusion-v2",
+ )
+
+
+SERVICE = TriageService()
+
+
+# ---------------------------------------------------------------------------
+# Happy-path band routing
+# ---------------------------------------------------------------------------
+
+class TestBandRouting:
+ def test_high_concern_for_strong_signal_with_symptoms(self) -> None:
+ result = SERVICE.assess(
+ _quality(),
+ _prediction(0.78),
+ SymptomInput(fatigue=True, dizziness=True, shortness_of_breath=True, poor_diet_low_iron=True),
+ )
+ assert result.band == "high_concern"
+
+ def test_moderate_risk_for_mid_signal_no_symptoms(self) -> None:
+ result = SERVICE.assess(_quality(), _prediction(0.52), SymptomInput())
+ assert result.band in {"moderate_risk", "high_concern"}
+
+ def test_low_risk_for_weak_signal_no_symptoms(self) -> None:
+ result = SERVICE.assess(_quality(), _prediction(0.18), SymptomInput())
+ assert result.band == "low_risk"
+
+ def test_symptoms_alone_cannot_override_quality_failure(self) -> None:
+ """Even with many symptoms, a quality failure must yield uncertain."""
+ heavy_symptoms = SymptomInput(
+ fatigue=True, dizziness=True, pale_skin=True, shortness_of_breath=True
+ )
+ result = SERVICE.assess(_quality(passed=False), None, heavy_symptoms)
+ assert result.band == "uncertain_retake_needed"
+
+
+# ---------------------------------------------------------------------------
+# Quality-failure paths
+# ---------------------------------------------------------------------------
+
+class TestQualityFailure:
+ def test_failed_quality_yields_uncertain(self) -> None:
+ result = SERVICE.assess(_quality(passed=False), None, SymptomInput(fatigue=True))
+ assert result.band == "uncertain_retake_needed"
+
+ def test_eye_not_visible_surfaces_in_summary(self) -> None:
+ quality = QualityAssessment(
+ passed=False,
+ blur_score=82.0,
+ brightness_score=0.2,
+ contrast_score=0.18,
+ framing_score=0.9,
+ issues=[
+ QualityIssue(
+ code="eye_not_visible",
+ severity="blocking",
+ title="Eye is not clearly visible",
+ message="Retake with the inner lower eyelid clearly visible.",
+ )
+ ],
+ )
+ result = SERVICE.assess(quality, None, SymptomInput())
+ assert result.band == "uncertain_retake_needed"
+ assert "inner eyelid" in result.summary.lower() or "eyelid" in result.summary.lower()
+
+ def test_blur_issue_surfaces_in_summary(self) -> None:
+ quality = _quality(
+ passed=False,
+ issues=[
+ {"code": "blur_detected", "severity": "blocking",
+ "title": "Image is blurry", "message": "Hold the camera steady and retake."}
+ ],
+ )
+ result = SERVICE.assess(quality, None, SymptomInput())
+ assert result.band == "uncertain_retake_needed"
+
+
+# ---------------------------------------------------------------------------
+# Score validity
+# ---------------------------------------------------------------------------
+
+@pytest.mark.parametrize("risk,sym_count", [
+ (0.1, 0), (0.35, 1), (0.5, 2), (0.7, 4), (0.95, 5),
+])
+def test_triage_score_always_in_unit_interval(risk: float, sym_count: int) -> None:
+ symptoms_on = list(SymptomInput.model_fields.keys())[:sym_count]
+ symptoms = SymptomInput(**{k: True for k in symptoms_on if k != "heavy_menstrual_bleeding"})
+ result = SERVICE.assess(_quality(), _prediction(risk), symptoms)
+ assert 0.0 <= result.score <= 1.0
+
+
+# ---------------------------------------------------------------------------
+# TriageResult computed properties
+# ---------------------------------------------------------------------------
+
+class TestTriageResultProperties:
+ def test_high_concern_requires_urgent_followup(self) -> None:
+ t = TriageResult(
+ band="high_concern", score=0.8, label="High concern",
+ summary="Urgent.", disclaimer="Screening only.",
+ )
+ assert t.requires_urgent_followup is True
+ assert t.requires_retake is False
+
+ def test_uncertain_requires_retake(self) -> None:
+ t = TriageResult(
+ band="uncertain_retake_needed", score=0.3, label="Uncertain",
+ summary="Retake needed.", disclaimer="Screening only.",
+ )
+ assert t.requires_retake is True
+ assert t.requires_urgent_followup is False
+
+ def test_low_risk_neither_urgent_nor_retake(self) -> None:
+ t = TriageResult(
+ band="low_risk", score=0.15, label="Low risk",
+ summary="Looking good.", disclaimer="Screening only.",
+ )
+ assert t.requires_urgent_followup is False
+ assert t.requires_retake is False
+
+
+# ---------------------------------------------------------------------------
+# Disclaimer is always present
+# ---------------------------------------------------------------------------
+
+def test_triage_result_always_has_disclaimer() -> None:
+ result = SERVICE.assess(_quality(), _prediction(0.5), SymptomInput())
+ assert len(result.disclaimer) > 20
+ assert "screening" in result.disclaimer.lower()
+
+
+def test_signal_breakdown_exposes_fusion_components() -> None:
+ quality = _quality()
+ prediction = _prediction(0.64, confidence=0.81, uncertainty=0.17)
+ symptoms = SymptomInput(fatigue=True, pale_skin=True)
+
+ breakdown = SERVICE.build_signal_breakdown(quality, prediction, symptoms)
+
+ assert breakdown.image_risk == 0.64
+ assert breakdown.symptom_score == pytest.approx(0.34)
+ assert breakdown.fused_score == pytest.approx((0.64 * 0.72) + (0.34 * 0.28))
+ assert breakdown.image_weight == 0.72
+ assert breakdown.symptom_weight == 0.28
+ assert breakdown.reliability_flag == "medium"