diff --git a/.gitignore b/.gitignore index 16143c0..bcbd562 100644 --- a/.gitignore +++ b/.gitignore @@ -169,3 +169,9 @@ cython_debug/ # and can be added to the global gitignore or merged into this file. For a more nuclear # option (not recommended) you can uncomment the following to ignore the entire idea folder. .idea/ + +# Prov +test_data +mlflow.db +mlartifacts +/example* diff --git a/README.md b/README.md index 43f0287..b366ae2 100644 --- a/README.md +++ b/README.md @@ -1 +1,426 @@ -# ML Kit \ No newline at end of file +# rationai.mlkit — Unified ML Provenance & Metrics Toolkit + +Automatic PROV-O-aware experiment tracking for pathology image classification, +with Lightning/Hydra integration and pathology-specific metric layers. + +All metadata (user, dataset, hardware, docker, git, environment, model architecture, +optimizer/scheduler settings, console output) is captured automatically via the +`@autolog` decorator and `ProvenanceCallback` — the user only writes their training loop. + +--- + +## Setup + +```bash +uv venv +source .venv/bin/activate +uv sync +``` + +Start the MLflow server: + +```bash +mlflow ui --host 127.0.0.1 --port 5000 # → http://localhost:5000 +``` + +--- + +## Quick start + +### 1. Register users and dataset (one-time setup) + +Before training, register researchers and datasets so provenance can reference them: + +```bash +# Edit example_provenance_setup.py with your own data, then run: +python example_provenance_setup.py +``` + +This creates runs in the `User_Registry` and `Dataset_Registry` MLflow experiments. + +### 2. Run a training experiment + +Use `@autolog` + `ProvenanceCallback` in a Hydra-based training script: + +```bash +python example_provenance_train.py +``` + +Check the MLflow UI at http://localhost:5000 to see the full provenance graph. + +See [example_provenance_setup.py](example_provenance_setup.py) and +[example_provenance_train.py](example_provenance_train.py) for complete working examples. + +--- + +## API reference + +### User registration + +Register a researcher into the `User_Registry` experiment: + +```python +from rationai.mlkit.provenance import register_new_user + +register_new_user( + username="researcher_01", + real_name="Jane Doe", + email="jane.doe@example.com", + organization="Example Org", + lead_name="John Smith", + lead_email="john.smith@example.com", +) +``` + +### Dataset registration + +Register a dataset (requires `manifest.csv` in the dataset directory): + +```python +from rationai.mlkit.provenance import register_dataset + +run_id = register_dataset( + dataset_dir="data/cohorts/pato_01", + dataset_name="pato_cohort_01", + version="2.0", +) +print(f"Registered as {run_id}") +``` + +Verify a dataset's file integrity: + +```python +from rationai.mlkit.provenance import verify_dataset + +result = verify_dataset(manifest_path="data/cohorts/pato_01/manifest.csv") +print(result["verified"]) # True if all files match +``` + +### Provenance — `@autolog` decorator (Hydra training scripts) + +Full auto-capture for Hydra-based training runs: + +```python +import hydra +from omegaconf import DictConfig +from rationai.mlkit import Trainer, autolog +from rationai.mlkit.lightning.loggers.mlflow import MLFlowLogger +from rationai.mlkit.lightning.callbacks import ProvenanceCallback + +@hydra.main(config_path=".", config_name="train_cfg", version_base=None) +@autolog +def main(config: DictConfig, logger: MLFlowLogger) -> None: + model = MyLightningModule() + data = MyDataModule(batch_size=config.batch_size) + + trainer = Trainer( + max_epochs=config.epochs, + logger=logger, + callbacks=[ + ProvenanceCallback(model_name="my_model_v1") + ], + ) + trainer.fit(model, datamodule=data) +``` + +| Auto-captured | Details | +|---|---| +| **User** | Resolved from `MLFLOW_USER` env → git config → linked to `User_Registry` run | +| **Dataset** | Latest `Dataset_Registry` run, file sizes verified | +| **Train/test split** | Manifest auto-discovered, CSVs saved as artifacts | +| **Model architecture** | Class name, param counts, per-layer summary | +| **Optimizer** | Type, lr, momentum, weight_decay, … | +| **Scheduler** | Type, step_size, gamma, milestones, … | +| **Hardware** | GPU/CPU/RAM/OS/Python version | +| **Docker** | Container ID, image name + hash (if in container) | +| **Git** | Commit, branch, remote URL | +| **Environment** | Frozen `requirements.txt` | +| **Console output** | stdout/stderr → `logs/console.log` artifact (ANSI-aware) | +| **PROV document** | OpenProvenance JSON → `provenance/prov.json` artifact | + +### Lightning callbacks + +#### ProvenanceCallback + +Drop-in callback that captures full PROV-O provenance for every training run: + +```python +from rationai.mlkit.lightning.callbacks import ProvenanceCallback + +callback = ProvenanceCallback( + model_name="resnet_v1", + experiment_name="Training_Pipeline", +) + +trainer = Trainer(callbacks=[callback], ...) +``` + +| Parameter | Default | Description | +|---|---|---| +| `model_name` | `None` | Model identifier (used in artifact paths) | +| `experiment_name` | `"Training_Pipeline"` | MLflow experiment name | +| `manifest_path` | auto-discover | Path to `manifest.csv` | +| `data_root` | auto-discover | Root directory of dataset | +| `test_size` | `0.2` | Test split fraction | +| `random_state` | `42` | Random seed for split | +| `fail_fast` | `True` | Stop training on verification failure | +| `strict` | `False` | Require all provenance fields | +| `register_model` | `True` | Register model architecture to MLflow | +| `register_optimizer` | `True` | Register optimizer config to MLflow | +| `register_scheduler` | `True` | Register scheduler config to MLflow | + +#### EnvironmentCallback + +Captures environment metadata (git, hardware, docker, environment freeze): + +```python +from rationai.mlkit.lightning.callbacks import EnvironmentCallback + +callback = EnvironmentCallback( + skip_hardware=False, # capture GPU/CPU/RAM info + snapshot_env=True, # freeze requirements.txt +) +``` + +#### DatasetVerificationCallback + +Verifies dataset integrity and optionally performs a stratified split: + +```python +from rationai.mlkit.lightning.callbacks import DatasetVerificationCallback + +callback = DatasetVerificationCallback( + test_size=0.2, # 0.0 to skip split + random_state=42, + fail_fast=True, # stop training if verification fails +) +``` + +### Metrics + +#### AggregatedMetricCollection — tile → slide aggregation + +Group tile-level predictions by slide and compute metrics at the slide level: + +```python +from torchmetrics import Accuracy +from rationai.mlkit import ( + AggregatedMetricCollection, + MaxAggregator, + MeanAggregator, +) + +agg = AggregatedMetricCollection( + metrics={"accuracy": Accuracy(task="binary")}, + aggregator=MaxAggregator(), +) + +for preds, targets, slide_ids in val_loader: + agg.update(preds, targets, keys=slide_ids) + +results = agg.compute() +# → {"key": ["slide1", "slide2"], "accuracy": [0.95, 0.87]} +``` + +Available aggregators: `MaxAggregator`, `MeanAggregator`, `TopKAggregator`, `MeanPoolMaxAggregator`. + +#### NestedMetricCollection — per-slide multiclass metrics + +Compute multiple torchmetrics per slide with class-level breakdowns: + +```python +from rationai.mlkit import NestedMetricCollection +from torchmetrics import Accuracy, Precision + +metrics = NestedMetricCollection( + metrics={ + "accuracy": Accuracy(task="multiclass", num_classes=3), + "precision": Precision(task="multiclass", num_classes=3, average=None), + }, + key_name="slide", + class_names=["benign", "low_grade", "high_grade"], +) + +metrics.update(preds, targets, keys=slide_ids) +result = metrics.compute() +``` + +### Stream capture — ANSI-aware stdout/stderr logging + +Captures console output (including progress bars and ANSI color codes) +without corrupting the log: + +```python +from rationai.mlkit import StreamCapture + +with StreamCapture(stream="stdout") as capture: + print("Hello!") + print("\033[92mGreen text\033[0m") # ANSI color + +text = capture.get_text() # Raw captured text +clean = capture.get_clean_text() # ANSI codes stripped +``` + +### Data utilities + +#### StratifiedBatchSampler + +Balanced class sampling across batches: + +```python +from rationai.mlkit import StratifiedBatchSampler + +sampler = StratifiedBatchSampler( + data_indices=[[0, 1, 2, 3], [4, 5, 6, 7]], # per-class indices + batch_size=4, +) +for batch in sampler: + train(model, batch) +``` + +#### MetaTiledSlides + +Load tile data from parquet or MLflow artifact URIs: + +```python +from rationai.mlkit import MetaTiledSlides + +dataset = MetaTiledSlides( + manifest_uri="s3://bucket/data/manifest.parquet", + tile_size=256, +) +``` + +### Lightning integration + +| Component | Import | Purpose | +|---|---|---| +| `Trainer` | `from rationai.mlkit import Trainer` | Lightning Trainer with MLflow checkpoint sync | +| `MLFlowLogger` | `from rationai.mlkit import MLFlowLogger` | Logger with git tags, stream capture, checkpoint sync | +| `MultiloaderLifecycle` | `from rationai.mlkit import MultiloaderLifecycle` | Per-dataloader callback hooks | +| `ProvenanceCallback` | `from rationai.mlkit.lightning.callbacks import ProvenanceCallback` | Full PROV-O provenance capture | +| `EnvironmentCallback` | `from rationai.mlkit.lightning.callbacks import EnvironmentCallback` | Environment metadata capture | +| `DatasetVerificationCallback` | `from rationai.mlkit.lightning.callbacks import DatasetVerificationCallback` | Dataset verification + split | +| `with_cli_args` | `from rationai.mlkit import with_cli_args` | Programmatic config injection (Hydra) | + +--- + +## Project structure + +```text +. +├── example_provenance_setup.py # Setup: register users & dataset +├── example_provenance_train.py # Training with @autolog + ProvenanceCallback +├── example_provenance_train_cfg.yaml # Hydra config for the training example +├── pyproject.toml # Project metadata + deps +├── test_data/ # Dummy datasets (gitignored) +└── rationai/ + └── mlkit/ + ├── __init__.py # Package exports (lazy loading) + ├── autolog.py # @autolog decorator for Hydra scripts + ├── with_cli_args.py # Programmatic config injection + ├── provenance/ # Dataset & user registration + │ ├── __init__.py + │ ├── register_dataset.py # register_dataset, verify_dataset + │ └── register_user.py # register_new_user + ├── stream/ # ANSI-aware console capture + │ ├── stream_capture.py + │ ├── stream_logger.py + │ └── stream_modifier.py + ├── metrics/ # Slide-level metric aggregation + │ ├── aggregated_metric_collection.py + │ ├── nested_metric_collection.py + │ ├── aggregators.py + │ └── lazy_metric_dict.py + ├── data/ # Data utilities + │ ├── shard_parquet.py + │ ├── samplers/ + │ │ └── stratified_batch_sampler.py + │ └── datasets/ + │ ├── meta_tiled_slides.py + │ ├── openslide_tiles_dataset.py + │ └── slides_tiles_loader.py + └── lightning/ # Lightning + Hydra integration + ├── trainer.py + ├── with_cli_args.py + ├── callbacks/ + │ ├── provenance.py # ProvenanceCallback + │ ├── environment.py # EnvironmentCallback + │ ├── dataset_verification.py # DatasetVerificationCallback + │ └── multiloader_lifecycle.py + └── loggers/ + └── mlflow.py # MLFlowLogger (checkpoint sync, git tags) +``` + +--- + +## MLflow experiments + +| Experiment | Purpose | +|---|---| +| `User_Registry` | Stores user identity runs (username, real name, org) | +| `Dataset_Registry` | Stores dataset manifest runs with file provenance | +| `Training_Pipeline` | Training runs with full auto-captured provenance | + +Cross-run tags (`user_run_id`, `dataset_run_id`) on training runs make +PROV graph reconstruction machine-readable. + +--- + +## PROV document (W3C PROV-O) + +Each training run emits a self-contained OpenProvenance JSON document at +`provenance/prov.json` (MLflow artifact). This document is compatible with +the `prov_mlflow` Java tool and +follows the W3C PROV-O standard. + +### Structure + +The document uses a bundle wrapper (`{"bundle": {"storage:": {...}}}`) +with 10 namespace prefixes and 7 sections: + +| Section | Purpose | +|---|---| +| `prefix` | Namespace URIs (`gen:`, `schema:`, `cpm:`, `prov:`, `sosa:`, …) | +| `entity` | Input data (dataset/WSI as `sosa:Sample`), metadata bundle (`cpm:BundleMetadata`) | +| `activity` | Training run (`schema:Action`) with hyperparameters, hardware, git commit; CPM wrapper (`cpm:mainActivity`) | +| `agent` | Researcher (`schema:Person`) with name, email, affiliation | +| `used` | Run activity → dataset entity | +| `wasAssociatedWith` | Run activity → agent | +| `wasGeneratedBy` | Metadata bundle ← run activity | + +### PROV elements + +| Element | ID Pattern | Type | Content | +|---|---|---|---| +| Agent | `gen:user_` | `schema:Person` | Name, email, affiliation | +| Dataset Entity | `gen:dataset_` | `sosa:Sample` | Input data (or WSI if path available) | +| Run Activity | `gen:run_` | `schema:Action` | Hyperparams, hardware, git, optimizer/scheduler settings | +| Meta Bundle | `meta:` | `cpm:BundleMetadata` | Full param/metric snapshot (model arch, layer summary, split info, final metrics) | +| Main Activity | `blank:TrainingRun_` | `cpm:mainActivity` | CPM wrapper linking meta bundle to run activity | + +### Compatibility + +The generated document matches the Java `prov_mlflow` +output format: + +- Bundle-wrapped JSON structure +- Qualified type annotations (`{"type": "prov:QUALIFIED_NAME", "$": "schema:Person"}`) +- Array-wrapped property values (Java convention) +- Blank node IDs for relationship references (`_:n0`, `_:n1`, …) +- CPM metadata entity with hardware, model, and metric details + +### Accessing the document + +```bash +# Via MLflow UI → Artifacts tab → provenance/prov.json + +# Via Python API +import mlflow, json +mlflow.set_tracking_uri("http://localhost:5000") +client = mlflow.tracking.MlflowClient() +run = client.search_runs("")[-1] +prov_path = client.download_artifacts(run.info.run_id, "provenance/prov.json") +with open(prov_path) as f: + prov = json.load(f) +``` diff --git a/pyproject.toml b/pyproject.toml index 8124b29..ced1cb3 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -20,6 +20,7 @@ dependencies = [ "pyarrow>=23.0.1", "rationai-masks", "ratiopath>=1.2.0", + "scikit-learn>=1.6.0", "torch>=2.9.1", ] diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index 43af683..6259b19 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -1,6 +1,79 @@ +"""rationai.mlkit — ML toolkit with provenance tracking.""" + +from typing import Any + from rationai.mlkit.autolog import autolog -from rationai.mlkit.lightning import Trainer -from rationai.mlkit.with_cli_args import with_cli_args +from rationai.mlkit.provenance.dataset import register_dataset +from rationai.mlkit.stream import StreamCapture, StreamLogger + + +__all__ = [ + "AggregatedMetricCollection", + "Aggregator", + "LazyMetricDict", + "MLFlowLogger", + "MaxAggregator", + "MeanAggregator", + "MeanPoolMaxAggregator", + "MetaTiledSlides", + "MultiloaderLifecycle", + "NestedMetricCollection", + "OpenSlideTilesDataset", + "PDMStratifiedBatchSampler", + "ProvenanceCallback", + "StratifiedBatchSampler", + "StreamCapture", + "StreamLogger", + "TopKAggregator", + "Trainer", + "autolog", + "register_dataset", + "with_cli_args", +] + + +def __getattr__(name: str) -> Any: + if name in ("Trainer", "MultiloaderLifecycle", "with_cli_args"): + import importlib + + _mod = importlib.import_module("rationai.mlkit.lightning") + return getattr(_mod, name) + + if name == "MLFlowLogger": + from rationai.mlkit.lightning.loggers.mlflow import MLFlowLogger + + return MLFlowLogger + + if name == "ProvenanceCallback": + from rationai.mlkit.lightning.callbacks import ProvenanceCallback + + return ProvenanceCallback + + if name in ( + "AggregatedMetricCollection", + "Aggregator", + "MaxAggregator", + "MeanAggregator", + "MeanPoolMaxAggregator", + "TopKAggregator", + "NestedMetricCollection", + "LazyMetricDict", + ): + import importlib + + _mod = importlib.import_module("rationai.mlkit.metrics") + return getattr(_mod, name) + + if name in ("StratifiedBatchSampler", "PDMStratifiedBatchSampler"): + import importlib + + _mod = importlib.import_module("rationai.mlkit.data.samplers") + return getattr(_mod, name) + + if name in ("MetaTiledSlides", "OpenSlideTilesDataset"): + import importlib + _mod = importlib.import_module("rationai.mlkit.data.datasets") + return getattr(_mod, name) -__all__ = ["Trainer", "autolog", "with_cli_args"] + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/rationai/mlkit/lightning/callbacks/__init__.py b/rationai/mlkit/lightning/callbacks/__init__.py index 914d283..65a01a4 100644 --- a/rationai/mlkit/lightning/callbacks/__init__.py +++ b/rationai/mlkit/lightning/callbacks/__init__.py @@ -1,6 +1,16 @@ +from rationai.mlkit.lightning.callbacks.dataset_verification import ( + DatasetVerificationCallback, +) +from rationai.mlkit.lightning.callbacks.environment import EnvironmentCallback from rationai.mlkit.lightning.callbacks.multiloader_lifecycle import ( MultiloaderLifecycle, ) +from rationai.mlkit.lightning.callbacks.provenance import ProvenanceCallback -__all__ = ["MultiloaderLifecycle"] +__all__ = [ + "DatasetVerificationCallback", + "EnvironmentCallback", + "MultiloaderLifecycle", + "ProvenanceCallback", +] diff --git a/rationai/mlkit/lightning/callbacks/dataset_verification.py b/rationai/mlkit/lightning/callbacks/dataset_verification.py new file mode 100644 index 0000000..0ac451c --- /dev/null +++ b/rationai/mlkit/lightning/callbacks/dataset_verification.py @@ -0,0 +1,194 @@ +"""Lightning callback that verifies the dataset against MLflow on trainer start. + +Optionally performs a stratified train/test split and logs it as an MLflow +artifact alongside verification results. + +Example:: + + from rationai.mlkit.lightning.callbacks import DatasetVerificationCallback + + # Verification only + trainer = Trainer(callbacks=[DatasetVerificationCallback()]) + + # Verification + train/test split + trainer = Trainer( + callbacks=[DatasetVerificationCallback(test_size=0.2, random_state=42)], + ) +""" + +from __future__ import annotations + +import logging +import os +import shutil +import uuid +from typing import Any + +import mlflow +from lightning.pytorch.callbacks import Callback + + +log = logging.getLogger(__name__) + + +class DatasetVerificationCallback(Callback): + """Run dataset verification (and optional train/test split) at training start. + + Auto-detects ``manifest.csv`` under ``data/``, looks up the latest + ``Dataset_Registry`` run, and checks that per-file sizes still match. + + Stores results on ``self`` so sibling callbacks (e.g. ``ProvenanceCallback``) + can read them without duplicating work: + + - ``_verification`` (dict | None) — verification result + - ``_split_data`` (dict | None) — train/test split data + + Args: + manifest_path: Path to manifest.csv (auto-detected if None). + test_size: Fraction of data for the test split. Set to 0 to skip splitting. + random_state: Random seed for train/test split. + fail_fast: Abort training if dataset verification fails. + """ + + def __init__( + self, + manifest_path: str | None = None, + test_size: float = 0.0, + random_state: int = 42, + fail_fast: bool = True, + ) -> None: + """Initialise the dataset verification callback. + + Args: + manifest_path: Path to manifest.csv (auto-detected if None). + test_size: Fraction of data for the test split. Set to 0 to skip splitting. + random_state: Random seed for train/test split. + fail_fast: Abort training if dataset verification fails. + """ + self._manifest_path = manifest_path + self.test_size = test_size + self.random_state = random_state + self.fail_fast = fail_fast + self._done = False + self._verification: dict[str, Any] | None = None + self._split_data: dict[str, Any] | None = None + + def on_fit_start(self, trainer: Any, pl_module: Any) -> None: + """Verify the dataset and optionally split into train/test. + + Looks up the latest ``Dataset_Registry`` run, checks file integrity, + logs verification params to MLflow, and (when ``test_size > 0``) + performs a stratified train/test split saved as an MLflow artifact. + """ + if self._done: + return + self._done = True + + from rationai.mlkit.provenance.dataset import ( + _detect_manifest, + _lookup_dataset_run, + _verify_dataset, + ) + from rationai.mlkit.provenance.dataset import ( + load_manifest as _load_manifest, + ) + + manifest_path = self._manifest_path + data_root = None + if manifest_path is None: + manifest_path, data_root = _detect_manifest() + + if manifest_path is None: + log.warning( + "[DatasetVerificationCallback] No manifest.csv found — skipping" + ) + return + + if data_root is None: + data_root = os.path.dirname(os.path.abspath(manifest_path)) + + # ── Verification ──────────────────────────────────────── + dataset_run_id = _lookup_dataset_run() + verification = _verify_dataset(manifest_path, data_root, dataset_run_id) + self._verification = verification + + for detail in verification.get("details", []): + log.info(f" [DatasetVerificationCallback] {detail}") + + # ── Log verification results ──────────────────────────── + if mlflow.active_run(): + mlflow.log_params( + { + "dataset_verified": verification["verified"], + "dataset_file_sizes_match": verification["file_sizes_match"] + is True, + "dataset_files_missing": verification["files_missing"], + "dataset_files_total": verification["files_total"], + } + ) + if verification["verified"]: + mlflow.set_tag("dataset_verification", "VERIFIED") + else: + mlflow.set_tag("dataset_verification", "MISMATCH") + mlflow.set_tag( + "dataset_verification_details", + "; ".join(verification.get("details", [])), + ) + + if self.fail_fast and not verification["verified"]: + raise RuntimeError( + "Dataset verification failed — aborting training.\n" + + "\n".join(f" {d}" for d in verification.get("details", [])) + ) + + # ── Train/test split (optional) ───────────────────────── + if self.test_size > 0: + import pandas as pd + from sklearn.model_selection import train_test_split + + samples = _load_manifest(manifest_path, data_root) + + train_samples, test_samples = train_test_split( + samples, + test_size=self.test_size, + random_state=self.random_state, + stratify=[s["label"] for s in samples], + ) + + self._split_data = { + "train": train_samples, + "test": test_samples, + "test_size": self.test_size, + "random_state": self.random_state, + } + + # Log split as artifact + if mlflow.active_run(): + split_dir = f"_mlflow_split_{uuid.uuid4().hex[:8]}" + os.makedirs(split_dir, exist_ok=True) + for subset_name, subset_samples in [ + ("train", train_samples), + ("test", test_samples), + ]: + split_file = os.path.join(split_dir, f"{subset_name}_split.csv") + pd.DataFrame(subset_samples).to_csv(split_file, index=False) + + mlflow.log_artifacts(split_dir, artifact_path="split") + shutil.rmtree(split_dir, ignore_errors=True) + + # Log split counts + train_labels = [s["label"] for s in train_samples] + test_labels = [s["label"] for s in test_samples] + mlflow.log_params( + { + "train_samples": len(train_samples), + "test_samples": len(test_samples), + "train_positive": sum(train_labels), + "train_negative": len(train_labels) - sum(train_labels), + "test_positive": sum(test_labels), + "test_negative": len(test_labels) - sum(test_labels), + "split_test_size": self.test_size, + "split_random_state": self.random_state, + "split_stratified": True, + } + ) diff --git a/rationai/mlkit/lightning/callbacks/environment.py b/rationai/mlkit/lightning/callbacks/environment.py new file mode 100644 index 0000000..3447bf7 --- /dev/null +++ b/rationai/mlkit/lightning/callbacks/environment.py @@ -0,0 +1,334 @@ +"""Lightning callback that captures environment provenance (hardware, docker, env snapshot). + +Extracted from ProvenanceCallback so users who only need environment metadata +don't have to pull in the full PROV machinery. + +Example:: + + from rationai.mlkit.lightning.callbacks import EnvironmentCallback + + trainer = Trainer( + callbacks=[EnvironmentCallback()], + logger=MLFlowLogger(...), + ) +""" + +from __future__ import annotations + +import contextlib +import hashlib +import logging +import os +import platform +import shutil +import subprocess +import uuid +from datetime import UTC, datetime +from typing import Any + +import mlflow +import pandas as pd +import torch +from lightning.pytorch.callbacks import Callback + + +log = logging.getLogger(__name__) + + +# ────────────────────────────────────────────── +# Helpers +# ────────────────────────────────────────────── + + +def _lookup_user_run() -> tuple[str | None, dict[str, str]]: + """Find the user run from User_Registry. Auto-detect username.""" + from rationai.mlkit.provenance.dataset import _lookup_experiment + + username = os.environ.get("MLFLOW_USER") + if not username: + with contextlib.suppress(subprocess.CalledProcessError): + username = ( + subprocess.check_output( + ["git", "config", "user.name"], + stderr=subprocess.DEVNULL, + ) + .decode() + .strip() + ) + if not username: + username = os.environ.get("USER", "unknown") + + exp_id = _lookup_experiment("User_Registry") + if exp_id is None: + return None, {} + + _runs_df = mlflow.search_runs(experiment_ids=[exp_id]) + runs_df: pd.DataFrame = _runs_df # search_runs may return RunList in old mlflow + if runs_df.empty: + return None, {} + + matched = runs_df[runs_df["tags.username"] == username] + if matched.empty: + matched = runs_df.head(1) + + row = matched.iloc[0] + run_obj = mlflow.get_run(row.run_id) + return row.run_id, dict(run_obj.data.tags) + + +def _detect_hardware() -> dict[str, str | int]: + """Detect CPU/GPU/hardware info.""" + info: dict[str, str | int] = {} + + if torch.cuda.is_available(): + info["gpu_name"] = torch.cuda.get_device_name(0) + info["gpu_count"] = torch.cuda.device_count() + cap = torch.cuda.get_device_capability(0) + info["gpu_compute_capability"] = f"{cap[0]}.{cap[1]}" + info["cuda_version"] = torch.version.cuda or "unknown" + else: + info["gpu_name"] = "none" + + info["cpu_count_logical"] = os.cpu_count() or 0 + info["os_platform"] = platform.platform() + info["python_version"] = platform.python_version() + + try: + import psutil + + mem = psutil.virtual_memory() + info["ram_total_gb"] = round(mem.total / 1e9, 1) + except ImportError: + pass + + return info + + +def _detect_docker() -> dict[str, str | bool]: + """Detect if running inside Docker and extract container info.""" + info: dict[str, str | bool] = {"docker": False} + + if os.path.exists("/.dockerenv"): + info["docker"] = True + + if not info["docker"]: + try: + with open("/proc/self/cgroup") as f: + for line in f: + for p in line.strip().split("/"): + if len(p) >= 12 and all( + c in "0123456789abcdef" for c in p[:12] + ): + info["docker"] = True + info["container_id_short"] = p[:12] + break + except FileNotFoundError: + pass + + if not info.get("container_id_short"): + try: + with open("/proc/self/mountinfo") as f: + for line in f: + for p in line.split(): + if len(p) == 64 and all(c in "0123456789abcdef" for c in p): + info["container_id_short"] = p[:12] + break + except FileNotFoundError: + pass + + if info["docker"]: + cid = str(info.get("container_id_short", "")) + if cid: + try: + result = subprocess.run( + ["docker", "inspect", "--format={{.Config.Image}}", cid], + capture_output=True, + text=True, + timeout=5, + ) + if result.returncode == 0 and result.stdout.strip(): + image = result.stdout.strip() + info["docker_image"] = image + info["docker_image_hash"] = hashlib.sha256( + image.encode() + ).hexdigest()[:16] + except (subprocess.TimeoutExpired, FileNotFoundError): + pass + + return info + + +def _snapshot_environment(artifact_dir: str) -> str: + """Freeze environment to *artifact_dir* and return the pip-freeze text.""" + req_path = os.path.join(artifact_dir, "requirements_frozen.txt") + with open(req_path, "w") as f: + subprocess.run(["uv", "pip", "freeze"], stdout=f, check=True) + + for src in ("pyproject.toml", "uv.lock"): + if os.path.exists(src): + shutil.copy2(src, os.path.join(artifact_dir, src)) + + with open(req_path) as f: + return f.read() + + +# ────────────────────────────────────────────── +# Callback +# ────────────────────────────────────────────── + + +class EnvironmentCallback(Callback): + """Capture hardware, docker, git, user, and environment snapshot at training start. + + Stores results on ``self`` so sibling callbacks (e.g. ``ProvenanceCallback``) + can read them without duplicating work. + + Attributes set after ``on_fit_start``: + - ``_git_commit``, ``_git_url``, ``_git_branch`` + - ``_hardware`` (dict) + - ``_docker`` (dict) + - ``_frozen_requirements`` (str | None) + - ``_user_run_id``, ``_user_tags`` + + Args: + skip_hardware: Skip hardware detection if True (MLflow system metrics + are already enabled). Auto-detected from trainer loggers by default. + snapshot_env: If True, freeze the environment to an MLflow artifact. + strict: If True, re-raise errors from optional steps instead of logging. + """ + + def __init__( + self, + skip_hardware: bool = False, + snapshot_env: bool = True, + strict: bool = False, + ) -> None: + """Initialise the environment callback. + + Args: + skip_hardware: Skip hardware detection if True. Auto-detected from + trainer loggers by default. + snapshot_env: If True, freeze the environment to an MLflow artifact. + strict: If True, re-raise errors from optional steps instead of logging. + """ + self.skip_hardware = skip_hardware + self.snapshot_env = snapshot_env + self.strict = strict + + # Populated during on_fit_start + self._git_commit: str = "unknown" + self._git_url: str = "unknown" + self._git_branch: str = "unknown" + self._hardware: dict[str, str | int] = {} + self._docker: dict[str, str | bool] = {} + self._frozen_requirements: str | None = None + self._user_run_id: str | None = None + self._user_tags: dict[str, str] = {} + self._temp_dirs: list[str] = [] + + def on_fit_start(self, trainer: Any, pl_module: Any) -> None: + """Capture environment metadata at the start of training.""" + if not mlflow.active_run(): + return + + # ── Git info (read from MLflow tags set by MLFlowLogger) ── + try: + run = mlflow.active_run() + git_tags: dict[str, str] = {} + if run and run.info and run.info.run_id: + client = mlflow.tracking.MlflowClient() + run_data = client.get_run(run.info.run_id) + git_tags = dict(run_data.data.tags) if run_data.data.tags else {} + self._git_commit = git_tags.get( + "mlflow.source.git.commit", git_tags.get("git.commit", "unknown") + ) + self._git_url = git_tags.get( + "mlflow.source.git.repoUrl", git_tags.get("git.repo_url", "unknown") + ) + self._git_branch = git_tags.get( + "mlflow.source.git.branch", git_tags.get("git.branch", "unknown") + ) + except Exception as e: + if self.strict: + raise + log.warning("[EnvironmentCallback] Git info failed: %s", e) + + # ── User lookup ───────────────────────────────────────── + try: + user_run_id, user_tags = _lookup_user_run() + self._user_run_id = user_run_id + self._user_tags = user_tags or {} + except Exception as e: + if self.strict: + raise + log.warning("[EnvironmentCallback] User lookup failed: %s", e) + + # ── Hardware (skip if MLflow system metrics are on) ───── + if not self.skip_hardware: + sys_metrics_on = any( + getattr(logger, "log_system_metrics", False) + for logger in trainer.loggers + ) + if sys_metrics_on: + log.info( + "[EnvironmentCallback] Skipping hardware — MLflow system metrics enabled" + ) + else: + try: + self._hardware = _detect_hardware() + except Exception as e: + if self.strict: + raise + log.warning( + "[EnvironmentCallback] Hardware detection failed: %s", e + ) + + # ── Docker detection ──────────────────────────────────── + try: + self._docker = _detect_docker() + except Exception as e: + if self.strict: + raise + log.warning("[EnvironmentCallback] Docker detection failed: %s", e) + + # ── Log tags (git + user) ─────────────────────────────── + env_tags: dict[str, str] = {} + if self._user_run_id: + env_tags["user_run_id"] = self._user_run_id + for key in ("username", "real_name", "organization"): + if key in self._user_tags: + env_tags[key] = self._user_tags[key] + + from rationai.mlkit.provenance.dataset import _lookup_dataset_run + + dataset_run_id = _lookup_dataset_run() + if dataset_run_id: + env_tags["dataset_run_id"] = dataset_run_id + + env_tags.update( + { + "git_commit": self._git_commit, + "git_url": self._git_url, + "git_branch": self._git_branch, + "prov_start_time": datetime.now(UTC).isoformat(), + } + ) + mlflow.set_tags(env_tags) + + # ── Log hardware + docker params ──────────────────────── + all_params: dict[str, str | float | int] = {**self._hardware, **self._docker} + if all_params: + mlflow.log_params(all_params) + + # ── Environment snapshot ──────────────────────────────── + if self.snapshot_env: + artifact_dir = f"_mlflow_env_{uuid.uuid4().hex[:8]}" + os.makedirs(artifact_dir, exist_ok=True) + self._temp_dirs.append(artifact_dir) + try: + self._frozen_requirements = _snapshot_environment(artifact_dir) + mlflow.log_artifacts(artifact_dir, artifact_path="environment") + except Exception as e: + if self.strict: + raise + log.warning("[EnvironmentCallback] Environment snapshot failed: %s", e) diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py new file mode 100644 index 0000000..18feef9 --- /dev/null +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -0,0 +1,613 @@ +"""Lightning callback that captures PROV-O provenance for MLflow runs. + +Slim callback that depends on sibling callbacks for environment and dataset +verification data: + + - :class:`~rationai.mlkit.lightning.callbacks.environment.EnvironmentCallback` + provides git info, hardware, docker, env snapshot, and user tags. + - :class:`~rationai.mlkit.lightning.callbacks.dataset_verification.DatasetVerificationCallback` + provides dataset verification and train/test split results. + +When used alone, it falls back to doing its own environment/verification work +so the user gets a single-drop-in experience. + +Example:: + + from rationai.mlkit.lightning.callbacks import ProvenanceCallback + + trainer = Trainer( + callbacks=[ProvenanceCallback(model_name="resnet_v1")], + logger=MLFlowLogger(...), + ) +""" + +from __future__ import annotations + +import json +import logging +import os +import shutil +import uuid +from datetime import UTC, datetime +from typing import Any + +import mlflow +import pandas as pd +from lightning.pytorch.callbacks import Callback + +# Import shared PROV helpers from prov.py to avoid duplication +from rationai.mlkit.provenance.common import ( + get_prov_prefixes as _get_prov_prefixes, +) +from rationai.mlkit.provenance.run import ( + build_training_run_prov as _build_prov_document, +) + + +log = logging.getLogger(__name__) + + +# ────────────────────────────────────────────── +# Model / Optimizer / Scheduler summaries +# ────────────────────────────────────────────── + + +def _model_summary(model: Any) -> dict[str, str | int]: + """Extract architecture details from a torch.nn.Module.""" + info: dict[str, str | int] = {} + + total_params = sum(p.numel() for p in model.parameters()) + trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) + info["total_parameters"] = total_params + info["trainable_parameters"] = trainable_params + + layer_lines = [] + for name, module in model.named_modules(): + if name == "": + continue + param_count = sum(p.numel() for p in module.parameters(recurse=False)) + children = len(list(module.children())) + layer_lines.append( + f"{name}({type(module).__name__}): params={param_count}, " + f"children={children}", + ) + + layer_summary: str = "\n".join(layer_lines[:20]) + if len(layer_lines) > 20: + layer_summary += f"\n... ({len(layer_lines)} layers total)" + info["layer_summary"] = layer_summary + + info["model_class"] = type(model).__name__ + return info + + +def _optimizer_summary(optimizer: Any) -> dict[str, str | float]: + """Extract optimizer settings from torch.optim.Optimizer.""" + info: dict[str, str | float] = {} + info["optimizer_type"] = type(optimizer).__name__ + for name, value in optimizer.defaults.items(): + if isinstance(value, (int, float, bool, str)): + info[f"opt_{name}"] = value + return info + + +def _scheduler_summary(scheduler: Any) -> dict[str, str | float]: + """Extract scheduler settings.""" + info: dict[str, str | float] = {} + if scheduler is None: + info["scheduler_type"] = "none" + return info + + info["scheduler_type"] = type(scheduler).__name__ + for attr in ( + "step_size", + "gamma", + "milestones", + "factor", + "patience", + "min_lr", + "T_max", + "eta_min", + ): + val = getattr(scheduler, attr, None) + if val is not None: + info[f"sch_{attr}"] = ( + str(list(val)) if isinstance(val, (list, tuple)) else str(val) + ) + + if hasattr(scheduler, "optimizer"): + for name, value in scheduler.optimizer.defaults.items(): + if isinstance(value, (int, float, bool)): + info[f"sch_opt_{name}"] = value + + return info + + +class ProvenanceCallback(Callback): + """Lightning callback that captures PROV document + run summary. + + Reads environment data from :class:`EnvironmentCallback` and dataset + verification/split data from :class:`DatasetVerificationCallback` when + present as sibling callbacks. When used alone, falls back to doing its + own environment/verification work so the user still gets a complete PROV + document. + + Args: + model_name: Identifier for this model (shown in run name). + experiment_name: MLflow experiment name (default: "Training_Pipeline"). + manifest_path: Path to manifest.csv (auto-detected if None). + data_root: Root directory for dataset files (auto-detected if None). + test_size: Fraction of data for the test split. + random_state: Random seed for train/test split. + fail_fast: Abort training if dataset verification fails. + strict: If True, re-raise errors from optional provenance steps. + register_model: If True, auto-log model summary from pl_module. + register_optimizer: Log optimizer config (or True to auto-detect). + register_scheduler: Log scheduler config (or True to auto-detect). + prov_prefixes: Optional override for PROV namespace prefixes. + """ + + def __init__( + self, + model_name: str | None = None, + experiment_name: str = "Training_Pipeline", + manifest_path: str | None = None, + data_root: str | None = None, + test_size: float = 0.2, + random_state: int = 42, + fail_fast: bool = True, + strict: bool = False, + register_model: bool = True, + register_optimizer: bool = True, + register_scheduler: bool = True, + prov_prefixes: dict[str, str] | None = None, + ) -> None: + """Initialise the provenance callback. + + Args: + model_name: Name of the model (defaults to ``MODEL_NAME`` env var or "model"). + experiment_name: MLflow experiment name for the training run. + manifest_path: Path to manifest.csv (auto-detected if None). + data_root: Root directory of the dataset (auto-detected if None). + test_size: Fraction of data for the test split. Set to 0 to skip splitting. + random_state: Random seed for train/test split. + fail_fast: Abort training if dataset verification fails. + strict: If True, re-raise errors from optional provenance steps. + register_model: If True, auto-log model summary from pl_module. + register_optimizer: Log optimizer config (or True to auto-detect). + register_scheduler: Log scheduler config (or True to auto-detect). + prov_prefixes: Optional override for PROV namespace prefixes. + """ + self.model_name = model_name or os.environ.get("MODEL_NAME", "model") + self.experiment_name = experiment_name + self.manifest_path = manifest_path + self.data_root = data_root + self.test_size = test_size + self.random_state = random_state + self.fail_fast = fail_fast + self.strict = strict + self.register_model = register_model + self.register_optimizer = register_optimizer + self.register_scheduler = register_scheduler + self._prov_prefixes = prov_prefixes + + # Internal state (populated by on_fit_start or sibling callbacks) + self._run_id: str | None = None + self._temp_dirs: list[str] = [] + self._split_data: dict[str, object] | None = None + self._verification: dict[str, object] | None = None + self._frozen_requirements: str | None = None + self._git_commit: str = "unknown" + self._git_url: str = "unknown" + self._git_branch: str = "unknown" + + # ── helpers ────────────────────────────────────────────── + + def _gather_from_siblings(self, trainer: Any) -> None: + """Read data already collected by sibling callbacks.""" + from rationai.mlkit.lightning.callbacks.dataset_verification import ( + DatasetVerificationCallback, + ) + from rationai.mlkit.lightning.callbacks.environment import EnvironmentCallback + + for cb in trainer.callbacks: + if isinstance(cb, EnvironmentCallback): + self._git_commit = getattr(cb, "_git_commit", "unknown") + self._git_url = getattr(cb, "_git_url", "unknown") + self._git_branch = getattr(cb, "_git_branch", "unknown") + self._frozen_requirements = getattr(cb, "_frozen_requirements", None) + elif isinstance(cb, DatasetVerificationCallback): + self._verification = getattr(cb, "_verification", None) + self._split_data = getattr(cb, "_split_data", None) + + def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: + """Do environment + verification work when no sibling callbacks exist.""" + from rationai.mlkit.lightning.callbacks.environment import ( + _detect_docker, + _detect_hardware, + _lookup_user_run, + _snapshot_environment, + ) + from rationai.mlkit.provenance.dataset import ( + _detect_manifest, + _lookup_dataset_run, + _verify_dataset, + ) + + # ── Git info (read from MLflow tags set by MLFlowLogger) ── + try: + run = mlflow.active_run() + run_tags: dict[str, str] = {} + if run and run.info and run.info.run_id: + client = mlflow.tracking.MlflowClient() + run_data = client.get_run(run.info.run_id) + run_tags = dict(run_data.data.tags) if run_data.data.tags else {} + self._git_commit = run_tags.get( + "mlflow.source.git.commit", run_tags.get("git.commit", "unknown") + ) + self._git_url = run_tags.get( + "mlflow.source.git.repoUrl", run_tags.get("git.repo_url", "unknown") + ) + self._git_branch = run_tags.get( + "mlflow.source.git.branch", run_tags.get("git.branch", "unknown") + ) + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] Git info failed: %s", e) + + # ── User lookup ───────────────────────────────────────── + try: + user_run_id, user_tags = _lookup_user_run() + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] User lookup failed: %s", e) + user_run_id, user_tags = None, {} + + # ── Hardware (skip if MLflow system metrics are on) ───── + sys_metrics_on = any( + getattr(logger, "log_system_metrics", False) for logger in trainer.loggers + ) + hardware = {} if sys_metrics_on else _detect_hardware() + docker = _detect_docker() + + # ── Dataset verification & split ──────────────────────── + manifest_path = self.manifest_path + data_root = self.data_root + + if manifest_path is None: + manifest_path, data_root = _detect_manifest() + elif data_root is None: + data_root = os.path.dirname(os.path.abspath(manifest_path)) + + if manifest_path and data_root: + from rationai.mlkit.provenance.dataset import ( + load_manifest, + ) + + # ── Verification (always) ──────────────────────────── + dataset_run_id = _lookup_dataset_run() + verification = _verify_dataset(manifest_path, data_root, dataset_run_id) + self._verification = verification or {} + verification_details: list[str] = verification.get("details", []) + + for detail in verification_details: + log.info(f" [ProvenanceCallback] {detail}") + + # Log verification results + if verification: + mlflow.log_params( + { + "dataset_verified": verification["verified"], + "dataset_file_sizes_match": verification["file_sizes_match"] + is True, + "dataset_files_missing": verification["files_missing"], + "dataset_files_total": verification["files_total"], + } + ) + if verification["verified"]: + mlflow.set_tag("dataset_verification", "VERIFIED") + else: + mlflow.set_tag("dataset_verification", "MISMATCH") + mlflow.set_tag( + "dataset_verification_details", + "; ".join(verification["details"]), + ) + + if self.fail_fast and not verification["verified"]: + raise RuntimeError( + "Dataset verification failed — aborting training.\n" + + "\n".join(f" {d}" for d in verification["details"]), + ) + + # ── Train/test split (only if test_size > 0) ───────── + if self.test_size > 0: + from sklearn.model_selection import train_test_split + + samples = load_manifest(manifest_path, data_root) + train_samples, test_samples = train_test_split( + samples, + test_size=self.test_size, + random_state=self.random_state, + stratify=[s["label"] for s in samples], + ) + + self._split_data = { + "train": train_samples, + "test": test_samples, + "test_size": self.test_size, + "random_state": self.random_state, + } + + # Log split as artifact + split_dir = f"_mlflow_split_{uuid.uuid4().hex[:8]}" + os.makedirs(split_dir, exist_ok=True) + self._temp_dirs.append(split_dir) + for subset_name, subset_samples in [ + ("train", train_samples), + ("test", test_samples), + ]: + split_file = os.path.join(split_dir, f"{subset_name}_split.csv") + pd.DataFrame(subset_samples).to_csv(split_file, index=False) + + mlflow.log_artifacts(split_dir, artifact_path="split") + + # Log split counts + train_labels = [s["label"] for s in train_samples] + test_labels = [s["label"] for s in test_samples] + mlflow.log_params( + { + "train_samples": len(train_samples), + "test_samples": len(test_samples), + "train_positive": sum(train_labels), + "train_negative": len(train_labels) - sum(train_labels), + "test_positive": sum(test_labels), + "test_negative": len(test_labels) - sum(test_labels), + } + ) + else: + log.warning( + "[ProvenanceCallback] No manifest.csv found — " + "train/test split not logged." + ) + + # ── Tags ──────────────────────────────────────────────── + tags: dict[str, str] = {} + if user_run_id: + tags["user_run_id"] = user_run_id + for key in ("username", "real_name", "organization"): + if key in user_tags: + tags[key] = user_tags[key] + + dataset_run_id = _lookup_dataset_run() + if dataset_run_id: + tags["dataset_run_id"] = dataset_run_id + + tags.update( + { + "git_commit": self._git_commit, + "git_url": self._git_url, + "git_branch": self._git_branch, + "prov_start_time": datetime.now(UTC).isoformat(), + } + ) + mlflow.set_tags(tags) + + # ── Params: hardware + docker + split config ──────────── + all_params: dict[str, str | float | int] = { + "model_name": self.model_name, + **hardware, + **docker, + "split_test_size": self.test_size, + "split_random_state": self.random_state, + "split_stratified": True, + } + mlflow.log_params(all_params) + + # ── Environment snapshot ──────────────────────────────── + artifact_dir = f"_mlflow_env_{uuid.uuid4().hex[:8]}" + os.makedirs(artifact_dir, exist_ok=True) + self._temp_dirs.append(artifact_dir) + try: + self._frozen_requirements = _snapshot_environment(artifact_dir) + mlflow.log_artifacts(artifact_dir, artifact_path="environment") + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] Environment snapshot failed: %s", e) + + # ── lightning hooks ─────────────────────────────────────── + + def _ensure_active_run(self, trainer: Any) -> str | None: + """Ensure MLflow has an active run by triggering the logger's experiment. + + Returns the run_id if successful, or None if no MLFlowLogger is present. + """ + try: + from rationai.mlkit.lightning.loggers.mlflow import MLFlowLogger + except ImportError: + # Standalone mlflow — rely on whatever active_run exists + run = mlflow.active_run() + return run.info.run_id if run else None + + for logger in trainer.loggers: + if isinstance(logger, MLFlowLogger): + # Access .experiment to trigger lazy init + active-run setup + _ = logger.experiment + self._run_id = logger.run_id + return logger.run_id + + run = mlflow.active_run() + return run.info.run_id if run else None + + def on_fit_start(self, trainer: Any, pl_module: Any) -> None: + """Gather environment/verification data from siblings or fall back.""" + from rationai.mlkit.lightning.callbacks.dataset_verification import ( + DatasetVerificationCallback, + ) + from rationai.mlkit.lightning.callbacks.environment import EnvironmentCallback + + # Ensure the MLFlowLogger has an active run before any fluent API calls + self._ensure_active_run(trainer) + + # Check if sibling callbacks are present + has_env = any(isinstance(cb, EnvironmentCallback) for cb in trainer.callbacks) + has_verify = any( + isinstance(cb, DatasetVerificationCallback) for cb in trainer.callbacks + ) + + if has_env or has_verify: + self._gather_from_siblings(trainer) + else: + # No siblings — do everything ourselves + self._fallback_on_fit_start(trainer, pl_module) + + def on_fit_end(self, trainer: Any, pl_module: Any) -> None: + """Log model/optimizer/scheduler summaries and PROV document.""" + _active_run = mlflow.active_run() + if _active_run: + run_id: str = _active_run.info.run_id + elif self._run_id: + run_id = self._run_id + else: + self._ensure_active_run(trainer) + if self._run_id: + run_id = self._run_id + else: + return + + # Get a Run object for metadata access + active_run = mlflow.get_run(run_id) + + # ── Model summary ─────────────────────────────────────── + if self.register_model and pl_module is not None: + try: + model_summary = _model_summary(pl_module) + mlflow.log_params(model_summary) + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] Model summary failed: %s", e) + + # ── Optimizer summary ─────────────────────────────────── + if self.register_optimizer and pl_module is not None: + try: + for opt in trainer.optimizers: + optimizer_info = _optimizer_summary(opt) + mlflow.log_params(optimizer_info) + break + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] Optimizer summary failed: %s", e) + + # ── Scheduler summary ─────────────────────────────────── + if self.register_scheduler and pl_module is not None: + try: + for sched in getattr(trainer, "lr_schedulers", []): + scheduler_info = _scheduler_summary(sched.get("scheduler")) + mlflow.log_params(scheduler_info) + break + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] Scheduler summary failed: %s", e) + + # ── PROV document + run summary ───────────────────────── + try: + run_data = mlflow.get_run(run_id) + params = {k: str(v) for k, v in run_data.data.params.items()} + metrics = {k: float(v) for k, v in run_data.data.metrics.items()} + tags = { + k: v + for k, v in run_data.data.tags.items() + if not k.startswith("mlflow.") + } + + # ── Run summary JSON ──────────────────────────────── + summary_dir = f"_mlflow_summary_{uuid.uuid4().hex[:8]}" + os.makedirs(summary_dir, exist_ok=True) + self._temp_dirs.append(summary_dir) + summary_path = os.path.join(summary_dir, "run_summary.json") + + summary = { + "model_name": self.model_name, + "params": dict(run_data.data.params), + "metrics": {k: float(v) for k, v in run_data.data.metrics.items()}, + "tags": tags, + "run_id": run_id, + "experiment_name": self.experiment_name, + "split": { + "test_size": self.test_size, + "random_state": self.random_state, + "stratified": True, + "train_count": len(self._split_data["train"]) + if isinstance(self._split_data, dict) + else 0, # type: ignore[arg-type] + "test_count": len(self._split_data["test"]) + if isinstance(self._split_data, dict) + else 0, # type: ignore[arg-type] + "train": self._split_data["train"] if self._split_data else None, + "test": self._split_data["test"] if self._split_data else None, + } + if self._split_data + else None, + "dataset_verification": self._verification, + "requirements": self._frozen_requirements, + "source": { + "git_commit": self._git_commit, + "git_branch": self._git_branch, + "git_remote": self._git_url, + }, + } + + with open(summary_path, "w") as f: + json.dump(summary, f, indent=2) + + mlflow.log_artifact(summary_path, artifact_path="provenance") + shutil.rmtree(summary_dir, ignore_errors=True) + + # ── PROV document (§9 — configurable prefixes) ────── + prov_doc = _build_prov_document( + run_id=run_id, + run_name=active_run.info.run_name or f"Training_{self.model_name}", + params=params, + metrics=metrics, + tags=tags, + start_time_ms=active_run.info.start_time, + end_time_ms=active_run.info.end_time, + split_data={ + "test_size": self.test_size, + "random_state": self.random_state, + "train": self._split_data["train"] if self._split_data else None, + "test": self._split_data["test"] if self._split_data else None, + } + if self._split_data + else None, + requirements=self._frozen_requirements, + verification=self._verification, + prov_prefixes=_get_prov_prefixes(self._prov_prefixes), + ) + + prov_dir = f"_mlflow_prov_{uuid.uuid4().hex[:8]}" + os.makedirs(prov_dir, exist_ok=True) + self._temp_dirs.append(prov_dir) + prov_path = os.path.join(prov_dir, "prov.json") + with open(prov_path, "w") as f: + json.dump(prov_doc, f, indent=2) + + mlflow.log_artifact(prov_path, artifact_path="provenance") + shutil.rmtree(prov_dir, ignore_errors=True) + + log.info("[ProvenanceCallback] Complete → %s", run_id) + except Exception as e: + if self.strict: + raise + log.warning( + "[ProvenanceCallback] Could not write provenance artifacts: %s", e + ) + + # ── Clean up temp dirs ────────────────────────────────── + for d in self._temp_dirs: + shutil.rmtree(d, ignore_errors=True) diff --git a/rationai/mlkit/lightning/loggers/mlflow.py b/rationai/mlkit/lightning/loggers/mlflow.py index 0f56026..cf4b6ac 100644 --- a/rationai/mlkit/lightning/loggers/mlflow.py +++ b/rationai/mlkit/lightning/loggers/mlflow.py @@ -59,7 +59,10 @@ def __init__( def experiment(self) -> MlflowClient: if not self._initialized: exp = super().experiment + # Resume the run in the fluent API context so callbacks can use + # mlflow.log_params, mlflow.log_artifacts, etc. mlflow.start_run(self.run_id, log_system_metrics=self.log_system_metrics) + self._initialized = True return exp return super().experiment diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py new file mode 100644 index 0000000..f2c2b6b --- /dev/null +++ b/rationai/mlkit/provenance/__init__.py @@ -0,0 +1,34 @@ +"""Provenance tracking - PROV-O-aware logging to MLflow. + +Submodules: + common - shared helpers (prefixes, IDs, timestamps) + user - build_user_prov + register_new_user + dataset - build_dataset_prov + register_dataset + verify_dataset + run - build_training_run_prov + +For automatic provenance capture with Lightning, use +:class:`~rationai.mlkit.lightning.callbacks.provenance.ProvenanceCallback`. +""" + +from __future__ import annotations + +from rationai.mlkit.provenance.dataset import ( + build_dataset_prov, + register_dataset, + verify_dataset, +) +from rationai.mlkit.provenance.run import build_training_run_prov +from rationai.mlkit.provenance.user import ( + build_user_prov, + register_new_user, +) + + +__all__ = [ + "build_dataset_prov", + "build_training_run_prov", + "build_user_prov", + "register_dataset", + "register_new_user", + "verify_dataset", +] diff --git a/rationai/mlkit/provenance/common.py b/rationai/mlkit/provenance/common.py new file mode 100644 index 0000000..9314d38 --- /dev/null +++ b/rationai/mlkit/provenance/common.py @@ -0,0 +1,80 @@ +"""Shared PROV-O helpers and namespace configuration. + +Used by all PROV document builders (user, dataset, training run). +""" + +from __future__ import annotations + +import datetime as _dt +import json +import os +import re +from typing import Any + + +# ────────────────────────────────────────────── +# OpenProvenance / CPM namespace URIs +# ────────────────────────────────────────────── + +_DEFAULT_PROV_PREFIXES: dict[str, str] = { + "storage": "http://localhost:8083/api/v1/documents/", + "meta": "http://localhost:8083/api/v1/documents/meta/", + "schema": "https://schema.org/", + "cpm": "https://www.commonprovenancemodel.org/cpm-namespace-v1-0/", + "blank": "https://openprovenance.org/blank/", + "xsd": "http://www.w3.org/2001/XMLSchema#", + "gen": "gen/", + "dct": "http://purl.org/dc/terms/", + "prov": "http://www.w3.org/ns/prov#", + "sosa": "http://www.w3.org/ns/sosa/", +} + + +def get_prov_prefixes(override: dict[str, str] | None = None) -> dict[str, str]: + """Return PROV prefix map. + + Priority: explicit override > ``PROV_BASE_URI`` env var > defaults. + """ + if override: + return {**_DEFAULT_PROV_PREFIXES, **override} + env_json = os.environ.get("PROV_BASE_URI", "") + if env_json: + try: + parsed = json.loads(env_json) + if not isinstance(parsed, dict): + raise TypeError("expected a JSON object") + merged = {**_DEFAULT_PROV_PREFIXES, **parsed} + return merged + except (json.JSONDecodeError, TypeError): + pass + return _DEFAULT_PROV_PREFIXES + + +# ────────────────────────────────────────────── +# Small helpers used inside PROV documents +# ────────────────────────────────────────────── + + +def _safe_id(name: str) -> str: + """Sanitise a name so it can be used as a PROV identifier fragment.""" + return re.sub(r"[^a-zA-Z0-9_]", "_", name) + + +def _qualified(prefix: str, local: str) -> str: + return f"{prefix}:{local}" + + +def _typed_value(value: Any) -> list[str]: + return [str(value)] + + +def _qualified_name(type_prefix: str, type_local: str) -> dict[str, str]: + return {"type": "prov:QUALIFIED_NAME", "$": f"{type_prefix}:{type_local}"} + + +def _iso_timestamp(ts_ms: int | None = None) -> str: + if ts_ms is not None: + dt = _dt.datetime.fromtimestamp(ts_ms / 1000, tz=_dt.UTC) + else: + dt = _dt.datetime.now(_dt.UTC) + return dt.strftime("%Y-%m-%dT%H:%M:%S.000+00:00") diff --git a/rationai/mlkit/provenance/dataset.py b/rationai/mlkit/provenance/dataset.py new file mode 100644 index 0000000..7a900f1 --- /dev/null +++ b/rationai/mlkit/provenance/dataset.py @@ -0,0 +1,485 @@ +"""Dataset registration, verification, and PROV-O document generation. + +Colocates ``build_dataset_prov`` (the PROV document builder) with the +MLflow registration entry points ``register_dataset`` and ``verify_dataset``. +""" + +from __future__ import annotations + +import json +import os +import shutil +import uuid +from typing import Any + +import mlflow +import pandas as pd + +from rationai.mlkit.provenance.common import ( + _iso_timestamp, + _qualified, + _qualified_name, + _safe_id, + _typed_value, + get_prov_prefixes, +) + + +# ────────────────────────────────────────────── +# PROV document builder +# ────────────────────────────────────────────── + + +def build_dataset_prov( + run_id: str, + dataset_name: str, + version: str, + dataset_root: str, + num_samples: int, + num_positive: int, + num_negative: int, + file_sizes: dict[str, int], + manifest_path: str | None = None, + prov_prefixes: dict[str, str] | None = None, +) -> dict[str, Any]: + """Build a PROV document for a dataset registration run. + + Produces a ``dataset`` entity (``sosa:Sample``) linked to the + registration activity and metadata bundle. + """ + prefixes = prov_prefixes or get_prov_prefixes() + + run_act_local = _safe_id(f"run_{run_id}") + run_act_id = _qualified("gen", run_act_local) + + ds_local = _safe_id(f"dataset_{dataset_name}_{version.replace('.', '_')}") + ds_id = _qualified("gen", ds_local) + + meta_local = run_id + meta_id = _qualified("meta", meta_local) + + main_act_local = f"DatasetReg_{run_id[:8]}" + main_act_id = _qualified("blank", main_act_local) + + entities: dict[str, dict[str, Any]] = {} + activities: dict[str, dict[str, Any]] = {} + used: dict[str, dict[str, Any]] = {} + was_generated_by: dict[str, dict[str, Any]] = {} + + rel_counter = [0] + + def _blank_rel_id() -> str: + rid = f"_:n{rel_counter[0]}" + rel_counter[0] += 1 + return rid + + now = _iso_timestamp() + + # ── DATASET ENTITY ─────────────────────────────────── + ds_props: dict[str, list[Any]] = { + "schema:name": _typed_value(dataset_name), + "prov:type": [_qualified_name("sosa", "Sample")], + "dct:description": _typed_value( + f"Dataset {dataset_name} v{version} ({num_samples} samples)", + ), + } + if manifest_path: + ds_props["schema:url"] = _typed_value(manifest_path) + entities[ds_id] = ds_props + + # ── ACTIVITY (the registration action) ──────────────── + run_activity: dict[str, Any] = {} + run_activity["prov:type"] = [_qualified_name("schema", "Action")] + run_activity["prov:startTime"] = [now] + run_activity["prov:endTime"] = [now] + run_activity["schema:name"] = _typed_value(f"Register dataset {dataset_name}") + run_activity["gen:dataset_name"] = _typed_value(dataset_name) + run_activity["gen:dataset_version"] = _typed_value(version) + activities[run_act_id] = run_activity + + # ── USED (activity consumed the dataset entity) ────── + used[_blank_rel_id()] = { + "prov:activity": run_act_id, + "prov:entity": ds_id, + } + + # ── CPM METADATA ENTITY ─────────────────────────────── + meta_entity: dict[str, list[Any]] = {} + meta_entity["prov:type"] = [_qualified_name("cpm", "BundleMetadata")] + meta_entity["gen:dataset_name"] = _typed_value(dataset_name) + meta_entity["gen:dataset_version"] = _typed_value(version) + meta_entity["gen:dataset_root"] = _typed_value(dataset_root) + meta_entity["gen:num_samples"] = _typed_value(str(num_samples)) + meta_entity["gen:num_positive"] = _typed_value(str(num_positive)) + meta_entity["gen:num_negative"] = _typed_value(str(num_negative)) + if manifest_path: + meta_entity["gen:manifest_path"] = _typed_value(manifest_path) + + file_sizes_str = json.dumps(file_sizes) + meta_entity["gen:file_sizes"] = [file_sizes_str] + entities[meta_id] = meta_entity + + # ── CPM MAIN ACTIVITY ──────────────────────────────── + main_activity: dict[str, Any] = {} + main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] + main_activity["cpm:referencedMetaBundleId"] = [ + {"type": "prov:QUALIFIED_NAME", "$": meta_id}, + ] + main_activity["dct:hasPart"] = [ + {"type": "prov:QUALIFIED_NAME", "$": run_act_id}, + ] + activities[main_act_id] = main_activity + + # ── RELATIONSHIPS ───────────────────────────────────── + was_generated_by[_blank_rel_id()] = { + "prov:entity": meta_id, + "prov:activity": run_act_id, + } + + # ── ASSEMBLE BUNDLE ─────────────────────────────────── + inner: dict[str, Any] = {"prefix": prefixes} + if entities: + inner["entity"] = entities + if activities: + inner["activity"] = activities + if used: + inner["used"] = used + if was_generated_by: + inner["wasGeneratedBy"] = was_generated_by + + bundle_key = f"storage:{run_id}" + return {"bundle": {bundle_key: inner}} + + +# ────────────────────────────────────────────── +# MLflow registration & verification helpers +# ────────────────────────────────────────────── + + +def _lookup_experiment(name: str) -> str | None: + """Return the MLflow experiment ID for *name*, or None.""" + exp = mlflow.get_experiment_by_name(name) + return exp.experiment_id if exp else None + + +def _lookup_dataset_run() -> str | None: + """Return the latest Dataset_Registry run ID. + + Falls back to the most recent run so that workflows with only one + registered dataset still work. + """ + exp_id = _lookup_experiment("Dataset_Registry") + if exp_id is None: + return None + + runs_df = mlflow.search_runs( + experiment_ids=[exp_id], + order_by=["start_time DESC"], + ) + if pd.DataFrame(runs_df).empty: + return None + + return pd.DataFrame(runs_df).iloc[0]["run_id"] + + +def _detect_manifest() -> tuple[str | None, str | None]: + """Walk data/ looking for manifest.csv. + + Returns (manifest_path, data_root) or (None, None). + """ + for root_dir in ("data", "test_data", "."): + for dirpath, _, filenames in os.walk(root_dir): + if "manifest.csv" in filenames: + return ( + os.path.join(dirpath, "manifest.csv"), + os.path.dirname( + os.path.abspath( + os.path.join(dirpath, "manifest.csv"), + ) + ), + ) + return None, None + + +# ────────────────────────────────────────────── +# Manifest loading (shared with ProvenanceCallback) +# ────────────────────────────────────────────── + + +def load_manifest(manifest_path: str, data_root: str) -> list[dict[str, Any]]: + """Load a manifest.csv and resolve WSI paths. + + Returns a list of dicts with keys ``path`` (absolute) and ``label``. + + Shared by ``register_dataset``, ``verify_dataset``, and + ``ProvenanceCallback`` to avoid duplicating the CSV iteration pattern. + """ + df = pd.read_csv(manifest_path) + samples: list[dict[str, Any]] = [] + for _, row in df.iterrows(): + rel = row["wsi_path"] + full = os.path.join(data_root, rel) if not os.path.isabs(rel) else rel + samples.append({"path": full, "label": int(row["cancer"])}) + return samples + + +# ────────────────────────────────────────────── +# Dataset verification +# ────────────────────────────────────────────── + + +def verify_dataset( + manifest_path: str | None = None, + data_root: str | None = None, +) -> dict[str, Any]: + """Public entry point — verify the current dataset against MLflow. + + Auto-detects the manifest if *manifest_path* is not given. + + Returns a dict with keys:: + + { + "verified": bool, + "file_sizes_match": bool | None, + "files_missing": int, + "files_total": int, + "details": list[str], + } + """ + if manifest_path is None: + manifest_path, data_root = _detect_manifest() + + if manifest_path is None: + return { + "verified": False, + "dataset_run_id": None, + "file_sizes_match": None, + "files_missing": 0, + "files_total": 0, + "details": ["No manifest.csv found — skipping verification"], + } + + if data_root is None: + data_root = os.path.dirname(os.path.abspath(manifest_path)) + + dataset_run_id = _lookup_dataset_run() + return _verify_dataset(manifest_path, data_root, dataset_run_id) + + +def _verify_dataset( + manifest_path: str, + data_root: str, + dataset_run_id: str | None, +) -> dict[str, Any]: + """Verify the current dataset against the registered version in MLflow. + + Checks: + 1. Per-file sizes match (file-level integrity) + 2. All WSI files exist on disk + + Returns a dict with verification results. + """ + result: dict[str, Any] = { + "verified": False, + "dataset_run_id": dataset_run_id, + "file_sizes_match": None, + "files_missing": 0, + "files_total": 0, + "details": [], + } + + if not dataset_run_id: + result["details"].append( + "No Dataset_Registry run found — skipping verification" + ) + return result + + # Fetch registered metadata + try: + reg_run = mlflow.get_run(dataset_run_id) + reg_tags = reg_run.data.tags + reg_file_sizes_str = reg_tags.get("file_sizes", "") + reg_file_sizes = json.loads(reg_file_sizes_str) if reg_file_sizes_str else {} + except Exception as e: + result["details"].append(f"Failed to fetch Dataset_Registry run: {e}") + return result + + samples = load_manifest(manifest_path, data_root) + curr_file_sizes: dict[str, int] = {} + for s in samples: + basename = os.path.basename(s["path"]) + if os.path.isfile(s["path"]): + curr_file_sizes[basename] = os.stat(s["path"]).st_size + else: + curr_file_sizes[basename] = -1 # missing + + manifest_match = set(curr_file_sizes) == set(reg_file_sizes) + if not manifest_match: + result["details"].append( + f"File manifest mismatch: expected {len(reg_file_sizes)} file(s), found {len(curr_file_sizes)} file(s)" + ) + sizes_match = all( + curr_file_sizes.get(k) == reg_file_sizes[k] for k in reg_file_sizes + ) + result["file_sizes_match"] = manifest_match and sizes_match + + if not result["file_sizes_match"]: + mismatched = [ + name + for name in reg_file_sizes + if curr_file_sizes.get(name) != reg_file_sizes[name] + ] + if mismatched: + result["details"].append( + f"File size mismatch on {len(mismatched)} file(s): " + + ", ".join(sorted(mismatched)[:5]) + + ("…" if len(mismatched) > 5 else ""), + ) + + missing = sum(1 for s in samples if not os.path.isfile(s["path"])) + result["files_total"] = len(samples) + result["files_missing"] = missing + + if missing > 0: + result["details"].append(f"{missing}/{len(samples)} WSI files missing on disk") + + result["verified"] = result["file_sizes_match"] and missing == 0 + + if result["verified"]: + result["details"].append("✅ Dataset verified — matches registered version") + else: + result["details"].append("❌ Dataset verification FAILED") + + return result + + +# ────────────────────────────────────────────── +# Hash-based registration (preferred) +# ────────────────────────────────────────────── + + +def register_dataset( + dataset_dir: str, + dataset_name: str | None = None, + version: str = "1.0.0", + experiment_name: str = "Dataset_Registry", +) -> str: + """Register a dataset in MLflow's Dataset_Registry experiment. + + Captures per-file metadata (size, last modified) and stores it as tags + so future training runs can inspect the exact state of the data. + + Args: + dataset_dir: Path to the dataset root (must contain manifest.csv). + dataset_name: Human-readable name (defaults to directory basename). + version: Dataset version string. + experiment_name: MLflow experiment for registration. + + Returns: + The run_id of the registration run. + + Example: + from rationai.mlkit.provenance import register_dataset + + run_id = register_dataset("data/pato_cohort_01", version="2.0") + print(f"Registered as {run_id}") + """ + dataset_dir = os.path.abspath(dataset_dir) + manifest_path = os.path.join(dataset_dir, "manifest.csv") + if not os.path.isfile(manifest_path): + raise FileNotFoundError( + f"No manifest.csv found in {dataset_dir}. " + "Dataset registration requires a manifest.csv file.", + ) + + if dataset_name is None: + dataset_name = os.path.basename(dataset_dir) + + samples = load_manifest(manifest_path, dataset_dir) + file_sizes = {} + file_mtimes = {} + + for s in samples: + basename = os.path.basename(s["path"]) + if os.path.isfile(s["path"]): + st = os.stat(s["path"]) + file_sizes[basename] = st.st_size + file_mtimes[basename] = st.st_mtime + else: + file_sizes[basename] = -1 + file_mtimes[basename] = -1 + + mlflow.set_experiment(experiment_name) + run = mlflow.start_run(run_name=f"Dataset_{dataset_name}_{version}") + run_id = run.info.run_id + run_active = True + + try: + mlflow.log_params( + { + "dataset_root": dataset_dir, + "num_samples": len(samples), + "num_positive": sum(1 for s in samples if s["label"] == 1), + "num_negative": sum(1 for s in samples if s["label"] == 0), + } + ) + + mlflow.set_tags( + { + "dataset_name": dataset_name, + "version": version, + "file_sizes": json.dumps(file_sizes), + "file_mtimes": json.dumps(file_mtimes), + } + ) + + # ── PROV-O document ──────────────────────────────── + prov_doc = build_dataset_prov( + run_id=run_id, + dataset_name=dataset_name, + version=version, + dataset_root=dataset_dir, + num_samples=len(samples), + num_positive=sum(1 for s in samples if s["label"] == 1), + num_negative=sum(1 for s in samples if s["label"] == 0), + file_sizes=file_sizes, + manifest_path=manifest_path, + ) + + prov_dir = f"_dataset_prov_{uuid.uuid4().hex[:8]}" + os.makedirs(prov_dir, exist_ok=True) + prov_path = os.path.join(prov_dir, "prov.json") + try: + with open(prov_path, "w") as f: + json.dump(prov_doc, f, indent=2) + mlflow.log_artifact(prov_path, artifact_path="provenance") + finally: + shutil.rmtree(prov_dir, ignore_errors=True) + + # ── Legacy dataset provenance JSON (backward compat) ── + legacy_prov_dir = f"_dataset_legacy_{uuid.uuid4().hex[:8]}" + os.makedirs(legacy_prov_dir, exist_ok=True) + legacy_path = os.path.join(legacy_prov_dir, "dataset_provenance.json") + try: + with open(legacy_path, "w") as f: + json.dump( + { + "dataset_name": dataset_name, + "version": version, + "dataset_root": dataset_dir, + "file_sizes": file_sizes, + "file_mtimes": file_mtimes, + "num_samples": len(samples), + }, + f, + indent=2, + ) + mlflow.log_artifact(legacy_path, artifact_path="provenance") + finally: + shutil.rmtree(legacy_prov_dir, ignore_errors=True) + finally: + if run_active: + mlflow.end_run() + + print(f" [register_dataset] {dataset_name} v{version} → run_id={run_id}") + return run_id diff --git a/rationai/mlkit/provenance/run.py b/rationai/mlkit/provenance/run.py new file mode 100644 index 0000000..a9797e4 --- /dev/null +++ b/rationai/mlkit/provenance/run.py @@ -0,0 +1,374 @@ +"""Training-run PROV-O document generation. + +Holds ``build_training_run_prov`` and the constants it needs for mapping +MLflow params/tags to PROV properties. +""" + +from __future__ import annotations + +import json +from typing import Any + +from rationai.mlkit.provenance.common import ( + _iso_timestamp, + _qualified, + _qualified_name, + _safe_id, + _typed_value, + get_prov_prefixes, +) + + +# ────────────────────────────────────────────── +# Hyperparameter keys that surface on the activity +# ────────────────────────────────────────────── + +_ACTIVITY_HP_KEYS: set[str] = { + "learning_rate", + "lr", + "batch_size", + "epochs", + "optimizer", + "loss_function", + "dropout", + "weight_decay", + "momentum", + "num_layers", + "hidden_size", + "embedding_dim", + "num_classes", + "patch_size", + "input_size", + "augmentations", +} + +_WSI_PARAM_KEYS: set[str] = { + "scanner", + "slide_id", + "wsi_id", + "patient_id", + "subject_id", + "institution", + "site", + "staining", + "slicing_method", +} + + +# ────────────────────────────────────────────── +# PROV document builder +# ────────────────────────────────────────────── + + +def build_training_run_prov( + run_id: str, + run_name: str, + params: dict[str, str], + metrics: dict[str, float], + tags: dict[str, str], + start_time_ms: int | None = None, + end_time_ms: int | None = None, + split_data: dict[str, object] | None = None, + requirements: str | None = None, + verification: dict[str, object] | None = None, + prov_prefixes: dict[str, str] | None = None, +) -> dict[str, object]: + """Build an OpenProvenance-compatible PROV document for a training run.""" + username = tags.get("username", tags.get("mlflow.user", "unknown")) + agent_local = _safe_id(f"user_{username}") + agent_id = _qualified("gen", agent_local) + + run_act_local = _safe_id(f"run_{run_id}") + run_act_id = _qualified("gen", run_act_local) + + meta_local = run_id + meta_id = _qualified("meta", meta_local) + + main_act_local = f"TrainingRun_{run_id[:8]}" + main_act_id = _qualified("blank", main_act_local) + + entities: dict[str, Any] = {} + activities: dict[str, Any] = {} + agents: dict[str, Any] = {} + used: dict[str, Any] = {} + was_associated_with: dict[str, Any] = {} + + rel_counter = [0] + + def _blank_rel_id() -> str: + rid = f"_:n{rel_counter[0]}" + rel_counter[0] += 1 + return rid + + # ── 1. AGENT ─────────────────────────────────────────── + agent_props: dict[str, Any] = {} + real_name = tags.get("real_name", username) + agent_props["schema:name"] = _typed_value(real_name) + email = tags.get("mlflow.source.git.user.email", f"{username}@unknown") + agent_props["schema:email"] = _typed_value(email) + org = tags.get("organization", "") + if org: + agent_props["schema:affiliation"] = _typed_value(org) + agent_props["prov:type"] = [_qualified_name("schema", "Person")] + agents[agent_id] = agent_props + + # ── 2. INPUT ENTITIES ────────────────────────────────── + image_path_candidates = ( + params.get("image_path") + or params.get("wsi_path") + or params.get("dataset_path") + or params.get("data_path") + or params.get("input_path") + ) + + if image_path_candidates: + wsi_local = _safe_id(f"wsi_{image_path_candidates}") + wsi_id = _qualified("gen", wsi_local) + wsi_props: dict[str, Any] = { + "schema:name": _typed_value(f"Input: {image_path_candidates}"), + "prov:type": [_qualified_name("sosa", "Sample")], + } + if "scanner" in params: + wsi_props["gen:scanner"] = _typed_value(params["scanner"]) + for pk, prov_key in [ + ("slide_id", "schema:identifier"), + ("wsi_id", "schema:identifier"), + ("patient_id", "gen:patient_pseudonym"), + ("subject_id", "gen:patient_pseudonym"), + ("institution", "gen:origin_institution"), + ("site", "gen:origin_institution"), + ("staining", "gen:staining_method"), + ("slicing_method", "gen:slicing_method"), + ]: + if pk in params: + wsi_props[prov_key] = _typed_value(params[pk]) + + entities[wsi_id] = wsi_props + used[_blank_rel_id()] = { + "prov:activity": run_act_id, + "prov:entity": wsi_id, + } + else: + train_count = params.get("train_samples", "0") + test_count = params.get("test_samples", "0") + ds_local = _safe_id(f"dataset_{run_id[:8]}") + ds_id = _qualified("gen", ds_local) + entities[ds_id] = { + "schema:name": _typed_value( + f"Training dataset ({train_count} train, {test_count} test)" + ), + "prov:type": [_qualified_name("sosa", "Sample")], + } + used[_blank_rel_id()] = { + "prov:activity": run_act_id, + "prov:entity": ds_id, + } + + # ── 3. RUN ACTIVITY ──────────────────────────────────── + run_activity: dict[str, Any] = {} + run_activity["prov:type"] = [_qualified_name("schema", "Action")] + run_activity["prov:startTime"] = [_iso_timestamp(start_time_ms)] + run_activity["prov:endTime"] = [_iso_timestamp(end_time_ms)] + run_activity["schema:name"] = _typed_value(run_name) + + exp_name = params.get("model_name", "") + if exp_name: + run_activity["gen:experiment_name"] = _typed_value(exp_name) + + if "model_class" in params: + run_activity["gen:model_config"] = _typed_value(params["model_class"]) + + git_commit = tags.get("git_commit", tags.get("mlflow.source.git.commit", "")) + if git_commit: + run_activity["schema:identifier"] = _typed_value(git_commit) + + for key in ("pretrained_model", "backbone", "feature_extractor"): + if key in params: + run_activity["gen:pretrained_model"] = _typed_value(params[key]) + + for key, prov_key in [ + ("dataset_name", "gen:dataset_name"), + ("dataset_version", "gen:dataset_version"), + ("data_split", "gen:data_split"), + ("split", "gen:data_split"), + ]: + if key in params: + run_activity[prov_key] = _typed_value(params[key]) + + for key in _ACTIVITY_HP_KEYS: + if key in params: + run_activity[f"gen:{key}"] = _typed_value(params[key]) + + for key, val in params.items(): + if key.startswith(("opt_", "sch_")): + clean = key.removeprefix("opt_").removeprefix("sch_") + if f"gen:{clean}" not in run_activity: + run_activity[f"gen:{clean}"] = _typed_value(val) + + for tag_key, prov_key in [ + ("mlflow.gpu.count", "gen:gpu_count"), + ("mlflow.gpu.names", "gen:gpu_names"), + ("mlflow.cpu.count", "gen:cpu_count"), + ("mlflow.memory_gb", "gen:memory_gb"), + ]: + if tag_key in tags: + run_activity[prov_key] = _typed_value(tags[tag_key]) + + for param_key, prov_key in [ + ("gpu_count", "gen:gpu_count"), + ("gpu_name", "gen:gpu_names"), + ("cpu_count_logical", "gen:cpu_count"), + ("ram_total_gb", "gen:memory_gb"), + ]: + if param_key in params and prov_key not in run_activity: + run_activity[prov_key] = _typed_value(params[param_key]) + + git_url = tags.get("git_url", tags.get("mlflow.source.git.remote", "")) + if git_url: + run_activity["gen:git_remote"] = _typed_value(git_url) + + source_name = tags.get("mlflow.source.name", "") + if source_name: + run_activity["gen:source_name"] = _typed_value(source_name) + + for key, prov_key in [ + ("segmentation", "gen:segmentation_config"), + ("model", "gen:model_config"), + ]: + if key in params: + run_activity[prov_key] = _typed_value(params[key]) + + activities[run_act_id] = run_activity + + # ── 4. CPM METADATA ENTITY ───────────────────────────── + meta_entity: dict[str, Any] = {} + meta_entity["prov:type"] = [_qualified_name("cpm", "BundleMetadata")] + org_val = tags.get("organization", "") + if org_val: + meta_entity["cpm:organization"] = _typed_value(org_val) + + skip_keys = ( + set(_ACTIVITY_HP_KEYS) + | _WSI_PARAM_KEYS + | { + "image_path", + "wsi_path", + "dataset_path", + "data_path", + "input_path", + "segmentation", + "model", + "pretrained_model", + "backbone", + "feature_extractor", + "dataset_name", + "dataset_version", + "data_split", + "split", + "scanner", + "slide_id", + "wsi_id", + "patient_id", + "subject_id", + "institution", + "site", + "staining", + "slicing_method", + } + ) + + for key, val in params.items(): + if key not in skip_keys: + safe_key = _safe_id(key) + meta_entity[f"gen:{safe_key}"] = _typed_value(val) + + for key, mval in metrics.items(): + safe_key = _safe_id(key) + meta_entity[f"gen:{safe_key}"] = _typed_value(mval) + + if split_data: + meta_entity["gen:split_test_size"] = _typed_value( + split_data.get("test_size", "0.2") + ) + meta_entity["gen:split_random_state"] = _typed_value( + str(split_data.get("random_state", "42")) + ) + meta_entity["gen:split_stratified"] = ["true"] + + if split_data.get("train"): + meta_entity["gen:split_train"] = _typed_value( + json.dumps(split_data["train"]) + ) + if split_data.get("test"): + meta_entity["gen:split_test"] = _typed_value(json.dumps(split_data["test"])) + + if requirements: + meta_entity["gen:requirements"] = [requirements] + + if verification: + meta_entity["gen:dataset_verified"] = [str(verification.get("verified", False))] + meta_entity["gen:dataset_run_id"] = [ + str(verification["dataset_run_id"]) + if verification.get("dataset_run_id") is not None + else "" + ] + fsm = verification.get("file_sizes_match") + if fsm is not None: + meta_entity["gen:file_sizes_match"] = [str(fsm)] + fm = verification.get("files_missing", 0) + ft = verification.get("files_total", 0) + meta_entity["gen:files_missing"] = [str(fm)] + meta_entity["gen:files_total"] = [str(ft)] + + for tag_key in ( + "mlflow.source.git.branch", + "mlflow.source.git.repo_url", + "mlflow.parentRunId", + "mlflow.note.content", + ): + if tag_key in tags: + safe_key = _safe_id(tag_key) + meta_entity[f"gen:{safe_key}"] = _typed_value(tags[tag_key]) + + entities[meta_id] = meta_entity + + was_generated_by: dict[str, Any] = {} + was_generated_by[_blank_rel_id()] = { + "prov:entity": meta_id, + "prov:activity": run_act_id, + } + + # ── 5. CPM MAIN ACTIVITY ─────────────────────────────── + main_activity: dict[str, Any] = {} + main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] + main_activity["cpm:referencedMetaBundleId"] = [ + {"type": "prov:QUALIFIED_NAME", "$": meta_id}, + ] + main_activity["dct:hasPart"] = [ + {"type": "prov:QUALIFIED_NAME", "$": run_act_id}, + ] + activities[main_act_id] = main_activity + + # ── 6. RELATIONSHIPS ─────────────────────────────────── + was_associated_with[_blank_rel_id()] = { + "prov:activity": run_act_id, + "prov:agent": agent_id, + } + + # ── 7. ASSEMBLE BUNDLE ───────────────────────────────── + inner: dict[str, object] = {"prefix": prov_prefixes or get_prov_prefixes()} + if entities: + inner["entity"] = entities + if activities: + inner["activity"] = activities + if agents: + inner["agent"] = agents + if was_associated_with: + inner["wasAssociatedWith"] = was_associated_with + if was_generated_by: + inner["wasGeneratedBy"] = was_generated_by + if used: + inner["used"] = used + + bundle_key = f"storage:{run_id}" + return {"bundle": {bundle_key: inner}} diff --git a/rationai/mlkit/provenance/user.py b/rationai/mlkit/provenance/user.py new file mode 100644 index 0000000..cfb81ca --- /dev/null +++ b/rationai/mlkit/provenance/user.py @@ -0,0 +1,232 @@ +"""User registration with PROV-O document generation. + +Colocates ``build_user_prov`` (the PROV document builder) and +``register_new_user`` (the MLflow registration entry point). +""" + +from __future__ import annotations + +import json +import os +import shutil +import uuid +from typing import Any + +import mlflow + +from rationai.mlkit.provenance.common import ( + _iso_timestamp, + _qualified, + _qualified_name, + _safe_id, + _typed_value, + get_prov_prefixes, +) + + +# ────────────────────────────────────────────── +# PROV document builder +# ────────────────────────────────────────────── + + +def build_user_prov( + run_id: str, + username: str, + real_name: str, + email: str, + organization: str, + lead_name: str | None = None, + lead_email: str | None = None, + prov_prefixes: dict[str, str] | None = None, +) -> dict[str, Any]: + """Build a PROV document for a user registration run. + + Produces an ``agent`` entity representing the researcher and links it + to a *registration activity* that generated the run's metadata bundle. + """ + prefixes = prov_prefixes or get_prov_prefixes() + + agent_local = _safe_id(f"user_{username}") + agent_id = _qualified("gen", agent_local) + + run_act_local = _safe_id(f"run_{run_id}") + run_act_id = _qualified("gen", run_act_local) + + meta_local = run_id + meta_id = _qualified("meta", meta_local) + + main_act_local = f"UserReg_{run_id[:8]}" + main_act_id = _qualified("blank", main_act_local) + + entities: dict[str, dict[str, Any]] = {} + activities: dict[str, dict[str, Any]] = {} + agents: dict[str, dict[str, list[Any]]] = {} + was_associated_with: dict[str, dict[str, str]] = {} + was_generated_by: dict[str, dict[str, str]] = {} + + rel_counter = [0] + + def _blank_rel_id() -> str: + rid = f"_:n{rel_counter[0]}" + rel_counter[0] += 1 + return rid + + now = _iso_timestamp() + + # ── AGENT ────────────────────────────────────────────── + agent_props: dict[str, list[Any]] = {} + agent_props["schema:name"] = _typed_value(real_name) + agent_props["schema:email"] = _typed_value(email) + if organization: + agent_props["schema:affiliation"] = _typed_value(organization) + agent_props["prov:type"] = [_qualified_name("schema", "Person")] + agents[agent_id] = agent_props + + # ── ACTIVITY (the registration action) ──────────────── + run_activity: dict[str, Any] = {} + run_activity["prov:type"] = [_qualified_name("schema", "Action")] + run_activity["prov:startTime"] = [now] + run_activity["prov:endTime"] = [now] + run_activity["schema:name"] = _typed_value(f"Register user {real_name}") + activities[run_act_id] = run_activity + + # ── CPM METADATA ENTITY ─────────────────────────────── + meta_entity: dict[str, list[Any]] = {} + meta_entity["prov:type"] = [_qualified_name("cpm", "BundleMetadata")] + meta_entity["gen:username"] = _typed_value(username) + meta_entity["gen:real_name"] = _typed_value(real_name) + meta_entity["gen:email"] = _typed_value(email) + if organization: + meta_entity["gen:organization"] = _typed_value(organization) + if lead_name: + meta_entity["gen:lead_name"] = _typed_value(lead_name) + if lead_email: + meta_entity["gen:lead_email"] = _typed_value(lead_email) + entities[meta_id] = meta_entity + + # ── CPM MAIN ACTIVITY ──────────────────────────────── + main_activity: dict[str, Any] = {} + main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] + main_activity["cpm:referencedMetaBundleId"] = [ + {"type": "prov:QUALIFIED_NAME", "$": meta_id}, + ] + main_activity["dct:hasPart"] = [ + {"type": "prov:QUALIFIED_NAME", "$": run_act_id}, + ] + activities[main_act_id] = main_activity + + # ── RELATIONSHIPS ───────────────────────────────────── + was_associated_with[_blank_rel_id()] = { + "prov:activity": run_act_id, + "prov:agent": agent_id, + } + was_generated_by[_blank_rel_id()] = { + "prov:entity": meta_id, + "prov:activity": run_act_id, + } + + # ── ASSEMBLE BUNDLE ─────────────────────────────────── + inner: dict[str, Any] = {"prefix": prefixes} + if entities: + inner["entity"] = entities + if activities: + inner["activity"] = activities + if agents: + inner["agent"] = agents + if was_associated_with: + inner["wasAssociatedWith"] = was_associated_with + if was_generated_by: + inner["wasGeneratedBy"] = was_generated_by + + bundle_key = f"storage:{run_id}" + return {"bundle": {bundle_key: inner}} + + +# ────────────────────────────────────────────── +# MLflow registration entry point +# ────────────────────────────────────────────── + + +def register_new_user( + username: str, + real_name: str, + email: str, + organization: str, + lead_name: str, + lead_email: str, +) -> str: + """Register a user and emit a PROV-O document as an artifact. + + Creates a run in the ``User_Registry`` experiment with identity tags + and a W3C PROV-O ``prov.json`` document compatible with the Java + ``prov_mlflow`` tool. + + Returns: + The MLflow run_id of the registration run. + """ + experiment_name = "User_Registry" + experiment_id: str | None + try: + experiment_id = mlflow.create_experiment(experiment_name) + except Exception: + exp = mlflow.get_experiment_by_name(experiment_name) + experiment_id = exp.experiment_id if exp else None + + with mlflow.start_run( + experiment_id=experiment_id, + run_name=f"User_{username}", + ) as run: + run_id = run.info.run_id + + mlflow.log_params( + { + "username": username, + "real_name": real_name, + "email": email, + "organization": organization, + "lead_name": lead_name, + "lead_email": lead_email, + } + ) + mlflow.set_tags( + { + "username": username, + "organization": organization, + } + ) + + # ── PROV-O document ──────────────────────────────── + prov_doc = build_user_prov( + run_id=run_id, + username=username, + real_name=real_name, + email=email, + organization=organization, + lead_name=lead_name, + lead_email=lead_email, + ) + + prov_dir = f"_user_prov_{uuid.uuid4().hex[:8]}" + os.makedirs(prov_dir, exist_ok=True) + try: + prov_path = os.path.join(prov_dir, "prov.json") + with open(prov_path, "w") as f: + json.dump(prov_doc, f, indent=2) + + mlflow.log_artifact(prov_path, artifact_path="provenance") + finally: + shutil.rmtree(prov_dir, ignore_errors=True) + + print(f" [register_new_user] {username} → run_id={run_id}") + return run_id + + +if __name__ == "__main__": + register_new_user( + username="researcher_01", + real_name="Jane Doe", + email="jane.doe@example.com", + organization="Example Org", + lead_name="John Smith", + lead_email="john.smith@example.com", + )