Skip to main content

Boilerplate-free, reproducible ML experiment workflows built on PyTorch Lightning and hydra-zen. Carved out of MIT-LL's responsible-ai-toolbox.

Project description

mushin logo

mushin

CI PyPI Python versions License: MIT

Docs: https://martinez-hub.github.io/mushin/

Boilerplate-free, reproducible machine-learning experiment workflows built on PyTorch Lightning and hydra-zen.

mushin is a standalone carve-out of the rai_toolbox.mushin subpackage from MIT Lincoln Laboratory's responsible-ai-toolbox. The upstream toolbox is no longer maintained (last release May 2023), but the mushin workflow layer still works against current versions of its dependencies. This package extracts just that layer so it can be maintained and used on its own.

Quickstart: run a sweep, get a dataset

Decorate your experiment with @mushin.sweep, sweep over parameters, and get the results back as a labeled xarray.Dataset — not rows in a dashboard you have to export. No subclassing, no callbacks:

import mushin

@mushin.sweep
def experiment(lr, seed):
    # ... train a model with this lr/seed, then evaluate it ...
    acc = ...  # your validation accuracy
    return dict(accuracy=acc)  # whatever you return becomes a data variable

ds = experiment.run(
    lr=mushin.multirun([0.01, 0.1, 1.0]),
    seed=mushin.multirun([0, 1, 2]),
)  # 9 runs, returned as a labeled xarray.Dataset
# <xarray.Dataset> Dimensions: (lr: 3, seed: 3)
#   Data variables: accuracy (lr, seed)

ds["accuracy"].mean("seed")   # average over seeds, per learning rate

Need the full tool — .failures, .plot(), provenance, custom analysis? Drop to experiment.workflow (the last-run instance), or subclass MultiRunMetricsWorkflow directly for advanced control (custom pre_task, jobs_post_process) — which is exactly what @mushin.sweep builds for you.

The full runnable version is in examples/sweep_to_dataset.py:

uv run python examples/sweep_to_dataset.py

The sweep layer is framework-agnostic — your task just returns a dict, so you can wrap scikit-learn, XGBoost, or anything and still get a labeled xarray.Dataset back (see examples/sklearn_sweep.py). The Lightning-specific conveniences (auto-tuning, HydraDDP, the compare batteries) assume PyTorch.

Compare methods, with statistics

Evaluate trained models on a standard battery and get a labeled dataset plus significance — metrics delegated to torchmetrics, statistics to scipy:

from mushin.benchmark import compare

result = compare(
    methods={"ours": [m0, m1, m2], "baseline": [b0, b1, b2]},  # one trained model per seed
    data=test_loader, task="classification", num_classes=10, test="welch",
)

result.summary()       # mean ± CI per method, with significance markers — paper-ready
result.comparisons     # tidy DataFrame: pairwise p-values + effect sizes
result.data            # the labeled xarray (method × seed) to slice and plot

Don't have the trained models in memory yet? Study runs the multi-seed training sweep (via Hydra) and feeds the results straight into compare — define → train → evaluate → report in one call:

from mushin import Study

study = Study(
    methods={"cnn": train_cnn, "mlp": train_mlp},   # train_fn(seed) -> checkpoint path
    load_fn=LitClassifier.load_from_checkpoint,       # path -> model
    seeds=[0, 1, 2], data=test_loader, num_classes=10, test="welch",
)
result = study.run()                                  # -> BenchmarkResult

# ...or compare checkpoints you already have, no training:
Study.from_checkpoints(
    checkpoints={"cnn": ["cnn_0.ckpt", ...], "mlp": ["mlp_0.ckpt", ...]},
    load_fn=LitClassifier.load_from_checkpoint,
    data=test_loader, num_classes=10, test="welch",
).run()

Compare LLM systems, with statistics

The same significance spine extends to LLM systems — the one eval setting where "is this difference real, or just sampling noise?" is most often skipped. Bring your own systems, data, and metric; mushin runs each across reproducible seeds and reports Holm-corrected significance:

from mushin.llm import compare_llms, llm_judge

result = compare_llms(
    systems={"prompt_a": system_a, "prompt_b": system_b},  # system(inputs, seed) -> outputs
    data=eval_data,                                         # [{"input": ..., "reference": ...}, ...]
    metric=llm_judge(my_judge, rubric="Is the answer correct?"),  # or any callable / torchmetric
    seeds=range(5), test="welch",
)
result.summary()       # same paper-ready table as the torch path

Systems can be plain callables or hydra-zen configs (instantiated once, reused across seeds); metrics can be a plain scorer, a torchmetrics text metric, or a named battery. An optional on-disk output cache makes reruns/resumes free. See the LLM evaluation guide.

What it provides

  • @mushin.sweep — the boilerplate-free core: decorate a task-style function and experiment.run(lr=multirun([...]), seed=multirun([...])) returns the labeled xarray.Dataset directly. Drop to experiment.workflow / experiment.workflow_cls for the full workflow, or subclass MultiRunMetricsWorkflow for advanced control.
  • benchmark.compare — run a standard metric battery (torchmetrics) across trained seeds and get a labeled dataset + significance (scipy): BenchmarkResult with .summary(), .comparisons, and .data.
  • register_task, get_task, list_tasks, Task — first-class, reusable evaluation tasks; compare and Study accept a task name or a Task. Built-in batteries: classification, segmentation, detection, regression, retrieval, image_quality, audio.
  • llm.compare_llms, llm.llm_judge — compare LLM systems across reproducible seeds with significance (callables or hydra-zen configs; plain / torchmetrics / judge metrics; optional output cache).
  • Study — orchestrate a multi-seed training sweep and route the trained models into compare, in one call; Study.from_checkpoints(...) for eval-only.
  • MultiRunMetricsWorkflow (plus its base mushin.workflows.BaseWorkflow and the mushin.workflows.RobustnessCurve variant) — declarative, reproducible experiment workflows that record configs, checkpoints, and metrics, and load results back as labeled xarray datasets.
  • Sweep resilience + provenancerun(on_error="nan") records a failed grid cell as NaN and keeps going; run(working_dir=..., resume=True) re-runs only the failed/missing cells and is durable across a hard process kill (OOM, SLURM preemption) — completed cells are never recomputed; a task can declare a mushin_resume parameter to resume its own training mid-run from the last checkpoint; compare/Study refuse statistics on an incomplete sweep (IncompleteSweepError) until you fix and resume; every run captures per-run provenance (mushin_provenance.json: git SHA, versions, config). See the resilience guide.
  • Out-of-process launchers — dispatch is stdlib-picklable, so sweeps run across cores or a scheduler with a standard Hydra launcher plugin: run(..., launcher="joblib") (or submitit). See the workflows guide.
  • tune_batch_size, tune_learning_rate — opt-in, reproducibility-preserving auto-tuning: find the batch size / LR once, pin it to a sidecar file, and reuse it — with an exact, hardware-independent effective batch (no drift).
  • MetricsCallback — a Lightning callback for capturing metrics.
  • HydraDDP — a Hydra/Lightning strategy for multi-GPU (DDP) launches.
  • multirun, hydra_list, load_experiment, load_from_checkpoint — helpers.

Analyze experiments from Claude Code (MCP)

mushin ships an optional read-only MCP server so Claude Code (or any MCP client) can load and analyze your completed runs — list experiments, summarize swept parameters and metrics, and inspect saved datasets — without launching anything.

pip install "mushin-py[mcp]"          # requires Python >= 3.10
claude mcp add mushin -- mushin-mcp --root ./outputs

See the MCP guide for the full tool list and example prompts.

Install

pip install mushin-py

Already use uv? uv pip install mushin-py (or uv add mushin-py inside a project) is faster.

Install name vs. import name: the PyPI distribution is mushin-py, but you import mushin (same pattern as scikit-learnsklearn).

Optional extras: viz (matplotlib, for plotting results) and netcdf (netCDF4) for the core; detection, image, and audio for those benchmark batteries; and mcp for the MCP server — e.g. pip install "mushin-py[viz]".

For a development environment (runtime deps + dev tooling), this project uses uv: uv sync.

Develop

uv run pytest tests/ --hypothesis-profile fast   # tests (DDP test needs >=2 GPUs)
uv run ruff check .                              # lint
uv run ruff format .                             # format
uv run codespell src tests                       # spell check

Or use the make shortcuts (make help to list them): make check runs lint + format-check + spell + tests (what CI runs); make test-py PYTHON=3.12 runs the suite on a specific Python version.

Supported Python versions: 3.10 – 3.14.

Relationship to upstream

This is a fork/extraction, not a replacement endorsed by MIT-LL. The configuration engine it depends on, hydra-zen, is actively maintained by the same group. See LICENSE.txt for attribution; the original MIT copyright is retained.

Project details


Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

mushin_py-0.7.0.tar.gz (118.0 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

mushin_py-0.7.0-py3-none-any.whl (93.0 kB view details)

Uploaded Python 3

File details

Details for the file mushin_py-0.7.0.tar.gz.

File metadata

  • Download URL: mushin_py-0.7.0.tar.gz
  • Upload date:
  • Size: 118.0 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.12

File hashes

Hashes for mushin_py-0.7.0.tar.gz
Algorithm Hash digest
SHA256 fc3f87af0b631ffc3ac11ac7d1475891dd6ade768d869b3f1063509b0e67b9e4
MD5 e5ae085e2fec2ea785d2c9656a951c0d
BLAKE2b-256 f597102e5c980c50469ad55fbaa9f40221531e6e083364c41320b14a4707aa7d

See more details on using hashes here.

Provenance

The following attestation bundles were made for mushin_py-0.7.0.tar.gz:

Publisher: publish.yml on martinez-hub/mushin

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file mushin_py-0.7.0-py3-none-any.whl.

File metadata

  • Download URL: mushin_py-0.7.0-py3-none-any.whl
  • Upload date:
  • Size: 93.0 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/6.1.0 CPython/3.13.12

File hashes

Hashes for mushin_py-0.7.0-py3-none-any.whl
Algorithm Hash digest
SHA256 ddb64dbe7b63678f4f3f212f223c0468b76f68997d047307882a3855e780938d
MD5 e87cc4c0655c57e19f18c1545c6ba10f
BLAKE2b-256 6188abad672f677702c7ee33020e17f62be36b1ff314da5da2a0b2a8399ebcb7

See more details on using hashes here.

Provenance

The following attestation bundles were made for mushin_py-0.7.0-py3-none-any.whl:

Publisher: publish.yml on martinez-hub/mushin

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Pingdom Monitoring Sentry Error logging StatusPage Status page