An end-to-end audit found the repo could not build, test or run as shipped. This fixes every finding, then adds a Cloud Run track so the demo costs about £1/month idle instead of ~£150. CI (red on its first run) - api: setuptools could not build the package (flat layout with app/ and alembic/) - web: missing @types/node; `vitest run` exited 1 with no test files - pipeline: the stub run needed a gitignored VCF, and no process had a stub block - ruff pinned, mypy configured, DB tests on real Postgres (pgserver locally, service in CI) ML serving (scores were meaningless) - the registered model now carries its own feature engineering and returns predict_proba, so serving sends raw columns and cannot drift from training - resolve by registry alias (stages are deprecated in MLflow 3) and record the real version; re-scoring upserts instead of failing on the unique constraint - ClinVar labels parsed from VEP's lowercase terms Pipeline - exact ref/alt recovered from a CHROM_POS_REF_ALT VCF ID; loading is idempotent - job status reaches running/failed/succeeded, so the UI stops polling dead jobs - DATABASE_URL travels in the environment or a Nextflow secret, never on a command line - VEP cache and plugins staged as inputs; the gcp profile runs tasks on Google Batch Deployment - the API serves /api (matching the ingress); the web app reads its API URL at runtime - migrations run in an init container under a Postgres advisory lock - terraform: custom VPC shared with Batch, private Cloud SQL, API enablement, Workload Identity bindings, Secret Manager, deletion protection - serverless track, now the default: Cloud Run services scaling to zero, a Cloud Run job for the Nextflow driver, and Neon or Cloud SQL behind one DATABASE_URL secret. GKE and Argo remain, behind -var deploy_kubernetes=true. See docs/cloud.md. Correctness and security - 409 on duplicate sample names, 422 on bad paging, natural chromosome ordering, wider VEP text columns, enum dropped on downgrade, the sample's assembly actually used - vcf_uri restricted to gs:// objects or files under the data root, blocking option injection - CORS restricted to configured origins; `make down` no longer deletes volumes Data - docs/data.md records the peer-reviewed, openly licensed sources (GIAB HG002, ClinVar, gnomAD) with citations and an honest evaluation plan; `make data` fetches a chr22 slice Verified: api 50 tests, ml 18, loader 16, web 12; ruff, mypy, svelte-check, terraform validate and both kustomize overlays clean.
108 lines
3.9 KiB
Python
108 lines
3.9 KiB
Python
"""Score a job's variants with the registered MLflow model.
|
|
|
|
The registered model is a pyfunc that owns its feature engineering (rarelens_ml.features travels
|
|
with it as model code) and returns P(pathogenic). Serving therefore only sends the raw columns
|
|
below and cannot drift from training.
|
|
"""
|
|
import asyncio
|
|
import uuid
|
|
from collections.abc import Sequence
|
|
from typing import Any
|
|
|
|
import mlflow
|
|
import pandas as pd
|
|
from mlflow import MlflowClient
|
|
from sqlalchemy import func, select
|
|
from sqlalchemy.dialects.postgresql import insert
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.config import settings
|
|
from app.models import Prediction, Variant
|
|
|
|
# Must match rarelens_ml.features.RAW_COLUMNS.
|
|
RAW_COLUMNS = ["impact", "consequence", "gnomad_af", "cadd_phred", "am_pathogenicity"]
|
|
CHUNK_SIZE = 5000
|
|
|
|
_models: dict[str, Any] = {} # model version -> loaded pyfunc
|
|
|
|
|
|
def load_model() -> tuple[Any, str]:
|
|
"""Return (model, version), from MODEL_URI if set, else from the registry alias."""
|
|
if settings.model_uri:
|
|
return _load_uri(settings.model_uri)
|
|
client = MlflowClient(tracking_uri=settings.mlflow_tracking_uri)
|
|
version = client.get_model_version_by_alias(settings.model_name, settings.model_alias).version
|
|
if version not in _models:
|
|
mlflow.set_tracking_uri(settings.mlflow_tracking_uri)
|
|
_models.clear()
|
|
_models[version] = mlflow.pyfunc.load_model(f"models:/{settings.model_name}/{version}")
|
|
return _models[version], version
|
|
|
|
|
|
def _load_uri(uri: str) -> tuple[Any, str]:
|
|
"""Load a model artifact directly (gs://...): no tracking server, nothing running when idle."""
|
|
version = uri.rstrip("/").rsplit("/", 1)[-1][:40] or "uri"
|
|
if version not in _models:
|
|
_models.clear()
|
|
_models[version] = mlflow.pyfunc.load_model(uri)
|
|
return _models[version], version
|
|
|
|
|
|
def raw_frame(variants: Sequence[Variant]) -> pd.DataFrame:
|
|
return pd.DataFrame(
|
|
{
|
|
"impact": [v.impact for v in variants],
|
|
"consequence": [v.consequence for v in variants],
|
|
"gnomad_af": [v.gnomad_af if v.gnomad_af is not None else float("nan") for v in variants],
|
|
"cadd_phred": [v.annotations.get("CADD_PHRED") for v in variants],
|
|
"am_pathogenicity": [v.annotations.get("am_pathogenicity") for v in variants],
|
|
},
|
|
columns=RAW_COLUMNS,
|
|
)
|
|
|
|
|
|
async def score_job(job_id: uuid.UUID, session: AsyncSession) -> tuple[int, str]:
|
|
model, version = await asyncio.to_thread(load_model)
|
|
scored = 0
|
|
last_id = 0
|
|
# Keyset pagination keeps memory flat for whole-genome jobs.
|
|
while True:
|
|
variants = (
|
|
await session.scalars(
|
|
select(Variant)
|
|
.where(Variant.job_id == job_id, Variant.id > last_id)
|
|
.order_by(Variant.id)
|
|
.limit(CHUNK_SIZE)
|
|
)
|
|
).all()
|
|
if not variants:
|
|
break
|
|
scores = await asyncio.to_thread(model.predict, raw_frame(variants))
|
|
stmt = insert(Prediction).values(
|
|
[
|
|
{
|
|
"variant_id": v.id,
|
|
"model_name": settings.model_name,
|
|
"model_version": version,
|
|
"score": float(s),
|
|
}
|
|
for v, s in zip(variants, scores, strict=True)
|
|
]
|
|
)
|
|
# Re-scoring (e.g. after a new model version) replaces the previous prediction.
|
|
await session.execute(
|
|
stmt.on_conflict_do_update(
|
|
index_elements=[Prediction.variant_id],
|
|
set_={
|
|
"model_name": stmt.excluded.model_name,
|
|
"model_version": stmt.excluded.model_version,
|
|
"score": stmt.excluded.score,
|
|
"created_at": func.now(),
|
|
},
|
|
)
|
|
)
|
|
scored += len(variants)
|
|
last_id = variants[-1].id
|
|
await session.commit()
|
|
return scored, version
|