From a324607f38673660371c696bb62d6f5bb79546b1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Mon, 20 Jul 2026 19:26:20 +0200 Subject: [PATCH 01/34] feat: prov implementaiton, WIP --- .gitignore | 4 + README.md | 435 ++++++++++++++++++- demo.py | 461 +++++++++++++++++++++ dummy_dataset_create.py | 146 +++++++ rationai/__init__.py | 1 + rationai/mlkit/__init__.py | 60 ++- rationai/mlkit/autolog.py | 101 +---- rationai/mlkit/data/__init__.py | 1 + rationai/mlkit/lightning/__init__.py | 13 +- rationai/mlkit/lightning/autolog.py | 106 +++++ rationai/mlkit/lightning/loggers/mlflow.py | 4 +- rationai/mlkit/lightning/with_cli_args.py | 57 +++ rationai/mlkit/stream/__init__.py | 4 +- rationai/mlkit/with_cli_args.py | 59 +-- tests/test_all.py | 378 +++++++++++++++++ user_to_mlflow.py | 37 ++ 16 files changed, 1701 insertions(+), 166 deletions(-) create mode 100644 demo.py create mode 100644 dummy_dataset_create.py create mode 100644 rationai/__init__.py create mode 100644 rationai/mlkit/data/__init__.py create mode 100644 rationai/mlkit/lightning/autolog.py create mode 100644 rationai/mlkit/lightning/with_cli_args.py create mode 100644 tests/test_all.py create mode 100644 user_to_mlflow.py diff --git a/.gitignore b/.gitignore index 16143c0..b16e130 100644 --- a/.gitignore +++ b/.gitignore @@ -169,3 +169,7 @@ 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 \ No newline at end of file diff --git a/README.md b/README.md index 43f0287..9a4586e 100644 --- a/README.md +++ b/README.md @@ -1 +1,434 @@ -# ML Kit \ No newline at end of file +# rationai.mlflow — 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 — the user only writes their training loop. + +--- + +## Setup + +```bash +uv venv +source .venv/bin/activate +uv pip install -r requirements.txt +``` + +Start the MLflow server: + +```bash +mlflow ui --host 0.0.0.0 --port 5000 # → http://localhost:5000 +``` + +--- + +## Quick start + +### 1. Run the demo + +The easiest way to see everything in action: + +```bash +# Full pipeline — uploads to localhost:5000 +python demo.py + +# With dummy datasets +python demo.py --datasets 3 + +# Push to a different MLflow server +python demo.py --uri http://your-server:5000 + +# Run unit tests instead +python demo.py --test +``` + +The demo exercises all components: + +| Step | Feature | What it shows | +|---|---|---| +| 1 | Dummy data creation | `data/dummy_dataset_*` with manifests | +| 2 | **StreamCapture** | ANSI-aware stdout/stderr capture | +| 3 | **AggregatedMetricCollection** | Tile → slide metric aggregation | +| 4 | **NestedMetricCollection** | Per-slide multiclass metrics | +| 5 | **StratifiedBatchSampler** | Balanced class batches | +| 6 | **Provenance** (`@autolog`) | Full training run with auto-captured provenance | +| 7 | **Lightning** (Trainer + MLFlowLogger) | Lightning training with full provenance tracking | + +Both steps 6 and 7 upload to MLflow with identical provenance depth: +model params, GPU/CPU info, optimizer config, train/test split stats, +environment freeze, console logs, and the PROV-O document. + +### 2. Create dummy data + +Generate test datasets (no shell script needed): + +```bash +# Default: 2 datasets × 50 WSIs each +python dummy_dataset_create.py + +# Custom: 3 datasets with 100 WSIs each +python dummy_dataset_create.py --datasets 3 --wsis-per-dataset 100 + +# Add more without deleting existing ones +python dummy_dataset_create.py -d 1 -w 30 --no-clean +``` + +This creates the following structure under `data/`: + +``` +data/ +├── dummy_dataset_1/ +│ ├── manifest.csv # patient_id, wsi_path (relative), cancer +│ └── wsis/ +│ ├── PAT_001.tiff +│ ├── PAT_002.tiff +│ └── ... +├── dummy_dataset_2/ +│ ├── manifest.csv +│ └── wsis/ +│ └── ... +``` + +**Options:** + +| Flag | Default | Description | +|---|---|---| +| `--datasets N` / `-d` | `2` | Number of dataset folders | +| `--wsis-per-dataset N` / `-w` | `50` | WSIs per dataset | +| `--data-dir DIR` | `data/` | Parent directory | +| `--seed N` | `42` | Reproducibility seed | +| `--no-clean` | off | Keep existing datasets, append new ones | +| `--img-size N` | `128` | Pixel size of dummy TIFF images | + +### 3. Register a user + +Edit the variables in `user_to_mlflow.py` and run: + +```bash +python user_to_mlflow.py +``` + +This creates a run in the **User_Registry** experiment with your identity tags. + +### 4. Register a dataset + +Edit the path/name variables in `dataset_to_mlflow.py` and run: + +```bash +python dataset_to_mlflow.py +``` + +This logs the manifest (enriched with file sizes & timestamps) as an MLflow +artifact under the **Dataset_Registry** experiment. + +### 5. Run a training experiment + +Edit `experiment.py` to define your model and training loop, then: + +```bash +python experiment.py +``` + +--- + +## API reference + +### Provenance — `@autolog` decorator + +Full auto-capture for plain PyTorch training runs: + +```python +from rationai.mlflow.provenance import autolog + +@autolog(model_name="my_model_v1", experiment_name="My_Experiment") +def train(run): + model = build_model() + run.register_model(model) + + optimizer = optim.SGD(model.parameters(), lr=1e-3, momentum=0.9) + run.register_optimizer(optimizer) + + scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.5) + run.register_scheduler(scheduler) + + for epoch in range(50): + loss = train_epoch(model, loader, optimizer) + run.log_metrics({"train_loss": loss}, step=epoch) + + run.save_model(model) + +if __name__ == "__main__": + train() +``` + +| Auto-captured | Details | +|---|---| +| **User** | Resolved from git config → linked to `User_Registry` run | +| **Dataset** | Latest `Dataset_Registry` run | +| **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 | +| **Git** | Commit, branch, remote URL | +| **Environment** | Frozen `requirements.txt` + `pyproject.toml` / `uv.lock` | +| **Console output** | stdout/stderr → `logs/console.log` artifact (ANSI-aware) | +| **PROV document** | OpenProvenance JSON → `provenance/prov.json` artifact | + +### Provenance + Lightning + +Wrap Lightning training in `@autolog` and pass the active run to `MLFlowLogger`: + +```python +import mlflow +from rationai.mlflow import Trainer, MLFlowLogger +from rationai.mlflow.provenance import autolog + +@autolog(model_name="my_lightning_model", experiment_name="My_Experiment") +def train(run): + model = MyLightningModule() + run.register_model(model) + run.register_optimizer(model.configure_optimizers()) + + # Reuse the @autolog run so Lightning logs to the same provenance-tracked run + logger = MLFlowLogger(experiment_name="My_Experiment", run_id=mlflow.active_run().info.run_id) + + trainer = Trainer(logger=logger, max_epochs=50) + trainer.fit(model, train_loader) + + run.save_model(model) +``` + +This gives you the same full provenance depth as plain PyTorch — GPU info, +model architecture, optimizer config, environment freeze, PROV document — +plus Lightning's native metric logging via `self.log()`. + +### 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.mlflow 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.mlflow 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.mlflow 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.mlflow 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.mlflow import MetaTiledSlides + +dataset = MetaTiledSlides( + manifest_uri="s3://bucket/data/manifest.parquet", + tile_size=256, +) +``` + +### Lightning integration + +| Component | Import | Purpose | +|---|---|---| +| `Trainer` | `from rationai.mlflow import Trainer` | Lightning Trainer with MLflow checkpoint sync | +| `MLFlowLogger` | `from rationai.mlflow import MLFlowLogger` | Logger with git tags, stream capture, checkpoint sync | +| `MultiloaderLifecycle` | `from rationai.mlflow import MultiloaderLifecycle` | Per-dataloader callback hooks | +| `lightning_autolog` | `from rationai.mlflow.lightning import autolog` | Lightning-specific autolog decorator | +| `with_cli_args` | `from rationai.mlflow.lightning import with_cli_args` | Programmatic config injection (Hydra) | + +--- + +## Project structure + +``` +. +├── demo.py # End-to-end demo (all components) +├── experiment.py # User training script (@autolog decorated) +├── dummy_dataset_create.py # Dummy data generator (CLI) +├── dataset_to_mlflow.py # Dataset registration script +├── user_to_mlflow.py # User registration script +├── requirements.txt # Python dependencies +├── pyproject.toml # Project metadata + deps +├── tests/ +│ └── test_all.py # Unit test suite +├── data/ # Dummy datasets (gitignored) +│ ├── dummy_dataset_1/ +│ │ ├── manifest.csv +│ │ └── wsis/ +│ └── ... +└── rationai/ + └── mlflow/ + ├── __init__.py # Package exports (lazy Lightning import) + ├── provenance.py # @autolog decorator + PROV-O engine + ├── 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 + │ ├── samplers/ + │ │ └── stratified_batch_sampler.py + │ └── datasets/ + │ └── meta_tiled_slides.py + └── lightning/ # Lightning + Hydra integration + ├── autolog.py + ├── trainer.py + ├── with_cli_args.py + ├── callbacks/ + │ └── 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](https://github.com/jiribuchta/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`](https://github.com/jiribuchta/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) +``` \ No newline at end of file diff --git a/demo.py b/demo.py new file mode 100644 index 0000000..8929644 --- /dev/null +++ b/demo.py @@ -0,0 +1,461 @@ +#!/usr/bin/env python3 +""" +Demo script for rationai.mlkit — full end-to-end pipeline. + +Creates dummy data, registers it in Dataset_Registry, trains a model with +provenance tracking (fail_fast=True), logs everything to MLflow, and prints +a summary of what was uploaded. + +Run: + python demo.py # full pipeline (local file store) + python demo.py --uri http://... # custom MLflow server +""" + +import argparse +import json +import os +import sys +from pathlib import Path + +os.environ["MLFLOW_ALLOW_FILE_STORE"] = "true" + +import logging +logging.getLogger("root").setLevel(logging.WARNING) +logging.getLogger("mlflow").setLevel(logging.ERROR) + + +# ────────────────────────────────────────────── +# Helpers +# ────────────────────────────────────────────── + +def sep(title=""): + print(f"\n{'='*60}") + if title: + print(f" {title}") + print(f"{'='*60}") + + +def sub(title): + print(f"\n{'─'*60}") + print(f" {title}") + print(f"{'─'*60}") + + +# ────────────────────────────────────────────── +# 1. Create dummy datasets +# ────────────────────────────────────────────── + +def step_create_datasets(n=2, wsis_per_ds=10): + sub("1. Creating dummy pathology datasets") + + from dummy_dataset_create import create_dummy_datasets + create_dummy_datasets( + num_datasets=n, + wsis_per_dataset=wsis_per_ds, + data_dir=Path("test_data"), + seed=42, + img_size=64, + clean=True, + ) + + data_dir = Path("test_data") + for d in sorted(data_dir.iterdir()): + if d.is_dir(): + manifest = d / "manifest.csv" + n_rows = sum(1 for _ in open(manifest)) - 1 if manifest.exists() else "?" + print(f" → {d.name}/ ({n_rows} samples)") + + +# ────────────────────────────────────────────── +# 2. Register dataset(s) in Dataset_Registry +# ────────────────────────────────────────────── + +def step_register_datasets(): + sub("2. Registering datasets in Dataset_Registry") + + from rationai.mlkit.provenance import register_dataset + + data_dir = Path("test_data") + for d in sorted(data_dir.iterdir()): + if d.is_dir(): + manifest = d / "manifest.csv" + if not manifest.exists(): + continue + ds_name = d.name + run_id = register_dataset( + dataset_dir=str(d), + dataset_name=ds_name, + ) + print(f" ✅ {ds_name}: registered (run={run_id[:8]})") + + +# ────────────────────────────────────────────── +# 3. Stream capture demo +# ────────────────────────────────────────────── + +def step_stream_capture(): + sub("3. StreamCapture + StreamModifier") + + import io + from rationai.mlkit import StreamCapture, StreamLogger, StreamModifier + + # StreamCapture — captures stdout into a logger + class _Buf(StreamLogger): + def __init__(self): + self._buf = io.StringIO() + def log_stream(self, text: str): + self._buf.write(text) + def get_value(self): + return self._buf.getvalue() + + logger = _Buf() + with StreamCapture(logger, streams=(sys.stdout,)): + print("Hello from stdout!") + print("\033[92mThis is green (ANSI)\033[0m") + + captured = logger.get_value() + has_ansi = "\x1b" in captured + print(f" Captured text : {captured.strip()!r}") + print(f" ANSI preserved: {'✅ (raw capture)' if has_ansi else '❌'}") + + # StreamModifier — injects side-effect logic into a stream's write + buf = io.StringIO() + side_log = [] + modifier = StreamModifier(stream=buf, id=42) + modifier.set_write(lambda s, iid: side_log.append(f"[{iid}] {s}")) + buf.write("hello") + modifier.teardown() + print(f" Side log : {side_log}") + print(f" Original buf : {buf.getvalue()!r}") + print(" ✅ StreamModifier works (side-effect injected before original write)") + + +# ────────────────────────────────────────────── +# 4. Metrics demo (AggregatedMetricCollection) +# — skipped if rationai.masks is not installed +# ────────────────────────────────────────────── + +def step_metrics(): + sub("4. AggregatedMetricCollection — tile → slide aggregation") + + try: + import torch + from torchmetrics import Accuracy + from rationai.mlkit import AggregatedMetricCollection, MaxAggregator, MeanAggregator + except ModuleNotFoundError as e: + if "rationai.masks" in str(e): + print(" ⊘ SKIPPED — rationai.masks (private dep) not installed") + return + raise + + preds = torch.tensor([0.1, 0.8, 0.3, 0.9]) + targets = torch.tensor([0, 1, 0, 1]) + keys = ["slide_A", "slide_A", "slide_B", "slide_B"] + + for name, agg in [("MaxAggregator", MaxAggregator()), ("MeanAggregator", MeanAggregator())]: + mc = AggregatedMetricCollection( + metrics={"accuracy": Accuracy(task="binary")}, + aggregator=agg, + ) + mc.update(preds, targets, keys) + result = mc.compute() + print(f" {name:20s}: accuracy = {result['accuracy'].item():.4f}") + + +# ────────────────────────────────────────────── +# 5. NestedMetricCollection demo +# — skipped if rationai.masks is not installed +# ────────────────────────────────────────────── + +def step_nested_metrics(): + sub("5. NestedMetricCollection — per-slide multiclass metrics") + + try: + import torch + from torchmetrics import Accuracy, Precision + from rationai.mlkit import NestedMetricCollection + except ModuleNotFoundError as e: + if "rationai.masks" in str(e): + print(" ⊘ SKIPPED — rationai.masks (private dep) not installed") + return + raise + + 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"], + ) + + preds = torch.tensor([[0.7, 0.2, 0.1], [0.1, 0.1, 0.8], [0.2, 0.6, 0.2]]) + targets = torch.tensor([0, 2, 1]) + keys = ["slide_1", "slide_1", "slide_2"] + + metrics.update(preds, targets, keys) + result = metrics.compute() + + print(f" Slides: {result['slide']}") + for k, v in result.items(): + if k != "slide": + print(f" {k:15s}: {v}") + + +# ────────────────────────────────────────────── +# 6. StratifiedBatchSampler demo +# ────────────────────────────────────────────── + +def step_sampler(): + sub("6. StratifiedBatchSampler — balanced class batches") + + from rationai.mlkit import StratifiedBatchSampler, PDMStratifiedBatchSampler + import pandas as pd + + # List-of-lists sampler + sampler = StratifiedBatchSampler( + data_indices=[[0, 1, 2, 3], [4, 5, 6, 7]], + batch_size=4, + ) + for i, batch in enumerate(sampler): + print(f" Batch {i}: {batch}") + print(f" Total batches: {len(sampler)}") + + # DataFrame-based sampler + df = pd.DataFrame({ + "idx": list(range(8)), + "label": [0, 0, 0, 1, 1, 1, 1, 0], + }) + pdm_sampler = PDMStratifiedBatchSampler(data=df, stratify_by="label", batch_size=4) + pdm_batches = list(pdm_sampler) + print(f" PDM sampler batches: {len(pdm_batches)}") + + +# ────────────────────────────────────────────── +# 7. Full training run with provenance autolog +# (fail_fast=True — dataset verification must pass) +# ────────────────────────────────────────────── + +def step_provenance(): + sub("7. Provenance — full training run, logged to MLflow (fail_fast=True)") + + import torch + import torch.nn as nn + import torch.optim as optim + from torch.utils.data import Dataset, DataLoader + from rationai.mlkit.provenance import autolog + + class _DummyDS(Dataset): + def __len__(self): + return 32 + def __getitem__(self, idx): + return torch.randn(64), torch.randint(0, 2, (1,)).item() + + @autolog(model_name="demo_model", experiment_name="Demo_Experiment", fail_fast=True) + def train(run): + model = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 2)) + run.register_model(model) + + optimizer = optim.Adam(model.parameters(), lr=0.01) + run.register_optimizer(optimizer) + + loader = DataLoader(_DummyDS(), batch_size=8) + criterion = nn.CrossEntropyLoss() + + for epoch in range(1, 5): + model.train() + total_loss = 0 + for bx, by in loader: + optimizer.zero_grad() + loss = criterion(model(bx), by) + loss.backward() + optimizer.step() + total_loss += loss.item() + + avg = total_loss / len(loader) + run.log_metrics({"train_loss": avg}, step=epoch) + print(f" Epoch {epoch}: loss={avg:.4f}") + + run.save_model(model) + + train() + + import mlflow as _mlf + while _mlf.active_run(): + _mlf.end_run() + + +# ────────────────────────────────────────────── +# 8. Lightning integration demo +# ────────────────────────────────────────────── + +def step_lightning(): + sub("8. Lightning — Trainer + MLFlowLogger (actual training)") + + import torch + import lightning as pl + import mlflow + from rationai.mlkit import Trainer, MLFlowLogger, MultiloaderLifecycle + + class _TinyModel(pl.LightningModule): + def __init__(self): + super().__init__() + self.net = torch.nn.Linear(8, 2) + def forward(self, x): + return self.net(x) + def training_step(self, batch, _): + x = torch.randn(4, 8, device=self.device) + loss = self.net(x).sum() + self.log("train_loss", loss) + return loss + def configure_optimizers(self): + return torch.optim.Adam(self.parameters(), lr=0.01) + + mlflow.set_experiment("Demo_Lightning") + run = mlflow.start_run(run_name="Training_demo_lightning_model") + + try: + model = _TinyModel() + optimizer = torch.optim.Adam(model.parameters(), lr=0.01) + + logger = MLFlowLogger(run_id=run.info.run_id) + trainer = Trainer( + logger=logger, + max_epochs=2, + enable_checkpointing=False, + enable_progress_bar=False, + enable_model_summary=False, + log_every_n_steps=1, + ) + + print(" Training a tiny Lightning model...") + trainer.fit(model, torch.utils.data.DataLoader( + torch.utils.data.TensorDataset(torch.randn(16, 8), torch.randint(0, 2, (16,))) + )) + print(" ✅ Training complete — check MLflow for logs") + finally: + mlflow.end_run() + + +# ────────────────────────────────────────────── +# 9. Summary — list all runs & artifacts in MLflow +# ────────────────────────────────────────────── + +def step_summary(): + sub("9. MLflow Summary") + + import mlflow + + tracking_uri = mlflow.get_tracking_uri() + print(f" Tracking URI: {tracking_uri}") + + client = mlflow.MlflowClient() + + for exp_name in ["Dataset_Registry", "Demo_Experiment", "Demo_Lightning"]: + exp = client.get_experiment_by_name(exp_name) + if not exp: + continue + runs = client.search_runs( + experiment_ids=[exp.experiment_id], + order_by=["attributes.start_time desc"], + ) + print(f"\n Experiment: {exp_name} ({len(runs)} run(s))") + + for r in runs[:5]: + params = dict(r.data.params) if r.data.params else {} + metrics = dict(r.data.metrics) if r.data.metrics else {} + artifacts = [a.path for a in client.list_artifacts(r.info.run_id)] + + print(f"\n Run: {r.info.run_name or r.info.run_id[:8]}") + print(f" Status : {r.info.status}") + if params: + print(f" Params : {params}") + if metrics: + print(f" Metrics: {metrics}") + + # Group artifacts by folder + folders = {} + for a in artifacts: + folder = a.split("/")[0] if "/" in a else "" + folders.setdefault(folder, []).append(a) + for folder, files in sorted(folders.items()): + print(f" [{folder}] {', '.join(os.path.basename(f) for f in files)}") + + # Check provenance summary + prov_path = None + for a in artifacts: + if "run_summary.json" in a: + prov_path = a + break + if prov_path: + local = client.download_artifacts(r.info.run_id, prov_path) + with open(local) as f: + summary = json.load(f) + print(f" Provenance keys: {', '.join(summary.keys())}") + + if "dataset_verification" in summary: + v = summary["dataset_verification"] + status = "✅" if v.get("verified") else "❌" + details = [d for d in v.get("details", []) if "FAILED" in d or "verified" in d] + print(f" Dataset verify: {status} {details[0] if details else '—'}") + + # Local file store hints + if tracking_uri.startswith("file://"): + db_path = Path(tracking_uri.replace("file://", "")) + print(f"\n 📂 Local data at: {db_path.resolve()}") + + +# ────────────────────────────────────────────── +# Main +# ────────────────────────────────────────────── + +def main(): + parser = argparse.ArgumentParser(description="rationai.mlkit — full pipeline demo") + parser.add_argument("--test", action="store_true", help="Run unit tests instead") + parser.add_argument("--uri", type=str, default=None, + help="MLflow tracking URI (default: local file store)") + args = parser.parse_args() + + import mlflow as _mlf + uri = args.uri or f"file:///tmp/mlkit_demo_mlruns_{os.getpid()}" + _mlf.set_tracking_uri(uri) + print(f"[*] MLflow tracking URI: {uri}") + + # Verify connectivity + try: + _client = _mlf.MlflowClient() + _ = _client.get_experiment_by_name("__ping__") + print(f"[*] Connected ✅\n") + except Exception as e: + print(f"[!] Warning: Could not connect to MLflow server: {e}\n") + + if args.test: + sys.path.insert(0, str(Path(__file__).resolve().parent)) + from tests.test_all import main as test_main + test_main() + return + + sep("rationai.mlkit — End-to-End Demo") + + # Step 1 & 2: create + register datasets (required for fail_fast=True) + step_create_datasets(n=2, wsis_per_ds=10) + step_register_datasets() + + # Steps 3-6: component demos + step_stream_capture() + step_metrics() + step_nested_metrics() + step_sampler() + + # Steps 7-8: training runs (fail_fast=True — verification must pass) + step_provenance() + step_lightning() + + # Step 9: summary + step_summary() + + sep("✅ Demo complete — all runs are in MLflow!") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/dummy_dataset_create.py b/dummy_dataset_create.py new file mode 100644 index 0000000..5254607 --- /dev/null +++ b/dummy_dataset_create.py @@ -0,0 +1,146 @@ +""" +Create dummy pathology datasets for local testing. + +Generates N datasets, each with M fake WSI images (small random TIFFs) +and a manifest.csv matching the real data format: + + patient_id,wsi_path,cancer + PAT_001,wsis/PAT_001.tiff,0 + +Usage examples: + # Default: 2 datasets x 50 WSIs each + python dummy_dataset_create.py + + # 3 datasets with 20 WSIs each + python dummy_dataset_create.py --datasets 3 --wsis-per-dataset 20 + + # 1 dataset with 100 WSIs (replaces existing data/) + python dummy_dataset_create.py --datasets 1 --wsis-per-dataset 100 + + # Keep existing datasets, just add more + python dummy_dataset_create.py --datasets 1 --wsis-per-dataset 30 --no-clean +""" + +import argparse +import csv +import io +import os +import shutil +import uuid +from pathlib import Path + +import numpy as np +from PIL import Image + + +def _parse_args(): + p = argparse.ArgumentParser( + description="Create dummy pathology datasets for local testing.", + ) + p.add_argument( + "--datasets", "-d", + type=int, default=2, + help="Number of dataset folders to create (default: 2)", + ) + p.add_argument( + "--wsis-per-dataset", "-w", + type=int, default=50, + help="Number of WSI images per dataset (default: 50)", + ) + p.add_argument( + "--data-dir", + type=str, default="test_data", + help="Parent directory for the datasets (default: test_data/)", + ) + p.add_argument( + "--seed", + type=int, default=42, + help="Random seed for reproducibility (default: 42)", + ) + p.add_argument( + "--no-clean", + action="store_true", + help="Skip removing existing dummy_dataset_* folders before creating new ones", + ) + p.add_argument( + "--img-size", + type=int, default=128, + help="Size of each dummy WSI image in pixels (default: 128x128)", + ) + return p.parse_args() + + +def _clean_existing(data_dir: Path): + """Remove old dummy_dataset_* folders from data/.""" + for entry in sorted(data_dir.iterdir()): + if entry.is_dir() and entry.name.startswith("dummy_dataset_"): + shutil.rmtree(entry) + print(f" Removed existing {entry.name}/") + + +def _generate_wsi(img_size: int, rng: np.random.Generator) -> bytes: + """Generate a small random TIFF image.""" + arr = rng.integers(0, 256, (img_size, img_size, 3), dtype=np.uint8) + img = Image.fromarray(arr) + buf = io.BytesIO() + img.save(buf, format="TIFF") + return buf.getvalue() + + +def create_dummy_datasets( + num_datasets: int = 2, + wsis_per_dataset: int = 50, + data_dir: Path = Path("test_data"), + seed: int = 42, + img_size: int = 128, + clean: bool = True, +): + rng = np.random.default_rng(seed) + + if clean and data_dir.exists(): + _clean_existing(data_dir) + + for ds_idx in range(num_datasets): + ds_name = f"dummy_dataset_{ds_idx + 1}" + ds_path = data_dir / ds_name + wsis_path = ds_path / "wsis" + wsis_path.mkdir(parents=True, exist_ok=True) + + rows = [] + for i in range(wsis_per_dataset): + patient_id = f"PAT_{uuid.uuid4().hex[:8]}" + tiff_name = f"{patient_id}.tiff" + cancer = int(rng.random() > 0.5) + wsi_path = wsis_path / tiff_name + + wsi_path.write_bytes(_generate_wsi(img_size, rng)) + rows.append((patient_id, str(wsi_path.relative_to(ds_path)), cancer)) + + manifest_path = ds_path / "manifest.csv" + with open(manifest_path, "w", newline="") as f: + writer = csv.writer(f) + writer.writerow(["patient_id", "wsi_path", "cancer"]) + writer.writerows(rows) + + print(f" Created {ds_name}/ ({len(rows)} samples)") + + +def main(): + args = _parse_args() + data_dir = Path(args.data_dir) + data_dir.mkdir(parents=True, exist_ok=True) + + print(f"[*] Creating {args.datasets} dataset(s) with {args.wsis_per_dataset} WSIs each...") + create_dummy_datasets( + num_datasets=args.datasets, + wsis_per_dataset=args.wsis_per_dataset, + data_dir=data_dir, + seed=args.seed, + img_size=args.img_size, + clean=not args.no_clean, + ) + print(f"[*] Done. Data in: {data_dir.resolve()}") + + +if __name__ == "__main__": + main() diff --git a/rationai/__init__.py b/rationai/__init__.py new file mode 100644 index 0000000..444020e --- /dev/null +++ b/rationai/__init__.py @@ -0,0 +1 @@ +# rationai namespace package diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index 43af683..9ce9e25 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -1,6 +1,58 @@ -from rationai.mlkit.autolog import autolog -from rationai.mlkit.lightning import Trainer -from rationai.mlkit.with_cli_args import with_cli_args +"""rationai.mlkit — ML toolkit with provenance tracking.""" +from rationai.mlkit.stream import StreamCapture, StreamLogger, StreamModifier +from rationai.mlkit.provenance import ( + autolog as provenance_autolog, + register_dataset, +) -__all__ = ["Trainer", "autolog", "with_cli_args"] +__all__ = [ + "StreamCapture", + "StreamLogger", + "StreamModifier", + "AggregatedMetricCollection", + "Aggregator", + "MaxAggregator", + "MeanAggregator", + "MeanPoolMaxAggregator", + "TopKAggregator", + "NestedMetricCollection", + "LazyMetricDict", + "StratifiedBatchSampler", + "PDMStratifiedBatchSampler", + "MetaTiledSlides", + "OpenSlideTilesDataset", + "Trainer", + "MLFlowLogger", + "MultiloaderLifecycle", + "autolog", + "with_cli_args", + "provenance_autolog", + "register_dataset", +] + + +def __getattr__(name): + if name in ("Trainer", "MLFlowLogger", "MultiloaderLifecycle", "autolog", "with_cli_args"): + import importlib + _mod = importlib.import_module("rationai.mlkit.lightning") + return getattr(_mod, name) + + 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) + + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/rationai/mlkit/autolog.py b/rationai/mlkit/autolog.py index 372e03a..dac19da 100644 --- a/rationai/mlkit/autolog.py +++ b/rationai/mlkit/autolog.py @@ -1,99 +1,2 @@ -import logging -import os -import tempfile -from collections.abc import Callable -from functools import partial, wraps -from pathlib import Path -from typing import overload - -import hydra -from hydra.core.hydra_config import HydraConfig -from omegaconf import DictConfig, OmegaConf - -from rationai.mlkit.lightning.loggers import MLFlowLogger -from rationai.mlkit.stream import StreamCapture - - -log = logging.getLogger(__name__) - - -WrapperT = Callable[[DictConfig], None] -FunctionT = Callable[[DictConfig, MLFlowLogger], None] - - -@overload -def autolog( - func: FunctionT, - *, - log_config: bool = True, - log_stream: bool = True, - log_hyperparams: bool = True, -) -> WrapperT: ... - - -@overload -def autolog( - func: None = None, - *, - log_config: bool = True, - log_stream: bool = True, - log_hyperparams: bool = True, -) -> Callable[[FunctionT], WrapperT]: ... - - -def autolog( - func: FunctionT | None = None, - *, - log_config: bool = True, - log_stream: bool = True, - log_hyperparams: bool = True, -) -> WrapperT | Callable[[FunctionT], WrapperT]: - """Decorator for automatic logging. - - Args: - func: The function to decorate - log_config: Whether to log the hydra configuration files - log_stream: Whether to log the std streams - log_hyperparams: Whether to log the hyperparameters defined in the - config.metadata.hyperparams. - """ - if func is None: - return partial(autolog, log_config=log_config, log_stream=log_stream) - - @wraps(func) - def wrapper(config: DictConfig) -> None: - logger: MLFlowLogger = hydra.utils.instantiate(config.logger) - - if log_config: - _log_config(config, logger) - - if ( - log_hyperparams - and hasattr(config, "metadata") - and hasattr(config.metadata, "hyperparams") - ): - logger.log_hyperparams(config.metadata.hyperparams) - - if log_stream: - with StreamCapture(logger): - return func(config, logger) - - return func(config, logger) - - return wrapper - - -def _log_config(config: DictConfig, logger: MLFlowLogger) -> None: - """Logs the hydra config.""" - with tempfile.TemporaryDirectory(dir=os.getcwd()) as tmp_dir_str: - tmp_dir = Path(tmp_dir_str) - with open(tmp_dir / "hydra.yaml", "w", encoding="utf-8") as file: - OmegaConf.save(HydraConfig.get(), file) - - with open(tmp_dir / "config.yaml", "w", encoding="utf-8") as file: - OmegaConf.save(config, file) - - with open(tmp_dir / "config-resolved.yaml", "w", encoding="utf-8") as file: - OmegaConf.save(config, file, resolve=True) - - logger.log_artifacts(tmp_dir_str, "configs") +from rationai.mlkit.lightning.autolog import * # noqa: F401, F403 +from rationai.mlkit.lightning.autolog import autolog # noqa: F401 diff --git a/rationai/mlkit/data/__init__.py b/rationai/mlkit/data/__init__.py new file mode 100644 index 0000000..33c9ec7 --- /dev/null +++ b/rationai/mlkit/data/__init__.py @@ -0,0 +1 @@ +# rationai.mlkit.data namespace diff --git a/rationai/mlkit/lightning/__init__.py b/rationai/mlkit/lightning/__init__.py index 7fc4eb6..90342bb 100644 --- a/rationai/mlkit/lightning/__init__.py +++ b/rationai/mlkit/lightning/__init__.py @@ -1,4 +1,13 @@ +from rationai.mlkit.lightning.autolog import autolog +from rationai.mlkit.lightning.callbacks import MultiloaderLifecycle +from rationai.mlkit.lightning.loggers import MLFlowLogger from rationai.mlkit.lightning.trainer import Trainer +from rationai.mlkit.lightning.with_cli_args import with_cli_args - -__all__ = ["Trainer"] +__all__ = [ + "Trainer", + "MLFlowLogger", + "MultiloaderLifecycle", + "autolog", + "with_cli_args", +] diff --git a/rationai/mlkit/lightning/autolog.py b/rationai/mlkit/lightning/autolog.py new file mode 100644 index 0000000..df2d2db --- /dev/null +++ b/rationai/mlkit/lightning/autolog.py @@ -0,0 +1,106 @@ +import logging +import os +import tempfile +from collections.abc import Callable +from functools import partial, wraps +from pathlib import Path +from typing import overload + +import hydra +from hydra.core.hydra_config import HydraConfig +from omegaconf import DictConfig, OmegaConf + +from rationai.mlkit.lightning.loggers import MLFlowLogger +from rationai.mlkit.stream import StreamCapture + + +log = logging.getLogger(__name__) + + +WrapperT = Callable[[DictConfig], None] +FunctionT = Callable[[DictConfig, MLFlowLogger], None] + + +@overload +def autolog( + func: FunctionT, + *, + log_config: bool = True, + log_stream: bool = True, + log_hyperparams: bool = True, +) -> WrapperT: ... + + +@overload +def autolog( + func: None = None, + *, + log_config: bool = True, + log_stream: bool = True, + log_hyperparams: bool = True, +) -> Callable[[FunctionT], WrapperT]: ... + + +def autolog( + func: FunctionT | None = None, + *, + log_config: bool = True, + log_stream: bool = True, + log_hyperparams: bool = True, +) -> WrapperT | Callable[[FunctionT], WrapperT]: + """Decorator for automatic logging. + + Args: + func: The function to decorate + log_config: Whether to log the hydra configuration files + log_stream: Whether to log the std streams + log_hyperparams: Whether to log the hyperparameters defined in the + config.metadata.hyperparams. + """ + if func is None: + return partial(autolog, log_config=log_config, log_stream=log_stream) + + @wraps(func) + def wrapper(config: DictConfig) -> None: + logger: MLFlowLogger = hydra.utils.instantiate(config.logger) + + if log_config: + _log_config(config, logger) + + if ( + log_hyperparams + and hasattr(config, "metadata") + and hasattr(config.metadata, "hyperparams") + ): + logger.log_hyperparams(config.metadata.hyperparams) + + if log_stream: + with StreamCapture(logger): + return func(config, logger) + + return func(config, logger) + + return wrapper + + +def _log_config(config: DictConfig, logger: MLFlowLogger) -> None: + """Logs the hydra config.""" + with tempfile.TemporaryDirectory(dir=os.getcwd()) as tmp_dir_str: + tmp_dir = Path(tmp_dir_str) + with open(tmp_dir / "hydra.yaml", "w", encoding="utf-8") as file: + OmegaConf.save(HydraConfig.get(), file) + + with open(tmp_dir / "config.yaml", "w", encoding="utf-8") as file: + OmegaConf.save(config, file) + + with open(tmp_dir / "config-resolved.yaml", "w", encoding="utf-8") as file: + OmegaConf.save(config, file, resolve=True) + + logger.log_artifacts(tmp_dir_str, "configs") + + +def __getattr__(name: str): + if name == "autolog_provenance": + from rationai.mlkit.provenance import autolog as _autolog + return _autolog + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/rationai/mlkit/lightning/loggers/mlflow.py b/rationai/mlkit/lightning/loggers/mlflow.py index 0f56026..ce33f30 100644 --- a/rationai/mlkit/lightning/loggers/mlflow.py +++ b/rationai/mlkit/lightning/loggers/mlflow.py @@ -59,7 +59,9 @@ def __init__( def experiment(self) -> MlflowClient: if not self._initialized: exp = super().experiment - mlflow.start_run(self.run_id, log_system_metrics=self.log_system_metrics) + # Only start a run if none is already active (e.g. from @autolog) + if not mlflow.active_run(): + mlflow.start_run(self.run_id, log_system_metrics=self.log_system_metrics) return exp return super().experiment diff --git a/rationai/mlkit/lightning/with_cli_args.py b/rationai/mlkit/lightning/with_cli_args.py new file mode 100644 index 0000000..ebddc0e --- /dev/null +++ b/rationai/mlkit/lightning/with_cli_args.py @@ -0,0 +1,57 @@ +import sys +from collections.abc import Callable +from functools import wraps +from typing import Any + + +def with_cli_args( + defaults: list[str] | None = None, overrides: list[str] | None = None +) -> Callable[[Callable[..., Any]], Callable[..., Any]]: + """Decorator to injects arguments into sys.argv. + + Args: + defaults: Arguments injected AFTER script name but BEFORE user args. + (Acts as defaults: User can override these). + overrides: Arguments injected AFTER user args. + (Acts as overrides: Forces value, User cannot override). + + Returns: + A decorator that modifies sys.argv for the duration of the decorated function. + + Examples: + >>> from rationai.mlkit import autolog, with_cli_args, MLFlowLogger + >>> from omegaconf import DictConfig + >>> import hydra + + >>> @with_cli_args(["+preprocessing=qc"]) + >>> @hydra.main(config_path="../configs", config_name="preprocessing", version_base=None) + >>> @autolog + >>> def main(config: DictConfig, logger: MLFlowLogger) -> None: + >>> pass + """ + prepend = defaults or [] + append = overrides or [] + + def decorator(func: Callable[..., Any]) -> Callable[..., Any]: + @wraps(func) + def wrapper(*args: Any, **kwargs: Any) -> Any: + # 1. Save original state + original_argv = sys.argv[:] + + # 2. Deconstruct existing argv + # sys.argv[0] is the script name + script_name = [sys.argv[0]] + user_provided_args = sys.argv[1:] + + # 3. Reconstruct: [Script] + [Start] + [User] + [End] + sys.argv = script_name + prepend + user_provided_args + append + + try: + return func(*args, **kwargs) + finally: + # 4. Restore original state guarantees safety + sys.argv = original_argv + + return wrapper + + return decorator diff --git a/rationai/mlkit/stream/__init__.py b/rationai/mlkit/stream/__init__.py index eeafc52..905a205 100644 --- a/rationai/mlkit/stream/__init__.py +++ b/rationai/mlkit/stream/__init__.py @@ -1,5 +1,5 @@ from rationai.mlkit.stream.stream_capture import StreamCapture from rationai.mlkit.stream.stream_logger import StreamLogger +from rationai.mlkit.stream.stream_modifier import StreamModifier - -__all__ = ["StreamCapture", "StreamLogger"] +__all__ = ["StreamCapture", "StreamLogger", "StreamModifier"] diff --git a/rationai/mlkit/with_cli_args.py b/rationai/mlkit/with_cli_args.py index ebddc0e..ba04856 100644 --- a/rationai/mlkit/with_cli_args.py +++ b/rationai/mlkit/with_cli_args.py @@ -1,57 +1,2 @@ -import sys -from collections.abc import Callable -from functools import wraps -from typing import Any - - -def with_cli_args( - defaults: list[str] | None = None, overrides: list[str] | None = None -) -> Callable[[Callable[..., Any]], Callable[..., Any]]: - """Decorator to injects arguments into sys.argv. - - Args: - defaults: Arguments injected AFTER script name but BEFORE user args. - (Acts as defaults: User can override these). - overrides: Arguments injected AFTER user args. - (Acts as overrides: Forces value, User cannot override). - - Returns: - A decorator that modifies sys.argv for the duration of the decorated function. - - Examples: - >>> from rationai.mlkit import autolog, with_cli_args, MLFlowLogger - >>> from omegaconf import DictConfig - >>> import hydra - - >>> @with_cli_args(["+preprocessing=qc"]) - >>> @hydra.main(config_path="../configs", config_name="preprocessing", version_base=None) - >>> @autolog - >>> def main(config: DictConfig, logger: MLFlowLogger) -> None: - >>> pass - """ - prepend = defaults or [] - append = overrides or [] - - def decorator(func: Callable[..., Any]) -> Callable[..., Any]: - @wraps(func) - def wrapper(*args: Any, **kwargs: Any) -> Any: - # 1. Save original state - original_argv = sys.argv[:] - - # 2. Deconstruct existing argv - # sys.argv[0] is the script name - script_name = [sys.argv[0]] - user_provided_args = sys.argv[1:] - - # 3. Reconstruct: [Script] + [Start] + [User] + [End] - sys.argv = script_name + prepend + user_provided_args + append - - try: - return func(*args, **kwargs) - finally: - # 4. Restore original state guarantees safety - sys.argv = original_argv - - return wrapper - - return decorator +from rationai.mlkit.lightning.with_cli_args import * # noqa: F401, F403 +from rationai.mlkit.lightning.with_cli_args import with_cli_args # noqa: F401 diff --git a/tests/test_all.py b/tests/test_all.py new file mode 100644 index 0000000..30a34ad --- /dev/null +++ b/tests/test_all.py @@ -0,0 +1,378 @@ +""" +End-to-end test suite for rationai.mlkit. + +Exercises every major component: + 1. Stream capture + StreamModifier + 2. AggregatedMetricCollection + aggregators (with torchmetrics) + 3. NestedMetricCollection + 4. StratifiedBatchSampler / PDMStratifiedBatchSampler + 5. Lightning: Trainer, MLFlowLogger, MultiloaderLifecycle, autolog, with_cli_args + 6. Provenance @autolog (full training run + provenance artifact verification) + +Run: python tests/test_all.py +""" + +import sys +import io +import os +import json +import tempfile +from pathlib import Path + +os.environ["MLFLOW_ALLOW_FILE_STORE"] = "true" + +sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) + +import torch +import mlflow + + +def section(title): + print(f"\n{'='*60}") + print(f" {title}") + print(f"{'='*60}\n") + + +# ────────────────────────────────────────────── +# 1. Stream Capture + StreamModifier +# ────────────────────────────────────────────── + +def test_stream_capture(): + section("1. StreamCapture + StreamModifier") + + from rationai.mlkit import StreamCapture, StreamLogger, StreamModifier + + class _TestLogger(StreamLogger): + def __init__(self): + self._buffer = io.StringIO() + def log_stream(self, text: str): + self._buffer.write(text) + def get_value(self): + return self._buffer.getvalue() + + logger = _TestLogger() + with StreamCapture(logger, streams=(sys.stdout,)): + print("hello") + print("\033[31mred text\033[0m") + + captured = logger.get_value() + assert "hello" in captured + assert "\x1b" in captured, "StreamCapture preserves ANSI codes (raw capture)" + print(f" Captured: {captured.strip()!r}") + print(" ✅ StreamCapture works (captures raw output including ANSI)") + + # StreamModifier — wraps a stream's write to inject side-effect + # logic before the original write fires. The callback receives (text, id) + # and can e.g. log or tag the output elsewhere. + buf = io.StringIO() + side_log = [] + modifier = StreamModifier(stream=buf, id=42) + modifier.set_write(lambda s, iid: side_log.append(f"[{iid}] {s}")) + buf.write("hello") + assert "[42] hello" in side_log, f"Side effect not called: {side_log}" + assert buf.getvalue() == "hello", f"Original write not called: {buf.getvalue()!r}" + modifier.teardown() + print(f" Side log: {side_log}") + print(f" Original buf: {buf.getvalue()!r}") + print(" ✅ StreamModifier works") + + +# ────────────────────────────────────────────── +# 2. AggregatedMetricCollection + aggregators +# ────────────────────────────────────────────── + +def test_aggregated_metrics(): + section("2. AggregatedMetricCollection") + + try: + from torchmetrics import Accuracy + from rationai.mlkit import ( + AggregatedMetricCollection, + MaxAggregator, + MeanAggregator, + ) + except ModuleNotFoundError as e: + if "rationai.masks" in str(e): + raise # handled by main as skip + raise + + preds = torch.tensor([0.1, 0.8, 0.3, 0.9]) + targets = torch.tensor([0, 1, 0, 1]) + keys = ["slide_A", "slide_A", "slide_B", "slide_B"] + + for name, agg in [ + ("MaxAggregator", MaxAggregator()), + ("MeanAggregator", MeanAggregator()), + ]: + mc = AggregatedMetricCollection( + metrics={"accuracy": Accuracy(task="binary")}, + aggregator=agg, + ) + mc.update(preds, targets, keys) + result = mc.compute() + print(f" {name:20s}: accuracy={result['accuracy'].item():.4f}") + + print(" ✅ AggregatedMetricCollection works") + + +# ────────────────────────────────────────────── +# 3. NestedMetricCollection +# ────────────────────────────────────────────── + +def test_nested_metrics(): + section("3. NestedMetricCollection") + + from torchmetrics import Accuracy, Precision + from rationai.mlkit import NestedMetricCollection + + 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"], + ) + + preds = torch.tensor([[0.7, 0.2, 0.1], [0.1, 0.1, 0.8], [0.2, 0.6, 0.2]]) + targets = torch.tensor([0, 2, 1]) + keys = ["slide_1", "slide_1", "slide_2"] + + metrics.update(preds, targets, keys) + result = metrics.compute() + + assert "slide" in result + print(f" Slides: {result['slide']}") + for k, v in result.items(): + if k != "slide": + print(f" {k:15s}: {v}") + print(" ✅ NestedMetricCollection works") + + +# ────────────────────────────────────────────── +# 4. StratifiedBatchSampler / PDMStratifiedBatchSampler +# ────────────────────────────────────────────── + +def test_samplers(): + section("4. Samplers") + + from rationai.mlkit import StratifiedBatchSampler + + sampler = StratifiedBatchSampler( + data_indices=[[0, 1, 2, 3], [4, 5, 6, 7]], + batch_size=4, + ) + batches = list(sampler) + print(f" StratifiedBatchSampler: {len(batches)} batches") + for i, batch in enumerate(batches): + print(f" Batch {i}: {batch}") + assert len(batches) == 2 + + # PDMStratifiedBatchSampler requires a DataFrame + from rationai.mlkit import PDMStratifiedBatchSampler + import pandas as pd + + df = pd.DataFrame({ + "idx": list(range(8)), + "label": [0, 0, 0, 1, 1, 1, 1, 0], + }) + pdm_sampler = PDMStratifiedBatchSampler( + data=df, + stratify_by="label", + batch_size=4, + ) + pdm_batches = list(pdm_sampler) + print(f" PDMStratifiedBatchSampler: {len(pdm_batches)} batches") + assert len(pdm_batches) >= 1 + + print(" ✅ Samplers work") + + +# ────────────────────────────────────────────── +# 5. Lightning imports + basic functionality +# ────────────────────────────────────────────── + +def test_lightning(): + section("5. Lightning (Trainer, MLFlowLogger, MultiloaderLifecycle, autolog, with_cli_args)") + + from rationai.mlkit import Trainer, MLFlowLogger, MultiloaderLifecycle, autolog, with_cli_args + print(f" Trainer: {Trainer}") + print(f" MLFlowLogger: {MLFlowLogger}") + print(f" MultiloaderLifecycle: {MultiloaderLifecycle}") + print(f" autolog: {autolog}") + print(f" with_cli_args: {with_cli_args}") + + import lightning as pl + + class _TinyModel(pl.LightningModule): + def __init__(self): + super().__init__() + self.net = torch.nn.Linear(8, 2) + def forward(self, x): + return self.net(x) + def training_step(self, batch, _): + x = torch.randn(4, 8, device=self.device) + loss = self.net(x).sum() + self.log("train_loss", loss) + return loss + def configure_optimizers(self): + return torch.optim.Adam(self.parameters(), lr=0.01) + + model = _TinyModel() + trainer = Trainer( + max_epochs=1, + enable_checkpointing=False, + enable_progress_bar=False, + enable_model_summary=False, + logger=False, + ) + ds = torch.utils.data.TensorDataset(torch.randn(8, 8), torch.randint(0, 2, (8,))) + trainer.fit(model, torch.utils.data.DataLoader(ds)) + print(" ✅ Lightning Trainer works") + + +# ────────────────────────────────────────────── +# 6. Provenance @autolog +# ────────────────────────────────────────────── + +def test_provenance(): + section("6. Provenance (@autolog + artifact verification)") + + import torch.nn as nn + import torch.optim as optim + from torch.utils.data import Dataset, DataLoader + from rationai.mlkit.provenance import autolog + + class _DummyDS(Dataset): + def __len__(self): + return 16 + def __getitem__(self, idx): + return torch.randn(32), torch.randint(0, 2, (1,)).item() + + with tempfile.TemporaryDirectory() as tmpdir: + mlflow.set_tracking_uri(f"file://{tmpdir}/mlruns") + + @autolog(model_name="test_model", experiment_name="Test_Provenance", fail_fast=False) + def train(run): + model = nn.Sequential(nn.Linear(32, 16), nn.ReLU(), nn.Linear(16, 2)) + run.register_model(model) + + optimizer = optim.Adam(model.parameters(), lr=0.01) + run.register_optimizer(optimizer) + + loader = DataLoader(_DummyDS(), batch_size=4) + criterion = nn.CrossEntropyLoss() + + for epoch in range(1, 3): + model.train() + total_loss = 0 + for bx, by in loader: + optimizer.zero_grad() + loss = criterion(model(bx), by) + loss.backward() + optimizer.step() + total_loss += loss.item() + + avg = total_loss / len(loader) + run.log_metrics({"train_loss": avg}, step=epoch) + print(f" Epoch {epoch}: loss={avg:.4f}") + + run.save_model(model) + + train() + + client = mlflow.MlflowClient() + exp = client.get_experiment_by_name("Test_Provenance") + runs = client.search_runs(experiment_ids=[exp.experiment_id]) + assert len(runs) >= 1, "Expected at least 1 run" + + r = runs[0] + artifacts = [a.path for a in client.list_artifacts(r.info.run_id)] + print(f" Artifacts: {artifacts}") + + prov_path = None + for a in artifacts: + if "provenance" in a.lower(): + prov_path = a.rstrip("/") + break + + if prov_path: + local_dir = client.download_artifacts(r.info.run_id, prov_path) + # download_artifacts returns a directory; find run_summary.json inside + import glob as _glob + summary_files = _glob.glob(f"{local_dir}/**/run_summary.json", recursive=True) + if not summary_files: + summary_files = [f for f in os.listdir(local_dir) if f.endswith(".json")] + summary_files = [os.path.join(local_dir, f) for f in summary_files] if summary_files else [] + if summary_files: + with open(summary_files[0]) as f: + summary = json.load(f) + print(f" Provenance keys: {list(summary.keys())}") + assert "model_name" in summary or "params" in summary or "metrics" in summary + else: + print(f" (Provenance dir contents: {os.listdir(local_dir)})") + else: + print(" (No provenance artifact found — checking run params/metrics)") + assert r.data.params, "Expected params in run" + + print(" ✅ Provenance @autolog works") + + +# ────────────────────────────────────────────── +# Backward compatibility +# ────────────────────────────────────────────── + +def test_backward_compat(): + section("7. Backward Compatibility (old import paths)") + + from rationai.mlkit import Trainer, autolog, with_cli_args + print(" from rationai.mlkit import Trainer, autolog, with_cli_args: ✅") + + from rationai.mlkit.autolog import autolog as _autolog + print(" from rationai.mlkit.autolog import autolog: ✅") + + from rationai.mlkit.with_cli_args import with_cli_args as _wca + print(" from rationai.mlkit.with_cli_args import with_cli_args: ✅") + + from rationai.mlkit.lightning.autolog import autolog as _la + print(" from rationai.mlkit.lightning.autolog import autolog: ✅") + + +# ────────────────────────────────────────────── +# Main +# ────────────────────────────────────────────── + +def main(): + section("rationai.mlkit — Test Suite") + + tests = [ + ("Stream Capture", test_stream_capture), + ("Aggregated Metrics", test_aggregated_metrics), + ("Nested Metrics", test_nested_metrics), + ("Samplers", test_samplers), + ("Lightning", test_lightning), + ("Provenance", test_provenance), + ("Backward Compat", test_backward_compat), + ] + + passed = failed = skipped = 0 + for name, fn in tests: + try: + fn() + passed += 1 + except ModuleNotFoundError as e: + print(f" ⊘ SKIPPED {name}: {e}") + skipped += 1 + except Exception as e: + print(f" ❌ {name}: {e}") + import traceback + traceback.print_exc() + failed += 1 + + section(f"Results: {passed} passed, {failed} failed, {skipped} skipped") + if failed: + sys.exit(1) + + +if __name__ == "__main__": + main() diff --git a/user_to_mlflow.py b/user_to_mlflow.py new file mode 100644 index 0000000..b4ca05b --- /dev/null +++ b/user_to_mlflow.py @@ -0,0 +1,37 @@ +import mlflow + +def register_new_user(username, real_name, email, organization, lead_name, lead_email): + """ + Registruje uživatele do MLflow experimentu 'User_Registry' s kompletními údaji. + """ + experiment_name = "User_Registry" + experiment = mlflow.get_experiment_by_name(experiment_name) + + if not experiment: + experiment_id = mlflow.create_experiment(experiment_name) + print(f"Vytvořen experiment: {experiment_name}") + else: + experiment_id = experiment.experiment_id + + # Vytvoření unikátního runu pro tohoto uživatele + with mlflow.start_run(experiment_id=experiment_id, run_name=f"User_{username}"): + mlflow.set_tags({ + "username": username, + "real_name": real_name, + "email": email, + "organization": organization, + "lead_name": lead_name, + "lead_email": lead_email + }) + print(f"Uživatel {real_name} ({username}) byl úspěšně zaregistrován.") + +if __name__ == "__main__": + # Příklad registrace + register_new_user( + username="jiribuchta", + real_name="Jiří Buchta", + email="524981@mail.muni.cz", + organization="RationAI", + lead_name="Tomáš Brázdil", + lead_email="brazdil@muni.cz" + ) \ No newline at end of file From d8be7f7658748d87c19190080f9851b2916f04cc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Mon, 20 Jul 2026 22:40:02 +0200 Subject: [PATCH 02/34] feat: added provenance and fixed README --- README.md | 67 +- rationai/mlkit/provenance.py | 1381 ++++++++++++++++++++++++++++++++++ 2 files changed, 1408 insertions(+), 40 deletions(-) create mode 100644 rationai/mlkit/provenance.py diff --git a/README.md b/README.md index 9a4586e..f57ed4c 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,4 @@ -# rationai.mlflow — Unified ML Provenance & Metrics Toolkit +# 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. @@ -14,7 +14,7 @@ optimizer/scheduler settings, console output) is captured automatically via the ```bash uv venv source .venv/bin/activate -uv pip install -r requirements.txt +uv sync ``` Start the MLflow server: @@ -35,9 +35,6 @@ The easiest way to see everything in action: # Full pipeline — uploads to localhost:5000 python demo.py -# With dummy datasets -python demo.py --datasets 3 - # Push to a different MLflow server python demo.py --uri http://your-server:5000 @@ -49,7 +46,7 @@ The demo exercises all components: | Step | Feature | What it shows | |---|---|---| -| 1 | Dummy data creation | `data/dummy_dataset_*` with manifests | +| 1 | Dummy data creation | `test_data/dummy_dataset_*` with manifests | | 2 | **StreamCapture** | ANSI-aware stdout/stderr capture | | 3 | **AggregatedMetricCollection** | Tile → slide metric aggregation | | 4 | **NestedMetricCollection** | Per-slide multiclass metrics | @@ -113,24 +110,15 @@ python user_to_mlflow.py This creates a run in the **User_Registry** experiment with your identity tags. -### 4. Register a dataset +### 4. Run a training experiment -Edit the path/name variables in `dataset_to_mlflow.py` and run: +Use the `@autolog` decorator from `rationai.mlkit.provenance` in your training script, then: ```bash -python dataset_to_mlflow.py +python your_experiment.py ``` -This logs the manifest (enriched with file sizes & timestamps) as an MLflow -artifact under the **Dataset_Registry** experiment. - -### 5. Run a training experiment - -Edit `experiment.py` to define your model and training loop, then: - -```bash -python experiment.py -``` +See the [API reference](#provenance--autolog-decorator) below for details. --- @@ -141,7 +129,7 @@ python experiment.py Full auto-capture for plain PyTorch training runs: ```python -from rationai.mlflow.provenance import autolog +from rationai.mlkit.provenance import autolog @autolog(model_name="my_model_v1", experiment_name="My_Experiment") def train(run): @@ -185,8 +173,8 @@ Wrap Lightning training in `@autolog` and pass the active run to `MLFlowLogger`: ```python import mlflow -from rationai.mlflow import Trainer, MLFlowLogger -from rationai.mlflow.provenance import autolog +from rationai.mlkit import Trainer, MLFlowLogger +from rationai.mlkit.provenance import autolog @autolog(model_name="my_lightning_model", experiment_name="My_Experiment") def train(run): @@ -215,7 +203,7 @@ Group tile-level predictions by slide and compute metrics at the slide level: ```python from torchmetrics import Accuracy -from rationai.mlflow import ( +from rationai.mlkit import ( AggregatedMetricCollection, MaxAggregator, MeanAggregator, @@ -240,7 +228,7 @@ Available aggregators: `MaxAggregator`, `MeanAggregator`, `TopKAggregator`, `Mea Compute multiple torchmetrics per slide with class-level breakdowns: ```python -from rationai.mlflow import NestedMetricCollection +from rationai.mlkit import NestedMetricCollection from torchmetrics import Accuracy, Precision metrics = NestedMetricCollection( @@ -262,7 +250,7 @@ Captures console output (including progress bars and ANSI color codes) without corrupting the log: ```python -from rationai.mlflow import StreamCapture +from rationai.mlkit import StreamCapture with StreamCapture(stream="stdout") as capture: print("Hello!") @@ -279,7 +267,7 @@ clean = capture.get_clean_text() # ANSI codes stripped Balanced class sampling across batches: ```python -from rationai.mlflow import StratifiedBatchSampler +from rationai.mlkit import StratifiedBatchSampler sampler = StratifiedBatchSampler( data_indices=[[0, 1, 2, 3], [4, 5, 6, 7]], # per-class indices @@ -294,7 +282,7 @@ for batch in sampler: Load tile data from parquet or MLflow artifact URIs: ```python -from rationai.mlflow import MetaTiledSlides +from rationai.mlkit import MetaTiledSlides dataset = MetaTiledSlides( manifest_uri="s3://bucket/data/manifest.parquet", @@ -306,11 +294,11 @@ dataset = MetaTiledSlides( | Component | Import | Purpose | |---|---|---| -| `Trainer` | `from rationai.mlflow import Trainer` | Lightning Trainer with MLflow checkpoint sync | -| `MLFlowLogger` | `from rationai.mlflow import MLFlowLogger` | Logger with git tags, stream capture, checkpoint sync | -| `MultiloaderLifecycle` | `from rationai.mlflow import MultiloaderLifecycle` | Per-dataloader callback hooks | -| `lightning_autolog` | `from rationai.mlflow.lightning import autolog` | Lightning-specific autolog decorator | -| `with_cli_args` | `from rationai.mlflow.lightning import with_cli_args` | Programmatic config injection (Hydra) | +| `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 | +| `lightning_autolog` | `from rationai.mlkit.lightning import autolog` | Lightning-specific autolog decorator | +| `with_cli_args` | `from rationai.mlkit.lightning import with_cli_args` | Programmatic config injection (Hydra) | --- @@ -319,23 +307,21 @@ dataset = MetaTiledSlides( ``` . ├── demo.py # End-to-end demo (all components) -├── experiment.py # User training script (@autolog decorated) ├── dummy_dataset_create.py # Dummy data generator (CLI) -├── dataset_to_mlflow.py # Dataset registration script ├── user_to_mlflow.py # User registration script -├── requirements.txt # Python dependencies ├── pyproject.toml # Project metadata + deps ├── tests/ │ └── test_all.py # Unit test suite -├── data/ # Dummy datasets (gitignored) +├── test_data/ # Dummy datasets (gitignored) │ ├── dummy_dataset_1/ │ │ ├── manifest.csv │ │ └── wsis/ │ └── ... └── rationai/ - └── mlflow/ + └── mlkit/ ├── __init__.py # Package exports (lazy Lightning import) - ├── provenance.py # @autolog decorator + PROV-O engine + ├── autolog.py # Re-exports lightning.autolog + ├── with_cli_args.py # Re-exports lightning.with_cli_args ├── stream/ # ANSI-aware console capture │ ├── stream_capture.py │ ├── stream_logger.py @@ -349,7 +335,8 @@ dataset = MetaTiledSlides( │ ├── samplers/ │ │ └── stratified_batch_sampler.py │ └── datasets/ - │ └── meta_tiled_slides.py + │ ├── meta_tiled_slides.py + │ └── openslide_tiles_dataset.py └── lightning/ # Lightning + Hydra integration ├── autolog.py ├── trainer.py @@ -431,4 +418,4 @@ 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) -``` \ No newline at end of file +``` diff --git a/rationai/mlkit/provenance.py b/rationai/mlkit/provenance.py new file mode 100644 index 0000000..801e213 --- /dev/null +++ b/rationai/mlkit/provenance.py @@ -0,0 +1,1381 @@ +""" +Automatic PROV-O-aware provenance logger for MLflow. + +Inspired by rationai.mlkit.autolog — the user decorates their training function +and all metadata is captured automatically. + +Usage: + + from provenance import autolog + + @autolog(model_name="resnet_baseline_v1") + def train(run): + run.log_params({"learning_rate": 1e-3}) + model = build_model() + run.register_model(model) + # ... training loop ... + run.save_model(model) + +Everything else (user, dataset, hardware, docker, git, environment, +train/test split, console output) is detected and logged automatically. +""" + +from __future__ import annotations + +import io +import os +import re +import json +import uuid +import shutil +import types +import platform +import hashlib +import subprocess +import contextlib +from datetime import datetime, timezone +from functools import partial, wraps +from collections.abc import Callable + +import mlflow +import torch +import pandas as pd + +# ────────────────────────────────────────────── +# OpenProvenance / CPM namespace URIs +# ────────────────────────────────────────────── + +_PROV_PREFIXES = { + "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/", +} + +# Hyperparameter keys that go on the activity vs. metadata entity +_ACTIVITY_HP_KEYS = { + "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", +} + +# Param keys that map to WSI entity properties +_WSI_PARAM_KEYS = { + "scanner", "slide_id", "wsi_id", "patient_id", "subject_id", + "institution", "site", "staining", "slicing_method", +} + + +# ────────────────────────────────────────────── +# Auto-detection helpers +# ────────────────────────────────────────────── + +def _get_git_info(): + """Return (commit, remote_url, branch) or ('unknown', ...) on failure.""" + try: + commit = subprocess.check_output( + ["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL, + ).decode().strip() + remote = subprocess.check_output( + ["git", "config", "--get", "remote.origin.url"], + stderr=subprocess.DEVNULL, + ).decode().strip() + branch = subprocess.check_output( + ["git", "rev-parse", "--abbrev-ref", "HEAD"], + stderr=subprocess.DEVNULL, + ).decode().strip() + except subprocess.CalledProcessError: + commit, remote, branch = "unknown", "unknown", "unknown" + return commit, remote, branch + + +def _lookup_experiment(name): + exp = mlflow.get_experiment_by_name(name) + return exp.experiment_id if exp else None + + +def _lookup_user_run(): + """Find the user run from User_Registry. Auto-detect username.""" + username = os.environ.get("MLFLOW_USER") + if not username: + try: + username = subprocess.check_output( + ["git", "config", "user.name"], stderr=subprocess.DEVNULL, + ).decode().strip() + except subprocess.CalledProcessError: + pass + if not username: + username = os.environ.get("USER", "unknown") + + exp_id = _lookup_experiment("User_Registry") + if exp_id is None: + return None, {} + + runs = mlflow.search_runs(experiment_ids=[exp_id]) + if runs.empty: + return None, {} + + matched = runs[runs["tags.username"] == username] + if matched.empty: + matched = runs.head(1) + + row = matched.iloc[0] + run = mlflow.get_run(row.run_id) + return row.run_id, dict(run.data.tags) + + +def _lookup_dataset_run(manifest_path: str | None = None): + """Return the Dataset_Registry run that matches the detected manifest. + + If *manifest_path* is given, looks for a registration run whose + ``manifest_hash`` tag matches the SHA-256 of that file. Falls back + to the latest Dataset_Registry run if no match is found (so that + runs 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 runs_df.empty: + return None + + # If we have a manifest, try to match by hash first + if manifest_path and os.path.isfile(manifest_path): + current_hash = _hash_manifest(manifest_path) + for _, row in runs_df.iterrows(): + reg_hash = row.get("tags.manifest_hash", "") + if reg_hash == current_hash: + return row["run_id"] + + # Fallback: latest run (backward compat when only one dataset registered) + return runs_df.iloc[0]["run_id"] + + +def _hash_manifest(manifest_path: str) -> str: + """Compute SHA-256 hash of a manifest CSV file.""" + h = hashlib.sha256() + with open(manifest_path, "rb") as f: + for chunk in iter(lambda: f.read(8192), b""): + h.update(chunk) + return h.hexdigest() + + +def _hash_samples(samples: list[dict]) -> str: + """Compute a deterministic hash over the set of samples (path+label pairs). + + Order-independent: sorts by path before hashing so that any reordering + of the manifest doesn't change the fingerprint. + """ + h = hashlib.sha256() + for s in sorted(samples, key=lambda x: x["path"]): + h.update(f"{s['path']}:{s['label']}\n".encode()) + return h.hexdigest() + + +def register_dataset( + dataset_dir: str, + dataset_name: str | None = None, + version: str = "1.0.0", + experiment_name: str = "Dataset_Registry", +): + """Register a dataset in MLflow's Dataset_Registry experiment. + + Computes a SHA-256 hash of the manifest and stores it as a tag so that + future training runs can verify they're using the exact same 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.mlflow.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) + + # Compute hashes + manifest_hash = _hash_manifest(manifest_path) + + # Read manifest and compute sample-level hash + df = pd.read_csv(manifest_path) + samples = [] + for _, row in df.iterrows(): + rel = row["wsi_path"] + full = os.path.join(dataset_dir, rel) if not os.path.isabs(rel) else rel + samples.append({"path": full, "label": int(row["cancer"])}) + samples_hash = _hash_samples(samples) + + # File-level hashes for each WSI + file_hashes = {} + for s in samples: + if os.path.isfile(s["path"]): + fh = hashlib.sha256() + with open(s["path"], "rb") as f: + for chunk in iter(lambda: f.read(8192), b""): + fh.update(chunk) + file_hashes[os.path.basename(s["path"])] = fh.hexdigest()[:16] + else: + file_hashes[os.path.basename(s["path"])] = "MISSING" + + # Register in MLflow + mlflow.set_experiment(experiment_name) + run = mlflow.start_run(run_name=f"Dataset_{dataset_name}_{version}") + run_id = run.info.run_id + + 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, + "manifest_hash": manifest_hash, + "samples_hash": samples_hash, + "file_hashes": json.dumps(file_hashes), + }) + + # Save manifest hash as an artifact for offline verification + prov_dir = f"_dataset_prov_{uuid.uuid4().hex[:8]}" + os.makedirs(prov_dir, exist_ok=True) + prov_path = os.path.join(prov_dir, "dataset_provenance.json") + with open(prov_path, "w") as f: + json.dump({ + "dataset_name": dataset_name, + "version": version, + "dataset_root": dataset_dir, + "manifest_hash": manifest_hash, + "samples_hash": samples_hash, + "file_hashes": file_hashes, + "num_samples": len(samples), + }, f, indent=2) + mlflow.log_artifact(prov_path, artifact_path="provenance") + shutil.rmtree(prov_dir, ignore_errors=True) + + mlflow.end_run() + print(f" [register_dataset] {dataset_name} v{version} → run_id={run_id}") + print(f" manifest_hash : {manifest_hash[:16]}…") + print(f" samples_hash : {samples_hash[:16]}…") + return run_id + + +def _verify_dataset( + manifest_path: str, + data_root: str, + dataset_run_id: str | None, +) -> dict: + """Verify the current dataset against the registered version in MLflow. + + Checks: + 1. Manifest hash matches (file-level integrity) + 2. Samples hash matches (content-level integrity, order-independent) + 3. All WSI files exist on disk + + Returns a dict with verification results. + """ + result: dict = { + "verified": False, + "dataset_run_id": dataset_run_id, + "manifest_hash_match": None, + "samples_hash_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 hashes + try: + reg_run = mlflow.get_run(dataset_run_id) + reg_tags = reg_run.data.tags + reg_manifest_hash = reg_tags.get("manifest_hash", "") + reg_samples_hash = reg_tags.get("samples_hash", "") + reg_file_hashes_str = reg_tags.get("file_hashes", "") + reg_file_hashes = json.loads(reg_file_hashes_str) if reg_file_hashes_str else {} + except Exception as e: + result["details"].append(f"Failed to fetch Dataset_Registry run: {e}") + return result + + # Compute current hashes + curr_manifest_hash = _hash_manifest(manifest_path) + result["manifest_hash_match"] = curr_manifest_hash == reg_manifest_hash + + if not result["manifest_hash_match"]: + result["details"].append( + f"Manifest hash mismatch: " + f"current={curr_manifest_hash[:16]}… registered={reg_manifest_hash[:16]}…" + ) + + # Compute samples hash + df = pd.read_csv(manifest_path) + samples = [] + 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"])}) + curr_samples_hash = _hash_samples(samples) + result["samples_hash_match"] = curr_samples_hash == reg_samples_hash + + if not result["samples_hash_match"]: + result["details"].append( + f"Samples hash mismatch: " + f"current={curr_samples_hash[:16]}… registered={reg_samples_hash[:16]}…" + ) + + # Check file existence + missing = 0 + for s in samples: + if not os.path.isfile(s["path"]): + missing += 1 + result["files_total"] = len(samples) + result["files_missing"] = missing + + if missing > 0: + result["details"].append(f"{missing}/{len(samples)} WSI files missing on disk") + + # Overall verdict + result["verified"] = ( + result["manifest_hash_match"] + and result["samples_hash_match"] + and missing == 0 + ) + + if result["verified"]: + result["details"].append("✅ Dataset verified — matches registered version") + else: + result["details"].append("❌ Dataset verification FAILED") + + return result + + +def _detect_hardware(): + 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(): + info: dict[str, str | bool] = {"docker": False} + + if os.path.exists("/.dockerenv"): + info["docker"] = True + + # Fallback: cgroup v1 / v2 + 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 + + # Try docker inspect if socket is available inside container + if info["docker"]: + cid = 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 _detect_manifest(): + """Walk data/ looking for manifest.csv.""" + 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 + + +def _snapshot_environment(artifact_dir): + """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)) + + # Return the frozen requirements text for embedding into provenance + with open(req_path) as f: + return f.read() + + +def _source_hash(source_file: str) -> str | None: + """Return SHA-256 of a source file (for reproducibility verification).""" + try: + h = hashlib.sha256() + with open(source_file, "rb") as f: + for chunk in iter(lambda: f.read(8192), b""): + h.update(chunk) + return h.hexdigest() + except (FileNotFoundError, PermissionError): + return None + + +def _prepare_split(manifest_path, data_root, test_size=0.2, random_state=42): + from sklearn.model_selection import train_test_split + + df = pd.read_csv(manifest_path) + samples = [] + 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"])}) + + train_s, test_s = train_test_split( + samples, + test_size=test_size, + random_state=random_state, + stratify=[s["label"] for s in samples], + ) + return list(train_s), list(test_s) + + +# ────────────────────────────────────────────── +# Model / optimizer / scheduler introspection +# ────────────────────────────────────────────── + +def _model_summary(model): + """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}" + ) + + info["layer_summary"] = "\n".join(layer_lines[:20]) + if len(layer_lines) > 20: + info["layer_summary"] += f"\n... ({len(layer_lines)} layers total)" + + info["model_class"] = type(model).__name__ + return info + + +def _optimizer_summary(optimizer): + """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): + """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}"] = list(val) if isinstance(val, (list, tuple)) else 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 + + +# ────────────────────────────────────────────── +# OpenProvenance / CPM PROV document builder +# ────────────────────────────────────────────── + +def _safe_id(name: str) -> str: + """Sanitise a string for use as a PROV identifier fragment.""" + return re.sub(r'[^a-zA-Z0-9_]', '_', name) + + +def _qualified(prefix: str, local: str) -> str: + """Return a qualified name like 'gen:run_abc123'.""" + return f"{prefix}:{local}" + + +def _typed_value(value, type_prefix="xsd", type_local="string") -> list: + """Wrap a string value as [value] — matching Java's array convention.""" + return [str(value)] + + +def _qualified_name(type_prefix: str, type_local: str) -> dict: + """Build a prov:QUALIFIED_NAME type descriptor.""" + return {"type": "prov:QUALIFIED_NAME", "$": f"{type_prefix}:{type_local}"} + + +def _iso_timestamp(ts_ms: int | None = None) -> str: + """Return an ISO-8601 timestamp string (ms since epoch or now).""" + if ts_ms is not None: + dt = datetime.fromtimestamp(ts_ms / 1000, tz=timezone.utc) + else: + dt = datetime.now(timezone.utc) + return dt.strftime("%Y-%m-%dT%H:%M:%S.000+00:00") + + +def _build_prov_document( + 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 | None = None, + requirements: str | None = None, + verification: dict | None = None, +) -> dict: + """Build an OpenProvenance-compatible PROV document dict. + + Structure mirrors the Java prov_mlflow output: + - bundle wrapper with storage: key + - prefix namespace declarations + - entity, activity, agent sections + - wasAssociatedWith, used relationship sections + """ + + # ── Derive identifiers ──────────────────────────────── + 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) + + # ── Collect sections ────────────────────────────────── + entities: dict[str, dict] = {} + activities: dict[str, dict] = {} + agents: dict[str, dict] = {} + used: dict[str, dict] = {} + was_associated_with: dict[str, dict] = {} + + rel_counter = [0] + + def _blank_rel_id() -> str: + rid = f"_:n{rel_counter[0]}" + rel_counter[0] += 1 + return rid + + # ── 1. AGENT (researcher) ───────────────────────────── + agent_props: dict[str, list] = {} + 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 (WSI / dataset samples) ───────── + # Collect unique sample paths from params or tags + sample_paths: list[str] = [] + for key in ("train_samples", "test_samples"): + if key in params: + pass # counts, not paths — skip + + # Try to find WSI-related params + 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 we have a manifest reference, create a single dataset entity + if image_path_candidates: + wsi_local = _safe_id(f"wsi_{image_path_candidates}") + wsi_id = _qualified("gen", wsi_local) + wsi_props: dict[str, list] = { + "schema:name": _typed_value(f"Input: {image_path_candidates}"), + "prov:type": [_qualified_name("sosa", "Sample")], + } + # Add optional WSI metadata from params + 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: + # Fallback: create a generic dataset entity from split info + 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 (the ML training) ───────────────── + run_activity: dict[str, object] = {} + 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) + + # Experiment name from tags + exp_name = params.get("model_name", "") + if exp_name: + run_activity["gen:experiment_name"] = _typed_value(exp_name) + + # Model config + if "model_class" in params: + run_activity["gen:model_config"] = _typed_value(params["model_class"]) + + # Git commit + git_commit = tags.get("git_commit", tags.get("mlflow.source.git.commit", "")) + if git_commit: + run_activity["schema:identifier"] = _typed_value(git_commit) + + # Backward-compatible model params + for key in ("pretrained_model", "backbone", "feature_extractor"): + if key in params: + run_activity["gen:pretrained_model"] = _typed_value(params[key]) + + # Dataset info + 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]) + + # Hyperparameters + for key in _ACTIVITY_HP_KEYS: + if key in params: + run_activity[f"gen:{key}"] = _typed_value(params[key]) + + # Also add opt_ and sch_ prefixed params as hyperparams (stripped) + # Strip ALL leading prefixes to avoid sch_opt_lr → gen:opt_lr pollution + for key, val in params.items(): + if key.startswith("opt_") or key.startswith("sch_"): + clean = key + while clean.startswith("opt_") or clean.startswith("sch_"): + if clean.startswith("opt_"): + clean = clean[4:] + elif clean.startswith("sch_"): + clean = clean[4:] + if f"gen:{clean}" not in run_activity: # avoid duplicates + run_activity[f"gen:{clean}"] = _typed_value(val) + + # Hardware from tags (mlflow.* convention) and our custom params + 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]) + + # Our custom hardware params + 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 remote / source + 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) + + # Segmentation / model params (if present) + 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, list] = {} + 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 already placed on the activity + 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", + } + + # Remaining params → metadata entity + 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) + + # Metrics → metadata entity + for key, val in metrics.items(): + safe_key = _safe_id(key) + meta_entity[f"gen:{safe_key}"] = _typed_value(val) + + # ── Reproducibility: embedded dataset splits ─────────── + if split_data: + # Embed as JSON string so the PROV document is self-contained + 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"]))[0]] + if split_data.get("test"): + meta_entity["gen:split_test"] = [_typed_value(json.dumps(split_data["test"]))[0]] + + # ── Reproducibility: frozen requirements ─────────────── + if requirements: + meta_entity["gen:requirements"] = [requirements] + + # ── Reproducibility: dataset verification ────────────── + if verification: + meta_entity["gen:dataset_verified"] = [str(verification.get("verified", False))] + meta_entity["gen:dataset_run_id"] = [verification.get("dataset_run_id", "")] + mh = verification.get("manifest_hash_match") + if mh is not None: + meta_entity["gen:manifest_hash_match"] = [str(mh)] + sh = verification.get("samples_hash_match") + if sh is not None: + meta_entity["gen:samples_hash_match"] = [str(sh)] + 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)] + + # Selected tags → metadata entity + 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 + + # wasGeneratedBy: meta entity ← run activity + was_generated_by: dict[str, dict] = {} + was_generated_by[_blank_rel_id()] = { + "prov:entity": meta_id, + "prov:activity": run_act_id, + } + + # ── 5. CPM MAIN ACTIVITY ────────────────────────────── + main_activity: dict[str, object] = {} + 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} + 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}} + + +# ────────────────────────────────────────────── +# Console stream capture (like mlkit.stream) +# ────────────────────────────────────────────── + +_CONSOLE_LOG_NAME = "console.log" + + +class _StreamCapture: + """Captures stdout/stderr and writes them to MLflow as console.log.""" + + def __init__(self, run_id: str): + self.run_id = run_id + self._buffer = io.StringIO() + self._originals = {} + + def __enter__(self): + import sys + for stream_name in ("stdout", "stderr"): + stream = getattr(sys, stream_name) + original_write = stream.write + self._originals[stream_name] = original_write + + def _make_wrapper(buf, orig): + def wrapper(text): + buf.write(text) + return orig(text) + return wrapper + + setattr(stream, "write", _make_wrapper(self._buffer, self._originals[stream_name])) + + def __exit__(self, *args): + import sys + for stream_name in ("stdout", "stderr"): + original = self._originals.get(stream_name) + if original is not None: + setattr(getattr(sys, stream_name), "write", original) + + text = self._buffer.getvalue() + if text.strip(): + try: + client = mlflow.tracking.MlflowClient() + log_path = os.path.join( + f"_mlflow_console_{uuid.uuid4().hex[:8]}", + _CONSOLE_LOG_NAME, + ) + os.makedirs(os.path.dirname(log_path), exist_ok=True) + with open(log_path, "w") as f: + f.write(text) + client.log_artifact(self.run_id, log_path, artifact_path="logs") + shutil.rmtree(os.path.dirname(log_path), ignore_errors=True) + except Exception: + pass # don't fail the run + + +# ────────────────────────────────────────────── +# Run helper object (injected into the user function) +# ────────────────────────────────────────────── + +class _Run: + """Handle passed to the user's training function. + + Provides logging methods and model/optimizer/scheduler registration. + """ + + def __init__(self, run_id: str): + self._run_id = run_id + self._model = None + self._optimizer_info: dict | None = None + self._scheduler_info: dict | None = None + self.train_paths: list[str] = [] + self.test_paths: list[str] = [] + self.train_labels: list[int] = [] + self.test_labels: list[int] = [] + self._split_data: dict | None = None # {"train": [...], "test": [...]} + + def _set_split(self, train_samples: list[dict], test_samples: list[dict]): + """Store the full split data for embedding into provenance.""" + self._split_data = {"train": train_samples, "test": test_samples} + + # ── Logging (forwarded to mlflow) ───────── + + def log_param(self, key, value): + mlflow.log_param(key, value) + + def log_params(self, params_dict): + mlflow.log_params(params_dict) + + def log_metric(self, key, value, step=None): + mlflow.log_metric(key, value, step=step) + + def log_metrics(self, metrics_dict, step=None): + mlflow.log_metrics(metrics_dict, step=step) + + def log_artifact(self, local_path, artifact_path=None): + mlflow.log_artifact(local_path, artifact_path=artifact_path) + + def log_artifacts(self, local_dir, artifact_path=None): + mlflow.log_artifacts(local_dir, artifact_path=artifact_path) + + # ── Registration (logged at the end) ─────── + + def register_model(self, model): + self._model = model + + def register_optimizer(self, optimizer): + self._optimizer_info = _optimizer_summary(optimizer) + + def register_scheduler(self, scheduler): + self._scheduler_info = _scheduler_summary(scheduler) + + # ── Model saving ────────────────────────── + + def save_model(self, model, name="model", **kwargs): + if "export_model" not in kwargs: + kwargs["export_model"] = False + mlflow.pytorch.log_model(model, name, **kwargs) + + +# ────────────────────────────────────────────── +# Decorator — the main entry point +# ────────────────────────────────────────────── + +def autolog( + model_name: str | None = None, + experiment_name: str = "Training_Pipeline", + test_size: float = 0.2, + random_state: int = 42, + log_stream: bool = True, + fail_fast: bool = True, +): + """Decorator for automatic provenance logging. + + All metadata (user, dataset, hardware, docker, git, environment, + train/test split, model architecture, optimizer, scheduler, console + output) is captured automatically. + + Args: + model_name: Identifier for this model (shown in run name). + Defaults to MODEL_NAME env var or "model". + experiment_name: MLflow experiment name. + test_size: Fraction of data for the test split. + random_state: Random seed for train/test split. + log_stream: Whether to capture stdout/stderr into console.log. + fail_fast: If True (default), abort the run with RuntimeError + when dataset verification fails. + + Example: + from provenance import autolog + + @autolog(model_name="resnet_baseline_v1") + def train(run): + run.log_params({"learning_rate": 1e-3}) + model = build_model() + run.register_model(model) + optimizer = optim.SGD(model.parameters(), lr=1e-3) + run.register_optimizer(optimizer) + for epoch in range(50): + loss = train_epoch(...) + run.log_metrics({"train_loss": loss}, step=epoch) + run.save_model(model) + + if __name__ == "__main__": + train() + """ + + def decorator(func: Callable[..., None]) -> Callable[[], None]: + @wraps(func) + def wrapper(): + _run_autolog( + func=func, + model_name=model_name or os.environ.get("MODEL_NAME", "model"), + experiment_name=experiment_name, + test_size=test_size, + random_state=random_state, + log_stream=log_stream, + fail_fast=fail_fast, + ) + + return wrapper + + return decorator + + +def _run_autolog(func, model_name, experiment_name, test_size, random_state, log_stream, fail_fast=True): + """Core autolog logic.""" + _temp_dirs: list[str] = [] + + # ── Auto-detect everything ────────────────────────────── + user_run_id, user_tags = _lookup_user_run() + manifest_path, data_root = _detect_manifest() + dataset_run_id = _lookup_dataset_run(manifest_path) + git_commit, git_url, git_branch = _get_git_info() + hardware = _detect_hardware() + docker = _detect_docker() + + # ── Start MLflow run ─────────────────────────────────── + mlflow.set_experiment(experiment_name) + ts = datetime.now().strftime("%Y%m%d_%H%M%S") + run_name = f"Training_{model_name}_{ts}" + + mlflow_run = mlflow.start_run(run_name=run_name) + run_id = mlflow_run.info.run_id + + # ── 1. Tags (PROV relationships) ─────────────────────── + 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] + + if dataset_run_id: + tags["dataset_run_id"] = dataset_run_id + + tags.update({ + "git_commit": git_commit, + "git_url": git_url, + "git_branch": git_branch, + "prov_start_time": datetime.now(timezone.utc).isoformat(), + }) + mlflow.set_tags(tags) + + # ── 2. Params: hardware + docker + split config ──────── + all_params: dict[str, str | float | int] = { + "model_name": model_name, + **hardware, + **docker, + "split_test_size": test_size, + "split_random_state": random_state, + "split_stratified": True, + } + mlflow.log_params(all_params) + + # ── 3. Environment snapshot ──────────────────────────── + artifact_dir = f"_mlflow_env_{uuid.uuid4().hex[:8]}" + os.makedirs(artifact_dir, exist_ok=True) + _temp_dirs.append(artifact_dir) + frozen_requirements: str | None = None + try: + frozen_requirements = _snapshot_environment(artifact_dir) + mlflow.log_artifacts(artifact_dir, artifact_path="environment") + except Exception: + pass + + # ── 4. Train/test split from manifest ────────────────── + run_handle = _Run(run_id) + + if manifest_path and data_root: + train_samples, test_samples = _prepare_split( + manifest_path, data_root, test_size, random_state + ) + + run_handle.train_paths = [s["path"] for s in train_samples] + run_handle.test_paths = [s["path"] for s in test_samples] + run_handle.train_labels = [s["label"] for s in train_samples] + run_handle.test_labels = [s["label"] for s in test_samples] + + # Store full split data for embedding into provenance + run_handle._set_split(train_samples, test_samples) + + mlflow.log_params({ + "train_samples": len(train_samples), + "test_samples": len(test_samples), + "train_positive": sum(run_handle.train_labels), + "train_negative": len(run_handle.train_labels) - sum(run_handle.train_labels), + "test_positive": sum(run_handle.test_labels), + "test_negative": len(run_handle.test_labels) - sum(run_handle.test_labels), + }) + + split_dir = f"_mlflow_split_{uuid.uuid4().hex[:8]}" + os.makedirs(split_dir, exist_ok=True) + _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") + + # ── 4b. Verify dataset against Dataset_Registry ──── + verification = _verify_dataset(manifest_path, data_root, dataset_run_id) + for detail in verification["details"]: + print(f" [autolog] {detail}") + + # Log verification results as params + tags + mlflow.log_params({ + "dataset_verified": verification["verified"], + "dataset_manifest_hash_match": verification["manifest_hash_match"] is True, + "dataset_samples_hash_match": verification["samples_hash_match"] is True, + "dataset_files_missing": verification["files_missing"], + "dataset_files_total": verification["files_total"], + }) + if verification["verified"]: + tags["dataset_verification"] = "VERIFIED" + else: + tags["dataset_verification"] = "MISMATCH" + tags["dataset_verification_details"] = "; ".join(verification["details"]) + mlflow.set_tags(tags) + + # ── Hard abort on mismatch ───────────────────── + if fail_fast and not verification["verified"]: + mlflow.end_run(status='FAILED') + raise RuntimeError( + "Dataset verification failed — aborting training.\n" + + " ".join(" " + d for d in verification["details"]) + ) + + else: + print("[autolog] WARNING: No manifest.csv found — " + "train/test split not logged.") + verification = None + + # ── 5. Run the user's training function ──────────────── + try: + if log_stream: + with _StreamCapture(run_id): + func(run_handle) + else: + func(run_handle) + except Exception: + # Still log what we can before re-raising + raise + finally: + # ── Log model/optimizer/scheduler (if registered) ── + if run_handle._model is not None: + mlflow.log_params(_model_summary(run_handle._model)) + + if run_handle._optimizer_info is not None: + mlflow.log_params(run_handle._optimizer_info) + if run_handle._scheduler_info is not None: + mlflow.log_params(run_handle._scheduler_info) + + # ── Write self-contained provenance JSON ─────────── + # Embeds splits + requirements so the run is reproducible + # without needing MLflow at all. + try: + run_data = mlflow.get_run(run_id) + summary_dir = f"_mlflow_summary_{uuid.uuid4().hex[:8]}" + os.makedirs(summary_dir, exist_ok=True) + summary_path = os.path.join(summary_dir, "run_summary.json") + + # Source file hash for verification + source_file = getattr(func, "__wrapped__", func).__code__.co_filename + source_hash = _source_hash(source_file) if source_file else None + + summary = { + "model_name": model_name, + "params": dict(run_data.data.params), + "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_id": run_id, + "run_name": run_name, + "experiment_name": experiment_name, + + # ── Reproducibility: dataset splits ───────── + "split": { + "test_size": test_size, + "random_state": random_state, + "stratified": True, + "train_count": len(run_handle.train_paths), + "test_count": len(run_handle.test_paths), + "train": run_handle._split_data["train"] if run_handle._split_data else None, + "test": run_handle._split_data["test"] if run_handle._split_data else None, + } if run_handle._split_data else None, + + # ── Reproducibility: dataset verification ─── + "dataset_verification": verification if verification is not None else None, + + # ── Reproducibility: frozen environment ───── + "requirements": frozen_requirements, + + # ── Source verification ───────────────────── + "source": { + "file": source_file, + "sha256": source_hash, + "git_commit": git_commit, + "git_branch": git_branch, + "git_remote": 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) + + # ── Build & log PROV document ─────────────────── + prov_doc = _build_prov_document( + run_id=run_id, + run_name=run_name, + 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.") + }, + start_time_ms=mlflow_run.info.start_time, + end_time_ms=mlflow_run.info.end_time, + split_data={ + "test_size": test_size, + "random_state": random_state, + "train": run_handle._split_data["train"] if run_handle._split_data else None, + "test": run_handle._split_data["test"] if run_handle._split_data else None, + } if run_handle._split_data else None, + requirements=frozen_requirements, + verification=verification, + ) + + prov_dir = f"_mlflow_prov_{uuid.uuid4().hex[:8]}" + os.makedirs(prov_dir, exist_ok=True) + _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) + + print(f"\n[autolog] Complete → {run_id}") + except Exception as e: + print(f"[autolog] WARNING: Could not write provenance artifacts: {e}") + + # ── Clean up temp dirs ───────────────────────────── + for d in _temp_dirs: + shutil.rmtree(d, ignore_errors=True) From 96eb1645cef077bef4953d2ac4e141c450678f99 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Mon, 20 Jul 2026 22:49:43 +0200 Subject: [PATCH 03/34] fix: demo and readme --- .gitignore | 3 ++- README.md | 2 +- demo.py | 23 ++++++++++++----------- 3 files changed, 15 insertions(+), 13 deletions(-) diff --git a/.gitignore b/.gitignore index b16e130..a8d875e 100644 --- a/.gitignore +++ b/.gitignore @@ -172,4 +172,5 @@ cython_debug/ # Prov test_data -mlflow.db \ No newline at end of file +mlflow.db +mlartifacts \ No newline at end of file diff --git a/README.md b/README.md index f57ed4c..c66577a 100644 --- a/README.md +++ b/README.md @@ -32,7 +32,7 @@ mlflow ui --host 0.0.0.0 --port 5000 # → http://localhost:5000 The easiest way to see everything in action: ```bash -# Full pipeline — uploads to localhost:5000 +# Full pipeline — uploads to http://localhost:5000 python demo.py # Push to a different MLflow server diff --git a/demo.py b/demo.py index 8929644..31e5a36 100644 --- a/demo.py +++ b/demo.py @@ -294,8 +294,8 @@ def step_lightning(): import torch import lightning as pl - import mlflow from rationai.mlkit import Trainer, MLFlowLogger, MultiloaderLifecycle + from rationai.mlkit.provenance import autolog class _TinyModel(pl.LightningModule): def __init__(self): @@ -311,14 +311,15 @@ def training_step(self, batch, _): def configure_optimizers(self): return torch.optim.Adam(self.parameters(), lr=0.01) - mlflow.set_experiment("Demo_Lightning") - run = mlflow.start_run(run_name="Training_demo_lightning_model") - - try: + @autolog(model_name="demo_lightning_model", experiment_name="Demo_Lightning", fail_fast=False) + def train(run): model = _TinyModel() + run.register_model(model) + optimizer = torch.optim.Adam(model.parameters(), lr=0.01) + run.register_optimizer(optimizer) - logger = MLFlowLogger(run_id=run.info.run_id) + logger = MLFlowLogger(run_id=run._run_id) trainer = Trainer( logger=logger, max_epochs=2, @@ -332,9 +333,9 @@ def configure_optimizers(self): trainer.fit(model, torch.utils.data.DataLoader( torch.utils.data.TensorDataset(torch.randn(16, 8), torch.randint(0, 2, (16,))) )) - print(" ✅ Training complete — check MLflow for logs") - finally: - mlflow.end_run() + run.save_model(model) + + train() # ────────────────────────────────────────────── @@ -413,11 +414,11 @@ def main(): parser = argparse.ArgumentParser(description="rationai.mlkit — full pipeline demo") parser.add_argument("--test", action="store_true", help="Run unit tests instead") parser.add_argument("--uri", type=str, default=None, - help="MLflow tracking URI (default: local file store)") + help="MLflow tracking URI (default: http://localhost:5000)") args = parser.parse_args() import mlflow as _mlf - uri = args.uri or f"file:///tmp/mlkit_demo_mlruns_{os.getpid()}" + uri = args.uri or "http://localhost:5000" _mlf.set_tracking_uri(uri) print(f"[*] MLflow tracking URI: {uri}") From 25bc3e07f9ed4eb2f4e8963a305e0d683c031014 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Mon, 20 Jul 2026 22:53:18 +0200 Subject: [PATCH 04/34] feat: added dataset registrator to mlflow --- dataset_to_mlflow.py | 57 ++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 57 insertions(+) create mode 100644 dataset_to_mlflow.py diff --git a/dataset_to_mlflow.py b/dataset_to_mlflow.py new file mode 100644 index 0000000..f2eca60 --- /dev/null +++ b/dataset_to_mlflow.py @@ -0,0 +1,57 @@ +import os +import pandas as pd +import mlflow +from datetime import datetime + +def register_dataset_as_provenance(manifest_path, dataset_root, dataset_name, version): + """ + Zaregistruje dataset do MLflow. + Používá parametry pro indexaci a artefakty jako zdroj pravdy (PROV-O připraveno). + """ + mlflow.set_experiment("Dataset_Registry") + + with mlflow.start_run(run_name=f"Dataset_{dataset_name}_v{version}") as run: + # 1. Indexace v MLflow (pro rychlé hledání/filtrování) + mlflow.set_tag("dataset_name", dataset_name) + mlflow.set_tag("version", version) + + # 2. Metadata pro tracking + mlflow.log_param("dataset_root", dataset_root) + + # 3. Zpracování manifestu a obohacení o metadata (size, mtime) + df = pd.read_csv(manifest_path) + metadata_list = [] + + for path in df['wsi_path']: + full_path = os.path.join(dataset_root, path) if not os.path.isabs(path) else path + if os.path.exists(full_path): + stat = os.stat(full_path) + metadata_list.append({"file_size": stat.st_size, "last_modified": stat.st_mtime}) + else: + metadata_list.append({"file_size": -1, "last_modified": -1}) + + df_enriched = pd.concat([df, pd.DataFrame(metadata_list)], axis=1) + + # 4. Uložení artefaktu (Zlatý zdroj pravdy) + # Toto CSV budeš později skenovat pro tvorbu JSON-LD (PROV-O) + provenance_file = "dataset_provenance.csv" + df_enriched.to_csv(provenance_file, index=False) + mlflow.log_artifact(provenance_file, artifact_path="provenance") + + # 5. Uložení odkazu do tagu (velmi důležité pro automatizaci!) + # Uložíme si, kde v artefaktech to CSV leží + mlflow.set_tag("manifest_uri", f"runs:/{run.info.run_id}/provenance/{provenance_file}") + + print(f"Dataset '{dataset_name}' (v{version}) úspěšně zaregistrován.") + print(f"Run ID: {run.info.run_id}") + + os.remove(provenance_file) + +if __name__ == "__main__": + # Příklad použití pro tvůj dataset + register_dataset_as_provenance( + manifest_path="data/dummy_dataset_1/manifest.csv", + dataset_root="data/dummy_dataset_1", + dataset_name="pato_cohort_01", + version="1.0.0" + ) \ No newline at end of file From a0b3fbffc2fa7336d84a81f4714bb28879875f90 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 21:58:51 +0200 Subject: [PATCH 05/34] fix: refactoring, made repository structure more readable and removed unneeded files --- demo.py | 462 ------------------ dummy_dataset_create.py | 146 ------ rationai/__init__.py | 1 - rationai/mlkit/autolog.py | 108 +++- rationai/mlkit/lightning/__init__.py | 4 - rationai/mlkit/lightning/autolog.py | 106 ---- rationai/mlkit/lightning/with_cli_args.py | 57 --- .../mlkit/provenance/dataset_to_mlflow.py | 0 rationai/mlkit/{ => provenance}/provenance.py | 0 .../mlkit/provenance/user_to_mlflow.py | 0 rationai/mlkit/with_cli_args.py | 59 ++- tests/test_all.py | 378 -------------- 12 files changed, 163 insertions(+), 1158 deletions(-) delete mode 100644 demo.py delete mode 100644 dummy_dataset_create.py delete mode 100644 rationai/__init__.py delete mode 100644 rationai/mlkit/lightning/autolog.py delete mode 100644 rationai/mlkit/lightning/with_cli_args.py rename dataset_to_mlflow.py => rationai/mlkit/provenance/dataset_to_mlflow.py (100%) rename rationai/mlkit/{ => provenance}/provenance.py (100%) rename user_to_mlflow.py => rationai/mlkit/provenance/user_to_mlflow.py (100%) delete mode 100644 tests/test_all.py diff --git a/demo.py b/demo.py deleted file mode 100644 index 31e5a36..0000000 --- a/demo.py +++ /dev/null @@ -1,462 +0,0 @@ -#!/usr/bin/env python3 -""" -Demo script for rationai.mlkit — full end-to-end pipeline. - -Creates dummy data, registers it in Dataset_Registry, trains a model with -provenance tracking (fail_fast=True), logs everything to MLflow, and prints -a summary of what was uploaded. - -Run: - python demo.py # full pipeline (local file store) - python demo.py --uri http://... # custom MLflow server -""" - -import argparse -import json -import os -import sys -from pathlib import Path - -os.environ["MLFLOW_ALLOW_FILE_STORE"] = "true" - -import logging -logging.getLogger("root").setLevel(logging.WARNING) -logging.getLogger("mlflow").setLevel(logging.ERROR) - - -# ────────────────────────────────────────────── -# Helpers -# ────────────────────────────────────────────── - -def sep(title=""): - print(f"\n{'='*60}") - if title: - print(f" {title}") - print(f"{'='*60}") - - -def sub(title): - print(f"\n{'─'*60}") - print(f" {title}") - print(f"{'─'*60}") - - -# ────────────────────────────────────────────── -# 1. Create dummy datasets -# ────────────────────────────────────────────── - -def step_create_datasets(n=2, wsis_per_ds=10): - sub("1. Creating dummy pathology datasets") - - from dummy_dataset_create import create_dummy_datasets - create_dummy_datasets( - num_datasets=n, - wsis_per_dataset=wsis_per_ds, - data_dir=Path("test_data"), - seed=42, - img_size=64, - clean=True, - ) - - data_dir = Path("test_data") - for d in sorted(data_dir.iterdir()): - if d.is_dir(): - manifest = d / "manifest.csv" - n_rows = sum(1 for _ in open(manifest)) - 1 if manifest.exists() else "?" - print(f" → {d.name}/ ({n_rows} samples)") - - -# ────────────────────────────────────────────── -# 2. Register dataset(s) in Dataset_Registry -# ────────────────────────────────────────────── - -def step_register_datasets(): - sub("2. Registering datasets in Dataset_Registry") - - from rationai.mlkit.provenance import register_dataset - - data_dir = Path("test_data") - for d in sorted(data_dir.iterdir()): - if d.is_dir(): - manifest = d / "manifest.csv" - if not manifest.exists(): - continue - ds_name = d.name - run_id = register_dataset( - dataset_dir=str(d), - dataset_name=ds_name, - ) - print(f" ✅ {ds_name}: registered (run={run_id[:8]})") - - -# ────────────────────────────────────────────── -# 3. Stream capture demo -# ────────────────────────────────────────────── - -def step_stream_capture(): - sub("3. StreamCapture + StreamModifier") - - import io - from rationai.mlkit import StreamCapture, StreamLogger, StreamModifier - - # StreamCapture — captures stdout into a logger - class _Buf(StreamLogger): - def __init__(self): - self._buf = io.StringIO() - def log_stream(self, text: str): - self._buf.write(text) - def get_value(self): - return self._buf.getvalue() - - logger = _Buf() - with StreamCapture(logger, streams=(sys.stdout,)): - print("Hello from stdout!") - print("\033[92mThis is green (ANSI)\033[0m") - - captured = logger.get_value() - has_ansi = "\x1b" in captured - print(f" Captured text : {captured.strip()!r}") - print(f" ANSI preserved: {'✅ (raw capture)' if has_ansi else '❌'}") - - # StreamModifier — injects side-effect logic into a stream's write - buf = io.StringIO() - side_log = [] - modifier = StreamModifier(stream=buf, id=42) - modifier.set_write(lambda s, iid: side_log.append(f"[{iid}] {s}")) - buf.write("hello") - modifier.teardown() - print(f" Side log : {side_log}") - print(f" Original buf : {buf.getvalue()!r}") - print(" ✅ StreamModifier works (side-effect injected before original write)") - - -# ────────────────────────────────────────────── -# 4. Metrics demo (AggregatedMetricCollection) -# — skipped if rationai.masks is not installed -# ────────────────────────────────────────────── - -def step_metrics(): - sub("4. AggregatedMetricCollection — tile → slide aggregation") - - try: - import torch - from torchmetrics import Accuracy - from rationai.mlkit import AggregatedMetricCollection, MaxAggregator, MeanAggregator - except ModuleNotFoundError as e: - if "rationai.masks" in str(e): - print(" ⊘ SKIPPED — rationai.masks (private dep) not installed") - return - raise - - preds = torch.tensor([0.1, 0.8, 0.3, 0.9]) - targets = torch.tensor([0, 1, 0, 1]) - keys = ["slide_A", "slide_A", "slide_B", "slide_B"] - - for name, agg in [("MaxAggregator", MaxAggregator()), ("MeanAggregator", MeanAggregator())]: - mc = AggregatedMetricCollection( - metrics={"accuracy": Accuracy(task="binary")}, - aggregator=agg, - ) - mc.update(preds, targets, keys) - result = mc.compute() - print(f" {name:20s}: accuracy = {result['accuracy'].item():.4f}") - - -# ────────────────────────────────────────────── -# 5. NestedMetricCollection demo -# — skipped if rationai.masks is not installed -# ────────────────────────────────────────────── - -def step_nested_metrics(): - sub("5. NestedMetricCollection — per-slide multiclass metrics") - - try: - import torch - from torchmetrics import Accuracy, Precision - from rationai.mlkit import NestedMetricCollection - except ModuleNotFoundError as e: - if "rationai.masks" in str(e): - print(" ⊘ SKIPPED — rationai.masks (private dep) not installed") - return - raise - - 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"], - ) - - preds = torch.tensor([[0.7, 0.2, 0.1], [0.1, 0.1, 0.8], [0.2, 0.6, 0.2]]) - targets = torch.tensor([0, 2, 1]) - keys = ["slide_1", "slide_1", "slide_2"] - - metrics.update(preds, targets, keys) - result = metrics.compute() - - print(f" Slides: {result['slide']}") - for k, v in result.items(): - if k != "slide": - print(f" {k:15s}: {v}") - - -# ────────────────────────────────────────────── -# 6. StratifiedBatchSampler demo -# ────────────────────────────────────────────── - -def step_sampler(): - sub("6. StratifiedBatchSampler — balanced class batches") - - from rationai.mlkit import StratifiedBatchSampler, PDMStratifiedBatchSampler - import pandas as pd - - # List-of-lists sampler - sampler = StratifiedBatchSampler( - data_indices=[[0, 1, 2, 3], [4, 5, 6, 7]], - batch_size=4, - ) - for i, batch in enumerate(sampler): - print(f" Batch {i}: {batch}") - print(f" Total batches: {len(sampler)}") - - # DataFrame-based sampler - df = pd.DataFrame({ - "idx": list(range(8)), - "label": [0, 0, 0, 1, 1, 1, 1, 0], - }) - pdm_sampler = PDMStratifiedBatchSampler(data=df, stratify_by="label", batch_size=4) - pdm_batches = list(pdm_sampler) - print(f" PDM sampler batches: {len(pdm_batches)}") - - -# ────────────────────────────────────────────── -# 7. Full training run with provenance autolog -# (fail_fast=True — dataset verification must pass) -# ────────────────────────────────────────────── - -def step_provenance(): - sub("7. Provenance — full training run, logged to MLflow (fail_fast=True)") - - import torch - import torch.nn as nn - import torch.optim as optim - from torch.utils.data import Dataset, DataLoader - from rationai.mlkit.provenance import autolog - - class _DummyDS(Dataset): - def __len__(self): - return 32 - def __getitem__(self, idx): - return torch.randn(64), torch.randint(0, 2, (1,)).item() - - @autolog(model_name="demo_model", experiment_name="Demo_Experiment", fail_fast=True) - def train(run): - model = nn.Sequential(nn.Linear(64, 32), nn.ReLU(), nn.Linear(32, 2)) - run.register_model(model) - - optimizer = optim.Adam(model.parameters(), lr=0.01) - run.register_optimizer(optimizer) - - loader = DataLoader(_DummyDS(), batch_size=8) - criterion = nn.CrossEntropyLoss() - - for epoch in range(1, 5): - model.train() - total_loss = 0 - for bx, by in loader: - optimizer.zero_grad() - loss = criterion(model(bx), by) - loss.backward() - optimizer.step() - total_loss += loss.item() - - avg = total_loss / len(loader) - run.log_metrics({"train_loss": avg}, step=epoch) - print(f" Epoch {epoch}: loss={avg:.4f}") - - run.save_model(model) - - train() - - import mlflow as _mlf - while _mlf.active_run(): - _mlf.end_run() - - -# ────────────────────────────────────────────── -# 8. Lightning integration demo -# ────────────────────────────────────────────── - -def step_lightning(): - sub("8. Lightning — Trainer + MLFlowLogger (actual training)") - - import torch - import lightning as pl - from rationai.mlkit import Trainer, MLFlowLogger, MultiloaderLifecycle - from rationai.mlkit.provenance import autolog - - class _TinyModel(pl.LightningModule): - def __init__(self): - super().__init__() - self.net = torch.nn.Linear(8, 2) - def forward(self, x): - return self.net(x) - def training_step(self, batch, _): - x = torch.randn(4, 8, device=self.device) - loss = self.net(x).sum() - self.log("train_loss", loss) - return loss - def configure_optimizers(self): - return torch.optim.Adam(self.parameters(), lr=0.01) - - @autolog(model_name="demo_lightning_model", experiment_name="Demo_Lightning", fail_fast=False) - def train(run): - model = _TinyModel() - run.register_model(model) - - optimizer = torch.optim.Adam(model.parameters(), lr=0.01) - run.register_optimizer(optimizer) - - logger = MLFlowLogger(run_id=run._run_id) - trainer = Trainer( - logger=logger, - max_epochs=2, - enable_checkpointing=False, - enable_progress_bar=False, - enable_model_summary=False, - log_every_n_steps=1, - ) - - print(" Training a tiny Lightning model...") - trainer.fit(model, torch.utils.data.DataLoader( - torch.utils.data.TensorDataset(torch.randn(16, 8), torch.randint(0, 2, (16,))) - )) - run.save_model(model) - - train() - - -# ────────────────────────────────────────────── -# 9. Summary — list all runs & artifacts in MLflow -# ────────────────────────────────────────────── - -def step_summary(): - sub("9. MLflow Summary") - - import mlflow - - tracking_uri = mlflow.get_tracking_uri() - print(f" Tracking URI: {tracking_uri}") - - client = mlflow.MlflowClient() - - for exp_name in ["Dataset_Registry", "Demo_Experiment", "Demo_Lightning"]: - exp = client.get_experiment_by_name(exp_name) - if not exp: - continue - runs = client.search_runs( - experiment_ids=[exp.experiment_id], - order_by=["attributes.start_time desc"], - ) - print(f"\n Experiment: {exp_name} ({len(runs)} run(s))") - - for r in runs[:5]: - params = dict(r.data.params) if r.data.params else {} - metrics = dict(r.data.metrics) if r.data.metrics else {} - artifacts = [a.path for a in client.list_artifacts(r.info.run_id)] - - print(f"\n Run: {r.info.run_name or r.info.run_id[:8]}") - print(f" Status : {r.info.status}") - if params: - print(f" Params : {params}") - if metrics: - print(f" Metrics: {metrics}") - - # Group artifacts by folder - folders = {} - for a in artifacts: - folder = a.split("/")[0] if "/" in a else "" - folders.setdefault(folder, []).append(a) - for folder, files in sorted(folders.items()): - print(f" [{folder}] {', '.join(os.path.basename(f) for f in files)}") - - # Check provenance summary - prov_path = None - for a in artifacts: - if "run_summary.json" in a: - prov_path = a - break - if prov_path: - local = client.download_artifacts(r.info.run_id, prov_path) - with open(local) as f: - summary = json.load(f) - print(f" Provenance keys: {', '.join(summary.keys())}") - - if "dataset_verification" in summary: - v = summary["dataset_verification"] - status = "✅" if v.get("verified") else "❌" - details = [d for d in v.get("details", []) if "FAILED" in d or "verified" in d] - print(f" Dataset verify: {status} {details[0] if details else '—'}") - - # Local file store hints - if tracking_uri.startswith("file://"): - db_path = Path(tracking_uri.replace("file://", "")) - print(f"\n 📂 Local data at: {db_path.resolve()}") - - -# ────────────────────────────────────────────── -# Main -# ────────────────────────────────────────────── - -def main(): - parser = argparse.ArgumentParser(description="rationai.mlkit — full pipeline demo") - parser.add_argument("--test", action="store_true", help="Run unit tests instead") - parser.add_argument("--uri", type=str, default=None, - help="MLflow tracking URI (default: http://localhost:5000)") - args = parser.parse_args() - - import mlflow as _mlf - uri = args.uri or "http://localhost:5000" - _mlf.set_tracking_uri(uri) - print(f"[*] MLflow tracking URI: {uri}") - - # Verify connectivity - try: - _client = _mlf.MlflowClient() - _ = _client.get_experiment_by_name("__ping__") - print(f"[*] Connected ✅\n") - except Exception as e: - print(f"[!] Warning: Could not connect to MLflow server: {e}\n") - - if args.test: - sys.path.insert(0, str(Path(__file__).resolve().parent)) - from tests.test_all import main as test_main - test_main() - return - - sep("rationai.mlkit — End-to-End Demo") - - # Step 1 & 2: create + register datasets (required for fail_fast=True) - step_create_datasets(n=2, wsis_per_ds=10) - step_register_datasets() - - # Steps 3-6: component demos - step_stream_capture() - step_metrics() - step_nested_metrics() - step_sampler() - - # Steps 7-8: training runs (fail_fast=True — verification must pass) - step_provenance() - step_lightning() - - # Step 9: summary - step_summary() - - sep("✅ Demo complete — all runs are in MLflow!") - - -if __name__ == "__main__": - main() \ No newline at end of file diff --git a/dummy_dataset_create.py b/dummy_dataset_create.py deleted file mode 100644 index 5254607..0000000 --- a/dummy_dataset_create.py +++ /dev/null @@ -1,146 +0,0 @@ -""" -Create dummy pathology datasets for local testing. - -Generates N datasets, each with M fake WSI images (small random TIFFs) -and a manifest.csv matching the real data format: - - patient_id,wsi_path,cancer - PAT_001,wsis/PAT_001.tiff,0 - -Usage examples: - # Default: 2 datasets x 50 WSIs each - python dummy_dataset_create.py - - # 3 datasets with 20 WSIs each - python dummy_dataset_create.py --datasets 3 --wsis-per-dataset 20 - - # 1 dataset with 100 WSIs (replaces existing data/) - python dummy_dataset_create.py --datasets 1 --wsis-per-dataset 100 - - # Keep existing datasets, just add more - python dummy_dataset_create.py --datasets 1 --wsis-per-dataset 30 --no-clean -""" - -import argparse -import csv -import io -import os -import shutil -import uuid -from pathlib import Path - -import numpy as np -from PIL import Image - - -def _parse_args(): - p = argparse.ArgumentParser( - description="Create dummy pathology datasets for local testing.", - ) - p.add_argument( - "--datasets", "-d", - type=int, default=2, - help="Number of dataset folders to create (default: 2)", - ) - p.add_argument( - "--wsis-per-dataset", "-w", - type=int, default=50, - help="Number of WSI images per dataset (default: 50)", - ) - p.add_argument( - "--data-dir", - type=str, default="test_data", - help="Parent directory for the datasets (default: test_data/)", - ) - p.add_argument( - "--seed", - type=int, default=42, - help="Random seed for reproducibility (default: 42)", - ) - p.add_argument( - "--no-clean", - action="store_true", - help="Skip removing existing dummy_dataset_* folders before creating new ones", - ) - p.add_argument( - "--img-size", - type=int, default=128, - help="Size of each dummy WSI image in pixels (default: 128x128)", - ) - return p.parse_args() - - -def _clean_existing(data_dir: Path): - """Remove old dummy_dataset_* folders from data/.""" - for entry in sorted(data_dir.iterdir()): - if entry.is_dir() and entry.name.startswith("dummy_dataset_"): - shutil.rmtree(entry) - print(f" Removed existing {entry.name}/") - - -def _generate_wsi(img_size: int, rng: np.random.Generator) -> bytes: - """Generate a small random TIFF image.""" - arr = rng.integers(0, 256, (img_size, img_size, 3), dtype=np.uint8) - img = Image.fromarray(arr) - buf = io.BytesIO() - img.save(buf, format="TIFF") - return buf.getvalue() - - -def create_dummy_datasets( - num_datasets: int = 2, - wsis_per_dataset: int = 50, - data_dir: Path = Path("test_data"), - seed: int = 42, - img_size: int = 128, - clean: bool = True, -): - rng = np.random.default_rng(seed) - - if clean and data_dir.exists(): - _clean_existing(data_dir) - - for ds_idx in range(num_datasets): - ds_name = f"dummy_dataset_{ds_idx + 1}" - ds_path = data_dir / ds_name - wsis_path = ds_path / "wsis" - wsis_path.mkdir(parents=True, exist_ok=True) - - rows = [] - for i in range(wsis_per_dataset): - patient_id = f"PAT_{uuid.uuid4().hex[:8]}" - tiff_name = f"{patient_id}.tiff" - cancer = int(rng.random() > 0.5) - wsi_path = wsis_path / tiff_name - - wsi_path.write_bytes(_generate_wsi(img_size, rng)) - rows.append((patient_id, str(wsi_path.relative_to(ds_path)), cancer)) - - manifest_path = ds_path / "manifest.csv" - with open(manifest_path, "w", newline="") as f: - writer = csv.writer(f) - writer.writerow(["patient_id", "wsi_path", "cancer"]) - writer.writerows(rows) - - print(f" Created {ds_name}/ ({len(rows)} samples)") - - -def main(): - args = _parse_args() - data_dir = Path(args.data_dir) - data_dir.mkdir(parents=True, exist_ok=True) - - print(f"[*] Creating {args.datasets} dataset(s) with {args.wsis_per_dataset} WSIs each...") - create_dummy_datasets( - num_datasets=args.datasets, - wsis_per_dataset=args.wsis_per_dataset, - data_dir=data_dir, - seed=args.seed, - img_size=args.img_size, - clean=not args.no_clean, - ) - print(f"[*] Done. Data in: {data_dir.resolve()}") - - -if __name__ == "__main__": - main() diff --git a/rationai/__init__.py b/rationai/__init__.py deleted file mode 100644 index 444020e..0000000 --- a/rationai/__init__.py +++ /dev/null @@ -1 +0,0 @@ -# rationai namespace package diff --git a/rationai/mlkit/autolog.py b/rationai/mlkit/autolog.py index dac19da..d86506d 100644 --- a/rationai/mlkit/autolog.py +++ b/rationai/mlkit/autolog.py @@ -1,2 +1,106 @@ -from rationai.mlkit.lightning.autolog import * # noqa: F401, F403 -from rationai.mlkit.lightning.autolog import autolog # noqa: F401 +import logging +import os +import tempfile +from collections.abc import Callable +from functools import partial, wraps +from pathlib import Path +from typing import overload + +import hydra +from hydra.core.hydra_config import HydraConfig +from omegaconf import DictConfig, OmegaConf + +from rationai.mlkit.lightning.loggers import MLFlowLogger +from rationai.mlkit.stream import StreamCapture + + +log = logging.getLogger(__name__) + + +WrapperT = Callable[[DictConfig], None] +FunctionT = Callable[[DictConfig, MLFlowLogger], None] + + +@overload +def autolog( + func: FunctionT, + *, + log_config: bool = True, + log_stream: bool = True, + log_hyperparams: bool = True, +) -> WrapperT: ... + + +@overload +def autolog( + func: None = None, + *, + log_config: bool = True, + log_stream: bool = True, + log_hyperparams: bool = True, +) -> Callable[[FunctionT], WrapperT]: ... + + +def autolog( + func: FunctionT | None = None, + *, + log_config: bool = True, + log_stream: bool = True, + log_hyperparams: bool = True, +) -> WrapperT | Callable[[FunctionT], WrapperT]: + """Decorator for automatic logging. + + Args: + func: The function to decorate + log_config: Whether to log the hydra configuration files + log_stream: Whether to log the std streams + log_hyperparams: Whether to log the hyperparameters defined in the + config.metadata.hyperparams. + """ + if func is None: + return partial(autolog, log_config=log_config, log_stream=log_stream) + + @wraps(func) + def wrapper(config: DictConfig) -> None: + logger: MLFlowLogger = hydra.utils.instantiate(config.logger) + + if log_config: + _log_config(config, logger) + + if ( + log_hyperparams + and hasattr(config, "metadata") + and hasattr(config.metadata, "hyperparams") + ): + logger.log_hyperparams(config.metadata.hyperparams) + + if log_stream: + with StreamCapture(logger): + return func(config, logger) + + return func(config, logger) + + return wrapper + + +def _log_config(config: DictConfig, logger: MLFlowLogger) -> None: + """Logs the hydra config.""" + with tempfile.TemporaryDirectory(dir=os.getcwd()) as tmp_dir_str: + tmp_dir = Path(tmp_dir_str) + with open(tmp_dir / "hydra.yaml", "w", encoding="utf-8") as file: + OmegaConf.save(HydraConfig.get(), file) + + with open(tmp_dir / "config.yaml", "w", encoding="utf-8") as file: + OmegaConf.save(config, file) + + with open(tmp_dir / "config-resolved.yaml", "w", encoding="utf-8") as file: + OmegaConf.save(config, file, resolve=True) + + logger.log_artifacts(tmp_dir_str, "configs") + + +def __getattr__(name: str): + if name == "autolog_provenance": + from rationai.mlkit.provenance.provenance import autolog as _autolog + return _autolog + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/rationai/mlkit/lightning/__init__.py b/rationai/mlkit/lightning/__init__.py index 90342bb..12abb54 100644 --- a/rationai/mlkit/lightning/__init__.py +++ b/rationai/mlkit/lightning/__init__.py @@ -1,13 +1,9 @@ -from rationai.mlkit.lightning.autolog import autolog from rationai.mlkit.lightning.callbacks import MultiloaderLifecycle from rationai.mlkit.lightning.loggers import MLFlowLogger from rationai.mlkit.lightning.trainer import Trainer -from rationai.mlkit.lightning.with_cli_args import with_cli_args __all__ = [ "Trainer", "MLFlowLogger", "MultiloaderLifecycle", - "autolog", - "with_cli_args", ] diff --git a/rationai/mlkit/lightning/autolog.py b/rationai/mlkit/lightning/autolog.py deleted file mode 100644 index df2d2db..0000000 --- a/rationai/mlkit/lightning/autolog.py +++ /dev/null @@ -1,106 +0,0 @@ -import logging -import os -import tempfile -from collections.abc import Callable -from functools import partial, wraps -from pathlib import Path -from typing import overload - -import hydra -from hydra.core.hydra_config import HydraConfig -from omegaconf import DictConfig, OmegaConf - -from rationai.mlkit.lightning.loggers import MLFlowLogger -from rationai.mlkit.stream import StreamCapture - - -log = logging.getLogger(__name__) - - -WrapperT = Callable[[DictConfig], None] -FunctionT = Callable[[DictConfig, MLFlowLogger], None] - - -@overload -def autolog( - func: FunctionT, - *, - log_config: bool = True, - log_stream: bool = True, - log_hyperparams: bool = True, -) -> WrapperT: ... - - -@overload -def autolog( - func: None = None, - *, - log_config: bool = True, - log_stream: bool = True, - log_hyperparams: bool = True, -) -> Callable[[FunctionT], WrapperT]: ... - - -def autolog( - func: FunctionT | None = None, - *, - log_config: bool = True, - log_stream: bool = True, - log_hyperparams: bool = True, -) -> WrapperT | Callable[[FunctionT], WrapperT]: - """Decorator for automatic logging. - - Args: - func: The function to decorate - log_config: Whether to log the hydra configuration files - log_stream: Whether to log the std streams - log_hyperparams: Whether to log the hyperparameters defined in the - config.metadata.hyperparams. - """ - if func is None: - return partial(autolog, log_config=log_config, log_stream=log_stream) - - @wraps(func) - def wrapper(config: DictConfig) -> None: - logger: MLFlowLogger = hydra.utils.instantiate(config.logger) - - if log_config: - _log_config(config, logger) - - if ( - log_hyperparams - and hasattr(config, "metadata") - and hasattr(config.metadata, "hyperparams") - ): - logger.log_hyperparams(config.metadata.hyperparams) - - if log_stream: - with StreamCapture(logger): - return func(config, logger) - - return func(config, logger) - - return wrapper - - -def _log_config(config: DictConfig, logger: MLFlowLogger) -> None: - """Logs the hydra config.""" - with tempfile.TemporaryDirectory(dir=os.getcwd()) as tmp_dir_str: - tmp_dir = Path(tmp_dir_str) - with open(tmp_dir / "hydra.yaml", "w", encoding="utf-8") as file: - OmegaConf.save(HydraConfig.get(), file) - - with open(tmp_dir / "config.yaml", "w", encoding="utf-8") as file: - OmegaConf.save(config, file) - - with open(tmp_dir / "config-resolved.yaml", "w", encoding="utf-8") as file: - OmegaConf.save(config, file, resolve=True) - - logger.log_artifacts(tmp_dir_str, "configs") - - -def __getattr__(name: str): - if name == "autolog_provenance": - from rationai.mlkit.provenance import autolog as _autolog - return _autolog - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/rationai/mlkit/lightning/with_cli_args.py b/rationai/mlkit/lightning/with_cli_args.py deleted file mode 100644 index ebddc0e..0000000 --- a/rationai/mlkit/lightning/with_cli_args.py +++ /dev/null @@ -1,57 +0,0 @@ -import sys -from collections.abc import Callable -from functools import wraps -from typing import Any - - -def with_cli_args( - defaults: list[str] | None = None, overrides: list[str] | None = None -) -> Callable[[Callable[..., Any]], Callable[..., Any]]: - """Decorator to injects arguments into sys.argv. - - Args: - defaults: Arguments injected AFTER script name but BEFORE user args. - (Acts as defaults: User can override these). - overrides: Arguments injected AFTER user args. - (Acts as overrides: Forces value, User cannot override). - - Returns: - A decorator that modifies sys.argv for the duration of the decorated function. - - Examples: - >>> from rationai.mlkit import autolog, with_cli_args, MLFlowLogger - >>> from omegaconf import DictConfig - >>> import hydra - - >>> @with_cli_args(["+preprocessing=qc"]) - >>> @hydra.main(config_path="../configs", config_name="preprocessing", version_base=None) - >>> @autolog - >>> def main(config: DictConfig, logger: MLFlowLogger) -> None: - >>> pass - """ - prepend = defaults or [] - append = overrides or [] - - def decorator(func: Callable[..., Any]) -> Callable[..., Any]: - @wraps(func) - def wrapper(*args: Any, **kwargs: Any) -> Any: - # 1. Save original state - original_argv = sys.argv[:] - - # 2. Deconstruct existing argv - # sys.argv[0] is the script name - script_name = [sys.argv[0]] - user_provided_args = sys.argv[1:] - - # 3. Reconstruct: [Script] + [Start] + [User] + [End] - sys.argv = script_name + prepend + user_provided_args + append - - try: - return func(*args, **kwargs) - finally: - # 4. Restore original state guarantees safety - sys.argv = original_argv - - return wrapper - - return decorator diff --git a/dataset_to_mlflow.py b/rationai/mlkit/provenance/dataset_to_mlflow.py similarity index 100% rename from dataset_to_mlflow.py rename to rationai/mlkit/provenance/dataset_to_mlflow.py diff --git a/rationai/mlkit/provenance.py b/rationai/mlkit/provenance/provenance.py similarity index 100% rename from rationai/mlkit/provenance.py rename to rationai/mlkit/provenance/provenance.py diff --git a/user_to_mlflow.py b/rationai/mlkit/provenance/user_to_mlflow.py similarity index 100% rename from user_to_mlflow.py rename to rationai/mlkit/provenance/user_to_mlflow.py diff --git a/rationai/mlkit/with_cli_args.py b/rationai/mlkit/with_cli_args.py index ba04856..ebddc0e 100644 --- a/rationai/mlkit/with_cli_args.py +++ b/rationai/mlkit/with_cli_args.py @@ -1,2 +1,57 @@ -from rationai.mlkit.lightning.with_cli_args import * # noqa: F401, F403 -from rationai.mlkit.lightning.with_cli_args import with_cli_args # noqa: F401 +import sys +from collections.abc import Callable +from functools import wraps +from typing import Any + + +def with_cli_args( + defaults: list[str] | None = None, overrides: list[str] | None = None +) -> Callable[[Callable[..., Any]], Callable[..., Any]]: + """Decorator to injects arguments into sys.argv. + + Args: + defaults: Arguments injected AFTER script name but BEFORE user args. + (Acts as defaults: User can override these). + overrides: Arguments injected AFTER user args. + (Acts as overrides: Forces value, User cannot override). + + Returns: + A decorator that modifies sys.argv for the duration of the decorated function. + + Examples: + >>> from rationai.mlkit import autolog, with_cli_args, MLFlowLogger + >>> from omegaconf import DictConfig + >>> import hydra + + >>> @with_cli_args(["+preprocessing=qc"]) + >>> @hydra.main(config_path="../configs", config_name="preprocessing", version_base=None) + >>> @autolog + >>> def main(config: DictConfig, logger: MLFlowLogger) -> None: + >>> pass + """ + prepend = defaults or [] + append = overrides or [] + + def decorator(func: Callable[..., Any]) -> Callable[..., Any]: + @wraps(func) + def wrapper(*args: Any, **kwargs: Any) -> Any: + # 1. Save original state + original_argv = sys.argv[:] + + # 2. Deconstruct existing argv + # sys.argv[0] is the script name + script_name = [sys.argv[0]] + user_provided_args = sys.argv[1:] + + # 3. Reconstruct: [Script] + [Start] + [User] + [End] + sys.argv = script_name + prepend + user_provided_args + append + + try: + return func(*args, **kwargs) + finally: + # 4. Restore original state guarantees safety + sys.argv = original_argv + + return wrapper + + return decorator diff --git a/tests/test_all.py b/tests/test_all.py deleted file mode 100644 index 30a34ad..0000000 --- a/tests/test_all.py +++ /dev/null @@ -1,378 +0,0 @@ -""" -End-to-end test suite for rationai.mlkit. - -Exercises every major component: - 1. Stream capture + StreamModifier - 2. AggregatedMetricCollection + aggregators (with torchmetrics) - 3. NestedMetricCollection - 4. StratifiedBatchSampler / PDMStratifiedBatchSampler - 5. Lightning: Trainer, MLFlowLogger, MultiloaderLifecycle, autolog, with_cli_args - 6. Provenance @autolog (full training run + provenance artifact verification) - -Run: python tests/test_all.py -""" - -import sys -import io -import os -import json -import tempfile -from pathlib import Path - -os.environ["MLFLOW_ALLOW_FILE_STORE"] = "true" - -sys.path.insert(0, str(Path(__file__).resolve().parent.parent)) - -import torch -import mlflow - - -def section(title): - print(f"\n{'='*60}") - print(f" {title}") - print(f"{'='*60}\n") - - -# ────────────────────────────────────────────── -# 1. Stream Capture + StreamModifier -# ────────────────────────────────────────────── - -def test_stream_capture(): - section("1. StreamCapture + StreamModifier") - - from rationai.mlkit import StreamCapture, StreamLogger, StreamModifier - - class _TestLogger(StreamLogger): - def __init__(self): - self._buffer = io.StringIO() - def log_stream(self, text: str): - self._buffer.write(text) - def get_value(self): - return self._buffer.getvalue() - - logger = _TestLogger() - with StreamCapture(logger, streams=(sys.stdout,)): - print("hello") - print("\033[31mred text\033[0m") - - captured = logger.get_value() - assert "hello" in captured - assert "\x1b" in captured, "StreamCapture preserves ANSI codes (raw capture)" - print(f" Captured: {captured.strip()!r}") - print(" ✅ StreamCapture works (captures raw output including ANSI)") - - # StreamModifier — wraps a stream's write to inject side-effect - # logic before the original write fires. The callback receives (text, id) - # and can e.g. log or tag the output elsewhere. - buf = io.StringIO() - side_log = [] - modifier = StreamModifier(stream=buf, id=42) - modifier.set_write(lambda s, iid: side_log.append(f"[{iid}] {s}")) - buf.write("hello") - assert "[42] hello" in side_log, f"Side effect not called: {side_log}" - assert buf.getvalue() == "hello", f"Original write not called: {buf.getvalue()!r}" - modifier.teardown() - print(f" Side log: {side_log}") - print(f" Original buf: {buf.getvalue()!r}") - print(" ✅ StreamModifier works") - - -# ────────────────────────────────────────────── -# 2. AggregatedMetricCollection + aggregators -# ────────────────────────────────────────────── - -def test_aggregated_metrics(): - section("2. AggregatedMetricCollection") - - try: - from torchmetrics import Accuracy - from rationai.mlkit import ( - AggregatedMetricCollection, - MaxAggregator, - MeanAggregator, - ) - except ModuleNotFoundError as e: - if "rationai.masks" in str(e): - raise # handled by main as skip - raise - - preds = torch.tensor([0.1, 0.8, 0.3, 0.9]) - targets = torch.tensor([0, 1, 0, 1]) - keys = ["slide_A", "slide_A", "slide_B", "slide_B"] - - for name, agg in [ - ("MaxAggregator", MaxAggregator()), - ("MeanAggregator", MeanAggregator()), - ]: - mc = AggregatedMetricCollection( - metrics={"accuracy": Accuracy(task="binary")}, - aggregator=agg, - ) - mc.update(preds, targets, keys) - result = mc.compute() - print(f" {name:20s}: accuracy={result['accuracy'].item():.4f}") - - print(" ✅ AggregatedMetricCollection works") - - -# ────────────────────────────────────────────── -# 3. NestedMetricCollection -# ────────────────────────────────────────────── - -def test_nested_metrics(): - section("3. NestedMetricCollection") - - from torchmetrics import Accuracy, Precision - from rationai.mlkit import NestedMetricCollection - - 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"], - ) - - preds = torch.tensor([[0.7, 0.2, 0.1], [0.1, 0.1, 0.8], [0.2, 0.6, 0.2]]) - targets = torch.tensor([0, 2, 1]) - keys = ["slide_1", "slide_1", "slide_2"] - - metrics.update(preds, targets, keys) - result = metrics.compute() - - assert "slide" in result - print(f" Slides: {result['slide']}") - for k, v in result.items(): - if k != "slide": - print(f" {k:15s}: {v}") - print(" ✅ NestedMetricCollection works") - - -# ────────────────────────────────────────────── -# 4. StratifiedBatchSampler / PDMStratifiedBatchSampler -# ────────────────────────────────────────────── - -def test_samplers(): - section("4. Samplers") - - from rationai.mlkit import StratifiedBatchSampler - - sampler = StratifiedBatchSampler( - data_indices=[[0, 1, 2, 3], [4, 5, 6, 7]], - batch_size=4, - ) - batches = list(sampler) - print(f" StratifiedBatchSampler: {len(batches)} batches") - for i, batch in enumerate(batches): - print(f" Batch {i}: {batch}") - assert len(batches) == 2 - - # PDMStratifiedBatchSampler requires a DataFrame - from rationai.mlkit import PDMStratifiedBatchSampler - import pandas as pd - - df = pd.DataFrame({ - "idx": list(range(8)), - "label": [0, 0, 0, 1, 1, 1, 1, 0], - }) - pdm_sampler = PDMStratifiedBatchSampler( - data=df, - stratify_by="label", - batch_size=4, - ) - pdm_batches = list(pdm_sampler) - print(f" PDMStratifiedBatchSampler: {len(pdm_batches)} batches") - assert len(pdm_batches) >= 1 - - print(" ✅ Samplers work") - - -# ────────────────────────────────────────────── -# 5. Lightning imports + basic functionality -# ────────────────────────────────────────────── - -def test_lightning(): - section("5. Lightning (Trainer, MLFlowLogger, MultiloaderLifecycle, autolog, with_cli_args)") - - from rationai.mlkit import Trainer, MLFlowLogger, MultiloaderLifecycle, autolog, with_cli_args - print(f" Trainer: {Trainer}") - print(f" MLFlowLogger: {MLFlowLogger}") - print(f" MultiloaderLifecycle: {MultiloaderLifecycle}") - print(f" autolog: {autolog}") - print(f" with_cli_args: {with_cli_args}") - - import lightning as pl - - class _TinyModel(pl.LightningModule): - def __init__(self): - super().__init__() - self.net = torch.nn.Linear(8, 2) - def forward(self, x): - return self.net(x) - def training_step(self, batch, _): - x = torch.randn(4, 8, device=self.device) - loss = self.net(x).sum() - self.log("train_loss", loss) - return loss - def configure_optimizers(self): - return torch.optim.Adam(self.parameters(), lr=0.01) - - model = _TinyModel() - trainer = Trainer( - max_epochs=1, - enable_checkpointing=False, - enable_progress_bar=False, - enable_model_summary=False, - logger=False, - ) - ds = torch.utils.data.TensorDataset(torch.randn(8, 8), torch.randint(0, 2, (8,))) - trainer.fit(model, torch.utils.data.DataLoader(ds)) - print(" ✅ Lightning Trainer works") - - -# ────────────────────────────────────────────── -# 6. Provenance @autolog -# ────────────────────────────────────────────── - -def test_provenance(): - section("6. Provenance (@autolog + artifact verification)") - - import torch.nn as nn - import torch.optim as optim - from torch.utils.data import Dataset, DataLoader - from rationai.mlkit.provenance import autolog - - class _DummyDS(Dataset): - def __len__(self): - return 16 - def __getitem__(self, idx): - return torch.randn(32), torch.randint(0, 2, (1,)).item() - - with tempfile.TemporaryDirectory() as tmpdir: - mlflow.set_tracking_uri(f"file://{tmpdir}/mlruns") - - @autolog(model_name="test_model", experiment_name="Test_Provenance", fail_fast=False) - def train(run): - model = nn.Sequential(nn.Linear(32, 16), nn.ReLU(), nn.Linear(16, 2)) - run.register_model(model) - - optimizer = optim.Adam(model.parameters(), lr=0.01) - run.register_optimizer(optimizer) - - loader = DataLoader(_DummyDS(), batch_size=4) - criterion = nn.CrossEntropyLoss() - - for epoch in range(1, 3): - model.train() - total_loss = 0 - for bx, by in loader: - optimizer.zero_grad() - loss = criterion(model(bx), by) - loss.backward() - optimizer.step() - total_loss += loss.item() - - avg = total_loss / len(loader) - run.log_metrics({"train_loss": avg}, step=epoch) - print(f" Epoch {epoch}: loss={avg:.4f}") - - run.save_model(model) - - train() - - client = mlflow.MlflowClient() - exp = client.get_experiment_by_name("Test_Provenance") - runs = client.search_runs(experiment_ids=[exp.experiment_id]) - assert len(runs) >= 1, "Expected at least 1 run" - - r = runs[0] - artifacts = [a.path for a in client.list_artifacts(r.info.run_id)] - print(f" Artifacts: {artifacts}") - - prov_path = None - for a in artifacts: - if "provenance" in a.lower(): - prov_path = a.rstrip("/") - break - - if prov_path: - local_dir = client.download_artifacts(r.info.run_id, prov_path) - # download_artifacts returns a directory; find run_summary.json inside - import glob as _glob - summary_files = _glob.glob(f"{local_dir}/**/run_summary.json", recursive=True) - if not summary_files: - summary_files = [f for f in os.listdir(local_dir) if f.endswith(".json")] - summary_files = [os.path.join(local_dir, f) for f in summary_files] if summary_files else [] - if summary_files: - with open(summary_files[0]) as f: - summary = json.load(f) - print(f" Provenance keys: {list(summary.keys())}") - assert "model_name" in summary or "params" in summary or "metrics" in summary - else: - print(f" (Provenance dir contents: {os.listdir(local_dir)})") - else: - print(" (No provenance artifact found — checking run params/metrics)") - assert r.data.params, "Expected params in run" - - print(" ✅ Provenance @autolog works") - - -# ────────────────────────────────────────────── -# Backward compatibility -# ────────────────────────────────────────────── - -def test_backward_compat(): - section("7. Backward Compatibility (old import paths)") - - from rationai.mlkit import Trainer, autolog, with_cli_args - print(" from rationai.mlkit import Trainer, autolog, with_cli_args: ✅") - - from rationai.mlkit.autolog import autolog as _autolog - print(" from rationai.mlkit.autolog import autolog: ✅") - - from rationai.mlkit.with_cli_args import with_cli_args as _wca - print(" from rationai.mlkit.with_cli_args import with_cli_args: ✅") - - from rationai.mlkit.lightning.autolog import autolog as _la - print(" from rationai.mlkit.lightning.autolog import autolog: ✅") - - -# ────────────────────────────────────────────── -# Main -# ────────────────────────────────────────────── - -def main(): - section("rationai.mlkit — Test Suite") - - tests = [ - ("Stream Capture", test_stream_capture), - ("Aggregated Metrics", test_aggregated_metrics), - ("Nested Metrics", test_nested_metrics), - ("Samplers", test_samplers), - ("Lightning", test_lightning), - ("Provenance", test_provenance), - ("Backward Compat", test_backward_compat), - ] - - passed = failed = skipped = 0 - for name, fn in tests: - try: - fn() - passed += 1 - except ModuleNotFoundError as e: - print(f" ⊘ SKIPPED {name}: {e}") - skipped += 1 - except Exception as e: - print(f" ❌ {name}: {e}") - import traceback - traceback.print_exc() - failed += 1 - - section(f"Results: {passed} passed, {failed} failed, {skipped} skipped") - if failed: - sys.exit(1) - - -if __name__ == "__main__": - main() From 61e21f004b654d1307e34a7039248e3fb1273748 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 21:59:21 +0200 Subject: [PATCH 06/34] feat: made folder for provenance --- rationai/mlkit/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index 9ce9e25..49dcdb1 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -1,7 +1,7 @@ """rationai.mlkit — ML toolkit with provenance tracking.""" from rationai.mlkit.stream import StreamCapture, StreamLogger, StreamModifier -from rationai.mlkit.provenance import ( +from rationai.mlkit.provenance.provenance import ( autolog as provenance_autolog, register_dataset, ) From f2ff9149394e73d5584403830e0820eaf3fe5b8a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 22:05:13 +0200 Subject: [PATCH 07/34] fix: trying to minimize file overwriting --- rationai/mlkit/lightning/__init__.py | 9 ++------- rationai/mlkit/lightning/loggers/mlflow.py | 4 +--- rationai/mlkit/stream/__init__.py | 3 +-- 3 files changed, 4 insertions(+), 12 deletions(-) diff --git a/rationai/mlkit/lightning/__init__.py b/rationai/mlkit/lightning/__init__.py index 12abb54..db3f03b 100644 --- a/rationai/mlkit/lightning/__init__.py +++ b/rationai/mlkit/lightning/__init__.py @@ -1,9 +1,4 @@ -from rationai.mlkit.lightning.callbacks import MultiloaderLifecycle -from rationai.mlkit.lightning.loggers import MLFlowLogger + from rationai.mlkit.lightning.trainer import Trainer -__all__ = [ - "Trainer", - "MLFlowLogger", - "MultiloaderLifecycle", -] +__all__ = ["Trainer"] diff --git a/rationai/mlkit/lightning/loggers/mlflow.py b/rationai/mlkit/lightning/loggers/mlflow.py index ce33f30..0f56026 100644 --- a/rationai/mlkit/lightning/loggers/mlflow.py +++ b/rationai/mlkit/lightning/loggers/mlflow.py @@ -59,9 +59,7 @@ def __init__( def experiment(self) -> MlflowClient: if not self._initialized: exp = super().experiment - # Only start a run if none is already active (e.g. from @autolog) - if not mlflow.active_run(): - mlflow.start_run(self.run_id, log_system_metrics=self.log_system_metrics) + mlflow.start_run(self.run_id, log_system_metrics=self.log_system_metrics) return exp return super().experiment diff --git a/rationai/mlkit/stream/__init__.py b/rationai/mlkit/stream/__init__.py index 905a205..6fcd614 100644 --- a/rationai/mlkit/stream/__init__.py +++ b/rationai/mlkit/stream/__init__.py @@ -1,5 +1,4 @@ from rationai.mlkit.stream.stream_capture import StreamCapture from rationai.mlkit.stream.stream_logger import StreamLogger -from rationai.mlkit.stream.stream_modifier import StreamModifier -__all__ = ["StreamCapture", "StreamLogger", "StreamModifier"] +__all__ = ["StreamCapture", "StreamLogger"] From ad7228a4c0f626335148ddf9c932b0d62c43be56 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 22:06:29 +0200 Subject: [PATCH 08/34] fix: too much newlines --- rationai/mlkit/lightning/loggers/__init__.py | 1 - 1 file changed, 1 deletion(-) diff --git a/rationai/mlkit/lightning/loggers/__init__.py b/rationai/mlkit/lightning/loggers/__init__.py index 77459b5..5d4ea67 100644 --- a/rationai/mlkit/lightning/loggers/__init__.py +++ b/rationai/mlkit/lightning/loggers/__init__.py @@ -1,4 +1,3 @@ from rationai.mlkit.lightning.loggers.mlflow import MLFlowLogger - __all__ = ["MLFlowLogger"] From 5a9406dd61e7e707af9afbe4513466877558352c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 22:08:03 +0200 Subject: [PATCH 09/34] fix: trying to fix line again --- rationai/mlkit/lightning/loggers/__init__.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/rationai/mlkit/lightning/loggers/__init__.py b/rationai/mlkit/lightning/loggers/__init__.py index 5d4ea67..f09a2a4 100644 --- a/rationai/mlkit/lightning/loggers/__init__.py +++ b/rationai/mlkit/lightning/loggers/__init__.py @@ -1,3 +1,4 @@ from rationai.mlkit.lightning.loggers.mlflow import MLFlowLogger -__all__ = ["MLFlowLogger"] + +__all__ = ["MLFlowLogger"] \ No newline at end of file From 1a4495b5a2cebe63fd19ef45f31519d7bc329238 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 22:10:42 +0200 Subject: [PATCH 10/34] fix: tried again --- rationai/mlkit/lightning/__init__.py | 1 - rationai/mlkit/lightning/loggers/__init__.py | 2 +- rationai/mlkit/stream/__init__.py | 1 + 3 files changed, 2 insertions(+), 2 deletions(-) diff --git a/rationai/mlkit/lightning/__init__.py b/rationai/mlkit/lightning/__init__.py index db3f03b..2378b6b 100644 --- a/rationai/mlkit/lightning/__init__.py +++ b/rationai/mlkit/lightning/__init__.py @@ -1,4 +1,3 @@ - from rationai.mlkit.lightning.trainer import Trainer __all__ = ["Trainer"] diff --git a/rationai/mlkit/lightning/loggers/__init__.py b/rationai/mlkit/lightning/loggers/__init__.py index f09a2a4..77459b5 100644 --- a/rationai/mlkit/lightning/loggers/__init__.py +++ b/rationai/mlkit/lightning/loggers/__init__.py @@ -1,4 +1,4 @@ from rationai.mlkit.lightning.loggers.mlflow import MLFlowLogger -__all__ = ["MLFlowLogger"] \ No newline at end of file +__all__ = ["MLFlowLogger"] diff --git a/rationai/mlkit/stream/__init__.py b/rationai/mlkit/stream/__init__.py index 6fcd614..eeafc52 100644 --- a/rationai/mlkit/stream/__init__.py +++ b/rationai/mlkit/stream/__init__.py @@ -1,4 +1,5 @@ from rationai.mlkit.stream.stream_capture import StreamCapture from rationai.mlkit.stream.stream_logger import StreamLogger + __all__ = ["StreamCapture", "StreamLogger"] From 28d5cf77f1d4841163dd94e16fc74e120da43e3a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 22:11:24 +0200 Subject: [PATCH 11/34] fix: now it is perfect --- rationai/mlkit/lightning/__init__.py | 1 + 1 file changed, 1 insertion(+) diff --git a/rationai/mlkit/lightning/__init__.py b/rationai/mlkit/lightning/__init__.py index 2378b6b..7fc4eb6 100644 --- a/rationai/mlkit/lightning/__init__.py +++ b/rationai/mlkit/lightning/__init__.py @@ -1,3 +1,4 @@ from rationai.mlkit.lightning.trainer import Trainer + __all__ = ["Trainer"] From 710f59728c716e6afcca6def90fde30870d74351 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 22:56:10 +0200 Subject: [PATCH 12/34] feat: provenance code refactor --- .gitignore | 3 +- rationai/mlkit/__init__.py | 9 +- rationai/mlkit/provenance/__init__.py | 37 +++ .../mlkit/provenance/dataset_to_mlflow.py | 57 ---- rationai/mlkit/provenance/provenance.py | 269 +---------------- rationai/mlkit/provenance/register_dataset.py | 279 ++++++++++++++++++ .../{user_to_mlflow.py => register_user.py} | 0 7 files changed, 332 insertions(+), 322 deletions(-) create mode 100644 rationai/mlkit/provenance/__init__.py delete mode 100644 rationai/mlkit/provenance/dataset_to_mlflow.py create mode 100644 rationai/mlkit/provenance/register_dataset.py rename rationai/mlkit/provenance/{user_to_mlflow.py => register_user.py} (100%) diff --git a/.gitignore b/.gitignore index a8d875e..6591481 100644 --- a/.gitignore +++ b/.gitignore @@ -173,4 +173,5 @@ cython_debug/ # Prov test_data mlflow.db -mlartifacts \ No newline at end of file +mlartifacts +example_provenance.py \ No newline at end of file diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index 49dcdb1..fc7168e 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -1,15 +1,12 @@ """rationai.mlkit — ML toolkit with provenance tracking.""" -from rationai.mlkit.stream import StreamCapture, StreamLogger, StreamModifier -from rationai.mlkit.provenance.provenance import ( - autolog as provenance_autolog, - register_dataset, -) +from rationai.mlkit.stream import StreamCapture, StreamLogger +from rationai.mlkit.provenance.provenance import autolog as provenance_autolog +from rationai.mlkit.provenance.register_dataset import register_dataset __all__ = [ "StreamCapture", "StreamLogger", - "StreamModifier", "AggregatedMetricCollection", "Aggregator", "MaxAggregator", diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py new file mode 100644 index 0000000..3f08927 --- /dev/null +++ b/rationai/mlkit/provenance/__init__.py @@ -0,0 +1,37 @@ +"""Provenance tracking — PROV-O-aware logging to MLflow. + +Submodules: + provenance – @autolog decorator, internal helpers + register_dataset – register_dataset (hash-based) and + register_dataset_as_provenance (legacy CSV-based) + register_user – register_new_user + +MLflow tracking URI defaults to ``http://localhost:5000``. Override with +the ``MLFLOW_TRACKING_URI`` environment variable. +""" + +from __future__ import annotations + +import os + +# ── Set default tracking URI before any mlflow import runs ─────────── +if "MLFLOW_TRACKING_URI" not in os.environ: + os.environ["MLFLOW_TRACKING_URI"] = "http://localhost:5000" + +# Now safe to import – all child modules will pick up the env var +from .provenance import autolog # noqa: E402 +from .register_dataset import ( # noqa: E402 + register_dataset, + register_dataset_as_provenance, +) +from .register_user import register_new_user # noqa: E402 + +__all__ = [ + # Core provenance + "autolog", + # Dataset registration + "register_dataset", + "register_dataset_as_provenance", + # User registration + "register_new_user", +] diff --git a/rationai/mlkit/provenance/dataset_to_mlflow.py b/rationai/mlkit/provenance/dataset_to_mlflow.py deleted file mode 100644 index f2eca60..0000000 --- a/rationai/mlkit/provenance/dataset_to_mlflow.py +++ /dev/null @@ -1,57 +0,0 @@ -import os -import pandas as pd -import mlflow -from datetime import datetime - -def register_dataset_as_provenance(manifest_path, dataset_root, dataset_name, version): - """ - Zaregistruje dataset do MLflow. - Používá parametry pro indexaci a artefakty jako zdroj pravdy (PROV-O připraveno). - """ - mlflow.set_experiment("Dataset_Registry") - - with mlflow.start_run(run_name=f"Dataset_{dataset_name}_v{version}") as run: - # 1. Indexace v MLflow (pro rychlé hledání/filtrování) - mlflow.set_tag("dataset_name", dataset_name) - mlflow.set_tag("version", version) - - # 2. Metadata pro tracking - mlflow.log_param("dataset_root", dataset_root) - - # 3. Zpracování manifestu a obohacení o metadata (size, mtime) - df = pd.read_csv(manifest_path) - metadata_list = [] - - for path in df['wsi_path']: - full_path = os.path.join(dataset_root, path) if not os.path.isabs(path) else path - if os.path.exists(full_path): - stat = os.stat(full_path) - metadata_list.append({"file_size": stat.st_size, "last_modified": stat.st_mtime}) - else: - metadata_list.append({"file_size": -1, "last_modified": -1}) - - df_enriched = pd.concat([df, pd.DataFrame(metadata_list)], axis=1) - - # 4. Uložení artefaktu (Zlatý zdroj pravdy) - # Toto CSV budeš později skenovat pro tvorbu JSON-LD (PROV-O) - provenance_file = "dataset_provenance.csv" - df_enriched.to_csv(provenance_file, index=False) - mlflow.log_artifact(provenance_file, artifact_path="provenance") - - # 5. Uložení odkazu do tagu (velmi důležité pro automatizaci!) - # Uložíme si, kde v artefaktech to CSV leží - mlflow.set_tag("manifest_uri", f"runs:/{run.info.run_id}/provenance/{provenance_file}") - - print(f"Dataset '{dataset_name}' (v{version}) úspěšně zaregistrován.") - print(f"Run ID: {run.info.run_id}") - - os.remove(provenance_file) - -if __name__ == "__main__": - # Příklad použití pro tvůj dataset - register_dataset_as_provenance( - manifest_path="data/dummy_dataset_1/manifest.csv", - dataset_root="data/dummy_dataset_1", - dataset_name="pato_cohort_01", - version="1.0.0" - ) \ No newline at end of file diff --git a/rationai/mlkit/provenance/provenance.py b/rationai/mlkit/provenance/provenance.py index 801e213..06be884 100644 --- a/rationai/mlkit/provenance/provenance.py +++ b/rationai/mlkit/provenance/provenance.py @@ -96,9 +96,12 @@ def _get_git_info(): return commit, remote, branch -def _lookup_experiment(name): - exp = mlflow.get_experiment_by_name(name) - return exp.experiment_id if exp else None +# ── Dataset helpers live in register_dataset.py ────────────────────── +from .register_dataset import ( # noqa: F401 + _lookup_dataset_run, + _lookup_experiment, + _verify_dataset, +) def _lookup_user_run(): @@ -131,253 +134,7 @@ def _lookup_user_run(): return row.run_id, dict(run.data.tags) -def _lookup_dataset_run(manifest_path: str | None = None): - """Return the Dataset_Registry run that matches the detected manifest. - - If *manifest_path* is given, looks for a registration run whose - ``manifest_hash`` tag matches the SHA-256 of that file. Falls back - to the latest Dataset_Registry run if no match is found (so that - runs 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 runs_df.empty: - return None - - # If we have a manifest, try to match by hash first - if manifest_path and os.path.isfile(manifest_path): - current_hash = _hash_manifest(manifest_path) - for _, row in runs_df.iterrows(): - reg_hash = row.get("tags.manifest_hash", "") - if reg_hash == current_hash: - return row["run_id"] - - # Fallback: latest run (backward compat when only one dataset registered) - return runs_df.iloc[0]["run_id"] - - -def _hash_manifest(manifest_path: str) -> str: - """Compute SHA-256 hash of a manifest CSV file.""" - h = hashlib.sha256() - with open(manifest_path, "rb") as f: - for chunk in iter(lambda: f.read(8192), b""): - h.update(chunk) - return h.hexdigest() - - -def _hash_samples(samples: list[dict]) -> str: - """Compute a deterministic hash over the set of samples (path+label pairs). - - Order-independent: sorts by path before hashing so that any reordering - of the manifest doesn't change the fingerprint. - """ - h = hashlib.sha256() - for s in sorted(samples, key=lambda x: x["path"]): - h.update(f"{s['path']}:{s['label']}\n".encode()) - return h.hexdigest() - - -def register_dataset( - dataset_dir: str, - dataset_name: str | None = None, - version: str = "1.0.0", - experiment_name: str = "Dataset_Registry", -): - """Register a dataset in MLflow's Dataset_Registry experiment. - - Computes a SHA-256 hash of the manifest and stores it as a tag so that - future training runs can verify they're using the exact same 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.mlflow.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) - - # Compute hashes - manifest_hash = _hash_manifest(manifest_path) - - # Read manifest and compute sample-level hash - df = pd.read_csv(manifest_path) - samples = [] - for _, row in df.iterrows(): - rel = row["wsi_path"] - full = os.path.join(dataset_dir, rel) if not os.path.isabs(rel) else rel - samples.append({"path": full, "label": int(row["cancer"])}) - samples_hash = _hash_samples(samples) - - # File-level hashes for each WSI - file_hashes = {} - for s in samples: - if os.path.isfile(s["path"]): - fh = hashlib.sha256() - with open(s["path"], "rb") as f: - for chunk in iter(lambda: f.read(8192), b""): - fh.update(chunk) - file_hashes[os.path.basename(s["path"])] = fh.hexdigest()[:16] - else: - file_hashes[os.path.basename(s["path"])] = "MISSING" - - # Register in MLflow - mlflow.set_experiment(experiment_name) - run = mlflow.start_run(run_name=f"Dataset_{dataset_name}_{version}") - run_id = run.info.run_id - - 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, - "manifest_hash": manifest_hash, - "samples_hash": samples_hash, - "file_hashes": json.dumps(file_hashes), - }) - - # Save manifest hash as an artifact for offline verification - prov_dir = f"_dataset_prov_{uuid.uuid4().hex[:8]}" - os.makedirs(prov_dir, exist_ok=True) - prov_path = os.path.join(prov_dir, "dataset_provenance.json") - with open(prov_path, "w") as f: - json.dump({ - "dataset_name": dataset_name, - "version": version, - "dataset_root": dataset_dir, - "manifest_hash": manifest_hash, - "samples_hash": samples_hash, - "file_hashes": file_hashes, - "num_samples": len(samples), - }, f, indent=2) - mlflow.log_artifact(prov_path, artifact_path="provenance") - shutil.rmtree(prov_dir, ignore_errors=True) - - mlflow.end_run() - print(f" [register_dataset] {dataset_name} v{version} → run_id={run_id}") - print(f" manifest_hash : {manifest_hash[:16]}…") - print(f" samples_hash : {samples_hash[:16]}…") - return run_id - - -def _verify_dataset( - manifest_path: str, - data_root: str, - dataset_run_id: str | None, -) -> dict: - """Verify the current dataset against the registered version in MLflow. - - Checks: - 1. Manifest hash matches (file-level integrity) - 2. Samples hash matches (content-level integrity, order-independent) - 3. All WSI files exist on disk - - Returns a dict with verification results. - """ - result: dict = { - "verified": False, - "dataset_run_id": dataset_run_id, - "manifest_hash_match": None, - "samples_hash_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 hashes - try: - reg_run = mlflow.get_run(dataset_run_id) - reg_tags = reg_run.data.tags - reg_manifest_hash = reg_tags.get("manifest_hash", "") - reg_samples_hash = reg_tags.get("samples_hash", "") - reg_file_hashes_str = reg_tags.get("file_hashes", "") - reg_file_hashes = json.loads(reg_file_hashes_str) if reg_file_hashes_str else {} - except Exception as e: - result["details"].append(f"Failed to fetch Dataset_Registry run: {e}") - return result - - # Compute current hashes - curr_manifest_hash = _hash_manifest(manifest_path) - result["manifest_hash_match"] = curr_manifest_hash == reg_manifest_hash - - if not result["manifest_hash_match"]: - result["details"].append( - f"Manifest hash mismatch: " - f"current={curr_manifest_hash[:16]}… registered={reg_manifest_hash[:16]}…" - ) - - # Compute samples hash - df = pd.read_csv(manifest_path) - samples = [] - 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"])}) - curr_samples_hash = _hash_samples(samples) - result["samples_hash_match"] = curr_samples_hash == reg_samples_hash - - if not result["samples_hash_match"]: - result["details"].append( - f"Samples hash mismatch: " - f"current={curr_samples_hash[:16]}… registered={reg_samples_hash[:16]}…" - ) - - # Check file existence - missing = 0 - for s in samples: - if not os.path.isfile(s["path"]): - missing += 1 - result["files_total"] = len(samples) - result["files_missing"] = missing - - if missing > 0: - result["details"].append(f"{missing}/{len(samples)} WSI files missing on disk") - - # Overall verdict - result["verified"] = ( - result["manifest_hash_match"] - and result["samples_hash_match"] - and missing == 0 - ) - - if result["verified"]: - result["details"].append("✅ Dataset verified — matches registered version") - else: - result["details"].append("❌ Dataset verification FAILED") - - return result +# ── Dataset verification lives in register_dataset.py ───────────────── def _detect_hardware(): @@ -881,12 +638,9 @@ def _blank_rel_id() -> str: if verification: meta_entity["gen:dataset_verified"] = [str(verification.get("verified", False))] meta_entity["gen:dataset_run_id"] = [verification.get("dataset_run_id", "")] - mh = verification.get("manifest_hash_match") - if mh is not None: - meta_entity["gen:manifest_hash_match"] = [str(mh)] - sh = verification.get("samples_hash_match") - if sh is not None: - meta_entity["gen:samples_hash_match"] = [str(sh)] + 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)] @@ -1237,8 +991,7 @@ def _run_autolog(func, model_name, experiment_name, test_size, random_state, log # Log verification results as params + tags mlflow.log_params({ "dataset_verified": verification["verified"], - "dataset_manifest_hash_match": verification["manifest_hash_match"] is True, - "dataset_samples_hash_match": verification["samples_hash_match"] is True, + "dataset_file_sizes_match": verification["file_sizes_match"] is True, "dataset_files_missing": verification["files_missing"], "dataset_files_total": verification["files_total"], }) diff --git a/rationai/mlkit/provenance/register_dataset.py b/rationai/mlkit/provenance/register_dataset.py new file mode 100644 index 0000000..18ab764 --- /dev/null +++ b/rationai/mlkit/provenance/register_dataset.py @@ -0,0 +1,279 @@ +"""Dataset registration, verification, and legacy CSV-based paths.""" + +from __future__ import annotations + +import json +import os +import shutil +import uuid + +import mlflow +import pandas as pd + + +# ── Internal helpers ──────────────────────────────────────────────────── + +def _lookup_experiment(name): + """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(manifest_path: str | None = None) -> 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 runs_df.empty: + return None + + return runs_df.iloc[0]["run_id"] + + +def _verify_dataset( + manifest_path: str, + data_root: str, + dataset_run_id: str | None, +) -> dict: + """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 = { + "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 + + # Read manifest and check current file sizes + df = pd.read_csv(manifest_path) + samples = [] + curr_file_sizes = {} + 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"])}) + basename = os.path.basename(full) + if os.path.isfile(full): + curr_file_sizes[basename] = os.stat(full).st_size + else: + curr_file_sizes[basename] = -1 # missing + + # Compare sizes + result["file_sizes_match"] = curr_file_sizes == reg_file_sizes + + 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 "") + ) + + # Check file existence + 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") + + # Overall verdict + 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", +): + """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) + + # Read manifest and collect per-file metadata (size, last_modified) + df = pd.read_csv(manifest_path) + samples = [] + file_sizes = {} + file_mtimes = {} + + for _, row in df.iterrows(): + rel = row["wsi_path"] + full = os.path.join(dataset_dir, rel) if not os.path.isabs(rel) else rel + samples.append({"path": full, "label": int(row["cancer"])}) + + basename = os.path.basename(full) + if os.path.isfile(full): + st = os.stat(full) + file_sizes[basename] = st.st_size + file_mtimes[basename] = st.st_mtime + else: + file_sizes[basename] = -1 + file_mtimes[basename] = -1 + + # Register in MLflow + mlflow.set_experiment(experiment_name) + run = mlflow.start_run(run_name=f"Dataset_{dataset_name}_{version}") + run_id = run.info.run_id + + 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), + }) + + # Save provenance JSON as an artifact for offline verification + prov_dir = f"_dataset_prov_{uuid.uuid4().hex[:8]}" + os.makedirs(prov_dir, exist_ok=True) + prov_path = os.path.join(prov_dir, "dataset_provenance.json") + with open(prov_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(prov_path, artifact_path="provenance") + shutil.rmtree(prov_dir, ignore_errors=True) + + mlflow.end_run() + print(f" [register_dataset] {dataset_name} v{version} → run_id={run_id}") + return run_id + + +# ── Legacy CSV-based registration (backward compat) ───────────────────── + +def register_dataset_as_provenance(manifest_path, dataset_root, dataset_name, version): + """Legacy CSV-based dataset registration. + + Stores an enriched manifest as an artifact with a ``manifest_uri`` tag. + Does NOT set ``manifest_hash`` / ``samples_hash`` tags — use + :func:`register_dataset` instead for hash-based verification support. + """ + mlflow.set_experiment("Dataset_Registry") + + with mlflow.start_run(run_name=f"Dataset_{dataset_name}_v{version}") as run: + # 1. Indexace v MLflow (pro rychlé hledání/filtrování) + mlflow.set_tag("dataset_name", dataset_name) + mlflow.set_tag("version", version) + + # 2. Metadata pro tracking + mlflow.log_param("dataset_root", dataset_root) + + # 3. Zpracování manifestu a obohacení o metadata (size, mtime) + df = pd.read_csv(manifest_path) + metadata_list = [] + + for path in df["wsi_path"]: + full_path = os.path.join(dataset_root, path) if not os.path.isabs(path) else path + if os.path.exists(full_path): + stat = os.stat(full_path) + metadata_list.append({"file_size": stat.st_size, "last_modified": stat.st_mtime}) + else: + metadata_list.append({"file_size": -1, "last_modified": -1}) + + df_enriched = pd.concat([df, pd.DataFrame(metadata_list)], axis=1) + + # 4. Uložení artefaktu (Zlatý zdroj pravdy) + provenance_file = "dataset_provenance.csv" + df_enriched.to_csv(provenance_file, index=False) + mlflow.log_artifact(provenance_file, artifact_path="provenance") + + # 5. Uložení odkazu do tagu (velmi důležité pro automatizaci!) + mlflow.set_tag("manifest_uri", f"runs:/{run.info.run_id}/provenance/{provenance_file}") + + print(f"Dataset '{dataset_name}' (v{version}) úspěšně zaregistrován.") + print(f"Run ID: {run.info.run_id}") + + os.remove(provenance_file) + + +if __name__ == "__main__": + # Příklad použití pro tvůj dataset + register_dataset_as_provenance( + manifest_path="data/dummy_dataset_1/manifest.csv", + dataset_root="data/dummy_dataset_1", + dataset_name="pato_cohort_01", + version="1.0.0" + ) diff --git a/rationai/mlkit/provenance/user_to_mlflow.py b/rationai/mlkit/provenance/register_user.py similarity index 100% rename from rationai/mlkit/provenance/user_to_mlflow.py rename to rationai/mlkit/provenance/register_user.py From 528d4f9a6f11b2d21eb8b14d2f705d50505064a9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 23:21:27 +0200 Subject: [PATCH 13/34] feat: removed provenance files and added them to lightning callbacks WIP --- rationai/mlkit/__init__.py | 7 +- rationai/mlkit/autolog.py | 7 - .../mlkit/lightning/callbacks/__init__.py | 6 +- .../callbacks/dataset_verification.py | 76 ++ .../callbacks}/provenance.py | 856 +++++++----------- rationai/mlkit/provenance/__init__.py | 24 +- rationai/mlkit/provenance/register_dataset.py | 50 + 7 files changed, 503 insertions(+), 523 deletions(-) create mode 100644 rationai/mlkit/lightning/callbacks/dataset_verification.py rename rationai/mlkit/{provenance => lightning/callbacks}/provenance.py (57%) diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index fc7168e..a907ff0 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -1,7 +1,6 @@ """rationai.mlkit — ML toolkit with provenance tracking.""" from rationai.mlkit.stream import StreamCapture, StreamLogger -from rationai.mlkit.provenance.provenance import autolog as provenance_autolog from rationai.mlkit.provenance.register_dataset import register_dataset __all__ = [ @@ -22,9 +21,9 @@ "Trainer", "MLFlowLogger", "MultiloaderLifecycle", + "ProvenanceCallback", "autolog", "with_cli_args", - "provenance_autolog", "register_dataset", ] @@ -35,6 +34,10 @@ def __getattr__(name): _mod = importlib.import_module("rationai.mlkit.lightning") return getattr(_mod, name) + if name == "ProvenanceCallback": + from rationai.mlkit.lightning.callbacks import ProvenanceCallback + return ProvenanceCallback + if name in ("AggregatedMetricCollection", "Aggregator", "MaxAggregator", "MeanAggregator", "MeanPoolMaxAggregator", "TopKAggregator", "NestedMetricCollection", "LazyMetricDict"): diff --git a/rationai/mlkit/autolog.py b/rationai/mlkit/autolog.py index d86506d..372e03a 100644 --- a/rationai/mlkit/autolog.py +++ b/rationai/mlkit/autolog.py @@ -97,10 +97,3 @@ def _log_config(config: DictConfig, logger: MLFlowLogger) -> None: OmegaConf.save(config, file, resolve=True) logger.log_artifacts(tmp_dir_str, "configs") - - -def __getattr__(name: str): - if name == "autolog_provenance": - from rationai.mlkit.provenance.provenance import autolog as _autolog - return _autolog - 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..b1f5529 100644 --- a/rationai/mlkit/lightning/callbacks/__init__.py +++ b/rationai/mlkit/lightning/callbacks/__init__.py @@ -1,6 +1,10 @@ +from rationai.mlkit.lightning.callbacks.dataset_verification import ( + DatasetVerificationCallback, +) from rationai.mlkit.lightning.callbacks.multiloader_lifecycle import ( MultiloaderLifecycle, ) +from rationai.mlkit.lightning.callbacks.provenance import ProvenanceCallback -__all__ = ["MultiloaderLifecycle"] +__all__ = ["DatasetVerificationCallback", "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..1df9043 --- /dev/null +++ b/rationai/mlkit/lightning/callbacks/dataset_verification.py @@ -0,0 +1,76 @@ +"""Lightning callback that verifies the dataset against MLflow on trainer start.""" + +from __future__ import annotations + +import mlflow + +from lightning.pytorch.callbacks import Callback + + +class DatasetVerificationCallback(Callback): + """Run dataset verification once at the start of training. + + Auto-detects ``manifest.csv`` under ``data/``, looks up the latest + ``Dataset_Registry`` run, and checks that per-file sizes still match. + + Logs verification results as MLflow params so they appear on the run + page alongside metrics and artifacts. + + Example:: + + from rationai.mlkit.lightning.callbacks import DatasetVerificationCallback + + trainer = Trainer( + callbacks=[DatasetVerificationCallback()], + logger=MLFlowLogger(...), + ) + """ + + def __init__(self, manifest_path: str | None = None): + self._manifest_path = manifest_path + self._done = False + + def on_fit_start(self, trainer, pl_module): # noqa: ARG002 + if self._done: + return + self._done = True + + # Import here so the callback doesn't require provenance as a hard dep + from rationai.mlkit.provenance.register_dataset import ( + _detect_manifest, + _lookup_dataset_run, + _verify_dataset, + ) + + manifest_path = self._manifest_path + data_root = None + if manifest_path is None: + manifest_path, data_root = _detect_manifest() + + if manifest_path is None: + print(" [DatasetVerificationCallback] No manifest.csv found — skipping") + return + + if data_root is None: + import os + data_root = os.path.dirname(os.path.abspath(manifest_path)) + + dataset_run_id = _lookup_dataset_run() + verification = _verify_dataset(manifest_path, data_root, dataset_run_id) + + for detail in verification["details"]: + print(f" [DatasetVerificationCallback] {detail}") + + # Log to the active MLflow run (if any) + active_run_id = mlflow.active_run().info.run_id if mlflow.active_run() else None + if active_run_id: + 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") diff --git a/rationai/mlkit/provenance/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py similarity index 57% rename from rationai/mlkit/provenance/provenance.py rename to rationai/mlkit/lightning/callbacks/provenance.py index 06be884..9225082 100644 --- a/rationai/mlkit/provenance/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -1,23 +1,17 @@ -""" -Automatic PROV-O-aware provenance logger for MLflow. - -Inspired by rationai.mlkit.autolog — the user decorates their training function -and all metadata is captured automatically. +"""Lightning callback that captures full PROV-O provenance for MLflow runs. -Usage: +Replaces the plain-PyTorch ``@autolog`` decorator from +``rationai.mlkit.provenance.provenance`` with a Lightning-native callback +that hooks into ``on_fit_start`` / ``on_fit_end``. - from provenance import autolog +Example:: - @autolog(model_name="resnet_baseline_v1") - def train(run): - run.log_params({"learning_rate": 1e-3}) - model = build_model() - run.register_model(model) - # ... training loop ... - run.save_model(model) + from rationai.mlkit.lightning.callbacks import ProvenanceCallback -Everything else (user, dataset, hardware, docker, git, environment, -train/test split, console output) is detected and logged automatically. + trainer = Trainer( + callbacks=[ProvenanceCallback(model_name="resnet_v1")], + logger=MLFlowLogger(...), + ) """ from __future__ import annotations @@ -28,18 +22,19 @@ def train(run): import json import uuid import shutil -import types import platform import hashlib import subprocess import contextlib from datetime import datetime, timezone -from functools import partial, wraps +from functools import partial from collections.abc import Callable import mlflow import torch import pandas as pd +from lightning.pytorch.callbacks import Callback + # ────────────────────────────────────────────── # OpenProvenance / CPM namespace URIs @@ -58,7 +53,6 @@ def train(run): "sosa": "http://www.w3.org/ns/sosa/", } -# Hyperparameter keys that go on the activity vs. metadata entity _ACTIVITY_HP_KEYS = { "learning_rate", "lr", "batch_size", "epochs", "optimizer", "loss_function", "dropout", "weight_decay", "momentum", @@ -66,7 +60,6 @@ def train(run): "patch_size", "input_size", "augmentations", } -# Param keys that map to WSI entity properties _WSI_PARAM_KEYS = { "scanner", "slide_id", "wsi_id", "patient_id", "subject_id", "institution", "site", "staining", "slicing_method", @@ -74,7 +67,7 @@ def train(run): # ────────────────────────────────────────────── -# Auto-detection helpers +# Helpers # ────────────────────────────────────────────── def _get_git_info(): @@ -96,16 +89,10 @@ def _get_git_info(): return commit, remote, branch -# ── Dataset helpers live in register_dataset.py ────────────────────── -from .register_dataset import ( # noqa: F401 - _lookup_dataset_run, - _lookup_experiment, - _verify_dataset, -) - - def _lookup_user_run(): """Find the user run from User_Registry. Auto-detect username.""" + from rationai.mlkit.provenance.register_dataset import _lookup_experiment + username = os.environ.get("MLFLOW_USER") if not username: try: @@ -134,9 +121,6 @@ def _lookup_user_run(): return row.run_id, dict(run.data.tags) -# ── Dataset verification lives in register_dataset.py ───────────────── - - def _detect_hardware(): info: dict[str, str | int] = {} @@ -169,7 +153,6 @@ def _detect_docker(): if os.path.exists("/.dockerenv"): info["docker"] = True - # Fallback: cgroup v1 / v2 if not info["docker"]: try: with open("/proc/self/cgroup") as f: @@ -197,7 +180,6 @@ def _detect_docker(): except FileNotFoundError: pass - # Try docker inspect if socket is available inside container if info["docker"]: cid = info.get("container_id_short", "") if cid: @@ -218,20 +200,6 @@ def _detect_docker(): return info -def _detect_manifest(): - """Walk data/ looking for manifest.csv.""" - 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 - - def _snapshot_environment(artifact_dir): """Freeze environment to *artifact_dir* and return the pip-freeze text.""" req_path = os.path.join(artifact_dir, "requirements_frozen.txt") @@ -242,46 +210,10 @@ def _snapshot_environment(artifact_dir): if os.path.exists(src): shutil.copy2(src, os.path.join(artifact_dir, src)) - # Return the frozen requirements text for embedding into provenance with open(req_path) as f: return f.read() -def _source_hash(source_file: str) -> str | None: - """Return SHA-256 of a source file (for reproducibility verification).""" - try: - h = hashlib.sha256() - with open(source_file, "rb") as f: - for chunk in iter(lambda: f.read(8192), b""): - h.update(chunk) - return h.hexdigest() - except (FileNotFoundError, PermissionError): - return None - - -def _prepare_split(manifest_path, data_root, test_size=0.2, random_state=42): - from sklearn.model_selection import train_test_split - - df = pd.read_csv(manifest_path) - samples = [] - 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"])}) - - train_s, test_s = train_test_split( - samples, - test_size=test_size, - random_state=random_state, - stratify=[s["label"] for s in samples], - ) - return list(train_s), list(test_s) - - -# ────────────────────────────────────────────── -# Model / optimizer / scheduler introspection -# ────────────────────────────────────────────── - def _model_summary(model): """Extract architecture details from a torch.nn.Module.""" info: dict[str, str | int] = {} @@ -345,31 +277,80 @@ def _scheduler_summary(scheduler): # ────────────────────────────────────────────── -# OpenProvenance / CPM PROV document builder +# Console stream capture +# ────────────────────────────────────────────── + +_CONSOLE_LOG_NAME = "console.log" + + +class _StreamCapture: + """Captures stdout/stderr and writes them to MLflow as console.log.""" + + def __init__(self, run_id: str): + self.run_id = run_id + self._buffer = io.StringIO() + self._originals = {} + + def __enter__(self): + import sys + for stream_name in ("stdout", "stderr"): + stream = getattr(sys, stream_name) + original_write = stream.write + self._originals[stream_name] = original_write + + def _make_wrapper(buf, orig): + def wrapper(text): + buf.write(text) + return orig(text) + return wrapper + + setattr(stream, "write", _make_wrapper(self._buffer, self._originals[stream_name])) + + def __exit__(self, *args): + import sys + for stream_name in ("stdout", "stderr"): + original = self._originals.get(stream_name) + if original is not None: + setattr(getattr(sys, stream_name), "write", original) + + text = self._buffer.getvalue() + if text.strip(): + try: + client = mlflow.tracking.MlflowClient() + log_path = os.path.join( + f"_mlflow_console_{uuid.uuid4().hex[:8]}", + _CONSOLE_LOG_NAME, + ) + os.makedirs(os.path.dirname(log_path), exist_ok=True) + with open(log_path, "w") as f: + f.write(text) + client.log_artifact(self.run_id, log_path, artifact_path="logs") + shutil.rmtree(os.path.dirname(log_path), ignore_errors=True) + except Exception: + pass + + +# ────────────────────────────────────────────── +# PROV document builder # ────────────────────────────────────────────── def _safe_id(name: str) -> str: - """Sanitise a string for use as a PROV identifier fragment.""" return re.sub(r'[^a-zA-Z0-9_]', '_', name) def _qualified(prefix: str, local: str) -> str: - """Return a qualified name like 'gen:run_abc123'.""" return f"{prefix}:{local}" def _typed_value(value, type_prefix="xsd", type_local="string") -> list: - """Wrap a string value as [value] — matching Java's array convention.""" return [str(value)] def _qualified_name(type_prefix: str, type_local: str) -> dict: - """Build a prov:QUALIFIED_NAME type descriptor.""" return {"type": "prov:QUALIFIED_NAME", "$": f"{type_prefix}:{type_local}"} def _iso_timestamp(ts_ms: int | None = None) -> str: - """Return an ISO-8601 timestamp string (ms since epoch or now).""" if ts_ms is not None: dt = datetime.fromtimestamp(ts_ms / 1000, tz=timezone.utc) else: @@ -389,16 +370,8 @@ def _build_prov_document( requirements: str | None = None, verification: dict | None = None, ) -> dict: - """Build an OpenProvenance-compatible PROV document dict. + """Build an OpenProvenance-compatible PROV document dict.""" - Structure mirrors the Java prov_mlflow output: - - bundle wrapper with storage: key - - prefix namespace declarations - - entity, activity, agent sections - - wasAssociatedWith, used relationship sections - """ - - # ── Derive identifiers ──────────────────────────────── username = tags.get("username", tags.get("mlflow.user", "unknown")) agent_local = _safe_id(f"user_{username}") agent_id = _qualified("gen", agent_local) @@ -412,7 +385,6 @@ def _build_prov_document( main_act_local = f"TrainingRun_{run_id[:8]}" main_act_id = _qualified("blank", main_act_local) - # ── Collect sections ────────────────────────────────── entities: dict[str, dict] = {} activities: dict[str, dict] = {} agents: dict[str, dict] = {} @@ -426,7 +398,7 @@ def _blank_rel_id() -> str: rel_counter[0] += 1 return rid - # ── 1. AGENT (researcher) ───────────────────────────── + # ── 1. AGENT ─────────────────────────────────────────── agent_props: dict[str, list] = {} real_name = tags.get("real_name", username) agent_props["schema:name"] = _typed_value(real_name) @@ -438,14 +410,7 @@ def _blank_rel_id() -> str: agent_props["prov:type"] = [_qualified_name("schema", "Person")] agents[agent_id] = agent_props - # ── 2. INPUT ENTITIES (WSI / dataset samples) ───────── - # Collect unique sample paths from params or tags - sample_paths: list[str] = [] - for key in ("train_samples", "test_samples"): - if key in params: - pass # counts, not paths — skip - - # Try to find WSI-related params + # ── 2. INPUT ENTITIES ────────────────────────────────── image_path_candidates = ( params.get("image_path") or params.get("wsi_path") or @@ -454,7 +419,6 @@ def _blank_rel_id() -> str: params.get("input_path") ) - # If we have a manifest reference, create a single dataset entity if image_path_candidates: wsi_local = _safe_id(f"wsi_{image_path_candidates}") wsi_id = _qualified("gen", wsi_local) @@ -462,7 +426,6 @@ def _blank_rel_id() -> str: "schema:name": _typed_value(f"Input: {image_path_candidates}"), "prov:type": [_qualified_name("sosa", "Sample")], } - # Add optional WSI metadata from params if "scanner" in params: wsi_props["gen:scanner"] = _typed_value(params["scanner"]) for pk, prov_key in [ @@ -484,7 +447,6 @@ def _blank_rel_id() -> str: "prov:entity": wsi_id, } else: - # Fallback: create a generic dataset entity from split info train_count = params.get("train_samples", "0") test_count = params.get("test_samples", "0") ds_local = _safe_id(f"dataset_{run_id[:8]}") @@ -498,33 +460,28 @@ def _blank_rel_id() -> str: "prov:entity": ds_id, } - # ── 3. RUN ACTIVITY (the ML training) ───────────────── + # ── 3. RUN ACTIVITY ──────────────────────────────────── run_activity: dict[str, object] = {} 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) - # Experiment name from tags exp_name = params.get("model_name", "") if exp_name: run_activity["gen:experiment_name"] = _typed_value(exp_name) - # Model config if "model_class" in params: run_activity["gen:model_config"] = _typed_value(params["model_class"]) - # Git commit git_commit = tags.get("git_commit", tags.get("mlflow.source.git.commit", "")) if git_commit: run_activity["schema:identifier"] = _typed_value(git_commit) - # Backward-compatible model params for key in ("pretrained_model", "backbone", "feature_extractor"): if key in params: run_activity["gen:pretrained_model"] = _typed_value(params[key]) - # Dataset info for key, prov_key in [ ("dataset_name", "gen:dataset_name"), ("dataset_version", "gen:dataset_version"), @@ -534,13 +491,10 @@ def _blank_rel_id() -> str: if key in params: run_activity[prov_key] = _typed_value(params[key]) - # Hyperparameters for key in _ACTIVITY_HP_KEYS: if key in params: run_activity[f"gen:{key}"] = _typed_value(params[key]) - # Also add opt_ and sch_ prefixed params as hyperparams (stripped) - # Strip ALL leading prefixes to avoid sch_opt_lr → gen:opt_lr pollution for key, val in params.items(): if key.startswith("opt_") or key.startswith("sch_"): clean = key @@ -549,10 +503,9 @@ def _blank_rel_id() -> str: clean = clean[4:] elif clean.startswith("sch_"): clean = clean[4:] - if f"gen:{clean}" not in run_activity: # avoid duplicates + if f"gen:{clean}" not in run_activity: run_activity[f"gen:{clean}"] = _typed_value(val) - # Hardware from tags (mlflow.* convention) and our custom params for tag_key, prov_key in [ ("mlflow.gpu.count", "gen:gpu_count"), ("mlflow.gpu.names", "gen:gpu_names"), @@ -562,7 +515,6 @@ def _blank_rel_id() -> str: if tag_key in tags: run_activity[prov_key] = _typed_value(tags[tag_key]) - # Our custom hardware params for param_key, prov_key in [ ("gpu_count", "gen:gpu_count"), ("gpu_name", "gen:gpu_names"), @@ -572,7 +524,6 @@ def _blank_rel_id() -> str: if param_key in params and prov_key not in run_activity: run_activity[prov_key] = _typed_value(params[param_key]) - # Git remote / source git_url = tags.get("git_url", tags.get("mlflow.source.git.remote", "")) if git_url: run_activity["gen:git_remote"] = _typed_value(git_url) @@ -581,7 +532,6 @@ def _blank_rel_id() -> str: if source_name: run_activity["gen:source_name"] = _typed_value(source_name) - # Segmentation / model params (if present) for key, prov_key in [ ("segmentation", "gen:segmentation_config"), ("model", "gen:model_config"), @@ -591,14 +541,13 @@ def _blank_rel_id() -> str: activities[run_act_id] = run_activity - # ── 4. CPM METADATA ENTITY ──────────────────────────── + # ── 4. CPM METADATA ENTITY ───────────────────────────── meta_entity: dict[str, list] = {} 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 already placed on the activity 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", @@ -607,20 +556,16 @@ def _blank_rel_id() -> str: "institution", "site", "staining", "slicing_method", } - # Remaining params → metadata entity 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) - # Metrics → metadata entity for key, val in metrics.items(): safe_key = _safe_id(key) meta_entity[f"gen:{safe_key}"] = _typed_value(val) - # ── Reproducibility: embedded dataset splits ─────────── if split_data: - # Embed as JSON string so the PROV document is self-contained 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"] @@ -630,11 +575,9 @@ def _blank_rel_id() -> str: if split_data.get("test"): meta_entity["gen:split_test"] = [_typed_value(json.dumps(split_data["test"]))[0]] - # ── Reproducibility: frozen requirements ─────────────── if requirements: meta_entity["gen:requirements"] = [requirements] - # ── Reproducibility: dataset verification ────────────── if verification: meta_entity["gen:dataset_verified"] = [str(verification.get("verified", False))] meta_entity["gen:dataset_run_id"] = [verification.get("dataset_run_id", "")] @@ -646,7 +589,6 @@ def _blank_rel_id() -> str: meta_entity["gen:files_missing"] = [str(fm)] meta_entity["gen:files_total"] = [str(ft)] - # Selected tags → metadata entity for tag_key in ( "mlflow.source.git.branch", "mlflow.source.git.repo_url", @@ -659,14 +601,13 @@ def _blank_rel_id() -> str: entities[meta_id] = meta_entity - # wasGeneratedBy: meta entity ← run activity was_generated_by: dict[str, dict] = {} was_generated_by[_blank_rel_id()] = { "prov:entity": meta_id, "prov:activity": run_act_id, } - # ── 5. CPM MAIN ACTIVITY ────────────────────────────── + # ── 5. CPM MAIN ACTIVITY ─────────────────────────────── main_activity: dict[str, object] = {} main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] main_activity["cpm:referencedMetaBundleId"] = [ @@ -677,13 +618,13 @@ def _blank_rel_id() -> str: ] activities[main_act_id] = main_activity - # ── 6. RELATIONSHIPS ────────────────────────────────── + # ── 6. RELATIONSHIPS ─────────────────────────────────── was_associated_with[_blank_rel_id()] = { "prov:activity": run_act_id, "prov:agent": agent_id, } - # ── 7. ASSEMBLE BUNDLE ──────────────────────────────── + # ── 7. ASSEMBLE BUNDLE ───────────────────────────────── inner: dict[str, object] = {"prefix": _PROV_PREFIXES} if entities: inner["entity"] = entities @@ -703,387 +644,289 @@ def _blank_rel_id() -> str: # ────────────────────────────────────────────── -# Console stream capture (like mlkit.stream) +# ProvenanceCallback # ────────────────────────────────────────────── -_CONSOLE_LOG_NAME = "console.log" - +class ProvenanceCallback(Callback): + """Lightning callback that captures full PROV-O provenance. -class _StreamCapture: - """Captures stdout/stderr and writes them to MLflow as console.log.""" - - def __init__(self, run_id: str): - self.run_id = run_id - self._buffer = io.StringIO() - self._originals = {} - - def __enter__(self): - import sys - for stream_name in ("stdout", "stderr"): - stream = getattr(sys, stream_name) - original_write = stream.write - self._originals[stream_name] = original_write - - def _make_wrapper(buf, orig): - def wrapper(text): - buf.write(text) - return orig(text) - return wrapper - - setattr(stream, "write", _make_wrapper(self._buffer, self._originals[stream_name])) - - def __exit__(self, *args): - import sys - for stream_name in ("stdout", "stderr"): - original = self._originals.get(stream_name) - if original is not None: - setattr(getattr(sys, stream_name), "write", original) - - text = self._buffer.getvalue() - if text.strip(): - try: - client = mlflow.tracking.MlflowClient() - log_path = os.path.join( - f"_mlflow_console_{uuid.uuid4().hex[:8]}", - _CONSOLE_LOG_NAME, - ) - os.makedirs(os.path.dirname(log_path), exist_ok=True) - with open(log_path, "w") as f: - f.write(text) - client.log_artifact(self.run_id, log_path, artifact_path="logs") - shutil.rmtree(os.path.dirname(log_path), ignore_errors=True) - except Exception: - pass # don't fail the run - - -# ────────────────────────────────────────────── -# Run helper object (injected into the user function) -# ────────────────────────────────────────────── - -class _Run: - """Handle passed to the user's training function. - - Provides logging methods and model/optimizer/scheduler registration. - """ - - def __init__(self, run_id: str): - self._run_id = run_id - self._model = None - self._optimizer_info: dict | None = None - self._scheduler_info: dict | None = None - self.train_paths: list[str] = [] - self.test_paths: list[str] = [] - self.train_labels: list[int] = [] - self.test_labels: list[int] = [] - self._split_data: dict | None = None # {"train": [...], "test": [...]} - - def _set_split(self, train_samples: list[dict], test_samples: list[dict]): - """Store the full split data for embedding into provenance.""" - self._split_data = {"train": train_samples, "test": test_samples} - - # ── Logging (forwarded to mlflow) ───────── - - def log_param(self, key, value): - mlflow.log_param(key, value) - - def log_params(self, params_dict): - mlflow.log_params(params_dict) - - def log_metric(self, key, value, step=None): - mlflow.log_metric(key, value, step=step) - - def log_metrics(self, metrics_dict, step=None): - mlflow.log_metrics(metrics_dict, step=step) - - def log_artifact(self, local_path, artifact_path=None): - mlflow.log_artifact(local_path, artifact_path=artifact_path) - - def log_artifacts(self, local_dir, artifact_path=None): - mlflow.log_artifacts(local_dir, artifact_path=artifact_path) - - # ── Registration (logged at the end) ─────── - - def register_model(self, model): - self._model = model - - def register_optimizer(self, optimizer): - self._optimizer_info = _optimizer_summary(optimizer) - - def register_scheduler(self, scheduler): - self._scheduler_info = _scheduler_summary(scheduler) - - # ── Model saving ────────────────────────── - - def save_model(self, model, name="model", **kwargs): - if "export_model" not in kwargs: - kwargs["export_model"] = False - mlflow.pytorch.log_model(model, name, **kwargs) - - -# ────────────────────────────────────────────── -# Decorator — the main entry point -# ────────────────────────────────────────────── - -def autolog( - model_name: str | None = None, - experiment_name: str = "Training_Pipeline", - test_size: float = 0.2, - random_state: int = 42, - log_stream: bool = True, - fail_fast: bool = True, -): - """Decorator for automatic provenance logging. - - All metadata (user, dataset, hardware, docker, git, environment, - train/test split, model architecture, optimizer, scheduler, console - output) is captured automatically. + Logs user tags, git info, hardware, docker detection, environment + snapshot, dataset verification, train/test split, model summary, + optimizer/scheduler summary, console output capture, and builds + a self-contained PROV document on training completion. Args: model_name: Identifier for this model (shown in run name). - Defaults to MODEL_NAME env var or "model". - experiment_name: MLflow experiment 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. log_stream: Whether to capture stdout/stderr into console.log. - fail_fast: If True (default), abort the run with RuntimeError - when dataset verification fails. - - Example: - from provenance import autolog - - @autolog(model_name="resnet_baseline_v1") - def train(run): - run.log_params({"learning_rate": 1e-3}) - model = build_model() - run.register_model(model) - optimizer = optim.SGD(model.parameters(), lr=1e-3) - run.register_optimizer(optimizer) - for epoch in range(50): - loss = train_epoch(...) - run.log_metrics({"train_loss": loss}, step=epoch) - run.save_model(model) - - if __name__ == "__main__": - train() + fail_fast: Abort training if dataset verification fails. + register_model: If True, auto-log model summary from pl_module. + register_optimizer: Pass an optimizer to log its config (or True to auto-detect). + register_scheduler: Pass a scheduler to log its config (or True to auto-detect). """ - def decorator(func: Callable[..., None]) -> Callable[[], None]: - @wraps(func) - def wrapper(): - _run_autolog( - func=func, - model_name=model_name or os.environ.get("MODEL_NAME", "model"), - experiment_name=experiment_name, - test_size=test_size, - random_state=random_state, - log_stream=log_stream, - fail_fast=fail_fast, + 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, + log_stream: bool = True, + fail_fast: bool = True, + register_model: bool = True, + register_optimizer: bool = True, + register_scheduler: bool = True, + ): + 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.log_stream = log_stream + self.fail_fast = fail_fast + self.register_model = register_model + self.register_optimizer = register_optimizer + self.register_scheduler = register_scheduler + + # Internal state + self._run_id: str | None = None + self._mlflow_run = None + self._temp_dirs: list[str] = [] + self._split_data: dict | None = None + self._verification: dict | None = None + self._frozen_requirements: str | None = None + self._optimizer_info: dict | None = None + self._scheduler_info: dict | None = None + self._git_commit: str = "unknown" + self._git_url: str = "unknown" + self._git_branch: str = "unknown" + + def on_fit_start(self, trainer, pl_module): # noqa: ARG002 + """Run all provenance setup at the start of training.""" + from rationai.mlkit.provenance.register_dataset import ( + _detect_manifest, + _lookup_dataset_run, + _verify_dataset, + ) + + # ── Auto-detect everything ────────────────────────────── + user_run_id, user_tags = _lookup_user_run() + git_commit, git_url, git_branch = _get_git_info() + self._git_commit = git_commit + self._git_url = git_url + self._git_branch = git_branch + + hardware = _detect_hardware() + docker = _detect_docker() + + # ── Dataset detection & split ─────────────────────────── + manifest_path = self.manifest_path + data_root = self.data_root + + if manifest_path is None: + manifest_path, data_root = _detect_manifest() + + if manifest_path and data_root: + dataset_run_id = _lookup_dataset_run(manifest_path) + verification = _verify_dataset(manifest_path, data_root, dataset_run_id) + self._verification = verification + + for detail in verification["details"]: + print(f" [ProvenanceCallback] {detail}") + + # Train/test split + from sklearn.model_selection import train_test_split + df = pd.read_csv(manifest_path) + samples = [] + 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"])}) + + 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], ) - return wrapper - - return decorator - - -def _run_autolog(func, model_name, experiment_name, test_size, random_state, log_stream, fail_fast=True): - """Core autolog logic.""" - _temp_dirs: list[str] = [] - - # ── Auto-detect everything ────────────────────────────── - user_run_id, user_tags = _lookup_user_run() - manifest_path, data_root = _detect_manifest() - dataset_run_id = _lookup_dataset_run(manifest_path) - git_commit, git_url, git_branch = _get_git_info() - hardware = _detect_hardware() - docker = _detect_docker() - - # ── Start MLflow run ─────────────────────────────────── - mlflow.set_experiment(experiment_name) - ts = datetime.now().strftime("%Y%m%d_%H%M%S") - run_name = f"Training_{model_name}_{ts}" - - mlflow_run = mlflow.start_run(run_name=run_name) - run_id = mlflow_run.info.run_id - - # ── 1. Tags (PROV relationships) ─────────────────────── - 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] - - if dataset_run_id: - tags["dataset_run_id"] = dataset_run_id - - tags.update({ - "git_commit": git_commit, - "git_url": git_url, - "git_branch": git_branch, - "prov_start_time": datetime.now(timezone.utc).isoformat(), - }) - mlflow.set_tags(tags) - - # ── 2. Params: hardware + docker + split config ──────── - all_params: dict[str, str | float | int] = { - "model_name": model_name, - **hardware, - **docker, - "split_test_size": test_size, - "split_random_state": random_state, - "split_stratified": True, - } - mlflow.log_params(all_params) + self._split_data = { + "train": train_samples, + "test": test_samples, + "test_size": self.test_size, + "random_state": self.random_state, + } - # ── 3. Environment snapshot ──────────────────────────── - artifact_dir = f"_mlflow_env_{uuid.uuid4().hex[:8]}" - os.makedirs(artifact_dir, exist_ok=True) - _temp_dirs.append(artifact_dir) - frozen_requirements: str | None = None - try: - frozen_requirements = _snapshot_environment(artifact_dir) - mlflow.log_artifacts(artifact_dir, artifact_path="environment") - except Exception: - pass + # 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), + }) + + # Log verification results + 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"]), + ) - # ── 4. Train/test split from manifest ────────────────── - run_handle = _Run(run_id) + 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"]) + ) + else: + print("[ProvenanceCallback] WARNING: 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] + + if dataset_run_id := _lookup_dataset_run(manifest_path): + tags["dataset_run_id"] = dataset_run_id + + tags.update({ + "git_commit": git_commit, + "git_url": git_url, + "git_branch": git_branch, + "prov_start_time": datetime.now(timezone.utc).isoformat(), + }) + mlflow.set_tags(tags) - if manifest_path and data_root: - train_samples, test_samples = _prepare_split( - manifest_path, data_root, test_size, random_state - ) + # ── 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) - run_handle.train_paths = [s["path"] for s in train_samples] - run_handle.test_paths = [s["path"] for s in test_samples] - run_handle.train_labels = [s["label"] for s in train_samples] - run_handle.test_labels = [s["label"] for s in test_samples] - - # Store full split data for embedding into provenance - run_handle._set_split(train_samples, test_samples) - - mlflow.log_params({ - "train_samples": len(train_samples), - "test_samples": len(test_samples), - "train_positive": sum(run_handle.train_labels), - "train_negative": len(run_handle.train_labels) - sum(run_handle.train_labels), - "test_positive": sum(run_handle.test_labels), - "test_negative": len(run_handle.test_labels) - sum(run_handle.test_labels), - }) + # ── 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: + pass - split_dir = f"_mlflow_split_{uuid.uuid4().hex[:8]}" - os.makedirs(split_dir, exist_ok=True) - _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") - - # ── 4b. Verify dataset against Dataset_Registry ──── - verification = _verify_dataset(manifest_path, data_root, dataset_run_id) - for detail in verification["details"]: - print(f" [autolog] {detail}") - - # Log verification results as params + tags - 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"]: - tags["dataset_verification"] = "VERIFIED" - else: - tags["dataset_verification"] = "MISMATCH" - tags["dataset_verification_details"] = "; ".join(verification["details"]) - mlflow.set_tags(tags) + def on_fit_end(self, trainer, pl_module): # noqa: ARG002 + """Log model/optimizer/scheduler summaries and PROV document.""" + active_run = mlflow.active_run() + if not active_run: + return + run_id = active_run.info.run_id - # ── Hard abort on mismatch ───────────────────── - if fail_fast and not verification["verified"]: - mlflow.end_run(status='FAILED') - raise RuntimeError( - "Dataset verification failed — aborting training.\n" - + " ".join(" " + d for d in verification["details"]) - ) + # ── 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: + pass - else: - print("[autolog] WARNING: No manifest.csv found — " - "train/test split not logged.") - verification = None + # ── Optimizer summary ─────────────────────────────────── + if self.register_optimizer and pl_module is not None: + try: + for opt in trainer.optimizers: + self._optimizer_info = _optimizer_summary(opt) + mlflow.log_params(self._optimizer_info) + break + except Exception: + pass - # ── 5. Run the user's training function ──────────────── - try: - if log_stream: - with _StreamCapture(run_id): - func(run_handle) - else: - func(run_handle) - except Exception: - # Still log what we can before re-raising - raise - finally: - # ── Log model/optimizer/scheduler (if registered) ── - if run_handle._model is not None: - mlflow.log_params(_model_summary(run_handle._model)) - - if run_handle._optimizer_info is not None: - mlflow.log_params(run_handle._optimizer_info) - if run_handle._scheduler_info is not None: - mlflow.log_params(run_handle._scheduler_info) - - # ── Write self-contained provenance JSON ─────────── - # Embeds splits + requirements so the run is reproducible - # without needing MLflow at all. + # ── Scheduler summary ─────────────────────────────────── + if self.register_scheduler and pl_module is not None: + try: + for sched in trainer.lr_schedulers: + self._scheduler_info = _scheduler_summary(sched.get("scheduler")) + mlflow.log_params(self._scheduler_info) + break + except Exception: + pass + + # ── 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") - # Source file hash for verification - source_file = getattr(func, "__wrapped__", func).__code__.co_filename - source_hash = _source_hash(source_file) if source_file else None - summary = { - "model_name": model_name, + "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": { - k: v for k, v in run_data.data.tags.items() - if not k.startswith("mlflow.") - }, + "tags": tags, "run_id": run_id, - "run_name": run_name, - "experiment_name": experiment_name, - - # ── Reproducibility: dataset splits ───────── + "experiment_name": self.experiment_name, "split": { - "test_size": test_size, - "random_state": random_state, + "test_size": self.test_size, + "random_state": self.random_state, "stratified": True, - "train_count": len(run_handle.train_paths), - "test_count": len(run_handle.test_paths), - "train": run_handle._split_data["train"] if run_handle._split_data else None, - "test": run_handle._split_data["test"] if run_handle._split_data else None, - } if run_handle._split_data else None, - - # ── Reproducibility: dataset verification ─── - "dataset_verification": verification if verification is not None else None, - - # ── Reproducibility: frozen environment ───── - "requirements": frozen_requirements, - - # ── Source verification ───────────────────── + "train_count": len(self._split_data["train"]) if self._split_data else 0, + "test_count": len(self._split_data["test"]) if self._split_data else 0, + "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": { - "file": source_file, - "sha256": source_hash, - "git_commit": git_commit, - "git_branch": git_branch, - "git_remote": git_url, + "git_commit": self._git_commit, + "git_branch": self._git_branch, + "git_remote": self._git_url, }, } @@ -1093,31 +936,28 @@ def _run_autolog(func, model_name, experiment_name, test_size, random_state, log mlflow.log_artifact(summary_path, artifact_path="provenance") shutil.rmtree(summary_dir, ignore_errors=True) - # ── Build & log PROV document ─────────────────── + # ── PROV document ─────────────────────────────────── prov_doc = _build_prov_document( run_id=run_id, - run_name=run_name, - 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.") - }, - start_time_ms=mlflow_run.info.start_time, - end_time_ms=mlflow_run.info.end_time, + 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": test_size, - "random_state": random_state, - "train": run_handle._split_data["train"] if run_handle._split_data else None, - "test": run_handle._split_data["test"] if run_handle._split_data else None, - } if run_handle._split_data else None, - requirements=frozen_requirements, - verification=verification, + "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_dir = f"_mlflow_prov_{uuid.uuid4().hex[:8]}" os.makedirs(prov_dir, exist_ok=True) - _temp_dirs.append(prov_dir) + 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) @@ -1125,10 +965,10 @@ def _run_autolog(func, model_name, experiment_name, test_size, random_state, log mlflow.log_artifact(prov_path, artifact_path="provenance") shutil.rmtree(prov_dir, ignore_errors=True) - print(f"\n[autolog] Complete → {run_id}") + print(f"\n[ProvenanceCallback] Complete → {run_id}") except Exception as e: - print(f"[autolog] WARNING: Could not write provenance artifacts: {e}") + print(f"[ProvenanceCallback] WARNING: Could not write provenance artifacts: {e}") - # ── Clean up temp dirs ───────────────────────────── - for d in _temp_dirs: + # ── Clean up temp dirs ────────────────────────────────── + for d in self._temp_dirs: shutil.rmtree(d, ignore_errors=True) diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py index 3f08927..dca4e6d 100644 --- a/rationai/mlkit/provenance/__init__.py +++ b/rationai/mlkit/provenance/__init__.py @@ -1,11 +1,14 @@ """Provenance tracking — PROV-O-aware logging to MLflow. Submodules: - provenance – @autolog decorator, internal helpers + provenance – internal helpers (lookup, verification) register_dataset – register_dataset (hash-based) and register_dataset_as_provenance (legacy CSV-based) register_user – register_new_user +For automatic provenance capture with Lightning, use +:class:`~rationai.mlkit.lightning.callbacks.provenance.ProvenanceCallback`. + MLflow tracking URI defaults to ``http://localhost:5000``. Override with the ``MLFLOW_TRACKING_URI`` environment variable. """ @@ -19,19 +22,30 @@ os.environ["MLFLOW_TRACKING_URI"] = "http://localhost:5000" # Now safe to import – all child modules will pick up the env var -from .provenance import autolog # noqa: E402 from .register_dataset import ( # noqa: E402 register_dataset, register_dataset_as_provenance, + verify_dataset, ) from .register_user import register_new_user # noqa: E402 __all__ = [ - # Core provenance - "autolog", - # Dataset registration + # Dataset "register_dataset", "register_dataset_as_provenance", + "verify_dataset", # User registration "register_new_user", ] + + +def __getattr__(name: str): + """Raise helpful error for removed ``autolog``.""" + if name == "autolog": + raise ImportError( + "provenance.autolog has been removed. " + "Use ProvenanceCallback instead:\n\n" + " from rationai.mlkit.lightning.callbacks import ProvenanceCallback\n" + " trainer = Trainer(callbacks=[ProvenanceCallback(model_name='...')], ...)\n" + ) + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/rationai/mlkit/provenance/register_dataset.py b/rationai/mlkit/provenance/register_dataset.py index 18ab764..5949c0d 100644 --- a/rationai/mlkit/provenance/register_dataset.py +++ b/rationai/mlkit/provenance/register_dataset.py @@ -39,6 +39,56 @@ def _lookup_dataset_run(manifest_path: str | None = None) -> str | None: return 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 + + +def verify_dataset( + manifest_path: str | None = None, + data_root: str | None = None, +) -> dict: + """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, From 6e15647481d6e9c2d7cc3885681af605814be589 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Tue, 21 Jul 2026 23:27:56 +0200 Subject: [PATCH 14/34] clean: removed unneded code --- .../mlkit/lightning/callbacks/provenance.py | 58 ------------------- 1 file changed, 58 deletions(-) diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index 9225082..3070144 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -16,7 +16,6 @@ from __future__ import annotations -import io import os import re import json @@ -25,10 +24,7 @@ import platform import hashlib import subprocess -import contextlib from datetime import datetime, timezone -from functools import partial -from collections.abc import Callable import mlflow import torch @@ -276,60 +272,6 @@ def _scheduler_summary(scheduler): return info -# ────────────────────────────────────────────── -# Console stream capture -# ────────────────────────────────────────────── - -_CONSOLE_LOG_NAME = "console.log" - - -class _StreamCapture: - """Captures stdout/stderr and writes them to MLflow as console.log.""" - - def __init__(self, run_id: str): - self.run_id = run_id - self._buffer = io.StringIO() - self._originals = {} - - def __enter__(self): - import sys - for stream_name in ("stdout", "stderr"): - stream = getattr(sys, stream_name) - original_write = stream.write - self._originals[stream_name] = original_write - - def _make_wrapper(buf, orig): - def wrapper(text): - buf.write(text) - return orig(text) - return wrapper - - setattr(stream, "write", _make_wrapper(self._buffer, self._originals[stream_name])) - - def __exit__(self, *args): - import sys - for stream_name in ("stdout", "stderr"): - original = self._originals.get(stream_name) - if original is not None: - setattr(getattr(sys, stream_name), "write", original) - - text = self._buffer.getvalue() - if text.strip(): - try: - client = mlflow.tracking.MlflowClient() - log_path = os.path.join( - f"_mlflow_console_{uuid.uuid4().hex[:8]}", - _CONSOLE_LOG_NAME, - ) - os.makedirs(os.path.dirname(log_path), exist_ok=True) - with open(log_path, "w") as f: - f.write(text) - client.log_artifact(self.run_id, log_path, artifact_path="logs") - shutil.rmtree(os.path.dirname(log_path), ignore_errors=True) - except Exception: - pass - - # ────────────────────────────────────────────── # PROV document builder # ────────────────────────────────────────────── From 617c792d7e5c7635fde14aca7a00560f4dc9fd38 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 10:21:49 +0200 Subject: [PATCH 15/34] feat: Introduce EnvironmentCallback for environment metadata capture - Added EnvironmentCallback to capture hardware, docker, git, user, and environment snapshot at training start. - Refactored ProvenanceCallback to depend on EnvironmentCallback for environment data, enhancing modularity. - Removed legacy CSV-based dataset registration from register_dataset.py for cleaner codebase. - Introduced load_manifest function to streamline manifest loading across multiple modules. - Updated documentation and logging for better clarity and error handling. --- .../mlkit/lightning/callbacks/__init__.py | 8 +- .../callbacks/dataset_verification.py | 136 +++++- .../mlkit/lightning/callbacks/environment.py | 308 +++++++++++++ .../mlkit/lightning/callbacks/provenance.py | 433 ++++++++---------- rationai/mlkit/provenance/__init__.py | 5 +- rationai/mlkit/provenance/register_dataset.py | 104 ++--- 6 files changed, 656 insertions(+), 338 deletions(-) create mode 100644 rationai/mlkit/lightning/callbacks/environment.py diff --git a/rationai/mlkit/lightning/callbacks/__init__.py b/rationai/mlkit/lightning/callbacks/__init__.py index b1f5529..65a01a4 100644 --- a/rationai/mlkit/lightning/callbacks/__init__.py +++ b/rationai/mlkit/lightning/callbacks/__init__.py @@ -1,10 +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__ = ["DatasetVerificationCallback", "MultiloaderLifecycle", "ProvenanceCallback"] +__all__ = [ + "DatasetVerificationCallback", + "EnvironmentCallback", + "MultiloaderLifecycle", + "ProvenanceCallback", +] diff --git a/rationai/mlkit/lightning/callbacks/dataset_verification.py b/rationai/mlkit/lightning/callbacks/dataset_verification.py index 1df9043..e9e7ba9 100644 --- a/rationai/mlkit/lightning/callbacks/dataset_verification.py +++ b/rationai/mlkit/lightning/callbacks/dataset_verification.py @@ -1,45 +1,79 @@ -"""Lightning callback that verifies the dataset against MLflow on trainer start.""" +"""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 mlflow +import logging +import os +import shutil +import uuid +import mlflow from lightning.pytorch.callbacks import Callback +log = logging.getLogger(__name__) + + class DatasetVerificationCallback(Callback): - """Run dataset verification once at the start of training. + """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. - Logs verification results as MLflow params so they appear on the run - page alongside metrics and artifacts. + Stores results on ``self`` so sibling callbacks (e.g. ``ProvenanceCallback``) + can read them without duplicating work: - Example:: + - ``_verification`` (dict | None) — verification result + - ``_split_data`` (dict | None) — train/test split data - from rationai.mlkit.lightning.callbacks import DatasetVerificationCallback - - trainer = Trainer( - callbacks=[DatasetVerificationCallback()], - logger=MLFlowLogger(...), - ) + 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): + def __init__( + self, + manifest_path: str | None = None, + test_size: float = 0.0, + random_state: int = 42, + fail_fast: bool = True, + ): 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 | None = None + self._split_data: dict | None = None def on_fit_start(self, trainer, pl_module): # noqa: ARG002 if self._done: return self._done = True - # Import here so the callback doesn't require provenance as a hard dep from rationai.mlkit.provenance.register_dataset import ( _detect_manifest, _lookup_dataset_run, _verify_dataset, + load_manifest as _load_manifest, ) manifest_path = self._manifest_path @@ -48,22 +82,22 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 manifest_path, data_root = _detect_manifest() if manifest_path is None: - print(" [DatasetVerificationCallback] No manifest.csv found — skipping") + log.warning("[DatasetVerificationCallback] No manifest.csv found — skipping") return if data_root is None: - import os 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["details"]: - print(f" [DatasetVerificationCallback] {detail}") + for detail in verification.get("details", []): + log.info(f" [DatasetVerificationCallback] {detail}") - # Log to the active MLflow run (if any) - active_run_id = mlflow.active_run().info.run_id if mlflow.active_run() else None - if active_run_id: + # ── Log verification results ──────────────────────────── + if mlflow.active_run(): mlflow.log_params({ "dataset_verified": verification["verified"], "dataset_file_sizes_match": verification["file_sizes_match"] is True, @@ -74,3 +108,63 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 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..09397b0 --- /dev/null +++ b/rationai/mlkit/lightning/callbacks/environment.py @@ -0,0 +1,308 @@ +"""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 hashlib +import json +import logging +import os +import platform +import shutil +import subprocess +import uuid +from datetime import datetime, timezone + +import mlflow +import torch +from lightning.pytorch.callbacks import Callback + + +log = logging.getLogger(__name__) + + +# ────────────────────────────────────────────── +# Helpers +# ────────────────────────────────────────────── + +def _lookup_user_run(): + """Find the user run from User_Registry. Auto-detect username.""" + from rationai.mlkit.provenance.register_dataset import _lookup_experiment + + username = os.environ.get("MLFLOW_USER") + if not username: + try: + username = subprocess.check_output( + ["git", "config", "user.name"], stderr=subprocess.DEVNULL, + ).decode().strip() + except subprocess.CalledProcessError: + pass + if not username: + username = os.environ.get("USER", "unknown") + + exp_id = _lookup_experiment("User_Registry") + if exp_id is None: + return None, {} + + runs = mlflow.search_runs(experiment_ids=[exp_id]) + if runs.empty: + return None, {} + + matched = runs[runs["tags.username"] == username] + if matched.empty: + matched = runs.head(1) + + row = matched.iloc[0] + run = mlflow.get_run(row.run_id) + return row.run_id, dict(run.data.tags) + + +def _detect_hardware(): + """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(): + """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 = 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, + ): + 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 = {} + self._docker: dict = {} + self._frozen_requirements: str | None = None + self._user_run_id: str | None = None + self._user_tags: dict = {} + self._temp_dirs: list[str] = [] + + def on_fit_start(self, trainer, pl_module): # noqa: ARG002 + """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() + if run and run.info and run.info.run_id: + client = mlflow.tracking.MlflowClient() + run_data = client.get_run(run.info.run_id) + tags = run_data.data.tags or {} + else: + tags = {} + self._git_commit = tags.get("mlflow.source.git.commit", + tags.get("git.commit", "unknown")) + self._git_url = tags.get("mlflow.source.git.repoUrl", + tags.get("git.repo_url", "unknown")) + self._git_branch = tags.get("mlflow.source.git.branch", + 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) ─────────────────────────────── + tags: dict[str, str] = {} + if self._user_run_id: + tags["user_run_id"] = self._user_run_id + for key in ("username", "real_name", "organization"): + if key in self._user_tags: + tags[key] = self._user_tags[key] + + from rationai.mlkit.provenance.register_dataset import _lookup_dataset_run + 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(timezone.utc).isoformat(), + }) + mlflow.set_tags(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 index 3070144..176914b 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -1,8 +1,15 @@ -"""Lightning callback that captures full PROV-O provenance for MLflow runs. +"""Lightning callback that captures PROV-O provenance for MLflow runs. -Replaces the plain-PyTorch ``@autolog`` decorator from -``rationai.mlkit.provenance.provenance`` with a Lightning-native callback -that hooks into ``on_fit_start`` / ``on_fit_end``. +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:: @@ -16,27 +23,28 @@ from __future__ import annotations +import json +import logging import os import re -import json -import uuid import shutil -import platform -import hashlib -import subprocess +import uuid from datetime import datetime, timezone import mlflow -import torch import pandas as pd +import torch from lightning.pytorch.callbacks import Callback +log = logging.getLogger(__name__) + + # ────────────────────────────────────────────── -# OpenProvenance / CPM namespace URIs +# OpenProvenance / CPM namespace URIs (§9 — configurable via env var) # ────────────────────────────────────────────── -_PROV_PREFIXES = { +_DEFAULT_PROV_PREFIXES = { "storage": "http://localhost:8083/api/v1/documents/", "meta": "http://localhost:8083/api/v1/documents/meta/", "schema": "https://schema.org/", @@ -49,6 +57,25 @@ "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. + The env var accepts a JSON object of ``{prefix: base_uri}`` pairs that + merge into (and override) the defaults. + """ + if override: + return {**_DEFAULT_PROV_PREFIXES, **override} + env_json = os.environ.get("PROV_BASE_URI", "") + if env_json: + try: + merged = {**_DEFAULT_PROV_PREFIXES, **json.loads(env_json)} + return merged + except json.JSONDecodeError as e: + log.warning(f"PROV_BASE_URI is not valid JSON: {e} — using defaults") + return _DEFAULT_PROV_PREFIXES + _ACTIVITY_HP_KEYS = { "learning_rate", "lr", "batch_size", "epochs", "optimizer", "loss_function", "dropout", "weight_decay", "momentum", @@ -63,153 +90,9 @@ # ────────────────────────────────────────────── -# Helpers +# Model / Optimizer / Scheduler summaries # ────────────────────────────────────────────── -def _get_git_info(): - """Return (commit, remote_url, branch) or ('unknown', ...) on failure.""" - try: - commit = subprocess.check_output( - ["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL, - ).decode().strip() - remote = subprocess.check_output( - ["git", "config", "--get", "remote.origin.url"], - stderr=subprocess.DEVNULL, - ).decode().strip() - branch = subprocess.check_output( - ["git", "rev-parse", "--abbrev-ref", "HEAD"], - stderr=subprocess.DEVNULL, - ).decode().strip() - except subprocess.CalledProcessError: - commit, remote, branch = "unknown", "unknown", "unknown" - return commit, remote, branch - - -def _lookup_user_run(): - """Find the user run from User_Registry. Auto-detect username.""" - from rationai.mlkit.provenance.register_dataset import _lookup_experiment - - username = os.environ.get("MLFLOW_USER") - if not username: - try: - username = subprocess.check_output( - ["git", "config", "user.name"], stderr=subprocess.DEVNULL, - ).decode().strip() - except subprocess.CalledProcessError: - pass - if not username: - username = os.environ.get("USER", "unknown") - - exp_id = _lookup_experiment("User_Registry") - if exp_id is None: - return None, {} - - runs = mlflow.search_runs(experiment_ids=[exp_id]) - if runs.empty: - return None, {} - - matched = runs[runs["tags.username"] == username] - if matched.empty: - matched = runs.head(1) - - row = matched.iloc[0] - run = mlflow.get_run(row.run_id) - return row.run_id, dict(run.data.tags) - - -def _detect_hardware(): - 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(): - 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 = 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): - """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() - - def _model_summary(model): """Extract architecture details from a torch.nn.Module.""" info: dict[str, str | int] = {} @@ -311,6 +194,7 @@ def _build_prov_document( split_data: dict | None = None, requirements: str | None = None, verification: dict | None = None, + prov_prefixes: dict[str, str] | None = None, ) -> dict: """Build an OpenProvenance-compatible PROV document dict.""" @@ -567,7 +451,7 @@ def _blank_rel_id() -> str: } # ── 7. ASSEMBLE BUNDLE ───────────────────────────────── - inner: dict[str, object] = {"prefix": _PROV_PREFIXES} + inner: dict[str, object] = {"prefix": prov_prefixes or _get_prov_prefixes()} if entities: inner["entity"] = entities if activities: @@ -586,16 +470,17 @@ def _blank_rel_id() -> str: # ────────────────────────────────────────────── -# ProvenanceCallback +# Slim ProvenanceCallback # ────────────────────────────────────────────── class ProvenanceCallback(Callback): - """Lightning callback that captures full PROV-O provenance. + """Lightning callback that captures PROV document + run summary. - Logs user tags, git info, hardware, docker detection, environment - snapshot, dataset verification, train/test split, model summary, - optimizer/scheduler summary, console output capture, and builds - a self-contained PROV document on training completion. + 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). @@ -604,11 +489,12 @@ class ProvenanceCallback(Callback): 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. - log_stream: Whether to capture stdout/stderr into console.log. 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: Pass an optimizer to log its config (or True to auto-detect). - register_scheduler: Pass a scheduler to log its config (or True to auto-detect). + 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__( @@ -619,11 +505,12 @@ def __init__( data_root: str | None = None, test_size: float = 0.2, random_state: int = 42, - log_stream: bool = True, 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, ): self.model_name = model_name or os.environ.get("MODEL_NAME", "model") self.experiment_name = experiment_name @@ -631,44 +518,88 @@ def __init__( self.data_root = data_root self.test_size = test_size self.random_state = random_state - self.log_stream = log_stream 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 + # Internal state (populated by on_fit_start or sibling callbacks) self._run_id: str | None = None - self._mlflow_run = None self._temp_dirs: list[str] = [] self._split_data: dict | None = None self._verification: dict | None = None self._frozen_requirements: str | None = None - self._optimizer_info: dict | None = None - self._scheduler_info: dict | None = None self._git_commit: str = "unknown" self._git_url: str = "unknown" self._git_branch: str = "unknown" - def on_fit_start(self, trainer, pl_module): # noqa: ARG002 - """Run all provenance setup at the start of training.""" + # ── helpers ────────────────────────────────────────────── + + def _gather_from_siblings(self, trainer) -> None: + """Read data already collected by sibling callbacks.""" + from rationai.mlkit.lightning.callbacks.environment import EnvironmentCallback + from rationai.mlkit.lightning.callbacks.dataset_verification import ( + DatasetVerificationCallback, + ) + + 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, pl_module): # noqa: ARG002 + """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.register_dataset import ( _detect_manifest, _lookup_dataset_run, _verify_dataset, ) - # ── Auto-detect everything ────────────────────────────── + # ── Git info (read from MLflow tags set by MLFlowLogger) ── + try: + run = mlflow.active_run() + if run and run.info and run.info.run_id: + client = mlflow.tracking.MlflowClient() + run_data = client.get_run(run.info.run_id) + tags = run_data.data.tags or {} + else: + tags = {} + self._git_commit = tags.get("mlflow.source.git.commit", + tags.get("git.commit", "unknown")) + self._git_url = tags.get("mlflow.source.git.repoUrl", + tags.get("git.repo_url", "unknown")) + self._git_branch = tags.get("mlflow.source.git.branch", + tags.get("git.branch", "unknown")) + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] Git info failed: %s", e) + + # ── User lookup ───────────────────────────────────────── user_run_id, user_tags = _lookup_user_run() - git_commit, git_url, git_branch = _get_git_info() - self._git_commit = git_commit - self._git_url = git_url - self._git_branch = git_branch - hardware = _detect_hardware() + # ── 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 detection & split ─────────────────────────── + # ── Dataset verification & split ──────────────────────── manifest_path = self.manifest_path data_root = self.data_root @@ -676,22 +607,19 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 manifest_path, data_root = _detect_manifest() if manifest_path and data_root: - dataset_run_id = _lookup_dataset_run(manifest_path) + from rationai.mlkit.provenance.register_dataset import ( + _load_manifest, + ) + from sklearn.model_selection import train_test_split + + dataset_run_id = _lookup_dataset_run() verification = _verify_dataset(manifest_path, data_root, dataset_run_id) - self._verification = verification + self._verification = verification or {} - for detail in verification["details"]: - print(f" [ProvenanceCallback] {detail}") - - # Train/test split - from sklearn.model_selection import train_test_split - df = pd.read_csv(manifest_path) - samples = [] - 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"])}) + for detail in self._verification.get("details", []): + log.info(f" [ProvenanceCallback] {detail}") + samples = _load_manifest(manifest_path, data_root) train_samples, test_samples = train_test_split( samples, test_size=self.test_size, @@ -732,29 +660,30 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 }) # Log verification results - 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"]) - ) + 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"]) + ) else: - print("[ProvenanceCallback] WARNING: No manifest.csv found — " - "train/test split not logged.") + log.warning("[ProvenanceCallback] No manifest.csv found — " + "train/test split not logged.") # ── Tags ──────────────────────────────────────────────── tags: dict[str, str] = {} @@ -764,13 +693,14 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 if key in user_tags: tags[key] = user_tags[key] - if dataset_run_id := _lookup_dataset_run(manifest_path): + dataset_run_id = _lookup_dataset_run() + if dataset_run_id: tags["dataset_run_id"] = dataset_run_id tags.update({ - "git_commit": git_commit, - "git_url": git_url, - "git_branch": git_branch, + "git_commit": self._git_commit, + "git_url": self._git_url, + "git_branch": self._git_branch, "prov_start_time": datetime.now(timezone.utc).isoformat(), }) mlflow.set_tags(tags) @@ -793,8 +723,30 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 try: self._frozen_requirements = _snapshot_environment(artifact_dir) mlflow.log_artifacts(artifact_dir, artifact_path="environment") - except Exception: - pass + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] Environment snapshot failed: %s", e) + + # ── lightning hooks ─────────────────────────────────────── + + def on_fit_start(self, trainer, pl_module): # noqa: ARG002 + """Gather environment/verification data from siblings or fall back.""" + # 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, pl_module): # noqa: ARG002 """Log model/optimizer/scheduler summaries and PROV document.""" @@ -808,28 +760,34 @@ def on_fit_end(self, trainer, pl_module): # noqa: ARG002 try: model_summary = _model_summary(pl_module) mlflow.log_params(model_summary) - except Exception: - pass + 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: - self._optimizer_info = _optimizer_summary(opt) - mlflow.log_params(self._optimizer_info) + optimizer_info = _optimizer_summary(opt) + mlflow.log_params(optimizer_info) break - except Exception: - pass + 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 trainer.lr_schedulers: - self._scheduler_info = _scheduler_summary(sched.get("scheduler")) - mlflow.log_params(self._scheduler_info) + scheduler_info = _scheduler_summary(sched.get("scheduler")) + mlflow.log_params(scheduler_info) break - except Exception: - pass + except Exception as e: + if self.strict: + raise + log.warning("[ProvenanceCallback] Scheduler summary failed: %s", e) # ── PROV document + run summary ───────────────────────── try: @@ -878,7 +836,7 @@ def on_fit_end(self, trainer, pl_module): # noqa: ARG002 mlflow.log_artifact(summary_path, artifact_path="provenance") shutil.rmtree(summary_dir, ignore_errors=True) - # ── PROV document ─────────────────────────────────── + # ── 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}", @@ -895,6 +853,7 @@ def on_fit_end(self, trainer, pl_module): # noqa: ARG002 } 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]}" @@ -907,9 +866,11 @@ def on_fit_end(self, trainer, pl_module): # noqa: ARG002 mlflow.log_artifact(prov_path, artifact_path="provenance") shutil.rmtree(prov_dir, ignore_errors=True) - print(f"\n[ProvenanceCallback] Complete → {run_id}") + log.info("[ProvenanceCallback] Complete → %s", run_id) except Exception as e: - print(f"[ProvenanceCallback] WARNING: Could not write provenance artifacts: {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: diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py index dca4e6d..bfd8762 100644 --- a/rationai/mlkit/provenance/__init__.py +++ b/rationai/mlkit/provenance/__init__.py @@ -2,8 +2,7 @@ Submodules: provenance – internal helpers (lookup, verification) - register_dataset – register_dataset (hash-based) and - register_dataset_as_provenance (legacy CSV-based) + register_dataset – register_dataset (hash-based) register_user – register_new_user For automatic provenance capture with Lightning, use @@ -24,7 +23,6 @@ # Now safe to import – all child modules will pick up the env var from .register_dataset import ( # noqa: E402 register_dataset, - register_dataset_as_provenance, verify_dataset, ) from .register_user import register_new_user # noqa: E402 @@ -32,7 +30,6 @@ __all__ = [ # Dataset "register_dataset", - "register_dataset_as_provenance", "verify_dataset", # User registration "register_new_user", diff --git a/rationai/mlkit/provenance/register_dataset.py b/rationai/mlkit/provenance/register_dataset.py index 5949c0d..12b6b08 100644 --- a/rationai/mlkit/provenance/register_dataset.py +++ b/rationai/mlkit/provenance/register_dataset.py @@ -19,7 +19,7 @@ def _lookup_experiment(name): return exp.experiment_id if exp else None -def _lookup_dataset_run(manifest_path: str | None = None) -> str | 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 @@ -56,6 +56,23 @@ def _detect_manifest() -> tuple[str | None, str | None]: return None, None +def load_manifest(manifest_path: str, data_root: str) -> list[dict]: + """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] = [] + 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 + + def verify_dataset( manifest_path: str | None = None, data_root: str | None = None, @@ -125,17 +142,12 @@ def _verify_dataset( result["details"].append(f"Failed to fetch Dataset_Registry run: {e}") return result - # Read manifest and check current file sizes - df = pd.read_csv(manifest_path) - samples = [] + samples = load_manifest(manifest_path, data_root) curr_file_sizes = {} - 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"])}) - basename = os.path.basename(full) - if os.path.isfile(full): - curr_file_sizes[basename] = os.stat(full).st_size + 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 @@ -212,20 +224,14 @@ def register_dataset( if dataset_name is None: dataset_name = os.path.basename(dataset_dir) - # Read manifest and collect per-file metadata (size, last_modified) - df = pd.read_csv(manifest_path) - samples = [] + samples = load_manifest(manifest_path, dataset_dir) file_sizes = {} file_mtimes = {} - for _, row in df.iterrows(): - rel = row["wsi_path"] - full = os.path.join(dataset_dir, rel) if not os.path.isabs(rel) else rel - samples.append({"path": full, "label": int(row["cancer"])}) - - basename = os.path.basename(full) - if os.path.isfile(full): - st = os.stat(full) + 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: @@ -272,58 +278,4 @@ def register_dataset( return run_id -# ── Legacy CSV-based registration (backward compat) ───────────────────── - -def register_dataset_as_provenance(manifest_path, dataset_root, dataset_name, version): - """Legacy CSV-based dataset registration. - - Stores an enriched manifest as an artifact with a ``manifest_uri`` tag. - Does NOT set ``manifest_hash`` / ``samples_hash`` tags — use - :func:`register_dataset` instead for hash-based verification support. - """ - mlflow.set_experiment("Dataset_Registry") - - with mlflow.start_run(run_name=f"Dataset_{dataset_name}_v{version}") as run: - # 1. Indexace v MLflow (pro rychlé hledání/filtrování) - mlflow.set_tag("dataset_name", dataset_name) - mlflow.set_tag("version", version) - - # 2. Metadata pro tracking - mlflow.log_param("dataset_root", dataset_root) - - # 3. Zpracování manifestu a obohacení o metadata (size, mtime) - df = pd.read_csv(manifest_path) - metadata_list = [] - for path in df["wsi_path"]: - full_path = os.path.join(dataset_root, path) if not os.path.isabs(path) else path - if os.path.exists(full_path): - stat = os.stat(full_path) - metadata_list.append({"file_size": stat.st_size, "last_modified": stat.st_mtime}) - else: - metadata_list.append({"file_size": -1, "last_modified": -1}) - - df_enriched = pd.concat([df, pd.DataFrame(metadata_list)], axis=1) - - # 4. Uložení artefaktu (Zlatý zdroj pravdy) - provenance_file = "dataset_provenance.csv" - df_enriched.to_csv(provenance_file, index=False) - mlflow.log_artifact(provenance_file, artifact_path="provenance") - - # 5. Uložení odkazu do tagu (velmi důležité pro automatizaci!) - mlflow.set_tag("manifest_uri", f"runs:/{run.info.run_id}/provenance/{provenance_file}") - - print(f"Dataset '{dataset_name}' (v{version}) úspěšně zaregistrován.") - print(f"Run ID: {run.info.run_id}") - - os.remove(provenance_file) - - -if __name__ == "__main__": - # Příklad použití pro tvůj dataset - register_dataset_as_provenance( - manifest_path="data/dummy_dataset_1/manifest.csv", - dataset_root="data/dummy_dataset_1", - dataset_name="pato_cohort_01", - version="1.0.0" - ) From d580ab6a565568776e6529a164a3ea42e9da6d62 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 10:45:43 +0200 Subject: [PATCH 16/34] fix: update manifest loading and improve MLFlow run handling --- rationai/mlkit/lightning/callbacks/provenance.py | 9 +++++++-- rationai/mlkit/lightning/loggers/mlflow.py | 15 ++++++++++++++- 2 files changed, 21 insertions(+), 3 deletions(-) diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index 176914b..a3296fa 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -608,7 +608,7 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 if manifest_path and data_root: from rationai.mlkit.provenance.register_dataset import ( - _load_manifest, + load_manifest, ) from sklearn.model_selection import train_test_split @@ -619,7 +619,7 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 for detail in self._verification.get("details", []): log.info(f" [ProvenanceCallback] {detail}") - samples = _load_manifest(manifest_path, data_root) + samples = load_manifest(manifest_path, data_root) train_samples, test_samples = train_test_split( samples, test_size=self.test_size, @@ -732,6 +732,11 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 def on_fit_start(self, trainer, pl_module): # noqa: ARG002 """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 + # Check if sibling callbacks are present has_env = any( isinstance(cb, EnvironmentCallback) diff --git a/rationai/mlkit/lightning/loggers/mlflow.py b/rationai/mlkit/lightning/loggers/mlflow.py index 0f56026..2b54893 100644 --- a/rationai/mlkit/lightning/loggers/mlflow.py +++ b/rationai/mlkit/lightning/loggers/mlflow.py @@ -59,7 +59,20 @@ def __init__( def experiment(self) -> MlflowClient: if not self._initialized: exp = super().experiment - mlflow.start_run(self.run_id, log_system_metrics=self.log_system_metrics) + # Resume the run created in __init__. + # We avoid ``mlflow.start_run`` here because its global fluent-API + # experiment state may have been clobbered by prior calls + # (register_dataset, register_new_user) that created runs in + # different experiments. Instead we set the run status directly + # via the client API. + client = mlflow.tracking.MlflowClient() + client.set_terminated_status( + self.run_id, + "RUNNING", + int(__import__("time").time() * 1000), + ) + self._active_run_id = self.run_id + self._initialized = True return exp return super().experiment From 2be0a0458ac1ba332276862bfd4ad69c3f1ef45a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 11:02:51 +0200 Subject: [PATCH 17/34] feat: enhance MLFlow run handling in ProvenanceCallback and MLFlowLogger --- .../mlkit/lightning/callbacks/provenance.py | 32 ++++++++++++++++++- rationai/mlkit/lightning/loggers/mlflow.py | 23 +++++++------ 2 files changed, 42 insertions(+), 13 deletions(-) diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index a3296fa..1400066 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -730,6 +730,28 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 # ── lightning hooks ─────────────────────────────────────── + def _ensure_active_run(self, trainer) -> 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, pl_module): # noqa: ARG002 """Gather environment/verification data from siblings or fall back.""" from rationai.mlkit.lightning.callbacks.dataset_verification import ( @@ -737,6 +759,9 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 ) 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) @@ -757,7 +782,12 @@ def on_fit_end(self, trainer, pl_module): # noqa: ARG002 """Log model/optimizer/scheduler summaries and PROV document.""" active_run = mlflow.active_run() if not active_run: - return + # Fallback: try to get run_id from the logger directly + self._ensure_active_run(trainer) + if self._run_id: + active_run = mlflow.tracking.MlflowClient().get_run(self._run_id) + else: + return run_id = active_run.info.run_id # ── Model summary ─────────────────────────────────────── diff --git a/rationai/mlkit/lightning/loggers/mlflow.py b/rationai/mlkit/lightning/loggers/mlflow.py index 2b54893..7a2d4bf 100644 --- a/rationai/mlkit/lightning/loggers/mlflow.py +++ b/rationai/mlkit/lightning/loggers/mlflow.py @@ -59,19 +59,18 @@ def __init__( def experiment(self) -> MlflowClient: if not self._initialized: exp = super().experiment - # Resume the run created in __init__. - # We avoid ``mlflow.start_run`` here because its global fluent-API - # experiment state may have been clobbered by prior calls - # (register_dataset, register_new_user) that created runs in - # different experiments. Instead we set the run status directly - # via the client API. + # Establish a fluent-API active run context so that callbacks + # using ``mlflow.log_artifact``, ``mlflow.log_params``, etc. + # work correctly. We push an ActiveRun onto the thread-local + # stack directly rather than calling ``mlflow.start_run()`` which + # would re-resolve the experiment name and potentially create a + # new run in the wrong experiment (prior calls like + # register_dataset may have clobbered global fluent state). client = mlflow.tracking.MlflowClient() - client.set_terminated_status( - self.run_id, - "RUNNING", - int(__import__("time").time() * 1000), - ) - self._active_run_id = self.run_id + run_info = client.get_run(self.run_id) + active = mlflow.tracking.fluent.ActiveRun(run_info) + stack = mlflow.tracking.fluent._active_run_stack.get() + stack.append(active) self._initialized = True return exp From 96f276ce6173e72b41bd9167d5643c47d5d46ed0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 11:24:20 +0200 Subject: [PATCH 18/34] fix: update .gitignore and refactor MLFlowLogger run initialization --- .gitignore | 2 +- rationai/mlkit/__init__.py | 5 +++-- rationai/mlkit/lightning/loggers/mlflow.py | 15 +++------------ 3 files changed, 7 insertions(+), 15 deletions(-) diff --git a/.gitignore b/.gitignore index 6591481..5611bec 100644 --- a/.gitignore +++ b/.gitignore @@ -174,4 +174,4 @@ cython_debug/ test_data mlflow.db mlartifacts -example_provenance.py \ No newline at end of file +example* \ No newline at end of file diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index a907ff0..241e3f5 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -1,7 +1,8 @@ """rationai.mlkit — ML toolkit with provenance tracking.""" -from rationai.mlkit.stream import StreamCapture, StreamLogger +from rationai.mlkit.autolog import autolog from rationai.mlkit.provenance.register_dataset import register_dataset +from rationai.mlkit.stream import StreamCapture, StreamLogger __all__ = [ "StreamCapture", @@ -29,7 +30,7 @@ def __getattr__(name): - if name in ("Trainer", "MLFlowLogger", "MultiloaderLifecycle", "autolog", "with_cli_args"): + if name in ("Trainer", "MLFlowLogger", "MultiloaderLifecycle", "with_cli_args"): import importlib _mod = importlib.import_module("rationai.mlkit.lightning") return getattr(_mod, name) diff --git a/rationai/mlkit/lightning/loggers/mlflow.py b/rationai/mlkit/lightning/loggers/mlflow.py index 7a2d4bf..cf4b6ac 100644 --- a/rationai/mlkit/lightning/loggers/mlflow.py +++ b/rationai/mlkit/lightning/loggers/mlflow.py @@ -59,18 +59,9 @@ def __init__( def experiment(self) -> MlflowClient: if not self._initialized: exp = super().experiment - # Establish a fluent-API active run context so that callbacks - # using ``mlflow.log_artifact``, ``mlflow.log_params``, etc. - # work correctly. We push an ActiveRun onto the thread-local - # stack directly rather than calling ``mlflow.start_run()`` which - # would re-resolve the experiment name and potentially create a - # new run in the wrong experiment (prior calls like - # register_dataset may have clobbered global fluent state). - client = mlflow.tracking.MlflowClient() - run_info = client.get_run(self.run_id) - active = mlflow.tracking.fluent.ActiveRun(run_info) - stack = mlflow.tracking.fluent._active_run_stack.get() - stack.append(active) + # 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 From 3ece4bdf23b89cb79d6fcbfcc2375fce2949af30 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 11:32:45 +0200 Subject: [PATCH 19/34] docs: update README for user and dataset registration process --- README.md | 278 +++++++++++++++++++++++++++--------------------------- 1 file changed, 140 insertions(+), 138 deletions(-) diff --git a/README.md b/README.md index c66577a..f9f07cc 100644 --- a/README.md +++ b/README.md @@ -5,7 +5,7 @@ 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 — the user only writes their training loop. +`@autolog` decorator and `ProvenanceCallback` — the user only writes their training loop. --- @@ -27,173 +27,171 @@ mlflow ui --host 0.0.0.0 --port 5000 # → http://localhost:5000 ## Quick start -### 1. Run the demo +### 1. Register users and dataset (one-time setup) -The easiest way to see everything in action: +Before training, register researchers and datasets so provenance can reference them: ```bash -# Full pipeline — uploads to http://localhost:5000 -python demo.py - -# Push to a different MLflow server -python demo.py --uri http://your-server:5000 - -# Run unit tests instead -python demo.py --test +# Edit example_provenance_setup.py with your own data, then run: +python example_provenance_setup.py ``` -The demo exercises all components: - -| Step | Feature | What it shows | -|---|---|---| -| 1 | Dummy data creation | `test_data/dummy_dataset_*` with manifests | -| 2 | **StreamCapture** | ANSI-aware stdout/stderr capture | -| 3 | **AggregatedMetricCollection** | Tile → slide metric aggregation | -| 4 | **NestedMetricCollection** | Per-slide multiclass metrics | -| 5 | **StratifiedBatchSampler** | Balanced class batches | -| 6 | **Provenance** (`@autolog`) | Full training run with auto-captured provenance | -| 7 | **Lightning** (Trainer + MLFlowLogger) | Lightning training with full provenance tracking | - -Both steps 6 and 7 upload to MLflow with identical provenance depth: -model params, GPU/CPU info, optimizer config, train/test split stats, -environment freeze, console logs, and the PROV-O document. +This creates runs in the `User_Registry` and `Dataset_Registry` MLflow experiments. -### 2. Create dummy data +### 2. Run a training experiment -Generate test datasets (no shell script needed): +Use `@autolog` + `ProvenanceCallback` in a Hydra-based training script: ```bash -# Default: 2 datasets × 50 WSIs each -python dummy_dataset_create.py - -# Custom: 3 datasets with 100 WSIs each -python dummy_dataset_create.py --datasets 3 --wsis-per-dataset 100 - -# Add more without deleting existing ones -python dummy_dataset_create.py -d 1 -w 30 --no-clean +python example_provenance_train.py ``` -This creates the following structure under `data/`: +Check the MLflow UI at http://localhost:5000 to see the full provenance graph. -``` -data/ -├── dummy_dataset_1/ -│ ├── manifest.csv # patient_id, wsi_path (relative), cancer -│ └── wsis/ -│ ├── PAT_001.tiff -│ ├── PAT_002.tiff -│ └── ... -├── dummy_dataset_2/ -│ ├── manifest.csv -│ └── wsis/ -│ └── ... -``` +See [example_provenance_setup.py](example_provenance_setup.py) and +[example_provenance_train.py](example_provenance_train.py) for complete working examples. -**Options:** +--- -| Flag | Default | Description | -|---|---|---| -| `--datasets N` / `-d` | `2` | Number of dataset folders | -| `--wsis-per-dataset N` / `-w` | `50` | WSIs per dataset | -| `--data-dir DIR` | `data/` | Parent directory | -| `--seed N` | `42` | Reproducibility seed | -| `--no-clean` | off | Keep existing datasets, append new ones | -| `--img-size N` | `128` | Pixel size of dummy TIFF images | +## API reference -### 3. Register a user +### User registration -Edit the variables in `user_to_mlflow.py` and run: +Register a researcher into the `User_Registry` experiment: -```bash -python user_to_mlflow.py +```python +from rationai.mlkit.provenance import register_new_user + +register_new_user( + username="jiribuchta", + real_name="Jiří Buchta", + email="524981@mail.muni.cz", + organization="RationAI", + lead_name="Tomáš Brázdil", + lead_email="brazdil@muni.cz", +) ``` -This creates a run in the **User_Registry** experiment with your identity tags. +### Dataset registration -### 4. Run a training experiment +Register a dataset (requires `manifest.csv` in the dataset directory): -Use the `@autolog` decorator from `rationai.mlkit.provenance` in your training script, then: +```python +from rationai.mlkit.provenance import register_dataset -```bash -python your_experiment.py +run_id = register_dataset( + dataset_dir="data/cohorts/pato_01", + dataset_name="pato_cohort_01", + version="2.0", +) +print(f"Registered as {run_id}") ``` -See the [API reference](#provenance--autolog-decorator) below for details. - ---- - -## API reference - -### Provenance — `@autolog` decorator - -Full auto-capture for plain PyTorch training runs: +Verify a dataset's file integrity: ```python -from rationai.mlkit.provenance import autolog +from rationai.mlkit.provenance import verify_dataset -@autolog(model_name="my_model_v1", experiment_name="My_Experiment") -def train(run): - model = build_model() - run.register_model(model) - - optimizer = optim.SGD(model.parameters(), lr=1e-3, momentum=0.9) - run.register_optimizer(optimizer) +result = verify_dataset(manifest_path="data/cohorts/pato_01/manifest.csv") +print(result["verified"]) # True if all files match +``` - scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.5) - run.register_scheduler(scheduler) +### Provenance — `@autolog` decorator (Hydra training scripts) - for epoch in range(50): - loss = train_epoch(model, loader, optimizer) - run.log_metrics({"train_loss": loss}, step=epoch) +Full auto-capture for Hydra-based training runs: - run.save_model(model) +```python +from rationai.mlkit import Trainer, autolog +from rationai.mlkit.lightning.callbacks import ProvenanceCallback -if __name__ == "__main__": - train() +@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 git config → linked to `User_Registry` run | -| **Dataset** | Latest `Dataset_Registry` run | +| **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 | +| **Docker** | Container ID, image name + hash (if in container) | | **Git** | Commit, branch, remote URL | -| **Environment** | Frozen `requirements.txt` + `pyproject.toml` / `uv.lock` | +| **Environment** | Frozen `requirements.txt` | | **Console output** | stdout/stderr → `logs/console.log` artifact (ANSI-aware) | | **PROV document** | OpenProvenance JSON → `provenance/prov.json` artifact | -### Provenance + Lightning +### Lightning callbacks + +#### ProvenanceCallback -Wrap Lightning training in `@autolog` and pass the active run to `MLFlowLogger`: +Drop-in callback that captures full PROV-O provenance for every training run: ```python -import mlflow -from rationai.mlkit import Trainer, MLFlowLogger -from rationai.mlkit.provenance import autolog +from rationai.mlkit.lightning.callbacks import ProvenanceCallback -@autolog(model_name="my_lightning_model", experiment_name="My_Experiment") -def train(run): - model = MyLightningModule() - run.register_model(model) - run.register_optimizer(model.configure_optimizers()) +callback = ProvenanceCallback( + model_name="resnet_v1", + experiment_name="Training_Pipeline", +) - # Reuse the @autolog run so Lightning logs to the same provenance-tracked run - logger = MLFlowLogger(experiment_name="My_Experiment", run_id=mlflow.active_run().info.run_id) +trainer = Trainer(callbacks=[callback], ...) +``` - trainer = Trainer(logger=logger, max_epochs=50) - trainer.fit(model, train_loader) +| 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 - run.save_model(model) +callback = EnvironmentCallback( + skip_hardware=False, # capture GPU/CPU/RAM info + snapshot_env=True, # freeze requirements.txt +) ``` -This gives you the same full provenance depth as plain PyTorch — GPU info, -model architecture, optimizer config, environment freeze, PROV document — -plus Lightning's native metric logging via `self.log()`. +#### 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 @@ -297,8 +295,10 @@ dataset = MetaTiledSlides( | `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 | -| `lightning_autolog` | `from rationai.mlkit.lightning import autolog` | Lightning-specific autolog decorator | -| `with_cli_args` | `from rationai.mlkit.lightning import with_cli_args` | Programmatic config injection (Hydra) | +| `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) | --- @@ -306,45 +306,47 @@ dataset = MetaTiledSlides( ``` . -├── demo.py # End-to-end demo (all components) -├── dummy_dataset_create.py # Dummy data generator (CLI) -├── user_to_mlflow.py # User registration script -├── pyproject.toml # Project metadata + deps -├── tests/ -│ └── test_all.py # Unit test suite -├── test_data/ # Dummy datasets (gitignored) -│ ├── dummy_dataset_1/ -│ │ ├── manifest.csv -│ │ └── wsis/ -│ └── ... +├── 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 Lightning import) - ├── autolog.py # Re-exports lightning.autolog - ├── with_cli_args.py # Re-exports lightning.with_cli_args - ├── stream/ # ANSI-aware console capture + ├── __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 + ├── metrics/ # Slide-level metric aggregation │ ├── aggregated_metric_collection.py │ ├── nested_metric_collection.py │ ├── aggregators.py │ └── lazy_metric_dict.py - ├── data/ # Data utilities + ├── data/ # Data utilities + │ ├── shard_parquet.py │ ├── samplers/ │ │ └── stratified_batch_sampler.py │ └── datasets/ │ ├── meta_tiled_slides.py - │ └── openslide_tiles_dataset.py - └── lightning/ # Lightning + Hydra integration - ├── autolog.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.py # MLFlowLogger (checkpoint sync, git tags) ``` --- From c27f54d86030e8d2ffa20b5d60e69f576a464e98 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 11:41:20 +0200 Subject: [PATCH 20/34] feat: implement PROV-O document generation for user and dataset registration --- rationai/mlkit/provenance/__init__.py | 13 +- rationai/mlkit/provenance/prov.py | 316 ++++++++++++++++++ rationai/mlkit/provenance/register_dataset.py | 31 +- rationai/mlkit/provenance/register_user.py | 78 ++++- 4 files changed, 416 insertions(+), 22 deletions(-) create mode 100644 rationai/mlkit/provenance/prov.py diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py index bfd8762..5728b71 100644 --- a/rationai/mlkit/provenance/__init__.py +++ b/rationai/mlkit/provenance/__init__.py @@ -1,9 +1,9 @@ """Provenance tracking — PROV-O-aware logging to MLflow. Submodules: - provenance – internal helpers (lookup, verification) - register_dataset – register_dataset (hash-based) - register_user – register_new_user + prov – PROV-O document builders (W3C PROV compatible) + register_dataset – register_dataset (hash-based, emits prov.json) + register_user – register_new_user (emits prov.json) For automatic provenance capture with Lightning, use :class:`~rationai.mlkit.lightning.callbacks.provenance.ProvenanceCallback`. @@ -21,6 +21,10 @@ os.environ["MLFLOW_TRACKING_URI"] = "http://localhost:5000" # Now safe to import – all child modules will pick up the env var +from .prov import ( # noqa: E402 + build_dataset_prov, + build_user_prov, +) from .register_dataset import ( # noqa: E402 register_dataset, verify_dataset, @@ -33,6 +37,9 @@ "verify_dataset", # User registration "register_new_user", + # PROV document builders + "build_user_prov", + "build_dataset_prov", ] diff --git a/rationai/mlkit/provenance/prov.py b/rationai/mlkit/provenance/prov.py new file mode 100644 index 0000000..7dc0bbf --- /dev/null +++ b/rationai/mlkit/provenance/prov.py @@ -0,0 +1,316 @@ +"""Shared PROV-O document builder. + +Produces OpenProvenance-compatible JSON bundles compatible with the +Java ``prov_mlflow`` tool and used by both registration runs and +training provenance callbacks. +""" + +from __future__ import annotations + +import json +import os +import re +from datetime import datetime, timezone +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: + merged = {**_DEFAULT_PROV_PREFIXES, **json.loads(env_json)} + return merged + except json.JSONDecodeError: + pass + return _DEFAULT_PROV_PREFIXES + + +# ────────────────────────────────────────────── +# Small helpers used inside the PROV document +# ────────────────────────────────────────────── + +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, type_prefix: str = "xsd", type_local: str = "string") -> list[str]: + return [str(value)] + + +def _qualified_name(type_prefix: str, type_local: str) -> dict: + 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 = datetime.fromtimestamp(ts_ms / 1000, tz=timezone.utc) + else: + dt = datetime.now(timezone.utc) + return dt.strftime("%Y-%m-%dT%H:%M:%S.000+00:00") + + +# ────────────────────────────────────────────── +# PROV document builders +# ────────────────────────────────────────────── + +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: + """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] = {} + activities: dict[str, dict] = {} + agents: dict[str, dict] = {} + was_associated_with: dict[str, dict] = {} + was_generated_by: dict[str, dict] = {} + + 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] = {} + 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, object] = {} + 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] = {} + 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, object] = {} + 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, object] = {"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}} + + +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: + """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] = {} + activities: dict[str, dict] = {} + used: dict[str, dict] = {} + was_generated_by: dict[str, dict] = {} + + 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] = { + "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, object] = {} + 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] = {} + 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) + + # Store per-file sizes as a single string value (matches prov_mlflow convention) + 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, object] = {} + 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, object] = {"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}} diff --git a/rationai/mlkit/provenance/register_dataset.py b/rationai/mlkit/provenance/register_dataset.py index 12b6b08..cb35bfd 100644 --- a/rationai/mlkit/provenance/register_dataset.py +++ b/rationai/mlkit/provenance/register_dataset.py @@ -257,11 +257,34 @@ def register_dataset( "file_mtimes": json.dumps(file_mtimes), }) - # Save provenance JSON as an artifact for offline verification + # ── PROV-O document (W3C PROV-O compatible) ──────────── + from rationai.mlkit.provenance.prov import build_dataset_prov # noqa: PLC0415 + + 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, "dataset_provenance.json") + 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) + + # ── Legacy dataset provenance JSON (kept for 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") + with open(legacy_path, "w") as f: json.dump({ "dataset_name": dataset_name, "version": version, @@ -270,8 +293,8 @@ def register_dataset( "file_mtimes": file_mtimes, "num_samples": len(samples), }, f, indent=2) - mlflow.log_artifact(prov_path, artifact_path="provenance") - shutil.rmtree(prov_dir, ignore_errors=True) + mlflow.log_artifact(legacy_path, artifact_path="provenance") + shutil.rmtree(legacy_prov_dir, ignore_errors=True) mlflow.end_run() print(f" [register_dataset] {dataset_name} v{version} → run_id={run_id}") diff --git a/rationai/mlkit/provenance/register_user.py b/rationai/mlkit/provenance/register_user.py index b4ca05b..33cd809 100644 --- a/rationai/mlkit/provenance/register_user.py +++ b/rationai/mlkit/provenance/register_user.py @@ -1,37 +1,85 @@ +"""Register a researcher into MLflow's User_Registry experiment.""" + +from __future__ import annotations + +import json +import os +import shutil +import uuid + import mlflow -def register_new_user(username, real_name, email, organization, lead_name, lead_email): - """ - Registruje uživatele do MLflow experimentu 'User_Registry' s kompletními údaji. + +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 = mlflow.get_experiment_by_name(experiment_name) - + if not experiment: - experiment_id = mlflow.create_experiment(experiment_name) - print(f"Vytvořen experiment: {experiment_name}") - else: - experiment_id = experiment.experiment_id + mlflow.create_experiment(experiment_name) + + with mlflow.start_run( + experiment_id=mlflow.get_experiment_by_name(experiment_name).experiment_id, + run_name=f"User_{username}", + ) as run: + run_id = run.info.run_id - # Vytvoření unikátního runu pro tohoto uživatele - with mlflow.start_run(experiment_id=experiment_id, run_name=f"User_{username}"): mlflow.set_tags({ "username": username, "real_name": real_name, "email": email, "organization": organization, "lead_name": lead_name, - "lead_email": lead_email + "lead_email": lead_email, }) - print(f"Uživatel {real_name} ({username}) byl úspěšně zaregistrován.") + + # ── PROV-O document ──────────────────────────────── + from rationai.mlkit.provenance.prov import build_user_prov # noqa: PLC0415 + + 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) + 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) + + print(f" [register_new_user] {real_name} ({username}) → run_id={run_id}") + return run_id + if __name__ == "__main__": - # Příklad registrace register_new_user( username="jiribuchta", real_name="Jiří Buchta", email="524981@mail.muni.cz", organization="RationAI", lead_name="Tomáš Brázdil", - lead_email="brazdil@muni.cz" - ) \ No newline at end of file + lead_email="brazdil@muni.cz", + ) From bb719a12bd180ee630c74ba4874c158de43382e2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 13:59:59 +0200 Subject: [PATCH 21/34] feat: enhance dataset and user registration with improved logging and error handling --- .gitignore | 2 +- README.md | 23 ++-- .../mlkit/lightning/callbacks/provenance.py | 2 +- rationai/mlkit/provenance/__init__.py | 10 +- rationai/mlkit/provenance/prov.py | 52 ++++----- rationai/mlkit/provenance/register_dataset.py | 108 ++++++++++-------- rationai/mlkit/provenance/register_user.py | 32 +++--- 7 files changed, 123 insertions(+), 106 deletions(-) diff --git a/.gitignore b/.gitignore index 5611bec..261446d 100644 --- a/.gitignore +++ b/.gitignore @@ -174,4 +174,4 @@ cython_debug/ test_data mlflow.db mlartifacts -example* \ No newline at end of file +/example* \ No newline at end of file diff --git a/README.md b/README.md index f9f07cc..b366ae2 100644 --- a/README.md +++ b/README.md @@ -20,7 +20,7 @@ uv sync Start the MLflow server: ```bash -mlflow ui --host 0.0.0.0 --port 5000 # → http://localhost:5000 +mlflow ui --host 127.0.0.1 --port 5000 # → http://localhost:5000 ``` --- @@ -63,12 +63,12 @@ Register a researcher into the `User_Registry` experiment: from rationai.mlkit.provenance import register_new_user register_new_user( - username="jiribuchta", - real_name="Jiří Buchta", - email="524981@mail.muni.cz", - organization="RationAI", - lead_name="Tomáš Brázdil", - lead_email="brazdil@muni.cz", + 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", ) ``` @@ -101,7 +101,10 @@ print(result["verified"]) # True if all files match 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) @@ -304,7 +307,7 @@ dataset = MetaTiledSlides( ## Project structure -``` +```text . ├── example_provenance_setup.py # Setup: register users & dataset ├── example_provenance_train.py # Training with @autolog + ProvenanceCallback @@ -368,7 +371,7 @@ PROV graph reconstruction machine-readable. Each training run emits a self-contained OpenProvenance JSON document at `provenance/prov.json` (MLflow artifact). This document is compatible with -the [prov_mlflow](https://github.com/jiribuchta/prov_mlflow) Java tool and +the `prov_mlflow` Java tool and follows the W3C PROV-O standard. ### Structure @@ -398,7 +401,7 @@ with 10 namespace prefixes and 7 sections: ### Compatibility -The generated document matches the Java [`prov_mlflow`](https://github.com/jiribuchta/prov_mlflow) +The generated document matches the Java `prov_mlflow` output format: - Bundle-wrapped JSON structure diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index 1400066..c7ac532 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -815,7 +815,7 @@ def on_fit_end(self, trainer, pl_module): # noqa: ARG002 # ── Scheduler summary ─────────────────────────────────── if self.register_scheduler and pl_module is not None: try: - for sched in trainer.lr_schedulers: + for sched in getattr(trainer, "lr_schedulers", []): scheduler_info = _scheduler_summary(sched.get("scheduler")) mlflow.log_params(scheduler_info) break diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py index 5728b71..d60266e 100644 --- a/rationai/mlkit/provenance/__init__.py +++ b/rationai/mlkit/provenance/__init__.py @@ -1,9 +1,9 @@ -"""Provenance tracking — PROV-O-aware logging to MLflow. +"""Provenance tracking - PROV-O-aware logging to MLflow. Submodules: - prov – PROV-O document builders (W3C PROV compatible) - register_dataset – register_dataset (hash-based, emits prov.json) - register_user – register_new_user (emits prov.json) + prov - PROV-O document builders (W3C PROV compatible) + register_dataset - register_dataset (hash-based, emits prov.json) + register_user - register_new_user (emits prov.json) For automatic provenance capture with Lightning, use :class:`~rationai.mlkit.lightning.callbacks.provenance.ProvenanceCallback`. @@ -20,7 +20,7 @@ if "MLFLOW_TRACKING_URI" not in os.environ: os.environ["MLFLOW_TRACKING_URI"] = "http://localhost:5000" -# Now safe to import – all child modules will pick up the env var +# Now safe to import - all child modules will pick up the env var from .prov import ( # noqa: E402 build_dataset_prov, build_user_prov, diff --git a/rationai/mlkit/provenance/prov.py b/rationai/mlkit/provenance/prov.py index 7dc0bbf..7964d6b 100644 --- a/rationai/mlkit/provenance/prov.py +++ b/rationai/mlkit/provenance/prov.py @@ -10,7 +10,7 @@ import json import os import re -from datetime import datetime, timezone +import datetime as _dt from typing import Any @@ -62,19 +62,19 @@ def _qualified(prefix: str, local: str) -> str: return f"{prefix}:{local}" -def _typed_value(value: Any, type_prefix: str = "xsd", type_local: str = "string") -> list[str]: +def _typed_value(value: Any) -> list[str]: return [str(value)] -def _qualified_name(type_prefix: str, type_local: str) -> dict: +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 = datetime.fromtimestamp(ts_ms / 1000, tz=timezone.utc) + dt = _dt.datetime.fromtimestamp(ts_ms / 1000, tz=_dt.UTC) else: - dt = datetime.now(timezone.utc) + dt = _dt.datetime.now(_dt.UTC) return dt.strftime("%Y-%m-%dT%H:%M:%S.000+00:00") @@ -91,7 +91,7 @@ def build_user_prov( lead_name: str | None = None, lead_email: str | None = None, prov_prefixes: dict[str, str] | None = None, -) -> dict: +) -> dict[str, Any]: """Build a PROV document for a user registration run. Produces an ``agent`` entity representing the researcher and links it @@ -111,11 +111,11 @@ def build_user_prov( main_act_local = f"UserReg_{run_id[:8]}" main_act_id = _qualified("blank", main_act_local) - entities: dict[str, dict] = {} - activities: dict[str, dict] = {} - agents: dict[str, dict] = {} - was_associated_with: dict[str, dict] = {} - was_generated_by: dict[str, dict] = {} + 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: @@ -126,7 +126,7 @@ def _blank_rel_id() -> str: now = _iso_timestamp() # ── AGENT ────────────────────────────────────────────── - agent_props: dict[str, list] = {} + agent_props: dict[str, list[Any]] = {} agent_props["schema:name"] = _typed_value(real_name) agent_props["schema:email"] = _typed_value(email) if organization: @@ -135,7 +135,7 @@ def _blank_rel_id() -> str: agents[agent_id] = agent_props # ── ACTIVITY (the registration action) ──────────────── - run_activity: dict[str, object] = {} + run_activity: dict[str, Any] = {} run_activity["prov:type"] = [_qualified_name("schema", "Action")] run_activity["prov:startTime"] = [now] run_activity["prov:endTime"] = [now] @@ -143,7 +143,7 @@ def _blank_rel_id() -> str: activities[run_act_id] = run_activity # ── CPM METADATA ENTITY ─────────────────────────────── - meta_entity: dict[str, list] = {} + 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) @@ -157,7 +157,7 @@ def _blank_rel_id() -> str: entities[meta_id] = meta_entity # ── CPM MAIN ACTIVITY ──────────────────────────────── - main_activity: dict[str, object] = {} + main_activity: dict[str, Any] = {} main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] main_activity["cpm:referencedMetaBundleId"] = [ {"type": "prov:QUALIFIED_NAME", "$": meta_id} @@ -178,7 +178,7 @@ def _blank_rel_id() -> str: } # ── ASSEMBLE BUNDLE ─────────────────────────────────── - inner: dict[str, object] = {"prefix": prefixes} + inner: dict[str, Any] = {"prefix": prefixes} if entities: inner["entity"] = entities if activities: @@ -205,7 +205,7 @@ def build_dataset_prov( file_sizes: dict[str, int], manifest_path: str | None = None, prov_prefixes: dict[str, str] | None = None, -) -> dict: +) -> dict[str, Any]: """Build a PROV document for a dataset registration run. Produces a ``dataset`` entity (``sosa:Sample``) linked to the @@ -225,10 +225,10 @@ def build_dataset_prov( main_act_local = f"DatasetReg_{run_id[:8]}" main_act_id = _qualified("blank", main_act_local) - entities: dict[str, dict] = {} - activities: dict[str, dict] = {} - used: dict[str, dict] = {} - was_generated_by: dict[str, dict] = {} + 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: @@ -239,7 +239,7 @@ def _blank_rel_id() -> str: now = _iso_timestamp() # ── DATASET ENTITY ─────────────────────────────────── - ds_props: dict[str, list] = { + ds_props: dict[str, list[Any]] = { "schema:name": _typed_value(dataset_name), "prov:type": [_qualified_name("sosa", "Sample")], "dct:description": _typed_value( @@ -251,7 +251,7 @@ def _blank_rel_id() -> str: entities[ds_id] = ds_props # ── ACTIVITY (the registration action) ──────────────── - run_activity: dict[str, object] = {} + run_activity: dict[str, Any] = {} run_activity["prov:type"] = [_qualified_name("schema", "Action")] run_activity["prov:startTime"] = [now] run_activity["prov:endTime"] = [now] @@ -267,7 +267,7 @@ def _blank_rel_id() -> str: } # ── CPM METADATA ENTITY ─────────────────────────────── - meta_entity: dict[str, list] = {} + 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) @@ -285,7 +285,7 @@ def _blank_rel_id() -> str: entities[meta_id] = meta_entity # ── CPM MAIN ACTIVITY ──────────────────────────────── - main_activity: dict[str, object] = {} + main_activity: dict[str, Any] = {} main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] main_activity["cpm:referencedMetaBundleId"] = [ {"type": "prov:QUALIFIED_NAME", "$": meta_id} @@ -302,7 +302,7 @@ def _blank_rel_id() -> str: } # ── ASSEMBLE BUNDLE ─────────────────────────────────── - inner: dict[str, object] = {"prefix": prefixes} + inner: dict[str, Any] = {"prefix": prefixes} if entities: inner["entity"] = entities if activities: diff --git a/rationai/mlkit/provenance/register_dataset.py b/rationai/mlkit/provenance/register_dataset.py index cb35bfd..47abea1 100644 --- a/rationai/mlkit/provenance/register_dataset.py +++ b/rationai/mlkit/provenance/register_dataset.py @@ -242,61 +242,69 @@ def register_dataset( 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 - 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 (W3C PROV-O compatible) ──────────── - from rationai.mlkit.provenance.prov import build_dataset_prov # noqa: PLC0415 - - 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, - ) + 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), + }) - 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") - 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) - - # ── Legacy dataset provenance JSON (kept for 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") - with open(legacy_path, "w") as f: - json.dump({ + mlflow.set_tags({ "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") - shutil.rmtree(legacy_prov_dir, ignore_errors=True) + "file_sizes": json.dumps(file_sizes), + "file_mtimes": json.dumps(file_mtimes), + }) + + # ── PROV-O document (W3C PROV-O compatible) ──────────── + from rationai.mlkit.provenance.prov import build_dataset_prov # noqa: PLC0415 + + 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 (kept for 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() - mlflow.end_run() print(f" [register_dataset] {dataset_name} v{version} → run_id={run_id}") return run_id diff --git a/rationai/mlkit/provenance/register_user.py b/rationai/mlkit/provenance/register_user.py index 33cd809..aea0f2f 100644 --- a/rationai/mlkit/provenance/register_user.py +++ b/rationai/mlkit/provenance/register_user.py @@ -28,18 +28,20 @@ def register_new_user( The MLflow run_id of the registration run. """ experiment_name = "User_Registry" - experiment = mlflow.get_experiment_by_name(experiment_name) - - if not experiment: - mlflow.create_experiment(experiment_name) + try: + experiment_id = mlflow.create_experiment(experiment_name) + except Exception: + # Already exists (possibly created concurrently) + exp = mlflow.get_experiment_by_name(experiment_name) + experiment_id = exp.experiment_id if exp else None with mlflow.start_run( - experiment_id=mlflow.get_experiment_by_name(experiment_name).experiment_id, + experiment_id=experiment_id, run_name=f"User_{username}", ) as run: run_id = run.info.run_id - mlflow.set_tags({ + mlflow.log_params({ "username": username, "real_name": real_name, "email": email, @@ -47,6 +49,10 @@ def register_new_user( "lead_name": lead_name, "lead_email": lead_email, }) + mlflow.set_tags({ + "username": username, + "organization": organization, + }) # ── PROV-O document ──────────────────────────────── from rationai.mlkit.provenance.prov import build_user_prov # noqa: PLC0415 @@ -70,16 +76,16 @@ def register_new_user( mlflow.log_artifact(prov_path, artifact_path="provenance") shutil.rmtree(prov_dir, ignore_errors=True) - print(f" [register_new_user] {real_name} ({username}) → run_id={run_id}") + print(f" [register_new_user] {username} → run_id={run_id}") return run_id if __name__ == "__main__": register_new_user( - username="jiribuchta", - real_name="Jiří Buchta", - email="524981@mail.muni.cz", - organization="RationAI", - lead_name="Tomáš Brázdil", - lead_email="brazdil@muni.cz", + 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", ) From e695e3ae8e4a9587b657a07143d976a10ed133ad Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 14:24:46 +0200 Subject: [PATCH 22/34] feat: enhance type hinting and improve code readability across multiple modules --- rationai/mlkit/__init__.py | 27 ++-- .../callbacks/dataset_verification.py | 9 +- .../mlkit/lightning/callbacks/environment.py | 70 ++++----- .../mlkit/lightning/callbacks/provenance.py | 136 +++++++++--------- rationai/mlkit/provenance/__init__.py | 20 +-- rationai/mlkit/provenance/prov.py | 14 +- rationai/mlkit/provenance/register_dataset.py | 29 ++-- rationai/mlkit/provenance/register_user.py | 3 +- 8 files changed, 160 insertions(+), 148 deletions(-) diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index 241e3f5..161ff09 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -1,35 +1,38 @@ """rationai.mlkit — ML toolkit with provenance tracking.""" +from typing import Any + from rationai.mlkit.autolog import autolog from rationai.mlkit.provenance.register_dataset import register_dataset from rationai.mlkit.stream import StreamCapture, StreamLogger + __all__ = [ - "StreamCapture", - "StreamLogger", "AggregatedMetricCollection", "Aggregator", + "LazyMetricDict", + "MLFlowLogger", "MaxAggregator", "MeanAggregator", "MeanPoolMaxAggregator", - "TopKAggregator", - "NestedMetricCollection", - "LazyMetricDict", - "StratifiedBatchSampler", - "PDMStratifiedBatchSampler", "MetaTiledSlides", - "OpenSlideTilesDataset", - "Trainer", - "MLFlowLogger", "MultiloaderLifecycle", + "NestedMetricCollection", + "OpenSlideTilesDataset", + "PDMStratifiedBatchSampler", "ProvenanceCallback", + "StratifiedBatchSampler", + "StreamCapture", + "StreamLogger", + "TopKAggregator", + "Trainer", "autolog", - "with_cli_args", "register_dataset", + "with_cli_args", ] -def __getattr__(name): +def __getattr__(name: str) -> Any: if name in ("Trainer", "MLFlowLogger", "MultiloaderLifecycle", "with_cli_args"): import importlib _mod = importlib.import_module("rationai.mlkit.lightning") diff --git a/rationai/mlkit/lightning/callbacks/dataset_verification.py b/rationai/mlkit/lightning/callbacks/dataset_verification.py index e9e7ba9..b113853 100644 --- a/rationai/mlkit/lightning/callbacks/dataset_verification.py +++ b/rationai/mlkit/lightning/callbacks/dataset_verification.py @@ -22,6 +22,7 @@ import os import shutil import uuid +from typing import Any import mlflow from lightning.pytorch.callbacks import Callback @@ -61,10 +62,10 @@ def __init__( self.random_state = random_state self.fail_fast = fail_fast self._done = False - self._verification: dict | None = None - self._split_data: dict | None = None + self._verification: dict[str, Any] | None = None + self._split_data: dict[str, Any] | None = None - def on_fit_start(self, trainer, pl_module): # noqa: ARG002 + def on_fit_start(self, trainer: Any, pl_module: Any) -> None: if self._done: return self._done = True @@ -73,6 +74,8 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 _detect_manifest, _lookup_dataset_run, _verify_dataset, + ) + from rationai.mlkit.provenance.register_dataset import ( load_manifest as _load_manifest, ) diff --git a/rationai/mlkit/lightning/callbacks/environment.py b/rationai/mlkit/lightning/callbacks/environment.py index 09397b0..5e15e0f 100644 --- a/rationai/mlkit/lightning/callbacks/environment.py +++ b/rationai/mlkit/lightning/callbacks/environment.py @@ -15,17 +15,19 @@ from __future__ import annotations +import contextlib import hashlib -import json import logging import os import platform import shutil import subprocess import uuid -from datetime import datetime, timezone +from datetime import UTC, datetime +from typing import Any import mlflow +import pandas as pd import torch from lightning.pytorch.callbacks import Callback @@ -37,18 +39,16 @@ # Helpers # ────────────────────────────────────────────── -def _lookup_user_run(): +def _lookup_user_run() -> tuple[str | None, dict[str, str]]: """Find the user run from User_Registry. Auto-detect username.""" from rationai.mlkit.provenance.register_dataset import _lookup_experiment username = os.environ.get("MLFLOW_USER") if not username: - try: + with contextlib.suppress(subprocess.CalledProcessError): username = subprocess.check_output( ["git", "config", "user.name"], stderr=subprocess.DEVNULL, ).decode().strip() - except subprocess.CalledProcessError: - pass if not username: username = os.environ.get("USER", "unknown") @@ -56,20 +56,21 @@ def _lookup_user_run(): if exp_id is None: return None, {} - runs = mlflow.search_runs(experiment_ids=[exp_id]) - if runs.empty: + _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[runs["tags.username"] == username] + matched = runs_df[runs_df["tags.username"] == username] if matched.empty: - matched = runs.head(1) + matched = runs_df.head(1) row = matched.iloc[0] - run = mlflow.get_run(row.run_id) - return row.run_id, dict(run.data.tags) + run_obj = mlflow.get_run(row.run_id) + return row.run_id, dict(run_obj.data.tags) -def _detect_hardware(): +def _detect_hardware() -> dict[str, str | int]: """Detect CPU/GPU/hardware info.""" info: dict[str, str | int] = {} @@ -96,7 +97,7 @@ def _detect_hardware(): return info -def _detect_docker(): +def _detect_docker() -> dict[str, str | bool]: """Detect if running inside Docker and extract container info.""" info: dict[str, str | bool] = {"docker": False} @@ -131,7 +132,7 @@ def _detect_docker(): pass if info["docker"]: - cid = info.get("container_id_short", "") + cid = str(info.get("container_id_short", "")) if cid: try: result = subprocess.run( @@ -202,14 +203,14 @@ def __init__( self._git_commit: str = "unknown" self._git_url: str = "unknown" self._git_branch: str = "unknown" - self._hardware: dict = {} - self._docker: dict = {} + 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 = {} + self._user_tags: dict[str, str] = {} self._temp_dirs: list[str] = [] - def on_fit_start(self, trainer, pl_module): # noqa: ARG002 + 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 @@ -217,18 +218,17 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 # ── 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) - tags = run_data.data.tags or {} - else: - tags = {} - self._git_commit = tags.get("mlflow.source.git.commit", - tags.get("git.commit", "unknown")) - self._git_url = tags.get("mlflow.source.git.repoUrl", - tags.get("git.repo_url", "unknown")) - self._git_branch = tags.get("mlflow.source.git.branch", - tags.get("git.branch", "unknown")) + 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 @@ -269,25 +269,25 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 log.warning("[EnvironmentCallback] Docker detection failed: %s", e) # ── Log tags (git + user) ─────────────────────────────── - tags: dict[str, str] = {} + env_tags: dict[str, str] = {} if self._user_run_id: - tags["user_run_id"] = 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: - tags[key] = self._user_tags[key] + env_tags[key] = self._user_tags[key] from rationai.mlkit.provenance.register_dataset import _lookup_dataset_run dataset_run_id = _lookup_dataset_run() if dataset_run_id: - tags["dataset_run_id"] = dataset_run_id + env_tags["dataset_run_id"] = dataset_run_id - tags.update({ + env_tags.update({ "git_commit": self._git_commit, "git_url": self._git_url, "git_branch": self._git_branch, - "prov_start_time": datetime.now(timezone.utc).isoformat(), + "prov_start_time": datetime.now(UTC).isoformat(), }) - mlflow.set_tags(tags) + mlflow.set_tags(env_tags) # ── Log hardware + docker params ──────────────────────── all_params: dict[str, str | float | int] = {**self._hardware, **self._docker} diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index c7ac532..8add38a 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -29,11 +29,11 @@ import re import shutil import uuid -from datetime import datetime, timezone +from datetime import UTC, datetime +from typing import Any import mlflow import pandas as pd -import torch from lightning.pytorch.callbacks import Callback @@ -93,7 +93,7 @@ def _get_prov_prefixes(override: dict[str, str] | None = None) -> dict[str, str] # Model / Optimizer / Scheduler summaries # ────────────────────────────────────────────── -def _model_summary(model): +def _model_summary(model: Any) -> dict[str, str | int]: """Extract architecture details from a torch.nn.Module.""" info: dict[str, str | int] = {} @@ -112,18 +112,19 @@ def _model_summary(model): children = len(list(module.children())) layer_lines.append( f"{name}({type(module).__name__}): params={param_count}, " - f"children={children}" + f"children={children}", ) - info["layer_summary"] = "\n".join(layer_lines[:20]) + layer_summary: str = "\n".join(layer_lines[:20]) if len(layer_lines) > 20: - info["layer_summary"] += f"\n... ({len(layer_lines)} layers total)" + 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): +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__ @@ -133,7 +134,7 @@ def _optimizer_summary(optimizer): return info -def _scheduler_summary(scheduler): +def _scheduler_summary(scheduler: Any) -> dict[str, str | float]: """Extract scheduler settings.""" info: dict[str, str | float] = {} if scheduler is None: @@ -145,7 +146,7 @@ def _scheduler_summary(scheduler): "patience", "min_lr", "T_max", "eta_min"): val = getattr(scheduler, attr, None) if val is not None: - info[f"sch_{attr}"] = list(val) if isinstance(val, (list, tuple)) else val + 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(): @@ -160,26 +161,26 @@ def _scheduler_summary(scheduler): # ────────────────────────────────────────────── def _safe_id(name: str) -> str: - return re.sub(r'[^a-zA-Z0-9_]', '_', name) + return re.sub(r"[^a-zA-Z0-9_]", "_", name) def _qualified(prefix: str, local: str) -> str: return f"{prefix}:{local}" -def _typed_value(value, type_prefix="xsd", type_local="string") -> list: +def _typed_value(value: object, type_prefix: str = "xsd", type_local: str = "string") -> list[str]: return [str(value)] -def _qualified_name(type_prefix: str, type_local: str) -> dict: +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 = datetime.fromtimestamp(ts_ms / 1000, tz=timezone.utc) + dt = datetime.fromtimestamp(ts_ms / 1000, tz=UTC) else: - dt = datetime.now(timezone.utc) + dt = datetime.now(UTC) return dt.strftime("%Y-%m-%dT%H:%M:%S.000+00:00") @@ -191,13 +192,12 @@ def _build_prov_document( tags: dict[str, str], start_time_ms: int | None = None, end_time_ms: int | None = None, - split_data: dict | None = None, + split_data: dict[str, object] | None = None, requirements: str | None = None, - verification: dict | None = None, + verification: dict[str, object] | None = None, prov_prefixes: dict[str, str] | None = None, -) -> dict: +) -> dict[str, object]: """Build an OpenProvenance-compatible PROV document dict.""" - username = tags.get("username", tags.get("mlflow.user", "unknown")) agent_local = _safe_id(f"user_{username}") agent_id = _qualified("gen", agent_local) @@ -211,11 +211,11 @@ def _build_prov_document( main_act_local = f"TrainingRun_{run_id[:8]}" main_act_id = _qualified("blank", main_act_local) - entities: dict[str, dict] = {} - activities: dict[str, dict] = {} - agents: dict[str, dict] = {} - used: dict[str, dict] = {} - was_associated_with: dict[str, dict] = {} + 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] @@ -225,7 +225,7 @@ def _blank_rel_id() -> str: return rid # ── 1. AGENT ─────────────────────────────────────────── - agent_props: dict[str, list] = {} + 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") @@ -248,7 +248,7 @@ def _blank_rel_id() -> str: if image_path_candidates: wsi_local = _safe_id(f"wsi_{image_path_candidates}") wsi_id = _qualified("gen", wsi_local) - wsi_props: dict[str, list] = { + wsi_props: dict[str, Any] = { "schema:name": _typed_value(f"Input: {image_path_candidates}"), "prov:type": [_qualified_name("sosa", "Sample")], } @@ -287,7 +287,7 @@ def _blank_rel_id() -> str: } # ── 3. RUN ACTIVITY ──────────────────────────────────── - run_activity: dict[str, object] = {} + 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)] @@ -322,12 +322,10 @@ def _blank_rel_id() -> str: run_activity[f"gen:{key}"] = _typed_value(params[key]) for key, val in params.items(): - if key.startswith("opt_") or key.startswith("sch_"): + if key.startswith(("opt_", "sch_")): clean = key - while clean.startswith("opt_") or clean.startswith("sch_"): - if clean.startswith("opt_"): - clean = clean[4:] - elif clean.startswith("sch_"): + while clean.startswith(("opt_", "sch_")): + if clean.startswith(("opt_", "sch_")): clean = clean[4:] if f"gen:{clean}" not in run_activity: run_activity[f"gen:{clean}"] = _typed_value(val) @@ -368,7 +366,7 @@ def _blank_rel_id() -> str: activities[run_act_id] = run_activity # ── 4. CPM METADATA ENTITY ───────────────────────────── - meta_entity: dict[str, list] = {} + meta_entity: dict[str, Any] = {} meta_entity["prov:type"] = [_qualified_name("cpm", "BundleMetadata")] org_val = tags.get("organization", "") if org_val: @@ -387,9 +385,9 @@ def _blank_rel_id() -> str: safe_key = _safe_id(key) meta_entity[f"gen:{safe_key}"] = _typed_value(val) - for key, val in metrics.items(): + for key, mval in metrics.items(): safe_key = _safe_id(key) - meta_entity[f"gen:{safe_key}"] = _typed_value(val) + 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")) @@ -427,20 +425,20 @@ def _blank_rel_id() -> str: entities[meta_id] = meta_entity - was_generated_by: dict[str, dict] = {} + 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, object] = {} + main_activity: dict[str, Any] = {} main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] main_activity["cpm:referencedMetaBundleId"] = [ - {"type": "prov:QUALIFIED_NAME", "$": meta_id} + {"type": "prov:QUALIFIED_NAME", "$": meta_id}, ] main_activity["dct:hasPart"] = [ - {"type": "prov:QUALIFIED_NAME", "$": run_act_id} + {"type": "prov:QUALIFIED_NAME", "$": run_act_id}, ] activities[main_act_id] = main_activity @@ -528,8 +526,8 @@ def __init__( # 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 | None = None - self._verification: dict | None = None + 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" @@ -537,12 +535,12 @@ def __init__( # ── helpers ────────────────────────────────────────────── - def _gather_from_siblings(self, trainer) -> None: + def _gather_from_siblings(self, trainer: Any) -> None: """Read data already collected by sibling callbacks.""" - from rationai.mlkit.lightning.callbacks.environment import EnvironmentCallback 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): @@ -554,7 +552,7 @@ def _gather_from_siblings(self, trainer) -> None: self._verification = getattr(cb, "_verification", None) self._split_data = getattr(cb, "_split_data", None) - def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 + 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, @@ -571,18 +569,17 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 # ── 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) - tags = run_data.data.tags or {} - else: - tags = {} - self._git_commit = tags.get("mlflow.source.git.commit", - tags.get("git.commit", "unknown")) - self._git_url = tags.get("mlflow.source.git.repoUrl", - tags.get("git.repo_url", "unknown")) - self._git_branch = tags.get("mlflow.source.git.branch", - tags.get("git.branch", "unknown")) + 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 @@ -607,16 +604,18 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 manifest_path, data_root = _detect_manifest() if manifest_path and data_root: + from sklearn.model_selection import train_test_split + from rationai.mlkit.provenance.register_dataset import ( load_manifest, ) - from sklearn.model_selection import train_test_split 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 self._verification.get("details", []): + for detail in verification_details: log.info(f" [ProvenanceCallback] {detail}") samples = load_manifest(manifest_path, data_root) @@ -679,7 +678,7 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 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"]) + + "\n".join(f" {d}" for d in verification["details"]), ) else: log.warning("[ProvenanceCallback] No manifest.csv found — " @@ -701,7 +700,7 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 "git_commit": self._git_commit, "git_url": self._git_url, "git_branch": self._git_branch, - "prov_start_time": datetime.now(timezone.utc).isoformat(), + "prov_start_time": datetime.now(UTC).isoformat(), }) mlflow.set_tags(tags) @@ -730,7 +729,7 @@ def _fallback_on_fit_start(self, trainer, pl_module): # noqa: ARG002 # ── lightning hooks ─────────────────────────────────────── - def _ensure_active_run(self, trainer) -> str | None: + 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. @@ -752,7 +751,7 @@ def _ensure_active_run(self, trainer) -> str | None: run = mlflow.active_run() return run.info.run_id if run else None - def on_fit_start(self, trainer, pl_module): # noqa: ARG002 + 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, @@ -778,17 +777,22 @@ def on_fit_start(self, trainer, pl_module): # noqa: ARG002 # No siblings — do everything ourselves self._fallback_on_fit_start(trainer, pl_module) - def on_fit_end(self, trainer, pl_module): # noqa: ARG002 + 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 not active_run: - # Fallback: try to get run_id from the logger directly + _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: - active_run = mlflow.tracking.MlflowClient().get_run(self._run_id) + run_id = self._run_id else: return - run_id = active_run.info.run_id + + # 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: @@ -851,8 +855,8 @@ def on_fit_end(self, trainer, pl_module): # noqa: ARG002 "test_size": self.test_size, "random_state": self.random_state, "stratified": True, - "train_count": len(self._split_data["train"]) if self._split_data else 0, - "test_count": len(self._split_data["test"]) if self._split_data else 0, + "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, diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py index d60266e..a289c4b 100644 --- a/rationai/mlkit/provenance/__init__.py +++ b/rationai/mlkit/provenance/__init__.py @@ -15,35 +15,35 @@ from __future__ import annotations import os +from typing import Any + # ── Set default tracking URI before any mlflow import runs ─────────── if "MLFLOW_TRACKING_URI" not in os.environ: os.environ["MLFLOW_TRACKING_URI"] = "http://localhost:5000" # Now safe to import - all child modules will pick up the env var -from .prov import ( # noqa: E402 +from rationai.mlkit.provenance.prov import ( build_dataset_prov, build_user_prov, ) -from .register_dataset import ( # noqa: E402 +from rationai.mlkit.provenance.register_dataset import ( register_dataset, verify_dataset, ) -from .register_user import register_new_user # noqa: E402 +from rationai.mlkit.provenance.register_user import register_new_user + __all__ = [ - # Dataset + "build_dataset_prov", + "build_user_prov", "register_dataset", - "verify_dataset", - # User registration "register_new_user", - # PROV document builders - "build_user_prov", - "build_dataset_prov", + "verify_dataset", ] -def __getattr__(name: str): +def __getattr__(name: str) -> Any: """Raise helpful error for removed ``autolog``.""" if name == "autolog": raise ImportError( diff --git a/rationai/mlkit/provenance/prov.py b/rationai/mlkit/provenance/prov.py index 7964d6b..338fb04 100644 --- a/rationai/mlkit/provenance/prov.py +++ b/rationai/mlkit/provenance/prov.py @@ -7,10 +7,10 @@ from __future__ import annotations +import datetime as _dt import json import os import re -import datetime as _dt from typing import Any @@ -55,7 +55,7 @@ def get_prov_prefixes(override: dict[str, str] | None = None) -> dict[str, str]: 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) + return re.sub(r"[^a-zA-Z0-9_]", "_", name) def _qualified(prefix: str, local: str) -> str: @@ -160,10 +160,10 @@ def _blank_rel_id() -> str: main_activity: dict[str, Any] = {} main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] main_activity["cpm:referencedMetaBundleId"] = [ - {"type": "prov:QUALIFIED_NAME", "$": meta_id} + {"type": "prov:QUALIFIED_NAME", "$": meta_id}, ] main_activity["dct:hasPart"] = [ - {"type": "prov:QUALIFIED_NAME", "$": run_act_id} + {"type": "prov:QUALIFIED_NAME", "$": run_act_id}, ] activities[main_act_id] = main_activity @@ -243,7 +243,7 @@ def _blank_rel_id() -> str: "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)" + f"Dataset {dataset_name} v{version} ({num_samples} samples)", ), } if manifest_path: @@ -288,10 +288,10 @@ def _blank_rel_id() -> str: main_activity: dict[str, Any] = {} main_activity["prov:type"] = [_qualified_name("cpm", "mainActivity")] main_activity["cpm:referencedMetaBundleId"] = [ - {"type": "prov:QUALIFIED_NAME", "$": meta_id} + {"type": "prov:QUALIFIED_NAME", "$": meta_id}, ] main_activity["dct:hasPart"] = [ - {"type": "prov:QUALIFIED_NAME", "$": run_act_id} + {"type": "prov:QUALIFIED_NAME", "$": run_act_id}, ] activities[main_act_id] = main_activity diff --git a/rationai/mlkit/provenance/register_dataset.py b/rationai/mlkit/provenance/register_dataset.py index 47abea1..bbe4f07 100644 --- a/rationai/mlkit/provenance/register_dataset.py +++ b/rationai/mlkit/provenance/register_dataset.py @@ -6,6 +6,7 @@ import os import shutil import uuid +from typing import Any import mlflow import pandas as pd @@ -13,7 +14,7 @@ # ── Internal helpers ──────────────────────────────────────────────────── -def _lookup_experiment(name): +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 @@ -33,10 +34,10 @@ def _lookup_dataset_run() -> str | None: experiment_ids=[exp_id], order_by=["start_time DESC"], ) - if runs_df.empty: + if pd.DataFrame(runs_df).empty: return None - return runs_df.iloc[0]["run_id"] + return pd.DataFrame(runs_df).iloc[0]["run_id"] def _detect_manifest() -> tuple[str | None, str | None]: @@ -50,13 +51,13 @@ def _detect_manifest() -> tuple[str | None, str | None]: return ( os.path.join(dirpath, "manifest.csv"), os.path.dirname(os.path.abspath( - os.path.join(dirpath, "manifest.csv") + os.path.join(dirpath, "manifest.csv"), )), ) return None, None -def load_manifest(manifest_path: str, data_root: str) -> list[dict]: +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``. @@ -65,7 +66,7 @@ def load_manifest(manifest_path: str, data_root: str) -> list[dict]: ``ProvenanceCallback`` to avoid duplicating the CSV iteration pattern. """ df = pd.read_csv(manifest_path) - samples: list[dict] = [] + 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 @@ -76,7 +77,7 @@ def load_manifest(manifest_path: str, data_root: str) -> list[dict]: def verify_dataset( manifest_path: str | None = None, data_root: str | None = None, -) -> dict: +) -> dict[str, Any]: """Public entry point — verify the current dataset against MLflow. Auto-detects the manifest if *manifest_path* is not given. @@ -110,7 +111,7 @@ def _verify_dataset( manifest_path: str, data_root: str, dataset_run_id: str | None, -) -> dict: +) -> dict[str, Any]: """Verify the current dataset against the registered version in MLflow. Checks: @@ -119,7 +120,7 @@ def _verify_dataset( Returns a dict with verification results. """ - result: dict = { + result: dict[str, Any] = { "verified": False, "dataset_run_id": dataset_run_id, "file_sizes_match": None, @@ -143,7 +144,7 @@ def _verify_dataset( return result samples = load_manifest(manifest_path, data_root) - curr_file_sizes = {} + curr_file_sizes: dict[str, int] = {} for s in samples: basename = os.path.basename(s["path"]) if os.path.isfile(s["path"]): @@ -163,7 +164,7 @@ def _verify_dataset( result["details"].append( f"File size mismatch on {len(mismatched)} file(s): " + ", ".join(sorted(mismatched)[:5]) - + ("…" if len(mismatched) > 5 else "") + + ("…" if len(mismatched) > 5 else ""), ) # Check file existence @@ -192,7 +193,7 @@ def register_dataset( 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 @@ -218,7 +219,7 @@ def register_dataset( if not os.path.isfile(manifest_path): raise FileNotFoundError( f"No manifest.csv found in {dataset_dir}. " - "Dataset registration requires a manifest.csv file." + "Dataset registration requires a manifest.csv file.", ) if dataset_name is None: @@ -260,7 +261,7 @@ def register_dataset( }) # ── PROV-O document (W3C PROV-O compatible) ──────────── - from rationai.mlkit.provenance.prov import build_dataset_prov # noqa: PLC0415 + from rationai.mlkit.provenance.prov import build_dataset_prov prov_doc = build_dataset_prov( run_id=run_id, diff --git a/rationai/mlkit/provenance/register_user.py b/rationai/mlkit/provenance/register_user.py index aea0f2f..3b5c1b7 100644 --- a/rationai/mlkit/provenance/register_user.py +++ b/rationai/mlkit/provenance/register_user.py @@ -28,6 +28,7 @@ def register_new_user( 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: @@ -55,7 +56,7 @@ def register_new_user( }) # ── PROV-O document ──────────────────────────────── - from rationai.mlkit.provenance.prov import build_user_prov # noqa: PLC0415 + from rationai.mlkit.provenance.prov import build_user_prov prov_doc = build_user_prov( run_id=run_id, From 1153c478d53b5c3baf8219a2b781fb7fab00a0ea Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 14:26:27 +0200 Subject: [PATCH 23/34] feat: adding docstrings --- .../callbacks/dataset_verification.py | 16 +++++++++++++++- .../mlkit/lightning/callbacks/environment.py | 10 +++++++++- .../mlkit/lightning/callbacks/provenance.py | 18 +++++++++++++++++- 3 files changed, 41 insertions(+), 3 deletions(-) diff --git a/rationai/mlkit/lightning/callbacks/dataset_verification.py b/rationai/mlkit/lightning/callbacks/dataset_verification.py index b113853..145627a 100644 --- a/rationai/mlkit/lightning/callbacks/dataset_verification.py +++ b/rationai/mlkit/lightning/callbacks/dataset_verification.py @@ -56,7 +56,15 @@ def __init__( 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 @@ -66,6 +74,12 @@ def __init__( 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 diff --git a/rationai/mlkit/lightning/callbacks/environment.py b/rationai/mlkit/lightning/callbacks/environment.py index 5e15e0f..da1772c 100644 --- a/rationai/mlkit/lightning/callbacks/environment.py +++ b/rationai/mlkit/lightning/callbacks/environment.py @@ -194,7 +194,15 @@ def __init__( 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 diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index 8add38a..eaa61a8 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -509,7 +509,23 @@ def __init__( 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 From c06811404a6c2df43762b7a54341f7144ecae883 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 14:48:36 +0200 Subject: [PATCH 24/34] feat: refactor __getattr__ to improve MLFlowLogger import handling --- rationai/mlkit/__init__.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index 161ff09..57375a1 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -33,11 +33,15 @@ def __getattr__(name: str) -> Any: - if name in ("Trainer", "MLFlowLogger", "MultiloaderLifecycle", "with_cli_args"): + 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 From 43ccd0cad30a2888bbbffed71050bb840996f579 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 15:06:45 +0200 Subject: [PATCH 25/34] feat: enhance error handling in _get_prov_prefixes and user lookup in ProvenanceCallback --- .../mlkit/lightning/callbacks/provenance.py | 17 ++++++++++++++--- 1 file changed, 14 insertions(+), 3 deletions(-) diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index eaa61a8..a2028de 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -70,9 +70,12 @@ def _get_prov_prefixes(override: dict[str, str] | None = None) -> dict[str, str] env_json = os.environ.get("PROV_BASE_URI", "") if env_json: try: - merged = {**_DEFAULT_PROV_PREFIXES, **json.loads(env_json)} + 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 as e: + except (json.JSONDecodeError, TypeError) as e: log.warning(f"PROV_BASE_URI is not valid JSON: {e} — using defaults") return _DEFAULT_PROV_PREFIXES @@ -602,7 +605,13 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: log.warning("[ProvenanceCallback] Git info failed: %s", e) # ── User lookup ───────────────────────────────────────── - user_run_id, user_tags = _lookup_user_run() + 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( @@ -618,6 +627,8 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: 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 sklearn.model_selection import train_test_split From 746dc0b662a86afd483170ee9eae0eab017df37b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 15:10:06 +0200 Subject: [PATCH 26/34] feat: improve code formatting and readability across multiple files --- rationai/mlkit/__init__.py | 19 +- .../callbacks/dataset_verification.py | 43 ++-- .../mlkit/lightning/callbacks/environment.py | 60 +++-- .../mlkit/lightning/callbacks/provenance.py | 227 ++++++++++++------ rationai/mlkit/provenance/prov.py | 4 + rationai/mlkit/provenance/register_dataset.py | 79 +++--- rationai/mlkit/provenance/register_user.py | 28 ++- 7 files changed, 303 insertions(+), 157 deletions(-) diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index 57375a1..51e3769 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -35,31 +35,44 @@ 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"): + 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) diff --git a/rationai/mlkit/lightning/callbacks/dataset_verification.py b/rationai/mlkit/lightning/callbacks/dataset_verification.py index 145627a..95ad9ae 100644 --- a/rationai/mlkit/lightning/callbacks/dataset_verification.py +++ b/rationai/mlkit/lightning/callbacks/dataset_verification.py @@ -99,7 +99,9 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: manifest_path, data_root = _detect_manifest() if manifest_path is None: - log.warning("[DatasetVerificationCallback] No manifest.csv found — skipping") + log.warning( + "[DatasetVerificationCallback] No manifest.csv found — skipping" + ) return if data_root is None: @@ -115,12 +117,15 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: # ── 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"], - }) + 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: @@ -174,14 +179,16 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: # 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, - }) + 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 index da1772c..62ad7e8 100644 --- a/rationai/mlkit/lightning/callbacks/environment.py +++ b/rationai/mlkit/lightning/callbacks/environment.py @@ -39,6 +39,7 @@ # 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.register_dataset import _lookup_experiment @@ -46,9 +47,14 @@ def _lookup_user_run() -> tuple[str | None, dict[str, str]]: 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() + username = ( + subprocess.check_output( + ["git", "config", "user.name"], + stderr=subprocess.DEVNULL, + ) + .decode() + .strip() + ) if not username: username = os.environ.get("USER", "unknown") @@ -89,6 +95,7 @@ def _detect_hardware() -> dict[str, str | int]: try: import psutil + mem = psutil.virtual_memory() info["ram_total_gb"] = round(mem.total / 1e9, 1) except ImportError: @@ -123,9 +130,7 @@ def _detect_docker() -> dict[str, str | bool]: 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 - ): + if len(p) == 64 and all(c in "0123456789abcdef" for c in p): info["container_id_short"] = p[:12] break except FileNotFoundError: @@ -137,7 +142,9 @@ def _detect_docker() -> dict[str, str | bool]: try: result = subprocess.run( ["docker", "inspect", "--format={{.Config.Image}}", cid], - capture_output=True, text=True, timeout=5, + capture_output=True, + text=True, + timeout=5, ) if result.returncode == 0 and result.stdout.strip(): image = result.stdout.strip() @@ -169,6 +176,7 @@ def _snapshot_environment(artifact_dir: str) -> str: # Callback # ────────────────────────────────────────────── + class EnvironmentCallback(Callback): """Capture hardware, docker, git, user, and environment snapshot at training start. @@ -231,12 +239,15 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: 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")) + 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 @@ -259,14 +270,18 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: for logger in trainer.loggers ) if sys_metrics_on: - log.info("[EnvironmentCallback] Skipping hardware — MLflow system metrics enabled") + 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) + log.warning( + "[EnvironmentCallback] Hardware detection failed: %s", e + ) # ── Docker detection ──────────────────────────────────── try: @@ -285,16 +300,19 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: env_tags[key] = self._user_tags[key] from rationai.mlkit.provenance.register_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(), - }) + 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 ──────────────────────── diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index a2028de..072bc7c 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -79,16 +79,36 @@ def _get_prov_prefixes(override: dict[str, str] | None = None) -> dict[str, str] log.warning(f"PROV_BASE_URI is not valid JSON: {e} — using defaults") return _DEFAULT_PROV_PREFIXES + _ACTIVITY_HP_KEYS = { - "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", + "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 = { - "scanner", "slide_id", "wsi_id", "patient_id", "subject_id", - "institution", "site", "staining", "slicing_method", + "scanner", + "slide_id", + "wsi_id", + "patient_id", + "subject_id", + "institution", + "site", + "staining", + "slicing_method", } @@ -96,14 +116,13 @@ def _get_prov_prefixes(override: dict[str, str] | None = None) -> dict[str, str] # 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 - ) + trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad) info["total_parameters"] = total_params info["trainable_parameters"] = trainable_params @@ -145,11 +164,21 @@ def _scheduler_summary(scheduler: Any) -> dict[str, str | float]: return info info["scheduler_type"] = type(scheduler).__name__ - for attr in ("step_size", "gamma", "milestones", "factor", - "patience", "min_lr", "T_max", "eta_min"): + 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) + 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(): @@ -163,6 +192,7 @@ def _scheduler_summary(scheduler: Any) -> dict[str, str | float]: # PROV document builder # ────────────────────────────────────────────── + def _safe_id(name: str) -> str: return re.sub(r"[^a-zA-Z0-9_]", "_", name) @@ -171,7 +201,9 @@ def _qualified(prefix: str, local: str) -> str: return f"{prefix}:{local}" -def _typed_value(value: object, type_prefix: str = "xsd", type_local: str = "string") -> list[str]: +def _typed_value( + value: object, type_prefix: str = "xsd", type_local: str = "string" +) -> list[str]: return [str(value)] @@ -241,11 +273,11 @@ def _blank_rel_id() -> str: # ── 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") + 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: @@ -281,7 +313,9 @@ def _blank_rel_id() -> str: 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)"), + "schema:name": _typed_value( + f"Training dataset ({train_count} train, {test_count} test)" + ), "prov:type": [_qualified_name("sosa", "Sample")], } used[_blank_rel_id()] = { @@ -375,13 +409,35 @@ def _blank_rel_id() -> str: 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", - } + 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: @@ -393,14 +449,22 @@ def _blank_rel_id() -> str: 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_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"]))[0]] + meta_entity["gen:split_train"] = [ + _typed_value(json.dumps(split_data["train"]))[0] + ] if split_data.get("test"): - meta_entity["gen:split_test"] = [_typed_value(json.dumps(split_data["test"]))[0]] + meta_entity["gen:split_test"] = [ + _typed_value(json.dumps(split_data["test"]))[0] + ] if requirements: meta_entity["gen:requirements"] = [requirements] @@ -474,6 +538,7 @@ def _blank_rel_id() -> str: # Slim ProvenanceCallback # ────────────────────────────────────────────── + class ProvenanceCallback(Callback): """Lightning callback that captures PROV document + run summary. @@ -593,12 +658,15 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: 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")) + 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 @@ -615,8 +683,7 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: # ── Hardware (skip if MLflow system metrics are on) ───── sys_metrics_on = any( - getattr(logger, "log_system_metrics", False) - for logger in trainer.loggers + getattr(logger, "log_system_metrics", False) for logger in trainer.loggers ) hardware = {} if sys_metrics_on else _detect_hardware() docker = _detect_docker() @@ -676,23 +743,28 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: # 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), - }) + 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), + } + ) # 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"], - }) + 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: @@ -708,8 +780,10 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: + "\n".join(f" {d}" for d in verification["details"]), ) else: - log.warning("[ProvenanceCallback] No manifest.csv found — " - "train/test split not logged.") + log.warning( + "[ProvenanceCallback] No manifest.csv found — " + "train/test split not logged." + ) # ── Tags ──────────────────────────────────────────────── tags: dict[str, str] = {} @@ -723,12 +797,14 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: 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(), - }) + 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 ──────────── @@ -789,13 +865,9 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: self._ensure_active_run(trainer) # Check if sibling callbacks are present - has_env = any( - isinstance(cb, EnvironmentCallback) - for cb in trainer.callbacks - ) + has_env = any(isinstance(cb, EnvironmentCallback) for cb in trainer.callbacks) has_verify = any( - isinstance(cb, DatasetVerificationCallback) - for cb in trainer.callbacks + isinstance(cb, DatasetVerificationCallback) for cb in trainer.callbacks ) if has_env or has_verify: @@ -861,7 +933,8 @@ def on_fit_end(self, trainer: Any, pl_module: Any) -> None: 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() + k: v + for k, v in run_data.data.tags.items() if not k.startswith("mlflow.") } @@ -882,11 +955,17 @@ def on_fit_end(self, trainer: Any, pl_module: Any) -> None: "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_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, + } + if self._split_data + else None, "dataset_verification": self._verification, "requirements": self._frozen_requirements, "source": { @@ -916,7 +995,9 @@ def on_fit_end(self, trainer: Any, pl_module: Any) -> None: "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, + } + if self._split_data + else None, requirements=self._frozen_requirements, verification=self._verification, prov_prefixes=_get_prov_prefixes(self._prov_prefixes), @@ -936,7 +1017,9 @@ def on_fit_end(self, trainer: Any, pl_module: Any) -> None: except Exception as e: if self.strict: raise - log.warning("[ProvenanceCallback] Could not write provenance artifacts: %s", e) + log.warning( + "[ProvenanceCallback] Could not write provenance artifacts: %s", e + ) # ── Clean up temp dirs ────────────────────────────────── for d in self._temp_dirs: diff --git a/rationai/mlkit/provenance/prov.py b/rationai/mlkit/provenance/prov.py index 338fb04..c261495 100644 --- a/rationai/mlkit/provenance/prov.py +++ b/rationai/mlkit/provenance/prov.py @@ -53,6 +53,7 @@ def get_prov_prefixes(override: dict[str, str] | None = None) -> dict[str, str]: # Small helpers used inside the PROV document # ────────────────────────────────────────────── + 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) @@ -82,6 +83,7 @@ def _iso_timestamp(ts_ms: int | None = None) -> str: # PROV document builders # ────────────────────────────────────────────── + def build_user_prov( run_id: str, username: str, @@ -118,6 +120,7 @@ def build_user_prov( 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 @@ -231,6 +234,7 @@ def build_dataset_prov( 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 diff --git a/rationai/mlkit/provenance/register_dataset.py b/rationai/mlkit/provenance/register_dataset.py index bbe4f07..27718f1 100644 --- a/rationai/mlkit/provenance/register_dataset.py +++ b/rationai/mlkit/provenance/register_dataset.py @@ -14,6 +14,7 @@ # ── Internal helpers ──────────────────────────────────────────────────── + def _lookup_experiment(name: str) -> str | None: """Return the MLflow experiment ID for *name*, or None.""" exp = mlflow.get_experiment_by_name(name) @@ -50,9 +51,11 @@ def _detect_manifest() -> tuple[str | None, str | None]: if "manifest.csv" in filenames: return ( os.path.join(dirpath, "manifest.csv"), - os.path.dirname(os.path.abspath( - os.path.join(dirpath, "manifest.csv"), - )), + os.path.dirname( + os.path.abspath( + os.path.join(dirpath, "manifest.csv"), + ) + ), ) return None, None @@ -84,8 +87,13 @@ def verify_dataset( Returns a dict with keys:: - {"verified": bool, "file_sizes_match": bool|None, - "files_missing": int, "files_total": int, "details": list[str]} + { + "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() @@ -130,7 +138,9 @@ def _verify_dataset( } if not dataset_run_id: - result["details"].append("No Dataset_Registry run found — skipping verification") + result["details"].append( + "No Dataset_Registry run found — skipping verification" + ) return result # Fetch registered metadata @@ -157,7 +167,8 @@ def _verify_dataset( if not result["file_sizes_match"]: mismatched = [ - name for name in reg_file_sizes + name + for name in reg_file_sizes if curr_file_sizes.get(name) != reg_file_sizes[name] ] if mismatched: @@ -188,6 +199,7 @@ def _verify_dataset( # ── Hash-based registration (preferred) ────────────────────────────────── + def register_dataset( dataset_dir: str, dataset_name: str | None = None, @@ -246,19 +258,23 @@ def register_dataset( 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), - }) + 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 (W3C PROV-O compatible) ──────────── from rationai.mlkit.provenance.prov import build_dataset_prov @@ -291,14 +307,18 @@ def register_dataset( 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) + 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) @@ -308,6 +328,3 @@ def register_dataset( print(f" [register_dataset] {dataset_name} v{version} → run_id={run_id}") return run_id - - - diff --git a/rationai/mlkit/provenance/register_user.py b/rationai/mlkit/provenance/register_user.py index 3b5c1b7..b6657ce 100644 --- a/rationai/mlkit/provenance/register_user.py +++ b/rationai/mlkit/provenance/register_user.py @@ -42,18 +42,22 @@ def register_new_user( ) 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, - }) + 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 ──────────────────────────────── from rationai.mlkit.provenance.prov import build_user_prov From 322f0a2f91410093e1882df65a09c0f047de6665 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 15:19:10 +0200 Subject: [PATCH 27/34] feat: enhance get_prov_prefixes to validate JSON input and improve error handling --- rationai/mlkit/provenance/prov.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/rationai/mlkit/provenance/prov.py b/rationai/mlkit/provenance/prov.py index c261495..c8c1511 100644 --- a/rationai/mlkit/provenance/prov.py +++ b/rationai/mlkit/provenance/prov.py @@ -42,9 +42,12 @@ def get_prov_prefixes(override: dict[str, str] | None = None) -> dict[str, str]: env_json = os.environ.get("PROV_BASE_URI", "") if env_json: try: - merged = {**_DEFAULT_PROV_PREFIXES, **json.loads(env_json)} + 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: + except (json.JSONDecodeError, TypeError): pass return _DEFAULT_PROV_PREFIXES From 72aff306cde184db38fbe71a932a061d4450b440 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 15:46:52 +0200 Subject: [PATCH 28/34] feat: add scikit-learn dependency and improve dataset verification logic --- pyproject.toml | 1 + .../mlkit/lightning/callbacks/provenance.py | 108 +++++++++--------- rationai/mlkit/provenance/register_dataset.py | 11 +- 3 files changed, 64 insertions(+), 56 deletions(-) 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/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index 072bc7c..0a01d49 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -360,10 +360,7 @@ def _blank_rel_id() -> str: for key, val in params.items(): if key.startswith(("opt_", "sch_")): - clean = key - while clean.startswith(("opt_", "sch_")): - if clean.startswith(("opt_", "sch_")): - clean = clean[4:] + clean = key.removeprefix("opt_").removeprefix("sch_") if f"gen:{clean}" not in run_activity: run_activity[f"gen:{clean}"] = _typed_value(val) @@ -458,13 +455,13 @@ def _blank_rel_id() -> str: meta_entity["gen:split_stratified"] = ["true"] if split_data.get("train"): - meta_entity["gen:split_train"] = [ - _typed_value(json.dumps(split_data["train"]))[0] - ] + 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"]))[0] - ] + meta_entity["gen:split_test"] = _typed_value( + json.dumps(split_data["test"]) + ) if requirements: meta_entity["gen:requirements"] = [requirements] @@ -698,12 +695,11 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: data_root = os.path.dirname(os.path.abspath(manifest_path)) if manifest_path and data_root: - from sklearn.model_selection import train_test_split - from rationai.mlkit.provenance.register_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 {} @@ -712,48 +708,6 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: for detail in verification_details: log.info(f" [ProvenanceCallback] {detail}") - 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), - } - ) - # Log verification results if verification: mlflow.log_params( @@ -779,6 +733,52 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: "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 — " diff --git a/rationai/mlkit/provenance/register_dataset.py b/rationai/mlkit/provenance/register_dataset.py index 27718f1..6ed42df 100644 --- a/rationai/mlkit/provenance/register_dataset.py +++ b/rationai/mlkit/provenance/register_dataset.py @@ -162,8 +162,15 @@ def _verify_dataset( else: curr_file_sizes[basename] = -1 # missing - # Compare sizes - result["file_sizes_match"] = curr_file_sizes == reg_file_sizes + # Compare sizes — must have the same keys AND the same sizes + if set(curr_file_sizes) != set(reg_file_sizes): + result["details"].append( + f"File manifest mismatch: expected {len(reg_file_sizes)} file(s), found {len(curr_file_sizes)} file(s)" + + ) + result["file_sizes_match"] = all( + curr_file_sizes.get(k) == reg_file_sizes[k] for k in reg_file_sizes + ) if not result["file_sizes_match"]: mismatched = [ From 4270cc6f92908af8e0bc2d0149bbd05ab35b51c7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 16:16:47 +0200 Subject: [PATCH 29/34] feat: Refactor provenance module: consolidate user registration and PROV-O document generation - Removed `prov.py` and `register_user.py` files. - Introduced `run.py` for training run PROV-O document generation. - Created `user.py` to handle user registration and PROV-O document building. - Updated `build_user_prov` and `register_new_user` functions to streamline user registration process. - Enhanced PROV document structure for better compatibility with MLflow and external tools. --- .gitignore | 2 +- rationai/mlkit/__init__.py | 2 +- .../callbacks/dataset_verification.py | 4 +- .../mlkit/lightning/callbacks/environment.py | 4 +- .../mlkit/lightning/callbacks/provenance.py | 431 +----------------- rationai/mlkit/provenance/__init__.py | 20 +- rationai/mlkit/provenance/common.py | 80 ++++ .../{register_dataset.py => dataset.py} | 178 +++++++- rationai/mlkit/provenance/prov.py | 323 ------------- rationai/mlkit/provenance/register_user.py | 96 ---- rationai/mlkit/provenance/run.py | 372 +++++++++++++++ rationai/mlkit/provenance/user.py | 230 ++++++++++ 12 files changed, 872 insertions(+), 870 deletions(-) create mode 100644 rationai/mlkit/provenance/common.py rename rationai/mlkit/provenance/{register_dataset.py => dataset.py} (59%) delete mode 100644 rationai/mlkit/provenance/prov.py delete mode 100644 rationai/mlkit/provenance/register_user.py create mode 100644 rationai/mlkit/provenance/run.py create mode 100644 rationai/mlkit/provenance/user.py diff --git a/.gitignore b/.gitignore index 261446d..bcbd562 100644 --- a/.gitignore +++ b/.gitignore @@ -174,4 +174,4 @@ cython_debug/ test_data mlflow.db mlartifacts -/example* \ No newline at end of file +/example* diff --git a/rationai/mlkit/__init__.py b/rationai/mlkit/__init__.py index 51e3769..6259b19 100644 --- a/rationai/mlkit/__init__.py +++ b/rationai/mlkit/__init__.py @@ -3,7 +3,7 @@ from typing import Any from rationai.mlkit.autolog import autolog -from rationai.mlkit.provenance.register_dataset import register_dataset +from rationai.mlkit.provenance.dataset import register_dataset from rationai.mlkit.stream import StreamCapture, StreamLogger diff --git a/rationai/mlkit/lightning/callbacks/dataset_verification.py b/rationai/mlkit/lightning/callbacks/dataset_verification.py index 95ad9ae..0ac451c 100644 --- a/rationai/mlkit/lightning/callbacks/dataset_verification.py +++ b/rationai/mlkit/lightning/callbacks/dataset_verification.py @@ -84,12 +84,12 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: return self._done = True - from rationai.mlkit.provenance.register_dataset import ( + from rationai.mlkit.provenance.dataset import ( _detect_manifest, _lookup_dataset_run, _verify_dataset, ) - from rationai.mlkit.provenance.register_dataset import ( + from rationai.mlkit.provenance.dataset import ( load_manifest as _load_manifest, ) diff --git a/rationai/mlkit/lightning/callbacks/environment.py b/rationai/mlkit/lightning/callbacks/environment.py index 62ad7e8..3447bf7 100644 --- a/rationai/mlkit/lightning/callbacks/environment.py +++ b/rationai/mlkit/lightning/callbacks/environment.py @@ -42,7 +42,7 @@ def _lookup_user_run() -> tuple[str | None, dict[str, str]]: """Find the user run from User_Registry. Auto-detect username.""" - from rationai.mlkit.provenance.register_dataset import _lookup_experiment + from rationai.mlkit.provenance.dataset import _lookup_experiment username = os.environ.get("MLFLOW_USER") if not username: @@ -299,7 +299,7 @@ def on_fit_start(self, trainer: Any, pl_module: Any) -> None: if key in self._user_tags: env_tags[key] = self._user_tags[key] - from rationai.mlkit.provenance.register_dataset import _lookup_dataset_run + from rationai.mlkit.provenance.dataset import _lookup_dataset_run dataset_run_id = _lookup_dataset_run() if dataset_run_id: diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index 0a01d49..5eba018 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -26,7 +26,6 @@ import json import logging import os -import re import shutil import uuid from datetime import UTC, datetime @@ -36,80 +35,16 @@ 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__) - - -# ────────────────────────────────────────────── -# OpenProvenance / CPM namespace URIs (§9 — configurable via env var) -# ────────────────────────────────────────────── -_DEFAULT_PROV_PREFIXES = { - "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. - The env var accepts a JSON object of ``{prefix: base_uri}`` pairs that - merge into (and override) the 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) as e: - log.warning(f"PROV_BASE_URI is not valid JSON: {e} — using defaults") - return _DEFAULT_PROV_PREFIXES - - -_ACTIVITY_HP_KEYS = { - "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 = { - "scanner", - "slide_id", - "wsi_id", - "patient_id", - "subject_id", - "institution", - "site", - "staining", - "slicing_method", -} +log = logging.getLogger(__name__) # ────────────────────────────────────────────── @@ -188,352 +123,6 @@ def _scheduler_summary(scheduler: Any) -> dict[str, str | float]: return info -# ────────────────────────────────────────────── -# PROV document builder -# ────────────────────────────────────────────── - - -def _safe_id(name: str) -> str: - return re.sub(r"[^a-zA-Z0-9_]", "_", name) - - -def _qualified(prefix: str, local: str) -> str: - return f"{prefix}:{local}" - - -def _typed_value( - value: object, type_prefix: str = "xsd", type_local: str = "string" -) -> 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 = datetime.fromtimestamp(ts_ms / 1000, tz=UTC) - else: - dt = datetime.now(UTC) - return dt.strftime("%Y-%m-%dT%H:%M:%S.000+00:00") - - -def _build_prov_document( - 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 dict.""" - 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"] = [verification.get("dataset_run_id", "")] - 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}} - - -# ────────────────────────────────────────────── -# Slim ProvenanceCallback -# ────────────────────────────────────────────── class ProvenanceCallback(Callback): @@ -641,7 +230,7 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: _lookup_user_run, _snapshot_environment, ) - from rationai.mlkit.provenance.register_dataset import ( + from rationai.mlkit.provenance.dataset import ( _detect_manifest, _lookup_dataset_run, _verify_dataset, @@ -695,7 +284,7 @@ def _fallback_on_fit_start(self, trainer: Any, pl_module: Any) -> None: data_root = os.path.dirname(os.path.abspath(manifest_path)) if manifest_path and data_root: - from rationai.mlkit.provenance.register_dataset import ( + from rationai.mlkit.provenance.dataset import ( load_manifest, ) diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py index a289c4b..d9debd6 100644 --- a/rationai/mlkit/provenance/__init__.py +++ b/rationai/mlkit/provenance/__init__.py @@ -1,9 +1,10 @@ """Provenance tracking - PROV-O-aware logging to MLflow. Submodules: - prov - PROV-O document builders (W3C PROV compatible) - register_dataset - register_dataset (hash-based, emits prov.json) - register_user - register_new_user (emits prov.json) + 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`. @@ -22,20 +23,21 @@ if "MLFLOW_TRACKING_URI" not in os.environ: os.environ["MLFLOW_TRACKING_URI"] = "http://localhost:5000" -# Now safe to import - all child modules will pick up the env var -from rationai.mlkit.provenance.prov import ( +from rationai.mlkit.provenance.dataset import ( build_dataset_prov, - build_user_prov, -) -from rationai.mlkit.provenance.register_dataset import ( register_dataset, verify_dataset, ) -from rationai.mlkit.provenance.register_user import register_new_user +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", 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/register_dataset.py b/rationai/mlkit/provenance/dataset.py similarity index 59% rename from rationai/mlkit/provenance/register_dataset.py rename to rationai/mlkit/provenance/dataset.py index 6ed42df..7a900f1 100644 --- a/rationai/mlkit/provenance/register_dataset.py +++ b/rationai/mlkit/provenance/dataset.py @@ -1,4 +1,8 @@ -"""Dataset registration, verification, and legacy CSV-based paths.""" +"""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 @@ -11,8 +15,145 @@ 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}} -# ── Internal helpers ──────────────────────────────────────────────────── + +# ────────────────────────────────────────────── +# MLflow registration & verification helpers +# ────────────────────────────────────────────── def _lookup_experiment(name: str) -> str | None: @@ -60,12 +201,17 @@ def _detect_manifest() -> tuple[str | None, str | None]: 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 + Shared by ``register_dataset``, ``verify_dataset``, and ``ProvenanceCallback`` to avoid duplicating the CSV iteration pattern. """ df = pd.read_csv(manifest_path) @@ -77,6 +223,11 @@ def load_manifest(manifest_path: str, data_root: str) -> list[dict[str, Any]]: return samples +# ────────────────────────────────────────────── +# Dataset verification +# ────────────────────────────────────────────── + + def verify_dataset( manifest_path: str | None = None, data_root: str | None = None, @@ -162,15 +313,15 @@ def _verify_dataset( else: curr_file_sizes[basename] = -1 # missing - # Compare sizes — must have the same keys AND the same sizes - if set(curr_file_sizes) != set(reg_file_sizes): + 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)" - ) - result["file_sizes_match"] = all( + 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 = [ @@ -185,7 +336,6 @@ def _verify_dataset( + ("…" if len(mismatched) > 5 else ""), ) - # Check file existence missing = sum(1 for s in samples if not os.path.isfile(s["path"])) result["files_total"] = len(samples) result["files_missing"] = missing @@ -193,7 +343,6 @@ def _verify_dataset( if missing > 0: result["details"].append(f"{missing}/{len(samples)} WSI files missing on disk") - # Overall verdict result["verified"] = result["file_sizes_match"] and missing == 0 if result["verified"]: @@ -204,7 +353,9 @@ def _verify_dataset( return result -# ── Hash-based registration (preferred) ────────────────────────────────── +# ────────────────────────────────────────────── +# Hash-based registration (preferred) +# ────────────────────────────────────────────── def register_dataset( @@ -258,7 +409,6 @@ def register_dataset( file_sizes[basename] = -1 file_mtimes[basename] = -1 - # Register in MLflow mlflow.set_experiment(experiment_name) run = mlflow.start_run(run_name=f"Dataset_{dataset_name}_{version}") run_id = run.info.run_id @@ -283,9 +433,7 @@ def register_dataset( } ) - # ── PROV-O document (W3C PROV-O compatible) ──────────── - from rationai.mlkit.provenance.prov import build_dataset_prov - + # ── PROV-O document ──────────────────────────────── prov_doc = build_dataset_prov( run_id=run_id, dataset_name=dataset_name, @@ -308,7 +456,7 @@ def register_dataset( finally: shutil.rmtree(prov_dir, ignore_errors=True) - # ── Legacy dataset provenance JSON (kept for backward compat) ── + # ── 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") diff --git a/rationai/mlkit/provenance/prov.py b/rationai/mlkit/provenance/prov.py deleted file mode 100644 index c8c1511..0000000 --- a/rationai/mlkit/provenance/prov.py +++ /dev/null @@ -1,323 +0,0 @@ -"""Shared PROV-O document builder. - -Produces OpenProvenance-compatible JSON bundles compatible with the -Java ``prov_mlflow`` tool and used by both registration runs and -training provenance callbacks. -""" - -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 the PROV document -# ────────────────────────────────────────────── - - -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") - - -# ────────────────────────────────────────────── -# PROV document builders -# ────────────────────────────────────────────── - - -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}} - - -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) - - # Store per-file sizes as a single string value (matches prov_mlflow convention) - 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}} diff --git a/rationai/mlkit/provenance/register_user.py b/rationai/mlkit/provenance/register_user.py deleted file mode 100644 index b6657ce..0000000 --- a/rationai/mlkit/provenance/register_user.py +++ /dev/null @@ -1,96 +0,0 @@ -"""Register a researcher into MLflow's User_Registry experiment.""" - -from __future__ import annotations - -import json -import os -import shutil -import uuid - -import mlflow - - -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: - # Already exists (possibly created concurrently) - 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 ──────────────────────────────── - from rationai.mlkit.provenance.prov import build_user_prov - - 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) - 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) - - 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", - ) diff --git a/rationai/mlkit/provenance/run.py b/rationai/mlkit/provenance/run.py new file mode 100644 index 0000000..a1f1d85 --- /dev/null +++ b/rationai/mlkit/provenance/run.py @@ -0,0 +1,372 @@ +"""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"] = [verification.get("dataset_run_id", "")] + 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..4e8d154 --- /dev/null +++ b/rationai/mlkit/provenance/user.py @@ -0,0 +1,230 @@ +"""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) + 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) + + 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", + ) From 4c2f9e93ecdaf2ae73ae1db45a6b9141f80a596a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 16:32:58 +0200 Subject: [PATCH 30/34] feat: remove default MLflow tracking URI setup from provenance module --- rationai/mlkit/provenance/__init__.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py index d9debd6..3a267fb 100644 --- a/rationai/mlkit/provenance/__init__.py +++ b/rationai/mlkit/provenance/__init__.py @@ -8,9 +8,6 @@ For automatic provenance capture with Lightning, use :class:`~rationai.mlkit.lightning.callbacks.provenance.ProvenanceCallback`. - -MLflow tracking URI defaults to ``http://localhost:5000``. Override with -the ``MLFLOW_TRACKING_URI`` environment variable. """ from __future__ import annotations @@ -18,11 +15,6 @@ import os from typing import Any - -# ── Set default tracking URI before any mlflow import runs ─────────── -if "MLFLOW_TRACKING_URI" not in os.environ: - os.environ["MLFLOW_TRACKING_URI"] = "http://localhost:5000" - from rationai.mlkit.provenance.dataset import ( build_dataset_prov, register_dataset, From 46494d9dfe9940aae6f935fc30ca6d528aa1cff4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 16:43:15 +0200 Subject: [PATCH 31/34] feat: remove unused imports and deprecated __getattr__ function from provenance module --- rationai/mlkit/provenance/__init__.py | 15 --------------- 1 file changed, 15 deletions(-) diff --git a/rationai/mlkit/provenance/__init__.py b/rationai/mlkit/provenance/__init__.py index 3a267fb..f2c2b6b 100644 --- a/rationai/mlkit/provenance/__init__.py +++ b/rationai/mlkit/provenance/__init__.py @@ -12,9 +12,6 @@ from __future__ import annotations -import os -from typing import Any - from rationai.mlkit.provenance.dataset import ( build_dataset_prov, register_dataset, @@ -35,15 +32,3 @@ "register_new_user", "verify_dataset", ] - - -def __getattr__(name: str) -> Any: - """Raise helpful error for removed ``autolog``.""" - if name == "autolog": - raise ImportError( - "provenance.autolog has been removed. " - "Use ProvenanceCallback instead:\n\n" - " from rationai.mlkit.lightning.callbacks import ProvenanceCallback\n" - " trainer = Trainer(callbacks=[ProvenanceCallback(model_name='...')], ...)\n" - ) - raise AttributeError(f"module {__name__!r} has no attribute {name!r}") From 8f417663ae409146e2db4d1f3e10fe27cbe818ef Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 16:44:19 +0200 Subject: [PATCH 32/34] refactor: clean up whitespace and improve readability in provenance callback and run module --- rationai/mlkit/lightning/callbacks/provenance.py | 2 -- rationai/mlkit/provenance/run.py | 4 +--- 2 files changed, 1 insertion(+), 5 deletions(-) diff --git a/rationai/mlkit/lightning/callbacks/provenance.py b/rationai/mlkit/lightning/callbacks/provenance.py index 5eba018..18feef9 100644 --- a/rationai/mlkit/lightning/callbacks/provenance.py +++ b/rationai/mlkit/lightning/callbacks/provenance.py @@ -123,8 +123,6 @@ def _scheduler_summary(scheduler: Any) -> dict[str, str | float]: return info - - class ProvenanceCallback(Callback): """Lightning callback that captures PROV document + run summary. diff --git a/rationai/mlkit/provenance/run.py b/rationai/mlkit/provenance/run.py index a1f1d85..e307498 100644 --- a/rationai/mlkit/provenance/run.py +++ b/rationai/mlkit/provenance/run.py @@ -300,9 +300,7 @@ def _blank_rel_id() -> str: json.dumps(split_data["train"]) ) if split_data.get("test"): - meta_entity["gen:split_test"] = _typed_value( - json.dumps(split_data["test"]) - ) + meta_entity["gen:split_test"] = _typed_value(json.dumps(split_data["test"])) if requirements: meta_entity["gen:requirements"] = [requirements] From 43dce6984377ca8cd986792a8ea8f8ee102b5778 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 17:22:47 +0200 Subject: [PATCH 33/34] fix: ensure dataset_run_id is properly formatted and handle file cleanup in user registration --- rationai/mlkit/provenance/run.py | 2 +- rationai/mlkit/provenance/user.py | 14 ++++++++------ 2 files changed, 9 insertions(+), 7 deletions(-) diff --git a/rationai/mlkit/provenance/run.py b/rationai/mlkit/provenance/run.py index e307498..884e218 100644 --- a/rationai/mlkit/provenance/run.py +++ b/rationai/mlkit/provenance/run.py @@ -307,7 +307,7 @@ def _blank_rel_id() -> str: if verification: meta_entity["gen:dataset_verified"] = [str(verification.get("verified", False))] - meta_entity["gen:dataset_run_id"] = [verification.get("dataset_run_id", "")] + 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)] diff --git a/rationai/mlkit/provenance/user.py b/rationai/mlkit/provenance/user.py index 4e8d154..cfb81ca 100644 --- a/rationai/mlkit/provenance/user.py +++ b/rationai/mlkit/provenance/user.py @@ -208,12 +208,14 @@ def register_new_user( prov_dir = f"_user_prov_{uuid.uuid4().hex[:8]}" os.makedirs(prov_dir, exist_ok=True) - 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) + 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 From e6c69a3dd81b1b453a40b3f38165ec1e00a0ba9c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ji=C5=99=C3=AD=20Buchta?= Date: Wed, 22 Jul 2026 17:23:14 +0200 Subject: [PATCH 34/34] fix: format dataset_run_id assignment for improved readability in build_training_run_prov function --- rationai/mlkit/provenance/run.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/rationai/mlkit/provenance/run.py b/rationai/mlkit/provenance/run.py index 884e218..a9797e4 100644 --- a/rationai/mlkit/provenance/run.py +++ b/rationai/mlkit/provenance/run.py @@ -307,7 +307,11 @@ def _blank_rel_id() -> str: 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 ""] + 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)]