diff --git a/README.md b/README.md index bdefeba7..4acb8bbc 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/apps/clinician_demo/__main__.py b/apps/clinician_demo/__main__.py index 0ba5b366..99f62609 100644 --- a/apps/clinician_demo/__main__.py +++ b/apps/clinician_demo/__main__.py @@ -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__) @@ -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) @@ -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: @@ -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, ) diff --git a/apps/clinician_demo/export_thresholds.py b/apps/clinician_demo/export_thresholds.py new file mode 100644 index 00000000..857d7dbd --- /dev/null +++ b/apps/clinician_demo/export_thresholds.py @@ -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 ``/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()) diff --git a/apps/clinician_demo/service.py b/apps/clinician_demo/service.py index f0a6e2a8..cce66689 100644 --- a/apps/clinician_demo/service.py +++ b/apps/clinician_demo/service.py @@ -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 @@ -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) diff --git a/apps/clinician_demo/thresholds.py b/apps/clinician_demo/thresholds.py index df8975f7..54502dff 100644 --- a/apps/clinician_demo/thresholds.py +++ b/apps/clinician_demo/thresholds.py @@ -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 @@ -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", ] diff --git a/docs/checkpoints.md b/docs/checkpoints.md new file mode 100644 index 00000000..7578bb99 --- /dev/null +++ b/docs/checkpoints.md @@ -0,0 +1,176 @@ +# Using a trained checkpoint + +This page is for people who receive a trained Odyssey checkpoint and want to +run it on their own data. It says what the checkpoint is, how it was +trained, and how to feed it data the way it expects. + +## Older versions of this repository + +Early versions of Odyssey had an `EHR-Mamba3` model, a default config +`odyssey/models/configs/ehr_mamba3.yaml` and a MEDS script +`scripts/meds/run_pipeline.sh`. These were removed in commit `3eef618`, when +the repository was rebuilt around the concept-bottleneck model. The Mamba-3 +backbone was then replaced by Mamba-2 (commits `f220d2c`, `febbef5`): its +kernels could not carry state correctly across chunks, which streaming +training needs. Code, configs and checkpoints from that era are not +compatible with the current repository. Use the current code and this page. + +## What a checkpoint directory holds + +A run directory written by `odyssey.training.train` holds everything needed +to run the model: + +| File | What it is | +| --- | --- | +| `checkpoint_best.pt` | weights (`model` key) from the step with the lowest validation loss; the other keys (`optimizer`, `step`, ...) are only for resuming training | +| `config.json` | the full training configuration (architecture, data source, task set) | +| `vocabulary.json` | the token vocabulary, built on the training split | +| `quantile_binner.json` | the value bins, fit on the training split | + +`odyssey.inference.run_inference.load_run` rebuilds the model from these +four files. It reads the architecture from the checkpoint's own weights +where the config could be ambiguous, so older run directories still load. + +## The MIMIC-IV checkpoint (`full_run_v10`) + +**Model.** Hybrid backbone: 8 blocks, each with a Mamba-2 branch and a +chunk-local attention branch run in parallel, hidden size 256. A concept +bottleneck reads the hidden state into 29 clinical concepts (task set `v3`; +mixture form, see the README). Heads on top: next-event forecasting, time +to the next event, and a discrete-time hazard for six events (ICU +admission, vasopressor start, acute kidney injury, Sepsis-3, death, 30-day +readmission). 32.2 million parameters. + +**Training.** One joint training run from random initialization. There is +no separate pretraining stage and no fine-tuning stage: the next-event, +timing, concept and hazard losses are all trained together from the start. +Two epochs over all 292 training shards; the best validation loss (2.056) +was at step 37,000. Trained at commit `cdbd4e7`; the registry row is in +[`experiments.md`](experiments.md) under `full_run_v10`. + +**Data.** MIMIC-IV 3.1, `hosp` and `icu` modules, extracted to MEDS with the +standard `meds-extract` tooling (see the README's "Data pipeline"). Subjects +are split train / tuning / held_out by the extraction's +`metadata/subject_splits.parquet`. The model saw the train and tuning +subjects. **Score only held_out subjects** if you report numbers on +MIMIC-IV; use the split file that comes with the checkpoint, not one from a +fresh extraction. + +**How well it does.** Held-out results (concept readout AUROCs, alert +AUROCs against a tuned gradient-boosted baseline) are in `alerts.json` and +`inference_results.json` beside the checkpoint, and in the paper. The +30-day readmission head is weak (AUROC about 0.59 at 7 days); do not use it +as a baseline. + +## Install + +Python 3.12 or later and [uv](https://github.com/astral-sh/uv): + +```bash +git clone https://github.com/VectorInstitute/odyssey.git +cd odyssey +uv sync --dev +``` + +That is enough to run the model on a CPU or an Apple-silicon GPU (MPS). +Without `mamba-ssm`, the hybrid backbone uses a pure-PyTorch version of its +layers (`odyssey/models/backbones/mamba_portable.py`) with the same weight +names, so the same checkpoint loads. It is for inference only. On an NVIDIA +GPU, install the CUDA kernels for speed (and for training): + +```bash +uv sync --extra cuda --no-build-isolation +``` + +On 12 patients (about 165,000 positions), the portable layers on Apple MPS +match the CUDA kernels to a mean risk difference of 0.0002 (largest 0.009), +and the top next-event forecast agrees at 99.8% of positions. + +## Prepare your data + +1. **Extract to MEDS** with the same tooling and spec the checkpoint was + trained on (for MIMIC-IV: `meds-extract-run spec=MIMIC-IV ...`, see the + README). Another source needs its own spec and its own codes; a + checkpoint trained on MIMIC-IV only knows MIMIC-IV codes. +2. **Do not rebuild the vocabulary or the value bins.** Tokenize with the + checkpoint's `vocabulary.json` and `quantile_binner.json`. Codes the + vocabulary does not know become an unknown token (with a coarser ICD + back-off for diagnoses), so many unknown codes mean the data does not + match the training extraction. +3. **Apply the run's own preprocessing.** `config.json` records it + (`source`, `normalize_medications`, `history_recap`); the example below + applies it from the loaded config. +4. **Sidecars are not model input.** `/sidecars/` (for example + microbiology for Sepsis-3) only feeds the labels used in evaluation. See + [`sidecars_and_task_sets.md`](sidecars_and_task_sets.md). + +## Run it + +**Evaluate on a shard directory** (forecasting, concept readout, +completeness), the same way the paper did: + +```bash +uv run python -m odyssey.inference.run_inference --run-dir \ + --held-out-shard-dir /data/held_out --output-json results.json +uv run python -m odyssey.inference.alerts --run-dir \ + --held-out-shard-dir /data/held_out --output-json alerts.json \ + --baseline-shard-dir /data/train # the GBM baseline is fit here +``` + +**Score one patient** and read the event risks and concept beliefs at every +position of the record: + +```python +import torch + +from odyssey.data.code_normalization import maybe_normalize +from odyssey.data.history_recap import maybe_history_recap +from odyssey.data.sequences import build_patient_sequence +from odyssey.data.value_binning import add_value_tokens +from odyssey.inference.patient_stream import risk_within, stream_patient +from odyssey.inference.run_inference import load_run +from odyssey.training.data import load_meds_subject +from odyssey.utils.device import default_device + +run_dir, shard, subject_id = "full_run_v10", "/data/held_out/0.parquet", 123 +device = default_device() # cuda, else Apple mps, else cpu +model, vocab, binner, config = load_run( + run_dir, device=device, checkpoint_path=f"{run_dir}/checkpoint_best.pt" +) +model.eval() + +# Prepare the record exactly as training did: the run's own normalization, +# value bins and vocabulary. +events = load_meds_subject(shard, subject_id) +events = maybe_normalize(events, enabled=config.normalize_medications, source=config.source) +events = maybe_history_recap(events, enabled=config.history_recap) +seq = build_patient_sequence(add_value_tokens(events, binner, source=config.source), vocab) + +heads = model.event_heads +risks, concepts = [], [] +for span in stream_patient(model, seq, device=device, chunk_size=config.chunk_size): + n = span.n_real + hazards = heads(span.fwd.features[0, :n]) # (positions, events, bins) + risks.append(risk_within(hazards, heads.edges, [8.0, 24.0, 72.0]).cpu()) + concepts.append(span.fwd.bottleneck.concept_probs[0, :n].cpu()) +risk = torch.cat(risks) # (positions, events, horizons): P(event within h) +concept_probs = torch.cat(concepts) # (positions, concepts) +for name, r in zip(heads.event_names, risk[-1, :, 1].tolist()): + print(f"{name}: {r:.3f} within 24 h") +``` + +`stream_patient` feeds the record in chunks and carries the recurrent state +between them, so a record of any length runs in constant memory. On an +Apple M4 an 11,000-event record takes about 4 seconds. + +**See it in a browser.** The clinician demo (`apps/clinician_demo`, see +[`clinician_demo.md`](clinician_demo.md)) replays one admission at a time +with the risks, concept beliefs, what-if edits and evidence. + +## Data governance + +A checkpoint trained on credentialed data (MIMIC-IV, eICU-CRD) is shared +only with people who hold that dataset's credentialed access and have +signed its data use agreement, and only for the purpose agreed. Do not +redistribute it. The `subject_splits.parquet` that comes with it lists +credentialed subject identifiers and is covered by the same terms. diff --git a/docs/clinician_demo.md b/docs/clinician_demo.md index 5d9bea3e..bffa8b29 100644 --- a/docs/clinician_demo.md +++ b/docs/clinician_demo.md @@ -51,7 +51,61 @@ The open demo extraction is made once with (the pipeline shells out to `MEDS_transform-stage`, so the venv must be on `PATH`). -## View it (laptop) +## Run it on a laptop (Apple silicon or CPU) + +The hybrid backbone needs the CUDA `mamba-ssm` kernels to train, but not to +run. Without `mamba-ssm`, `EHRHybridBackbone` builds from +`odyssey/models/backbones/mamba_portable.py`: the same layers in plain +PyTorch, with the same parameter names, so a GPU checkpoint loads +unchanged. `--device auto` (the default) picks CUDA, then Apple MPS, then +CPU. On an M4, MPS traces a 6,000-entry stay in about 10 s. + +Checked against the GPU on 12 open-demo patients (about 165,000 positions, +full_run_v10): risk differs by 0.0002 on average (99th percentile 0.002, +largest 0.009), concept beliefs by at most 0.006, and the top next-event +forecast agrees at 99.8% of positions. The GPU's TF32 and Triton +accumulation order explain the gap; CPU and MPS agree with each other more +closely than either agrees with the GPU. + +Only open mode belongs on a laptop: credentialed patient data stays on the +GPU host. Copy these files from the GPU host: + +- From the run directory: `checkpoint_best.pt`, `config.json`, + `vocabulary.json`, `quantile_binner.json`, and the scorecard aggregates + `alerts.json`, `alerts_cis.json`, `inference_results.json`. +- The alert lines, as aggregates. The demo sets them from the patient-level + `alerts_rows.parquet`, which stays on the host. Export them once per run; + the demo reads the export when the rows file is absent: + + ```bash + .venv/bin/python -m apps.clinician_demo.export_thresholds --run-dir ~/runs/full_run_v10 + ``` + + This writes `demo_thresholds_aggregate.json` into the run directory. +- The open MIMIC-IV demo extraction (`data/`, `metadata/`, `sidecars/`). +- The model's split for the 100 demo patients only, so the UI can say + which ones it trained on: + + ```bash + .venv/bin/python -c " + import polars as pl + ids = pl.scan_parquet('$HOME/data/mimiciv_demo_meds/data/**/*.parquet').select('subject_id').unique().collect() + pl.read_parquet('$HOME/data/mimiciv_3.1_v1/metadata/subject_splits.parquet').join(ids, on='subject_id').write_parquet('demo_subject_splits.parquet')" + ``` + +Then, from the repository root on the laptop: + +```bash +.venv/bin/python -m apps.clinician_demo --data-mode open --port 8766 \ + --run-dir /runs/full_run_v10 \ + --data-dir /mimiciv_demo_meds/data \ + --metadata-dir /mimiciv_demo_meds/metadata \ + --splits /demo_subject_splits.parquet +``` + +and open http://localhost:8766. No tunnel is needed. + +## View it from a laptop (server on the GPU host) ```bash gcloud compute ssh odyssey-cbm-a100 --zone us-central1-f \ diff --git a/odyssey/data/streaming.py b/odyssey/data/streaming.py index 385ea694..3f3ff784 100644 --- a/odyssey/data/streaming.py +++ b/odyssey/data/streaming.py @@ -70,6 +70,7 @@ from odyssey.data.sequences import NO_VISIT, PatientSequence from odyssey.data.types import AuxiliaryInputs, ClinicalSequenceBatch from odyssey.data.vocabulary import PAD_ID +from odyssey.utils.device import to_device # Sentinel `subject_ids` value at padding positions. @@ -423,10 +424,11 @@ def move_to_device(chunk: _MovableT, device: str) -> _MovableT: Works for :class:`StreamingChunk` and its nested batch/aux tuples without depending on their exact field lists, so a new field added to - any of them needs no matching change here. + any of them needs no matching change here. Tensors move with + :func:`~odyssey.utils.device.to_device` (float64 becomes float32 on MPS). """ if isinstance(chunk, torch.Tensor): - return chunk.to(device) # type: ignore[return-value] + return to_device(chunk, device) # type: ignore[return-value] if isinstance(chunk, tuple) and hasattr(chunk, "_fields"): # NamedTuple return type(chunk)(*(move_to_device(v, device) for v in chunk)) return chunk diff --git a/odyssey/inference/alerts.py b/odyssey/inference/alerts.py index c3231b2d..16fa37e1 100644 --- a/odyssey/inference/alerts.py +++ b/odyssey/inference/alerts.py @@ -75,7 +75,12 @@ from odyssey.data.packed_context import PackedContextSampler from odyssey.data.sequences import BIRTH_CODE from odyssey.data.sidecars import activate_sidecars -from odyssey.data.streaming import NO_SUBJECT, PackedLaneSampler, StreamingChunk +from odyssey.data.streaming import ( + NO_SUBJECT, + PackedLaneSampler, + StreamingChunk, + move_to_device, +) from odyssey.data.value_binning import ( VALUE_Z_CLIP, QuantileBinner, @@ -101,7 +106,6 @@ merge_event_times, shard_paths, ) -from odyssey.training.train import _move_chunk_to_device logger = logging.getLogger(__name__) @@ -489,7 +493,7 @@ def collect_model_scores( landmark_state: LandmarkState | None = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 if packed: landmark_state = None # see the docstring: never carried fwd = model.forward_with_features( @@ -721,7 +725,7 @@ def emit( state = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 fwd = model.forward_with_features( chunk.batch, state=state, reset_mask=chunk.reset_mask ) diff --git a/odyssey/inference/concept_attribution.py b/odyssey/inference/concept_attribution.py index 5770660d..5f40acea 100644 --- a/odyssey/inference/concept_attribution.py +++ b/odyssey/inference/concept_attribution.py @@ -46,7 +46,7 @@ from odyssey.data.code_normalization import maybe_normalize from odyssey.data.history_recap import maybe_history_recap from odyssey.data.sidecars import activate_sidecars -from odyssey.data.streaming import PackedLaneSampler +from odyssey.data.streaming import PackedLaneSampler, move_to_device from odyssey.data.value_binning import add_value_tokens from odyssey.data.vocabulary import Vocabulary from odyssey.inference.legacy_concept_pins import resolve_concepts_for_run @@ -54,7 +54,6 @@ from odyssey.models.concept_bottleneck import ConceptBottleneck from odyssey.models.sequence_model import ConceptBottleneckSequenceModel from odyssey.training.data import iter_patient_sequences, load_meds_shards -from odyssey.training.train import _move_chunk_to_device logger = logging.getLogger(__name__) @@ -184,7 +183,7 @@ def run_streaming_attribution( # noqa: PLR0915 -- one linear scoring pass state = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 hidden, state = model.backbone( chunk.batch, state=state, reset_mask=chunk.reset_mask ) @@ -297,7 +296,7 @@ def mean_concept_directions( state = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 hidden, state = model.backbone( chunk.batch, state=state, reset_mask=chunk.reset_mask ) diff --git a/odyssey/inference/counterfactual.py b/odyssey/inference/counterfactual.py index c36c7ba8..506c3ed2 100644 --- a/odyssey/inference/counterfactual.py +++ b/odyssey/inference/counterfactual.py @@ -50,6 +50,7 @@ from odyssey.inference.patient_stream import risk_within, stream_patient from odyssey.models.sequence_model import SequenceModel from odyssey.training.data import load_meds_shards +from odyssey.utils.device import default_device logger = logging.getLogger(__name__) @@ -560,7 +561,7 @@ def _main() -> None: parser.add_argument("--keep-per-subject", action="store_true") args = parser.parse_args() - device = "cuda" if torch.cuda.is_available() else "cpu" + device = default_device() run_dir = Path(args.run_dir) model, vocab, binner, config = load_run( run_dir, diff --git a/odyssey/inference/embedding_probe.py b/odyssey/inference/embedding_probe.py index 9dfaf39f..19f06490 100644 --- a/odyssey/inference/embedding_probe.py +++ b/odyssey/inference/embedding_probe.py @@ -25,13 +25,12 @@ import torch from odyssey.data.alert_events import AlertEvent, EventTimes -from odyssey.data.streaming import PackedLaneSampler +from odyssey.data.streaming import PackedLaneSampler, move_to_device from odyssey.data.vocabulary import Vocabulary from odyssey.inference.alerts import IndexRow, LandmarkState, _select_index_positions from odyssey.inference.alerts import outcome_at_horizon as _outcome_at_horizon from odyssey.models.sequence_model import ConceptBottleneckSequenceModel from odyssey.training.data import iter_patient_sequences -from odyssey.training.train import _move_chunk_to_device Key = tuple[int, int, float] @@ -75,7 +74,7 @@ def collect_embeddings( landmark_state: LandmarkState | None = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 hidden_states, state = model.backbone( chunk.batch, state=state, reset_mask=chunk.reset_mask ) diff --git a/odyssey/inference/interventions.py b/odyssey/inference/interventions.py index 38b2c90f..c6835303 100644 --- a/odyssey/inference/interventions.py +++ b/odyssey/inference/interventions.py @@ -108,7 +108,12 @@ from odyssey.data.code_normalization import maybe_normalize from odyssey.data.history_recap import maybe_history_recap from odyssey.data.sidecars import activate_sidecars -from odyssey.data.streaming import NO_SUBJECT, PackedLaneSampler, StreamingChunk +from odyssey.data.streaming import ( + NO_SUBJECT, + PackedLaneSampler, + StreamingChunk, + move_to_device, +) from odyssey.data.value_binning import add_value_tokens from odyssey.data.vocabulary import PAD_ID, Vocabulary from odyssey.inference.alerts import ( @@ -149,7 +154,6 @@ load_meds_shards, ) from odyssey.training.running_labels import position_running_labels -from odyssey.training.train import _move_chunk_to_device logger = logging.getLogger(__name__) @@ -726,7 +730,7 @@ def run_streaming_intervention( # noqa: PLR0912, PLR0915 -- one linear scoring state = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 intervention = ( None if mode in CALIBRATED_MODES # built below, from the model's own probs diff --git a/odyssey/inference/leakage.py b/odyssey/inference/leakage.py index c1f31ce6..b2de0c95 100644 --- a/odyssey/inference/leakage.py +++ b/odyssey/inference/leakage.py @@ -98,7 +98,7 @@ from torch import nn from odyssey.data.sequences import PatientSequence -from odyssey.data.streaming import PackedLaneSampler +from odyssey.data.streaming import PackedLaneSampler, move_to_device from odyssey.data.value_binning import add_value_tokens from odyssey.data.vocabulary import PAD_ID, Vocabulary from odyssey.inference.run_inference import ( @@ -116,7 +116,6 @@ ) from odyssey.training.data import iter_patient_sequences from odyssey.training.running_labels import position_running_labels -from odyssey.training.train import _move_chunk_to_device logger = logging.getLogger(__name__) @@ -284,7 +283,7 @@ def bank_valid( state = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 fwd = model.forward_with_features( chunk.batch, state=state, reset_mask=chunk.reset_mask ) diff --git a/odyssey/inference/patient_stream.py b/odyssey/inference/patient_stream.py index ffe47c4a..993fc31d 100644 --- a/odyssey/inference/patient_stream.py +++ b/odyssey/inference/patient_stream.py @@ -18,10 +18,9 @@ import torch from odyssey.data.sequences import PatientSequence -from odyssey.data.streaming import NO_SUBJECT, PackedLaneSampler +from odyssey.data.streaming import NO_SUBJECT, PackedLaneSampler, move_to_device from odyssey.models.sequence_model import ForwardWithFeatures, SequenceModel from odyssey.models.time_to_event import probability_within -from odyssey.training.train import _move_chunk_to_device @dataclass(frozen=True) @@ -65,7 +64,7 @@ def stream_patient( state = None offset = 0 for raw_chunk in sampler: - chunk = _move_chunk_to_device(raw_chunk, device) + chunk = move_to_device(raw_chunk, device) with torch.no_grad(): fwd = model.forward_with_features( chunk.batch, state=state, reset_mask=chunk.reset_mask diff --git a/odyssey/inference/rollouts.py b/odyssey/inference/rollouts.py index 70965669..f03a6e2c 100644 --- a/odyssey/inference/rollouts.py +++ b/odyssey/inference/rollouts.py @@ -41,12 +41,11 @@ from odyssey.data.alert_events import AlertEvent from odyssey.data.sequences import PatientSequence -from odyssey.data.streaming import PackedLaneSampler +from odyssey.data.streaming import PackedLaneSampler, move_to_device from odyssey.data.types import AuxiliaryInputs, ClinicalSequenceBatch from odyssey.data.vocabulary import PAD_ID, Vocabulary, code_type from odyssey.models.sequence_model import SequenceModel from odyssey.models.time_to_event import survival_curve -from odyssey.training.train import _move_chunk_to_device logger = logging.getLogger(__name__) @@ -241,7 +240,7 @@ def rollout_from_position( for chunk in PackedLaneSampler( iter([prefix]), num_lanes=1, chunk_size=chunk_size, reset_prob=0.0 ): - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 fwd = model.forward_with_features( chunk.batch, state=state, reset_mask=chunk.reset_mask ) diff --git a/odyssey/inference/run_inference.py b/odyssey/inference/run_inference.py index 72e28bfd..8028239b 100644 --- a/odyssey/inference/run_inference.py +++ b/odyssey/inference/run_inference.py @@ -38,7 +38,7 @@ from odyssey.data.sequences import PatientSequence from odyssey.data.sidecars import activate_sidecars from odyssey.data.signal_panel import SIGNAL_PANEL, SignalPanelResolver -from odyssey.data.streaming import PackedLaneSampler, StreamingChunk +from odyssey.data.streaming import PackedLaneSampler, StreamingChunk, move_to_device from odyssey.data.value_binning import QuantileBinner, add_value_tokens from odyssey.data.vocabulary import Vocabulary, code_type from odyssey.inference.legacy_concept_pins import ( @@ -92,10 +92,10 @@ ) from odyssey.training.train import ( TrainingConfig, - _move_chunk_to_device, build_model, hazard_event_names_for, ) +from odyssey.utils.device import default_device from odyssey.utils.env_fingerprint import verify_run_provenance @@ -988,7 +988,7 @@ def run_streaming_inference( state = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 fwd = model.forward_with_features( chunk.batch, state=state, reset_mask=chunk.reset_mask ) @@ -1156,7 +1156,7 @@ def evaluate_run( checkpoint_path: str | Path | None = None, ) -> InferenceResults: """End-to-end: load a trained run, score it against a held-out split.""" - device = device or ("cuda" if torch.cuda.is_available() else "cpu") + device = device or (default_device()) model, vocab, binner, config = load_run( run_dir, device=device, checkpoint_path=checkpoint_path ) diff --git a/odyssey/inference/time_head_probe.py b/odyssey/inference/time_head_probe.py index 9f3b713b..50e6fbd2 100644 --- a/odyssey/inference/time_head_probe.py +++ b/odyssey/inference/time_head_probe.py @@ -54,6 +54,7 @@ from torch import nn from odyssey.data.sequences import PatientSequence +from odyssey.data.streaming import move_to_device from odyssey.data.value_binning import add_value_tokens from odyssey.data.vocabulary import Vocabulary from odyssey.models.sequence_model import SequenceModel @@ -63,7 +64,6 @@ gap_to_bin, ) from odyssey.training.data import iter_patient_sequences -from odyssey.training.train import _move_chunk_to_device logger = logging.getLogger(__name__) @@ -153,7 +153,7 @@ def collect_feature_bank( state = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 fwd = model.forward_with_features( chunk.batch, state=state, reset_mask=chunk.reset_mask ) diff --git a/odyssey/models/backbones/hybrid.py b/odyssey/models/backbones/hybrid.py index 88985f97..8beda5f5 100644 --- a/odyssey/models/backbones/hybrid.py +++ b/odyssey/models/backbones/hybrid.py @@ -16,7 +16,10 @@ spirit, not a reproduction of their exact published architecture, since the paper does not give enough implementation detail to reproduce exactly. -Requires `mamba-ssm`, which needs a CUDA/`nvcc` build. The block stack is +Training requires `mamba-ssm`, which needs a CUDA/`nvcc` build. Without +it (CPU, Apple MPS) the backbone builds from the pure-PyTorch layers in +:mod:`odyssey.models.backbones.mamba_portable` instead: same parameter +names, so a GPU checkpoint loads unchanged; inference only. The block stack is built directly (each :class:`HybridBlock` constructed by hand) rather than through ``mamba_ssm``'s high-level ``MixerModel`` dispatcher: that dispatcher only builds a sequential stack of single-mixer blocks, with no @@ -83,6 +86,7 @@ from torch import nn from odyssey.data.types import ClinicalSequenceBatch +from odyssey.models.backbones import mamba_portable from odyssey.models.backbones.base import ( SequenceBackbone, TimeAwareState, @@ -362,6 +366,36 @@ def forward( # noqa: PLR0912, PLR0915 return Mamba2WithState +def _mixer_classes(*, portable: bool) -> tuple[Any, Any, Any]: + """Return ``(Mamba2WithState, MHA, RMSNorm)`` for this environment. + + ``portable=False`` imports the CUDA ``mamba-ssm`` modules (deferred: + importing them needs a CUDA build). ``portable=True`` returns the + pure-PyTorch stand-ins from :mod:`mamba_portable`, which have the same + parameter names, so one checkpoint loads in either environment. + """ + if portable: + return ( + mamba_portable.Mamba2WithState, + mamba_portable.MHA, + mamba_portable.RMSNorm, + ) + from mamba_ssm.modules.mamba2 import Mamba2 # noqa: PLC0415 + from mamba_ssm.modules.mha import MHA # noqa: PLC0415 + from mamba_ssm.ops.triton.layer_norm import RMSNorm # noqa: PLC0415 + + return _make_mamba2_with_state_cls(Mamba2), MHA, RMSNorm + + +def _inference_params_cls(*, portable: bool) -> Any: # noqa: ANN401 + """Return the ``InferenceParams`` class matching :func:`_mixer_classes`.""" + if portable: + return mamba_portable.InferenceParams + from mamba_ssm.utils.generation import InferenceParams # noqa: PLC0415 + + return InferenceParams + + class MergeAttention(nn.Module): """Learned fusion of two per-position branch outputs via a small attention. @@ -474,20 +508,13 @@ def __init__( # noqa: PLR0917 multi-head attention); set it lower for grouped-query attention, as Nemotron-H does (entry 03, Section 03). """ - try: - # Deferred: mamba-ssm needs CUDA. See the module docstring. - from mamba_ssm.modules.mamba2 import Mamba2 # noqa: PLC0415 - from mamba_ssm.modules.mha import MHA # noqa: PLC0415 - from mamba_ssm.ops.triton.layer_norm import RMSNorm # noqa: PLC0415 - except ImportError as exc: - raise ImportError( - "EHRHybridBackbone requires mamba-ssm, which needs a CUDA " - "build: `uv sync --extra cuda --no-build-isolation`. Use " - "odyssey.models.backbones.tiny_gru.TinyGRUBackbone for " - "CPU development instead." - ) from exc - super().__init__() + # mamba-ssm needs CUDA; without it (CPU, Apple MPS) the pure-PyTorch + # port stands in. See _mixer_classes. + self.portable = not mamba_portable.mamba_ssm_available() + mamba2_with_state_cls, MHA, RMSNorm = _mixer_classes( # noqa: N806 + portable=self.portable + ) self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers @@ -498,8 +525,6 @@ def __init__( # noqa: PLR0917 **embedding_kwargs, ) - mamba2_with_state_cls = _make_mamba2_with_state_cls(Mamba2) - def _make_block(layer_idx: int) -> HybridBlock: mamba_cls = partial( mamba2_with_state_cls, @@ -539,7 +564,7 @@ def forward( this, but its batch-dimension semantics haven't been validated against this backbone. """ - from mamba_ssm.utils.generation import InferenceParams # noqa: PLC0415 + InferenceParams = _inference_params_cls(portable=self.portable) # noqa: N806 typed_state: MambaStateDict | None = ( None if state is None else cast(HybridState, state.recurrent).mamba_states diff --git a/odyssey/models/backbones/mamba_portable.py b/odyssey/models/backbones/mamba_portable.py new file mode 100644 index 00000000..46f98f4b --- /dev/null +++ b/odyssey/models/backbones/mamba_portable.py @@ -0,0 +1,374 @@ +"""Pure-PyTorch stand-ins for the ``mamba-ssm`` modules the hybrid backbone uses. + +``mamba-ssm`` needs CUDA and Triton, so :class:`EHRHybridBackbone` cannot be +built on a laptop (CPU or Apple MPS). This module re-implements the four +pieces the backbone takes from it -- ``Mamba2`` (with the carried-state fix +of :func:`odyssey.models.backbones.hybrid._make_mamba2_with_state_cls`), +``MHA``, the plain ``RMSNorm`` and ``InferenceParams`` -- in plain PyTorch, +following the reference (``*_ref``) code paths of ``mamba-ssm`` 2.3.0. + +Parameter and buffer names match the upstream modules exactly, so a +checkpoint trained on the GPU loads with ``load_state_dict`` unchanged. + +Scope: inference only. The chunked SSD scan is the "minimal SSD" algorithm +(Listing 1 of the Mamba-2 paper) computed in float32; it is not memory- +or speed-optimized and has no backward kernel worth training with. Only +the configuration the backbone uses is supported: ``ngroups=1`` style +grouping, ``rmsnorm=True``, ``norm_before_gate=False``, no ``d_mlp`` +split, no rotary embedding or conv in attention, no varlen +(``cu_seqlens``/``seq_idx``) inputs and no single-token ``step()`` decode. +""" + +from dataclasses import dataclass, field +from typing import Any + +import torch +import torch.nn.functional as F # noqa: N812 +from torch import nn + + +@dataclass +class InferenceParams: + """The fields of ``mamba_ssm.utils.generation.InferenceParams`` used here.""" + + max_seqlen: int + max_batch_size: int + seqlen_offset: int = 0 + batch_size_offset: int = 0 + key_value_memory_dict: dict[int, Any] = field(default_factory=dict) + lengths_per_sample: torch.Tensor | None = None + + +class RMSNorm(nn.Module): + """``mamba_ssm.ops.triton.layer_norm.RMSNorm`` without the fused kernel.""" + + def __init__(self, hidden_size: int, eps: float = 1e-5) -> None: + """Initialize with a unit weight, as upstream does.""" + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(hidden_size)) + self.register_parameter("bias", None) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Normalize by the root mean square of the last dimension.""" + dtype = x.dtype + xf = x.float() + rstd = torch.rsqrt(xf.square().mean(dim=-1, keepdim=True) + self.eps) + return (xf * rstd * self.weight.float()).to(dtype) + + +class RMSNormGated(nn.Module): + """``mamba_ssm.ops.triton.layernorm_gated.RMSNorm`` (``rms_norm_ref`` path).""" + + def __init__( + self, + hidden_size: int, + eps: float = 1e-5, + group_size: int | None = None, + norm_before_gate: bool = False, + ) -> None: + """Initialize with a unit weight, as upstream does.""" + super().__init__() + self.eps = eps + self.weight = nn.Parameter(torch.ones(hidden_size)) + self.register_parameter("bias", None) + self.group_size = group_size + self.norm_before_gate = norm_before_gate + + def forward(self, x: torch.Tensor, z: torch.Tensor | None = None) -> torch.Tensor: + """Return ``norm(x * silu(z))`` (or ``norm(x) * silu(z)``), in float32.""" + dtype = x.dtype + xf = x.float() + zf = z.float() if z is not None else None + if zf is not None and not self.norm_before_gate: + xf = xf * F.silu(zf) + group = self.group_size or xf.shape[-1] + xg = xf.reshape(*xf.shape[:-1], -1, group) + rstd = torch.rsqrt(xg.square().mean(dim=-1, keepdim=True) + self.eps) + out = (xg * rstd).flatten(-2) * self.weight.float() + if zf is not None and self.norm_before_gate: + out = out * F.silu(zf) + result: torch.Tensor = out.to(dtype) + return result + + +def _segsum(x: torch.Tensor) -> torch.Tensor: + """Stable segment sum: ``out[..., i, j] = sum(x[..., j+1 : i+1])``, -inf above.""" + t = x.size(-1) + x = x.unsqueeze(-1).expand(*x.shape, t) + lower = torch.tril(torch.ones(t, t, device=x.device, dtype=torch.bool), diagonal=-1) + x = x.masked_fill(~lower, 0) + out = torch.cumsum(x, dim=-2) + keep = torch.tril(torch.ones(t, t, device=x.device, dtype=torch.bool), diagonal=0) + return out.masked_fill(~keep, -torch.inf) + + +def ssd_chunk_scan( # noqa: PLR0913, PLR0917 + x: torch.Tensor, + dt: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + c: torch.Tensor, + chunk_size: int, + d: torch.Tensor, + dt_bias: torch.Tensor, + initial_states: torch.Tensor | None, +) -> tuple[torch.Tensor, torch.Tensor]: + """Chunked SSD scan, the math of ``mamba_chunk_scan_combined`` (z=None). + + Shapes: ``x`` (b, l, h, p), ``dt`` (b, l, h), ``a`` (h,), ``b``/``c`` + (b, l, g, n), ``d`` (h,), ``dt_bias`` (h,), ``initial_states`` + (b, h, p, n). Returns ``(y, final_state)`` with ``y`` (b, l, h, p). + Computed in float32 throughout; ``dt`` gets ``softplus(dt + dt_bias)``. + """ + batch, seqlen, nheads, _ = x.shape + out_dtype = x.dtype + ngroups = b.shape[2] + dt = F.softplus(dt.float() + dt_bias.float()) + xf = x.float() + bf = b.float().repeat_interleave(nheads // ngroups, dim=2) + cf = c.float().repeat_interleave(nheads // ngroups, dim=2) + pad = (-seqlen) % chunk_size + if pad: + # Zero dt on the padding keeps the state unchanged (decay exp(0)=1, + # input dt*x = 0), so the final state is the true last-token state. + dt = F.pad(dt, (0, 0, 0, pad)) + xf = F.pad(xf, (0, 0, 0, 0, 0, pad)) + bf = F.pad(bf, (0, 0, 0, 0, 0, pad)) + cf = F.pad(cf, (0, 0, 0, 0, 0, pad)) + n_chunks = (seqlen + pad) // chunk_size + + def chunked(t: torch.Tensor) -> torch.Tensor: + return t.reshape(batch, n_chunks, chunk_size, *t.shape[2:]) + + xdt = chunked(xf * dt.unsqueeze(-1)) # (b, c, l, h, p) + a_dt = chunked(dt * a.float()).permute(0, 3, 1, 2) # (b, h, c, l) + bc, cc = chunked(bf), chunked(cf) # (b, c, l, h, n) + a_cumsum = torch.cumsum(a_dt, dim=-1) + + decay = torch.exp(_segsum(a_dt)) # (b, h, c, l, s) + y_diag = torch.einsum("bclhn,bcshn,bhcls,bcshp->bclhp", cc, bc, decay, xdt) + + decay_states = torch.exp(a_cumsum[..., -1:] - a_cumsum) + states = torch.einsum("bclhn,bhcl,bclhp->bchpn", bc, decay_states, xdt) + if initial_states is None: + initial_states = torch.zeros_like(states[:, 0]) + states = torch.cat([initial_states.float().unsqueeze(1), states], dim=1) + decay_chunk = torch.exp(_segsum(F.pad(a_cumsum[..., -1], (1, 0)))) + new_states = torch.einsum("bhzc,bchpn->bzhpn", decay_chunk, states) + states, final_state = new_states[:, :-1], new_states[:, -1] + + y_off = torch.einsum("bclhn,bchpn,bhcl->bclhp", cc, states, torch.exp(a_cumsum)) + y = (y_diag + y_off).reshape(batch, n_chunks * chunk_size, nheads, -1)[:, :seqlen] + y = y + xf[:, :seqlen] * d.float().view(1, 1, nheads, 1) + return y.to(out_dtype), final_state + + +class Mamba2WithState(nn.Module): + """``Mamba2`` + the backbone's carried-state fix, in plain PyTorch. + + Same parameters as ``mamba_ssm.modules.mamba2.Mamba2`` (``in_proj``, + ``conv1d``, ``dt_bias``, ``A_log``, ``D``, ``norm``, ``out_proj``) and + the same ``(conv_state, ssm_state)`` cache layout, so a chunk's carried + state continues into the next chunk exactly as on the GPU. + """ + + def __init__( # noqa: PLR0913 + self, + d_model: int, + *, + d_state: int = 128, + d_conv: int = 4, + expand: int = 2, + headdim: int = 64, + ngroups: int = 1, + chunk_size: int = 256, + layer_idx: int | None = None, + ) -> None: + """Build the layer with upstream's shapes; values come from the checkpoint.""" + super().__init__() + self.d_model = d_model + self.d_state = d_state + self.d_conv = d_conv + self.d_inner = expand * d_model + self.d_ssm = self.d_inner + self.headdim = headdim + self.ngroups = ngroups + self.nheads = self.d_ssm // headdim + self.chunk_size = chunk_size + self.layer_idx = layer_idx + d_in_proj = 2 * self.d_inner + 2 * ngroups * d_state + self.nheads + self.in_proj = nn.Linear(d_model, d_in_proj, bias=False) + conv_dim = self.d_ssm + 2 * ngroups * d_state + self.conv1d = nn.Conv1d( + conv_dim, + conv_dim, + kernel_size=d_conv, + groups=conv_dim, + padding=d_conv - 1, + bias=True, + ) + self.act = nn.SiLU() + self.dt_bias = nn.Parameter(torch.zeros(self.nheads)) + self.A_log = nn.Parameter(torch.zeros(self.nheads)) + self.D = nn.Parameter(torch.ones(self.nheads)) + self.norm = RMSNormGated( + self.d_ssm, + eps=1e-5, + group_size=self.d_ssm // ngroups, + norm_before_gate=False, + ) + self.out_proj = nn.Linear(self.d_inner, d_model, bias=False) + + def _get_states_from_cache( + self, inference_params: InferenceParams, batch_size: int + ) -> tuple[torch.Tensor, torch.Tensor]: + if self.layer_idx is None: + raise ValueError("a carried state needs layer_idx") + layer = self.layer_idx + cache = inference_params.key_value_memory_dict + if layer not in cache: + conv_state = torch.zeros( + batch_size, + self.d_conv, + self.conv1d.weight.shape[0], + device=self.conv1d.weight.device, + dtype=self.conv1d.weight.dtype, + ).transpose(1, 2) + ssm_state = torch.zeros( + batch_size, + self.nheads, + self.headdim, + self.d_state, + device=self.in_proj.weight.device, + dtype=self.in_proj.weight.dtype, + ) + cache[layer] = (conv_state, ssm_state) + states: tuple[torch.Tensor, torch.Tensor] = cache[layer] + return states + + def forward( + self, u: torch.Tensor, inference_params: InferenceParams | None = None + ) -> torch.Tensor: + """Prefill one chunk ``u`` (b, l, d), reading and writing the carried state.""" + batch, seqlen, _ = u.shape + conv_state = ssm_state = incoming_conv = None + if inference_params is not None: + if inference_params.seqlen_offset > 0: + raise NotImplementedError("single-token decode is not ported") + conv_state, ssm_state = self._get_states_from_cache(inference_params, batch) + incoming_conv = conv_state.clone() + + zxbcdt = self.in_proj(u) + z, xbc, dt = torch.split( + zxbcdt, + [self.d_ssm, self.d_ssm + 2 * self.ngroups * self.d_state, self.nheads], + dim=-1, + ) + xbc_t = xbc.transpose(1, 2) # (b, d, l) + if conv_state is not None: + conv_state.copy_(F.pad(xbc_t, (self.d_conv - xbc_t.shape[-1], 0))) + if incoming_conv is not None and self.d_conv > 1: + padded = torch.cat( + [incoming_conv[:, :, -(self.d_conv - 1) :], xbc_t], dim=-1 + ) + conv_out = F.conv1d( + padded, self.conv1d.weight, self.conv1d.bias, groups=padded.shape[1] + ) + else: + conv_out = self.conv1d(xbc_t)[:, :, : -(self.d_conv - 1)] + xbc = self.act(conv_out.transpose(1, 2)) + x, b, c = torch.split( + xbc, + [self.d_ssm, self.ngroups * self.d_state, self.ngroups * self.d_state], + dim=-1, + ) + y, final_state = ssd_chunk_scan( + x.reshape(batch, seqlen, self.nheads, self.headdim), + dt, + -torch.exp(self.A_log.float()), + b.reshape(batch, seqlen, self.ngroups, self.d_state), + c.reshape(batch, seqlen, self.ngroups, self.d_state), + self.chunk_size, + self.D, + self.dt_bias, + ssm_state, + ) + if ssm_state is not None: + ssm_state.copy_(final_state) + y = self.norm(y.flatten(-2), z) + out: torch.Tensor = self.out_proj(y) + return out + + +class MHA(nn.Module): + """``mamba_ssm.modules.mha.MHA`` for the configuration the backbone uses. + + Causal self-attention with optional grouped-query heads; no rotary + embedding, no conv, no MLP split, and no KV cache (the backbone always + passes ``inference_params=None`` to attention). + """ + + def __init__( + self, + embed_dim: int, + num_heads: int, + num_heads_kv: int | None = None, + causal: bool = True, + layer_idx: int | None = None, + ) -> None: + """Build the layer with upstream's shapes; values come from the checkpoint.""" + super().__init__() + self.num_heads = num_heads + self.num_heads_kv = num_heads_kv if num_heads_kv is not None else num_heads + self.head_dim = embed_dim // num_heads + self.causal = causal + self.layer_idx = layer_idx + qkv_dim = self.head_dim * (self.num_heads + 2 * self.num_heads_kv) + self.in_proj = nn.Linear(embed_dim, qkv_dim, bias=True) + self.out_proj = nn.Linear(self.head_dim * num_heads, embed_dim, bias=True) + + def forward(self, x: torch.Tensor, inference_params: Any = None) -> torch.Tensor: # noqa: ANN401 + """Attend over the chunk ``x`` (b, l, d).""" + if inference_params is not None: + raise NotImplementedError("attention KV caching is not ported") + qkv = self.in_proj(x) + q, kv = qkv.split( + [self.num_heads * self.head_dim, 2 * self.num_heads_kv * self.head_dim], + dim=-1, + ) + q = q.reshape(*q.shape[:-1], self.num_heads, self.head_dim) + k, v = kv.reshape(*kv.shape[:-1], 2, self.num_heads_kv, self.head_dim).unbind( + dim=-3 + ) + repeats = self.num_heads // self.num_heads_kv + k = k.repeat_interleave(repeats, dim=2) + v = v.repeat_interleave(repeats, dim=2) + context = F.scaled_dot_product_attention( + q.transpose(1, 2), + k.transpose(1, 2), + v.transpose(1, 2), + is_causal=self.causal, + ).transpose(1, 2) + out: torch.Tensor = self.out_proj(context.flatten(-2)) + return out + + +def mamba_ssm_available() -> bool: + """Return whether the CUDA ``mamba-ssm`` build can be imported.""" + try: + import mamba_ssm # noqa: F401, PLC0415 + except ImportError: + return False + return True + + +__all__ = [ + "MHA", + "InferenceParams", + "Mamba2WithState", + "RMSNorm", + "RMSNormGated", + "mamba_ssm_available", + "ssd_chunk_scan", +] diff --git a/odyssey/training/train.py b/odyssey/training/train.py index ead2e4b4..f72a573e 100644 --- a/odyssey/training/train.py +++ b/odyssey/training/train.py @@ -53,7 +53,6 @@ from pathlib import Path from typing import ( Any, - TypeVar, ) import polars as pl @@ -66,7 +65,7 @@ from odyssey.data.packed_context import PackedContextSampler from odyssey.data.sequences import PatientSequence from odyssey.data.sidecars import activate_sidecars, active_sidecar_names -from odyssey.data.streaming import PackedLaneSampler, StreamingChunk +from odyssey.data.streaming import PackedLaneSampler, StreamingChunk, move_to_device from odyssey.data.value_binning import CLIP_TAIL, QuantileBinner, add_value_tokens from odyssey.data.vocabulary import PAD_ID, Vocabulary from odyssey.inference.legacy_concept_pins import write_run_pins @@ -529,25 +528,6 @@ def close(self) -> None: self._file.close() -_Movable = TypeVar("_Movable") - - -def _move_chunk_to_device(chunk: _Movable, device: str) -> _Movable: - """Move every tensor field of a (possibly nested) NamedTuple to ``device``. - - Works for :class:`~odyssey.data.streaming.StreamingChunk` and its - nested :class:`~odyssey.data.types.ClinicalSequenceBatch`/ - :class:`~odyssey.data.types.AuxiliaryInputs` without depending on - their exact field lists, so a new field added to any of them doesn't - need a matching change here. - """ - if isinstance(chunk, torch.Tensor): - return chunk.to(device) # type: ignore[return-value] - if isinstance(chunk, tuple) and hasattr(chunk, "_fields"): # NamedTuple - return type(chunk)(*(_move_chunk_to_device(v, device) for v in chunk)) - return chunk - - def _detach_state(state: TimeAwareState) -> TimeAwareState: """Truncate BPTT across chunks for a backbone's carried recurrent state. @@ -1026,7 +1006,7 @@ def evaluate_streaming( for i, chunk in enumerate(sampler): if max_chunks is not None and i >= max_chunks: break - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 event_targets = ( event_hazard_targets(chunk, event_tables) if event_tables is not None @@ -1642,7 +1622,7 @@ def make_tuning_sampler() -> StreamingSampler: steps_this_epoch = steps_into_epoch for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 event_targets = ( event_hazard_targets(chunk, train_event_tables) if train_event_tables is not None diff --git a/odyssey/utils/device.py b/odyssey/utils/device.py new file mode 100644 index 00000000..c980a134 --- /dev/null +++ b/odyssey/utils/device.py @@ -0,0 +1,37 @@ +"""Torch device selection and transfer for CUDA, Apple MPS and CPU. + +Inference runs on an NVIDIA GPU (with the ``mamba-ssm`` kernels) or on a +laptop (Apple MPS or CPU, with :mod:`odyssey.models.backbones.mamba_portable`). +These helpers keep the device-specific rules in one place. +""" + +import torch + + +def default_device() -> str: + """Return the best available device: ``cuda``, then ``mps``, then ``cpu``.""" + if torch.cuda.is_available(): + return "cuda" + if torch.backends.mps.is_available(): + return "mps" + return "cpu" + + +def resolve_device(device: str | None) -> str: + """Map ``None`` or ``"auto"`` to :func:`default_device`; pass others through.""" + return default_device() if device in {None, "auto"} else str(device) + + +def to_device(tensor: torch.Tensor, device: str | torch.device) -> torch.Tensor: + """Move ``tensor`` to ``device``, casting float64 to float32 on MPS. + + MPS has no float64. The only float64 model inputs are time stamps in + hours, which the model casts to float32 before use, so the cast loses + nothing the model would see. + """ + if tensor.dtype == torch.float64 and torch.device(device).type == "mps": + return tensor.to(device=device, dtype=torch.float32) + return tensor.to(device) + + +__all__ = ["default_device", "resolve_device", "to_device"] diff --git a/odyssey/utils/env_fingerprint.py b/odyssey/utils/env_fingerprint.py index 6c45ff8b..09f36190 100644 --- a/odyssey/utils/env_fingerprint.py +++ b/odyssey/utils/env_fingerprint.py @@ -129,11 +129,11 @@ def numeric_canary( (fixed generator seed, fixed shapes), in eval mode, and records robust statistics of the logits. Any kernel or weight change moves them. """ + from odyssey.data.streaming import move_to_device # noqa: PLC0415 from odyssey.data.types import ( # noqa: PLC0415 AuxiliaryInputs, ClinicalSequenceBatch, ) - from odyssey.training.train import _move_chunk_to_device # noqa: PLC0415 g = torch.Generator().manual_seed(12345) lanes, t = 2, 64 @@ -152,7 +152,7 @@ def numeric_canary( ), ), ) - batch = _move_chunk_to_device(batch, device) + batch = move_to_device(batch, device) was_training = model.training model.eval() with torch.no_grad(): diff --git a/scripts/probe_channel.py b/scripts/probe_channel.py index c4febf7a..4f6159a4 100644 --- a/scripts/probe_channel.py +++ b/scripts/probe_channel.py @@ -99,7 +99,7 @@ from odyssey.data.history_recap import maybe_history_recap from odyssey.data.sequences import PatientSequence from odyssey.data.sidecars import activate_sidecars -from odyssey.data.streaming import PackedLaneSampler +from odyssey.data.streaming import PackedLaneSampler, move_to_device from odyssey.data.value_binning import QuantileBinner, add_value_tokens from odyssey.data.vocabulary import PAD_ID, Vocabulary from odyssey.inference.leakage import ( @@ -127,7 +127,7 @@ ) from odyssey.training.running_labels import position_running_labels from odyssey.training.shard_stream import shard_paths -from odyssey.training.train import TrainingConfig, _move_chunk_to_device +from odyssey.training.train import TrainingConfig logger = logging.getLogger("probe_channel") @@ -318,7 +318,7 @@ def collect_channel_bank( # noqa: PLR0915 -- one linear streaming pass state = None with torch.no_grad(): for chunk in sampler: - chunk = _move_chunk_to_device(chunk, device) # noqa: PLW2901 + chunk = move_to_device(chunk, device) # noqa: PLW2901 hidden, state = model.backbone( chunk.batch, state=state, reset_mask=chunk.reset_mask ) diff --git a/tests/apps/clinician_demo/test_main_and_layering.py b/tests/apps/clinician_demo/test_main_and_layering.py index 6a11d91e..dfca9739 100644 --- a/tests/apps/clinician_demo/test_main_and_layering.py +++ b/tests/apps/clinician_demo/test_main_and_layering.py @@ -25,6 +25,7 @@ "thresholds", "showcase", "scorecard", + "export_thresholds", ] @@ -72,6 +73,16 @@ def test_every_flag_reaches_the_config() -> None: ) +def test_device_auto_resolves_to_an_available_device( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr("odyssey.utils.device.default_device", lambda: "mps") + config, _ = parse_args(["--run-dir", "/r", "--data-dir", "/d"]) + assert config.device == "mps" + config, _ = parse_args(["--run-dir", "/r", "--data-dir", "/d", "--device", "auto"]) + assert config.device == "mps" + + @pytest.mark.parametrize( "argv", [ diff --git a/tests/apps/clinician_demo/test_thresholds.py b/tests/apps/clinician_demo/test_thresholds.py index 41d54af6..0e008e95 100644 --- a/tests/apps/clinician_demo/test_thresholds.py +++ b/tests/apps/clinician_demo/test_thresholds.py @@ -6,12 +6,18 @@ import polars as pl import pytest +from apps.clinician_demo.export_thresholds import main as export_main from apps.clinician_demo.thresholds import ( + 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, ) @@ -149,3 +155,76 @@ def test_unreadable_cache_is_ignored_and_rewritten(tmp_path: Path) -> None: rows_path, cache, ["acute_kidney_injury"], [24.0], 0.2 ) assert len(points) == 1 and json.loads(cache.read_text())["points"] + + +def test_export_then_load_round_trips_the_alert_lines(tmp_path: Path) -> None: + rows_path, out = tmp_path / "alerts_rows.parquet", tmp_path / "agg.json" + _rows().write_parquet(rows_path) + exported = export_operating_points(rows_path, out, [24.0, 72.0], 0.2) + payload = json.loads(out.read_text()) + assert payload["alert_rate"] == 0.2 + assert {p["event"] for p in payload["points"]} == {"acute_kidney_injury", "death"} + # aggregates only: nothing row-level travels with the file + assert all("subject_id" not in p and "visit_id" not in p for p in payload["points"]) + loaded = load_aggregate_operating_points( + out, ["acute_kidney_injury", "death"], [24.0], 0.2 + ) + assert loaded == exported + + +def test_loaded_alert_lines_are_filtered_to_the_request(tmp_path: Path) -> None: + rows_path, out = tmp_path / "alerts_rows.parquet", tmp_path / "agg.json" + _rows().write_parquet(rows_path) + export_operating_points(rows_path, out, [24.0], 0.2) + only_death = load_aggregate_operating_points(out, ["death"], [24.0], 0.2) + assert [(p.event, p.horizon_hours) for p in only_death] == [("death", 24.0)] + assert load_aggregate_operating_points(out, ["death"], [8.0], 0.2) == [] + + +def test_aggregate_file_absent_or_at_another_rate_gives_no_lines( + tmp_path: Path, caplog: pytest.LogCaptureFixture +) -> None: + out = tmp_path / "agg.json" + assert load_aggregate_operating_points(out, ["death"], [24.0], 0.05) == [] + rows_path = tmp_path / "alerts_rows.parquet" + _rows().write_parquet(rows_path) + export_operating_points(rows_path, out, [24.0], 0.2) + assert load_aggregate_operating_points(out, ["death"], [24.0], 0.05) == [] + assert "alert rate" in caplog.text + + +def test_a_run_uses_its_rows_when_present_else_its_export(tmp_path: Path) -> None: + gpu_run, laptop_run = tmp_path / "gpu", tmp_path / "laptop" + gpu_run.mkdir() + laptop_run.mkdir() + _rows().write_parquet(gpu_run / ALERTS_ROWS_FILENAME) + cache = tmp_path / "cache" / "thresholds.json" + from_rows = operating_points_for_run(gpu_run, cache, ["death"], [24.0], 0.2) + assert cache.exists() and [p.event for p in from_rows] == ["death"] + + assert operating_points_for_run(laptop_run, cache, ["death"], [24.0], 0.2) == [] + export_operating_points( + gpu_run / ALERTS_ROWS_FILENAME, + laptop_run / AGGREGATE_THRESHOLDS_FILENAME, + [24.0], + 0.2, + ) + from_export = operating_points_for_run(laptop_run, cache, ["death"], [24.0], 0.2) + assert from_export == from_rows + + +def test_export_command_writes_next_to_the_checkpoint( + tmp_path: Path, capsys: pytest.CaptureFixture[str] +) -> None: + _rows().write_parquet(tmp_path / ALERTS_ROWS_FILENAME) + assert export_main(["--run-dir", str(tmp_path), "--alert-rate", "0.2"]) == 0 + payload = json.loads((tmp_path / AGGREGATE_THRESHOLDS_FILENAME).read_text()) + assert payload["alert_rate"] == 0.2 + # the demo's horizons: rows here carry 24 h only + assert {p["horizon_hours"] for p in payload["points"]} == {24.0} + assert "wrote 2 alert lines" in capsys.readouterr().out + + +def test_export_command_refuses_a_run_without_rows(tmp_path: Path) -> None: + with pytest.raises(SystemExit): + export_main(["--run-dir", str(tmp_path)]) diff --git a/tests/odyssey/models/backbones/hybrid_helpers.py b/tests/odyssey/models/backbones/hybrid_helpers.py new file mode 100644 index 00000000..d7f3f335 --- /dev/null +++ b/tests/odyssey/models/backbones/hybrid_helpers.py @@ -0,0 +1,51 @@ +"""Shared synthetic batches and backbones for the hybrid-backbone tests.""" + +import torch + +from odyssey.data.types import AuxiliaryInputs, ClinicalSequenceBatch +from odyssey.models.backbones.hybrid import EHRHybridBackbone + + +VOCAB_SIZE = 40 +HIDDEN_SIZE = 64 # divisible by both MAMBA_HEADDIM and ATTN_NUM_HEADS + + +def make_batch( + batch: int, seq_len: int, device: str = "cpu", seed: int = 0 +) -> ClinicalSequenceBatch: + """Return a random batch with increasing time stamps, as real records have.""" + gen = torch.Generator().manual_seed(seed) + + def ints(high: int) -> torch.Tensor: + return torch.randint(0, high, (batch, seq_len), generator=gen) + + times = torch.rand(batch, seq_len, generator=gen).cumsum(dim=1) * 10 + aux = AuxiliaryInputs( + type_ids=ints(9), + time_stamps=times.double(), + ages=torch.rand(batch, seq_len, generator=gen) * 90, + visit_orders=ints(5), + visit_segments=ints(3), + ) + batch_cpu = ClinicalSequenceBatch(concept_ids=ints(VOCAB_SIZE - 1) + 1, aux=aux) + return ClinicalSequenceBatch( + concept_ids=batch_cpu.concept_ids.to(device), + aux=AuxiliaryInputs( + *(t.to(device) if t is not None else None for t in batch_cpu.aux) + ), + ) + + +def make_backbone( + chunk_size: int = 16, num_hidden_layers: int = 2 +) -> EHRHybridBackbone: + """Return a small hybrid backbone; the environment picks the mixers.""" + return EHRHybridBackbone( + vocab_size=VOCAB_SIZE, + hidden_size=HIDDEN_SIZE, + num_hidden_layers=num_hidden_layers, + mamba_state_size=16, + mamba_headdim=32, + mamba_chunk_size=chunk_size, + attn_num_heads=8, + ) diff --git a/tests/odyssey/models/backbones/test_hybrid_portable.py b/tests/odyssey/models/backbones/test_hybrid_portable.py new file mode 100644 index 00000000..97163a38 --- /dev/null +++ b/tests/odyssey/models/backbones/test_hybrid_portable.py @@ -0,0 +1,116 @@ +"""The hybrid backbone built from the pure-PyTorch mixers (CPU, any host). + +``mamba_ssm_available`` is forced to ``False`` so these run the portable +path even on a GPU host that has ``mamba-ssm`` installed. +""" + +import pytest +import torch + +from odyssey.models.backbones import mamba_portable +from odyssey.models.backbones.base import TimeAwareState +from tests.odyssey.models.backbones.hybrid_helpers import ( + HIDDEN_SIZE, + make_backbone, + make_batch, +) + + +@pytest.fixture(autouse=True) +def _portable(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(mamba_portable, "mamba_ssm_available", lambda: False) + + +def _slice(batch, start: int, stop: int): # type: ignore[no-untyped-def] + return type(batch)( + concept_ids=batch.concept_ids[:, start:stop], + aux=type(batch.aux)( + *(t[:, start:stop] if t is not None else None for t in batch.aux) + ), + ) + + +def test_builds_without_mamba_ssm_and_uses_the_portable_mixers() -> None: + backbone = make_backbone() + assert backbone.portable + block = backbone.layers[0] + assert isinstance(block.mamba, mamba_portable.Mamba2WithState) + assert isinstance(block.attn, mamba_portable.MHA) + assert isinstance(backbone.norm_f, mamba_portable.RMSNorm) + + +def test_forward_shapes_and_carried_state() -> None: + torch.manual_seed(0) + backbone = make_backbone().eval() + hidden, state = backbone(make_batch(2, 12)) + assert hidden.shape == (2, 12, HIDDEN_SIZE) + assert isinstance(state, TimeAwareState) + assert set(state.recurrent.mamba_states) == {0, 1} + + +def test_mamba_state_carries_across_chunks() -> None: + """Chunked streaming equals one pass on the Mamba side. + + Attention has no cross-chunk memory by design (see hybrid's module + docstring), so this compares a backbone whose attention branch is + silenced: what remains must match exactly. + """ + torch.manual_seed(0) + backbone = make_backbone().eval() + with torch.no_grad(): + for block in backbone.layers: + block.attn.out_proj.weight.zero_() + block.attn.out_proj.bias.zero_() + block.merge.value_proj.weight.data[:] = torch.eye(HIDDEN_SIZE) + block.merge.value_proj.bias.zero_() + block.merge.key_proj.weight.zero_() + block.merge.key_proj.bias.zero_() + batch = make_batch(2, 20) + with torch.no_grad(): + whole, _ = backbone(batch) + first, state = backbone(_slice(batch, 0, 7)) + second, _ = backbone(_slice(batch, 7, 20), state=state) + # Equal keys give the two branches equal weight, so attention's zero + # output halves the fused value on both paths alike. + torch.testing.assert_close( + torch.cat([first, second], 1), whole, atol=1e-4, rtol=1e-4 + ) + + +def test_carried_state_is_not_mutated_by_the_next_chunk() -> None: + torch.manual_seed(0) + backbone = make_backbone().eval() + batch = make_batch(1, 10) + with torch.no_grad(): + _, state = backbone(_slice(batch, 0, 5)) + before = { + k: tuple(t.clone() for t in v) + for k, v in state.recurrent.mamba_states.items() + } + backbone(_slice(batch, 5, 10), state=state) + for layer, tensors in state.recurrent.mamba_states.items(): + for kept, now in zip(before[layer], tensors, strict=True): + assert torch.equal(kept, now) + + +def test_a_state_dict_round_trips_between_two_portable_backbones() -> None: + torch.manual_seed(0) + source = make_backbone().eval() + target = make_backbone().eval() + target.load_state_dict(source.state_dict()) + batch = make_batch(2, 9) + with torch.no_grad(): + torch.testing.assert_close(source(batch)[0], target(batch)[0]) + + +@pytest.mark.skipif(not torch.backends.mps.is_available(), reason="needs Apple MPS") +def test_runs_on_apple_mps_and_matches_cpu() -> None: + from odyssey.data.streaming import move_to_device # noqa: PLC0415 + + torch.manual_seed(0) + backbone = make_backbone().eval() + batch = make_batch(2, 24) + with torch.no_grad(): + cpu_out, _ = backbone(batch) + mps_out, _ = backbone.to("mps")(move_to_device(batch, "mps")) + torch.testing.assert_close(mps_out.cpu(), cpu_out, atol=1e-4, rtol=1e-4) diff --git a/tests/odyssey/models/backbones/test_mamba_portable.py b/tests/odyssey/models/backbones/test_mamba_portable.py new file mode 100644 index 00000000..ac3ba690 --- /dev/null +++ b/tests/odyssey/models/backbones/test_mamba_portable.py @@ -0,0 +1,203 @@ +"""The pure-PyTorch ``mamba-ssm`` stand-ins (CPU; see the GPU parity test too). + +The reference for the scan is the SSM recurrence itself, stepped one token +at a time; the references for the norms and attention are their textbook +formulas. Parameter names are pinned to ``mamba-ssm`` 2.3.0's so a GPU +checkpoint loads unchanged. +""" + +import pytest +import torch +import torch.nn.functional as F # noqa: N812 + +from odyssey.models.backbones.mamba_portable import ( + MHA, + InferenceParams, + Mamba2WithState, + RMSNorm, + RMSNormGated, + ssd_chunk_scan, +) + + +def _naive_scan( # noqa: PLR0913, PLR0917 + x: torch.Tensor, + dt: torch.Tensor, + a: torch.Tensor, + b: torch.Tensor, + c: torch.Tensor, + d: torch.Tensor, + dt_bias: torch.Tensor, + state: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor]: + """``h_t = exp(dt_t a) h_{t-1} + dt_t x_t b_t^T``; ``y_t = h_t c_t + d x_t``.""" + heads = x.shape[2] + b = b.repeat_interleave(heads // b.shape[2], dim=2) + c = c.repeat_interleave(heads // c.shape[2], dim=2) + dt = F.softplus(dt + dt_bias) + ys = [] + for t in range(x.shape[1]): + decay = torch.exp(dt[:, t] * a)[:, :, None, None] + update = dt[:, t, :, None, None] * x[:, t, :, :, None] * b[:, t, :, None, :] + state = decay * state + update + ys.append(torch.einsum("bhpn,bhn->bhp", state, c[:, t]) + d[:, None] * x[:, t]) + return torch.stack(ys, dim=1), state + + +def _scan_inputs( + seqlen: int, ngroups: int = 1, seed: int = 0 +) -> tuple[torch.Tensor, ...]: + gen = torch.Generator().manual_seed(seed) + batch, heads, headdim, dstate = 2, 4, 3, 5 + + def randn(*shape: int) -> torch.Tensor: + return torch.randn(*shape, generator=gen) + + return ( + randn(batch, seqlen, heads, headdim), + randn(batch, seqlen, heads), + -torch.rand(heads, generator=gen) * 2, + randn(batch, seqlen, ngroups, dstate), + randn(batch, seqlen, ngroups, dstate), + randn(heads), + randn(heads), + randn(batch, heads, headdim, dstate), + ) + + +@pytest.mark.parametrize( + ("seqlen", "chunk_size", "ngroups"), + [(16, 8, 1), (13, 8, 1), (5, 8, 1), (13, 4, 2)], +) +def test_chunk_scan_matches_the_stepwise_recurrence( + seqlen: int, chunk_size: int, ngroups: int +) -> None: + x, dt, a, b, c, d, dt_bias, state = _scan_inputs(seqlen, ngroups) + y, final = ssd_chunk_scan(x, dt, a, b, c, chunk_size, d, dt_bias, state) + y_ref, final_ref = _naive_scan(x, dt, a, b, c, d, dt_bias, state) + torch.testing.assert_close(y, y_ref, atol=1e-4, rtol=1e-4) + torch.testing.assert_close(final, final_ref, atol=1e-4, rtol=1e-4) + + +def test_no_initial_state_means_a_zero_state() -> None: + x, dt, a, b, c, d, dt_bias, state = _scan_inputs(9) + y_none, final_none = ssd_chunk_scan(x, dt, a, b, c, 4, d, dt_bias, None) + y_zero, final_zero = ssd_chunk_scan( + x, dt, a, b, c, 4, d, dt_bias, torch.zeros_like(state) + ) + torch.testing.assert_close(y_none, y_zero) + torch.testing.assert_close(final_none, final_zero) + + +def _mixer(chunk_size: int = 8, seed: int = 0) -> Mamba2WithState: + torch.manual_seed(seed) + mixer = Mamba2WithState( + 32, d_state=8, headdim=8, chunk_size=chunk_size, layer_idx=0 + ) + with torch.no_grad(): + for param in mixer.parameters(): + param.normal_(std=0.2) + return mixer.eval() + + +def test_carried_state_continues_a_split_sequence_exactly() -> None: + mixer = _mixer() + u = torch.randn(2, 23, 32) + whole = mixer(u, inference_params=InferenceParams(23, 2)) + params = InferenceParams(23, 2) + parts = [ + mixer(u[:, s], inference_params=params) + for s in (slice(0, 2), slice(2, 11), slice(11, 23)) + ] + torch.testing.assert_close(torch.cat(parts, dim=1), whole, atol=1e-5, rtol=1e-5) + + +def test_output_does_not_depend_on_the_scan_chunk_size() -> None: + u = torch.randn(1, 19, 32) + small, large = _mixer(chunk_size=4), _mixer(chunk_size=64) + torch.testing.assert_close(small(u), large(u), atol=1e-5, rtol=1e-5) + + +def test_single_token_decode_is_refused() -> None: + params = InferenceParams(4, 1, seqlen_offset=1) + with pytest.raises(NotImplementedError, match="decode"): + _mixer()(torch.randn(1, 1, 32), inference_params=params) + + +def test_a_carried_state_needs_a_layer_index() -> None: + mixer = Mamba2WithState(32, d_state=8, headdim=8) + with pytest.raises(ValueError, match="layer_idx"): + mixer(torch.randn(1, 3, 32), inference_params=InferenceParams(3, 1)) + + +def test_rms_norm_matches_its_formula() -> None: + norm = RMSNorm(6, eps=1e-5) + with torch.no_grad(): + norm.weight.copy_(torch.arange(1.0, 7.0)) + x = torch.randn(3, 6) + expected = x / torch.sqrt(x.pow(2).mean(-1, keepdim=True) + 1e-5) * norm.weight + torch.testing.assert_close(norm(x), expected) + + +@pytest.mark.parametrize("norm_before_gate", [False, True]) +def test_gated_rms_norm_gates_on_the_documented_side(norm_before_gate: bool) -> None: + norm = RMSNormGated(8, eps=1e-5, group_size=4, norm_before_gate=norm_before_gate) + x, z = torch.randn(2, 8), torch.randn(2, 8) + + def grouped_rms(t: torch.Tensor) -> torch.Tensor: + g = t.reshape(2, 2, 4) + return (g / torch.sqrt(g.pow(2).mean(-1, keepdim=True) + 1e-5)).reshape(2, 8) + + gate = F.silu(z) + expected = grouped_rms(x) * gate if norm_before_gate else grouped_rms(x * gate) + torch.testing.assert_close(norm(x, z), expected) + + +@pytest.mark.parametrize("num_heads_kv", [4, 2, 1]) +def test_attention_is_causal_multi_head_with_grouped_kv(num_heads_kv: int) -> None: + torch.manual_seed(0) + attn = MHA(16, num_heads=4, num_heads_kv=num_heads_kv).eval() + x = torch.randn(2, 7, 16) + head_dim = 4 + q, k, v = attn.in_proj(x).split( + [4 * head_dim, num_heads_kv * head_dim, num_heads_kv * head_dim], dim=-1 + ) + q = q.reshape(2, 7, 4, head_dim).transpose(1, 2) + k = k.reshape(2, 7, num_heads_kv, head_dim).transpose(1, 2) + v = v.reshape(2, 7, num_heads_kv, head_dim).transpose(1, 2) + k = k.repeat_interleave(4 // num_heads_kv, dim=1) + v = v.repeat_interleave(4 // num_heads_kv, dim=1) + scores = q @ k.transpose(-1, -2) / head_dim**0.5 + scores = scores.masked_fill(torch.ones(7, 7).triu(1).bool(), float("-inf")) + context = (scores.softmax(-1) @ v).transpose(1, 2).reshape(2, 7, 16) + torch.testing.assert_close(attn(x), attn.out_proj(context), atol=1e-5, rtol=1e-5) + + +def test_attention_refuses_a_kv_cache() -> None: + with pytest.raises(NotImplementedError): + MHA(16, num_heads=4)(torch.randn(1, 2, 16), inference_params=object()) + + +def test_parameter_names_and_shapes_match_mamba_ssm() -> None: + """Pinned to mamba-ssm 2.3.0 so GPU checkpoints load without remapping.""" + d_model, d_state, headdim = 32, 8, 8 + d_inner, nheads = 2 * d_model, 2 * d_model // headdim + conv_dim = d_inner + 2 * d_state + mixer = Mamba2WithState(d_model, d_state=d_state, headdim=headdim) + assert {k: tuple(v.shape) for k, v in mixer.state_dict().items()} == { + "in_proj.weight": (2 * d_inner + 2 * d_state + nheads, d_model), + "conv1d.weight": (conv_dim, 1, 4), + "conv1d.bias": (conv_dim,), + "dt_bias": (nheads,), + "A_log": (nheads,), + "D": (nheads,), + "norm.weight": (d_inner,), + "out_proj.weight": (d_model, d_inner), + } + assert set(MHA(16, num_heads=4).state_dict()) == { + "in_proj.weight", + "in_proj.bias", + "out_proj.weight", + "out_proj.bias", + } + assert set(RMSNorm(8).state_dict()) == {"weight"} diff --git a/tests/odyssey/models/backbones/test_mamba_portable_gpu.py b/tests/odyssey/models/backbones/test_mamba_portable_gpu.py new file mode 100644 index 00000000..856551f7 --- /dev/null +++ b/tests/odyssey/models/backbones/test_mamba_portable_gpu.py @@ -0,0 +1,64 @@ +"""GPU parity: the pure-PyTorch backbone against the ``mamba-ssm`` kernels. + +Builds the hybrid backbone twice from one state dict -- once with the CUDA +mixers, once with :mod:`mamba_portable` -- and checks that both give the +same hidden states on the same chunked stream. Auto-skips without +``mamba-ssm`` and a CUDA device. Tolerances allow for TF32 matmuls and the +Triton kernels' own accumulation order. +""" + +import pytest +import torch + + +pytest.importorskip("mamba_ssm", reason="mamba-ssm not installed (needs CUDA)") +cuda_required = pytest.mark.skipif( + not torch.cuda.is_available(), reason="requires a CUDA device" +) + +from odyssey.data.streaming import move_to_device # noqa: E402 +from odyssey.models.backbones import mamba_portable # noqa: E402 +from tests.odyssey.models.backbones.hybrid_helpers import ( # noqa: E402 + make_backbone, + make_batch, +) + + +def _pair(monkeypatch: pytest.MonkeyPatch, device: str): # type: ignore[no-untyped-def] + torch.manual_seed(0) + cuda_backbone = make_backbone().to("cuda").eval() + assert not cuda_backbone.portable + with monkeypatch.context() as patch: + patch.setattr(mamba_portable, "mamba_ssm_available", lambda: False) + portable = make_backbone().eval() + assert portable.portable + portable.load_state_dict(cuda_backbone.state_dict()) + return cuda_backbone, portable.to(device) + + +@cuda_required +@pytest.mark.parametrize("device", ["cuda", "cpu"]) +def test_portable_backbone_matches_the_cuda_kernels( + monkeypatch: pytest.MonkeyPatch, device: str +) -> None: + cuda_backbone, portable = _pair(monkeypatch, device) + batch = make_batch(2, 48) + halves = [(0, 20), (20, 48)] + with torch.no_grad(): + state_cuda = state_port = None + for start, stop in halves: + chunk = type(batch)( + concept_ids=batch.concept_ids[:, start:stop], + aux=type(batch.aux)( + *(t[:, start:stop] if t is not None else None for t in batch.aux) + ), + ) + out_cuda, state_cuda = cuda_backbone( + move_to_device(chunk, "cuda"), state=state_cuda + ) + out_port, state_port = portable( + move_to_device(chunk, device), state=state_port + ) + torch.testing.assert_close( + out_port.cpu(), out_cuda.cpu(), atol=2e-3, rtol=2e-3 + ) diff --git a/tests/odyssey/models/test_sequence_model.py b/tests/odyssey/models/test_sequence_model.py index 0ba4a89c..daf007fa 100644 --- a/tests/odyssey/models/test_sequence_model.py +++ b/tests/odyssey/models/test_sequence_model.py @@ -313,28 +313,3 @@ def test_synthetic_training_reduces_next_token_and_concept_loss() -> None: assert task_losses[-1] < task_losses[0] * 0.8 assert concept_losses[-1] < concept_losses[0] * 0.5 - - -# --------------------------------------------------------------------------- -# The real hybrid backbone can't be executed here (CUDA-only); confirm the -# import guard fails helpfully instead of with an opaque ImportError. -# --------------------------------------------------------------------------- - - -def test_ehr_hybrid_backbone_raises_helpful_error_without_cuda() -> None: - """Can't test the real backbone's forward pass without a GPU here. - - This instead validates that, absent `mamba-ssm`, the import guard - raises a clear and actionable error rather than an opaque one. - """ - try: - import mamba_ssm # noqa: F401, PLC0415 - - pytest.skip("mamba-ssm is installed here; the guard path isn't exercised") - except ImportError: - pass - - from odyssey.models.backbones.hybrid import EHRHybridBackbone # noqa: PLC0415 - - with pytest.raises(ImportError, match="mamba-ssm"): - EHRHybridBackbone(vocab_size=10, hidden_size=8, num_hidden_layers=1) diff --git a/tests/odyssey/training/test_train.py b/tests/odyssey/training/test_train.py index f08614da..fbcfab20 100644 --- a/tests/odyssey/training/test_train.py +++ b/tests/odyssey/training/test_train.py @@ -16,7 +16,6 @@ from odyssey.data.alert_events import hazard_events_for from odyssey.data.sidecars import activate_sidecars from odyssey.data.streaming import PackedLaneSampler -from odyssey.data.types import AuxiliaryInputs, ClinicalSequenceBatch from odyssey.data.vocabulary import Vocabulary from odyssey.inference.legacy_concept_pins import HAZARD_NUM_BINS from odyssey.models.backbones.base import TimeAwareState @@ -37,7 +36,6 @@ _batch_config_fields, _combined_val_loss, _detach_state, - _move_chunk_to_device, build_model, build_objective, evaluate_streaming, @@ -155,37 +153,6 @@ def test_loss_logger_appends_to_existing_file(tmp_path: Path) -> None: assert len(lines) == 2 -# --------------------------------------------------------------------------- -# _move_chunk_to_device -# --------------------------------------------------------------------------- - - -def test_move_chunk_to_device_preserves_structure_and_values() -> None: - batch = ClinicalSequenceBatch( - concept_ids=torch.tensor([[1, 2]]), - aux=AuxiliaryInputs( - type_ids=torch.tensor([[0, 1]]), - time_stamps=torch.tensor([[0.0, 1.0]]), - ages=torch.tensor([[30.0, 30.0]]), - visit_orders=torch.tensor([[0, 0]]), - visit_segments=torch.tensor([[0, 0]]), - ), - ) - - moved = _move_chunk_to_device(batch, "cpu") - - assert isinstance(moved, ClinicalSequenceBatch) - assert isinstance(moved.aux, AuxiliaryInputs) - assert torch.equal(moved.concept_ids, batch.concept_ids) - assert torch.equal(moved.aux.time_stamps, batch.aux.time_stamps) - - -def test_move_chunk_to_device_passes_through_non_tensor_values() -> None: - assert _move_chunk_to_device(5, "cpu") == 5 - assert _move_chunk_to_device("x", "cpu") == "x" - assert _move_chunk_to_device(None, "cpu") is None - - # --------------------------------------------------------------------------- # _detach_state # --------------------------------------------------------------------------- diff --git a/tests/odyssey/utils/test_device.py b/tests/odyssey/utils/test_device.py new file mode 100644 index 00000000..0c083cb7 --- /dev/null +++ b/tests/odyssey/utils/test_device.py @@ -0,0 +1,43 @@ +"""Device selection and transfer across CUDA, Apple MPS and CPU.""" + +import pytest +import torch + +from odyssey.utils import device as device_mod +from odyssey.utils.device import default_device, resolve_device, to_device + + +@pytest.mark.parametrize( + ("cuda", "mps", "expected"), + [(True, True, "cuda"), (False, True, "mps"), (False, False, "cpu")], +) +def test_default_device_prefers_cuda_then_mps_then_cpu( + monkeypatch: pytest.MonkeyPatch, cuda: bool, mps: bool, expected: str +) -> None: + monkeypatch.setattr(device_mod.torch.cuda, "is_available", lambda: cuda) + monkeypatch.setattr(device_mod.torch.backends.mps, "is_available", lambda: mps) + assert default_device() == expected + + +def test_resolve_device_maps_auto_and_none_and_keeps_explicit_choices() -> None: + assert resolve_device("auto") == default_device() + assert resolve_device(None) == default_device() + assert resolve_device("cpu") == "cpu" + assert resolve_device("cuda:1") == "cuda:1" + + +def test_to_device_keeps_dtypes_off_mps() -> None: + stamps = torch.tensor([1.0, 2.5], dtype=torch.float64) + moved = to_device(stamps, "cpu") + assert moved.dtype == torch.float64 and torch.equal(moved, stamps) + ids = torch.tensor([1, 2]) + assert to_device(ids, torch.device("cpu")).dtype == torch.long + + +@pytest.mark.skipif(not torch.backends.mps.is_available(), reason="needs Apple MPS") +def test_to_device_casts_float64_to_float32_on_mps_only() -> None: + stamps = torch.tensor([123.25, 456.5], dtype=torch.float64) + moved = to_device(stamps, "mps") + assert moved.device.type == "mps" and moved.dtype == torch.float32 + assert torch.equal(moved.cpu(), stamps.float()) + assert to_device(torch.tensor([1, 2]), "mps").dtype == torch.long