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)} +
+
+
+ + + + + {steps_html} +
+
Recommended next steps
+
+
+ + + + +
+ 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"