Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,7 @@ The hybrid backbone depends on `mamba-ssm`, which needs CUDA/`nvcc` to build. On
uv sync --extra cuda --no-build-isolation
```

CPU/MPS development uses a lightweight stand-in backbone so the concept module and the harness can be built and tested without a GPU. The notes-sidecar text pipeline (`odyssey/text/`) needs `uv sync --extra text` (see `docs/sidecars_and_task_sets.md`).
Without `mamba-ssm` (CPU, Apple silicon), the hybrid backbone runs on a pure-PyTorch version of its layers with the same weight names, so a GPU-trained checkpoint loads and runs for inference; training still needs CUDA. To use a trained checkpoint, see [`docs/checkpoints.md`](docs/checkpoints.md). The notes-sidecar text pipeline (`odyssey/text/`) needs `uv sync --extra text` (see `docs/sidecars_and_task_sets.md`).

## Data pipeline

Expand Down
10 changes: 7 additions & 3 deletions apps/clinician_demo/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
from pathlib import Path

from apps.clinician_demo.config import DATA_MODES, DemoConfig
from odyssey.utils.device import resolve_device


logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -53,7 +54,9 @@ def parse_args(argv: list[str] | None = None) -> tuple[DemoConfig, bool]:
parser.add_argument("--alert-rate", type=float, default=0.05)
parser.add_argument("--max-shards", type=int, default=None)
parser.add_argument("--cache-dir", type=Path, default=None)
parser.add_argument("--device", default="cuda")
parser.add_argument(
"--device", default="auto", help="cuda, mps, cpu, or auto (first available)"
)
parser.add_argument("--no-warmup", action="store_true")
parser.add_argument("--self-check", action="store_true")
args = parser.parse_args(argv)
Expand All @@ -69,7 +72,7 @@ def parse_args(argv: list[str] | None = None) -> tuple[DemoConfig, bool]:
alert_rate=args.alert_rate,
max_shards=args.max_shards,
cache_dir=args.cache_dir.expanduser() if args.cache_dir else None,
device=args.device,
device=resolve_device(args.device),
warmup=not args.no_warmup,
)
except ValueError as exc:
Expand Down Expand Up @@ -102,9 +105,10 @@ def main(argv: list[str] | None = None) -> int:
service.warm_up()
server = make_server(service, config.host, config.port)
logger.info(
"serving %s (%s mode) on http://%s:%d -- open an SSH tunnel to this port",
"serving %s (%s mode) on %s at http://%s:%d (over an SSH tunnel from a GPU host)",
config.run_name,
config.data_mode,
config.device,
config.host,
config.port,
)
Expand Down
42 changes: 42 additions & 0 deletions apps/clinician_demo/export_thresholds.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
"""Export a run's alert lines as aggregates, for serving the demo elsewhere.

The demo sets its alert lines from the run's patient-level
``alerts_rows.parquet``, which stays on the GPU host. Run this there once
per run; copy the resulting JSON next to the checkpoint and the demo uses
it on a host without the rows (for example a laptop)::

python -m apps.clinician_demo.export_thresholds --run-dir ~/runs/full_run_v10
"""

import argparse
import sys
from pathlib import Path

from apps.clinician_demo.config import HORIZONS_HOURS
from apps.clinician_demo.thresholds import (
AGGREGATE_THRESHOLDS_FILENAME,
ALERTS_ROWS_FILENAME,
export_operating_points,
)


def main(argv: list[str] | None = None) -> int:
"""Write ``<run-dir>/demo_thresholds_aggregate.json``; return the exit code."""
parser = argparse.ArgumentParser(
prog="python -m apps.clinician_demo.export_thresholds", description=__doc__
)
parser.add_argument("--run-dir", type=Path, required=True)
parser.add_argument("--alert-rate", type=float, default=0.05)
args = parser.parse_args(argv)
run_dir = args.run_dir.expanduser()
rows = run_dir / ALERTS_ROWS_FILENAME
if not rows.exists():
parser.error(f"no {ALERTS_ROWS_FILENAME} in {run_dir}")
out = run_dir / AGGREGATE_THRESHOLDS_FILENAME
points = export_operating_points(rows, out, HORIZONS_HOURS, args.alert_rate)
print(f"wrote {len(points)} alert lines to {out}")
return 0


if __name__ == "__main__":
sys.exit(main())
18 changes: 7 additions & 11 deletions apps/clinician_demo/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -71,7 +71,7 @@
from apps.clinician_demo.thresholds import (
ALERTS_ROWS_FILENAME,
horizon_key,
load_or_compute_operating_points,
operating_points_for_run,
)
from apps.clinician_demo.whatif import PRESETS, parse_edit_requests, run_whatif
from odyssey.data.alert_events import alert_events_for
Expand Down Expand Up @@ -296,16 +296,12 @@ def from_config(cls, config: DemoConfig) -> "DemoService":
describe=admission_label,
)
rows_path = config.run_dir / ALERTS_ROWS_FILENAME
points = (
load_or_compute_operating_points(
rows_path,
config.resolved_cache_dir / "thresholds.json",
ctx.events,
config.horizons,
config.alert_rate,
)
if rows_path.exists()
else []
points = operating_points_for_run(
config.run_dir,
config.resolved_cache_dir / "thresholds.json",
ctx.events,
config.horizons,
config.alert_rate,
)
concepts = [
ConceptInfo(d.name, concept_label(d.name), d.description, None)
Expand Down
87 changes: 87 additions & 0 deletions apps/clinician_demo/thresholds.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,9 @@
logger = logging.getLogger(__name__)

ALERTS_ROWS_FILENAME = "alerts_rows.parquet"
# Alert lines exported from the GPU host (aggregates only), for running the
# demo where the patient-level ``alerts_rows.parquet`` is not present.
AGGREGATE_THRESHOLDS_FILENAME = "demo_thresholds_aggregate.json"
CACHE_VERSION = 1


Expand Down Expand Up @@ -186,11 +189,95 @@ def load_or_compute_operating_points(
return points


def export_operating_points(
rows_path: str | Path,
out_path: str | Path,
horizons: Sequence[float],
alert_rate: float,
) -> list[OperatingPoint]:
"""Write the alert lines for every event in ``rows_path`` to ``out_path``.

Runs where the patient-level row file lives (the GPU host). The output
holds aggregates only, so it can travel with the checkpoint to a host
without the rows; :func:`load_aggregate_operating_points` reads it.
"""
rows_path = Path(rows_path)
events = sorted(
pl.scan_parquet(rows_path).select(pl.col("event").unique()).collect()["event"]
)
points = compute_operating_points(rows_path, events, horizons, alert_rate)
Path(out_path).write_text(
json.dumps({"alert_rate": alert_rate, "points": to_jsonable(points)}, indent=1)
)
return points


def load_aggregate_operating_points(
path: str | Path,
events: Sequence[str],
horizons: Sequence[float],
alert_rate: float,
) -> list[OperatingPoint]:
"""Read exported alert lines; empty when the file is absent or does not fit.

The file holds ``{"alert_rate": ..., "points": [...]}`` as written on the
GPU host from :func:`compute_operating_points`. Points are kept only for
the requested events and horizons, and only if the alert rate matches.
"""
path = Path(path)
if not path.exists():
return []
payload = json.loads(path.read_text())
if payload.get("alert_rate") != alert_rate:
logger.warning(
"[thresholds] %s was exported at alert rate %s, not %s; no alert lines",
path,
payload.get("alert_rate"),
alert_rate,
)
return []
wanted = {(e, float(h)) for e in events for h in horizons}
return [
OperatingPoint(**p)
for p in payload["points"]
if (p["event"], float(p["horizon_hours"])) in wanted
]


def operating_points_for_run(
run_dir: str | Path,
cache_path: str | Path,
events: Sequence[str],
horizons: Sequence[float],
alert_rate: float,
) -> list[OperatingPoint]:
"""Return a run's alert lines from its rows if present, else its export.

On the GPU host the patient-level ``alerts_rows.parquet`` is read (and
the result cached in ``cache_path``). Elsewhere the exported aggregates
(:data:`AGGREGATE_THRESHOLDS_FILENAME`) are used; with neither, the demo
runs without alert lines.
"""
run_dir = Path(run_dir)
rows_path = run_dir / ALERTS_ROWS_FILENAME
if rows_path.exists():
return load_or_compute_operating_points(
rows_path, cache_path, events, horizons, alert_rate
)
return load_aggregate_operating_points(
run_dir / AGGREGATE_THRESHOLDS_FILENAME, events, horizons, alert_rate
)


__all__ = [
"AGGREGATE_THRESHOLDS_FILENAME",
"ALERTS_ROWS_FILENAME",
"compute_operating_points",
"export_operating_points",
"flag_threshold",
"horizon_key",
"load_aggregate_operating_points",
"load_or_compute_operating_points",
"operating_point",
"operating_points_for_run",
]
Loading
Loading