mnn-correct provides mutual-nearest-neighbour-based batch correction for
AnnData objects. The package is centered on a stateful MNNCorrector that
supports three steps:
fit()estimates per-batch displacement models.correct()applies the fitted correction to the training dataset.project()propagates an already learned batch displacement onto new cells based on a batch column in the queryAnnData.
By default, fit() updates the current MNNCorrector in place and returns
None. Pass return_corrector=True if you want it to return the fitted
corrector instance for chaining.
Compared with mnnpy, mnn-correct uses a different correction mechanism and
workflow. mnnpy follows the classic pattern of finding MNN pairs, averaging
pairwise correction vectors, and smoothing them with a Gaussian kernel
controlled by sigma, with optional variance adjustment and biological
subspace correction. By contrast, mnn-correct treats MNN-supported query
cells as anchors, estimates their displacements, and propagates those
displacements across the full batch with a configurable weighted KNN graph
(default: jaccard_square). It also retains fitted batch-specific state so the
learned correction can later be projected onto new cells from already seen
batches.
Install the package with pip in your existing Python environment:
pip install .Core runtime dependencies are listed in pyproject.toml.
The corrector expects a shared representation for all cells that should be aligned. In practice this is usually something like:
adata.obsm["X_pca"]adata.obsm["X_scVI"]adata.obsm["X_harmony"]
If no representation is supplied, the corrector can fit a PCA model directly on
adata.X by passing use_rep=None. This PCA fallback is only used for the
current correction run and does not support future projection.
During fitting, the model:
- identifies mutual nearest neighbour pairs between a reference and query batch,
- estimates anchor-cell displacement vectors,
- propagates those displacements to all cells in the query batch,
- stores batch-specific projection state for later reuse when a reusable representation was supplied.
from mnn_correct import MNNCorrector
corrector = MNNCorrector(
k_mnn=10,
k_propagate=20,
weighting_scheme="jaccard_square",
)
corrector.fit(
adata,
batch_key="batch",
batch_order=["reference", "query_a", "query_b"],
use_rep="X_scVI",
)
# Optional chaining-friendly form:
same_corrector = corrector.fit(
adata,
batch_key="batch",
batch_order=["reference", "query_a", "query_b"],
use_rep="X_scVI",
return_corrector=True,
)corrector.correct(adata)
corrected = adata.obsm["X_scVI_mnn_corrected"]new_adata.obs["batch"] = ["query_b", "query_b", "reference", ...]
corrector.project(
new_adata,
batch_key="batch",
)
projected = new_adata.obsm["X_scVI_mnn_corrected"]project() validates that every batch category in new_adata.obs[batch_key]
was seen during fit(). Cells from the initial sequential batch or the fixed
reference batch are projected as an identity mapping. If you need to apply the
learned correction back to the original training dataset, use correct()
instead.
If fit() was run with use_rep=None, project() is unavailable because the
PCA fallback is not stored as a reusable projection model.
In sequential mode, batches are corrected one after another. Each corrected batch is merged into the growing reference before the next round.
corrector.fit(
adata,
batch_key="batch",
batch_order=["batch0", "batch1", "batch2"],
use_rep="X_pca",
)This is useful when the batches form a natural progression or when no single batch should be privileged as the sole reference.
In fixed-reference mode, each non-reference batch is corrected independently against one chosen reference batch.
corrector.fit(
adata,
batch_key="batch",
reference="atlas",
use_rep="X_scVI",
)This is useful when one batch serves as the canonical reference, such as an atlas or a well-curated control dataset.
Use mnn_correct() when you already have separate reference and query
AnnData objects.
from mnn_correct import mnn_correct
corrector = mnn_correct(
adata_ref,
adata_query,
use_rep="X_scVI",
batch_label="query_batch",
)
corrected_query = adata_query.obsm["X_scVI_mnn_corrected"]The reference object is left unchanged. The corrected embedding is written only to the query object.
Use mnn_correct_adata() when all batches are already contained in a single
AnnData object.
from mnn_correct import mnn_correct_adata
_, corrector = mnn_correct_adata(
adata,
batch_key="batch",
batch_order=["batch0", "batch1", "batch2"],
use_rep="X_pca",
)By default, corrected embeddings are written to:
"{use_rep}_mnn_corrected"whenuse_repis provided"X_pca_mnn_corrected"whenuse_rep=None
You can override this with key_added in fit(), correct(), project(),
mnn_correct(), or mnn_correct_adata().
fit()does not modify the inputAnnData.fit()updates theMNNCorrectorin place and returnsNoneunlessreturn_corrector=Trueis passed.correct()applies only to the same fitted dataset and validates that both the cells and source representation match the fitted state.project()is for new cells whose batch assignments are stored in a column ofadata.obs, and every batch in that column must already be known fromfit().- Projection state is stored per batch label in
corrector.projection_data_only when a reusable representation was supplied. - The initial sequential batch or fixed reference batch projects as an identity
mapping and is stored separately from
corrector.projection_data_. - If
use_rep=None,fit()falls back to PCA for the current correction and forcesstore_for_projection=False.
k_mnn: number of neighbors used to identify MNN pairs.k_propagate: number of neighbors used to smooth anchor-cell displacement to the full batch.weighting_scheme: how propagation edges are weighted. The default is"jaccard_square".use_rep: which latent embedding to correct.
For contributor setup, this repository also supports uv with the checked-in
uv.lock file:
uv sync --frozen --extra devIf uv needs the pinned Python version first:
uv python install 3.10.19Run the package commands in that environment with uv run, for example:
uv run pytestRun the test suite:
uv run pytestRun linting:
uv run ruff check src testsRun type checking:
uv run mypy srcMIT