Metadata-Version: 2.4
Name: bernn
Version: 1.0.6
Summary: Batch Effect Removal Neural Networks for Tandem Mass Spectrometry
Home-page: https://github.com/spell00/BERNN_MSMS
Author: Simon Pelletier
Author-email: 
License: MIT
Classifier: Programming Language :: Python :: 3
Classifier: Programming Language :: Python :: 3.11
Classifier: Programming Language :: Python :: 3.12
Classifier: Programming Language :: Python :: 3.13
Classifier: License :: OSI Approved :: MIT License
Classifier: Operating System :: OS Independent
Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
Classifier: Topic :: Scientific/Engineering :: Bio-Informatics
Requires-Python: >=3.10
Description-Content-Type: text/markdown
Requires-Dist: scikit-learn>=1.6.0
Requires-Dist: pandas>=2.2.0
Requires-Dist: scikit-optimize>=0.9.0
Requires-Dist: matplotlib>=3.7.0
Requires-Dist: seaborn>=0.12.2
Requires-Dist: tabulate>=0.9.0
Requires-Dist: scipy>=1.11.0
Requires-Dist: tqdm
Requires-Dist: joblib>=1.3.0
Requires-Dist: psutil>=5.9.4
Requires-Dist: scikit-image>=0.21.0
Requires-Dist: nibabel
Requires-Dist: mpmath>=1.3.0
Requires-Dist: patsy>=0.5.3
Requires-Dist: umap-learn>=0.5.3
Requires-Dist: shapely
Requires-Dist: numba>=0.58.0
Requires-Dist: openpyxl>=3.0.10
Requires-Dist: xgboost>=1.7.0
Requires-Dist: importlib-metadata>=6.0.0
Requires-Dist: threadpoolctl>=3.1.0
Requires-Dist: protobuf<7,>=6.31.1
Requires-Dist: requests<3.0.0,>=2.31.0
Requires-Dist: PyYAML>=6.0.1
Requires-Dist: python-dateutil>=2.8.2
Requires-Dist: nbformat>=5.9.2
Requires-Dist: statsmodels
Provides-Extra: minimal
Provides-Extra: core-extended
Requires-Dist: shap; extra == "core-extended"
Requires-Dist: pytest; extra == "core-extended"
Requires-Dist: pytest-cov; extra == "core-extended"
Requires-Dist: cython>=0.29.21; extra == "core-extended"
Requires-Dist: FuzzyTM>=0.4.0; extra == "core-extended"
Requires-Dist: blosc2<3.0.0,>=2.0.0; extra == "core-extended"
Requires-Dist: llvmlite>=0.40.1; extra == "core-extended"
Requires-Dist: pycombat; extra == "core-extended"
Provides-Extra: deep-learning
Requires-Dist: torch>=2.1.0; extra == "deep-learning"
Requires-Dist: torchvision>=0.16.0; extra == "deep-learning"
Requires-Dist: torch-geometric; extra == "deep-learning"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "deep-learning"
Requires-Dist: typing-extensions>=4.9.0; extra == "deep-learning"
Requires-Dist: numpy<2.3,>=1.24; extra == "deep-learning"
Requires-Dist: six>=1.16.0; extra == "deep-learning"
Provides-Extra: experiment-tracking
Requires-Dist: tensorboardX; extra == "experiment-tracking"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "experiment-tracking"
Requires-Dist: sqlalchemy>=2.0.0; extra == "experiment-tracking"
Requires-Dist: urllib3>=1.26.7; extra == "experiment-tracking"
Provides-Extra: notebooks
Requires-Dist: notebook>=7.0.0; extra == "notebooks"
Requires-Dist: ipywidgets>=8.0.0; extra == "notebooks"
Requires-Dist: jupyterlab>=4.0.0; extra == "notebooks"
Provides-Extra: tools
Requires-Dist: packaging>=21.0; extra == "tools"
Requires-Dist: python-dateutil>=2.8.2; extra == "tools"
Requires-Dist: PyYAML>=6.0.1; extra == "tools"
Requires-Dist: optuna>=3.0.0; extra == "tools"
Provides-Extra: python311-plus
Requires-Dist: torch>=2.1.0; extra == "python311-plus"
Requires-Dist: torchvision>=0.16.0; extra == "python311-plus"
Requires-Dist: torch-geometric; extra == "python311-plus"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "python311-plus"
Requires-Dist: typing-extensions>=4.9.0; extra == "python311-plus"
Requires-Dist: numpy<2.3,>=1.24; extra == "python311-plus"
Requires-Dist: six>=1.16.0; extra == "python311-plus"
Requires-Dist: tensorboardX; extra == "python311-plus"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "python311-plus"
Requires-Dist: sqlalchemy>=2.0.0; extra == "python311-plus"
Requires-Dist: urllib3>=1.26.7; extra == "python311-plus"
Provides-Extra: tools-with-ax
Requires-Dist: packaging>=21.0; extra == "tools-with-ax"
Requires-Dist: python-dateutil>=2.8.2; extra == "tools-with-ax"
Requires-Dist: PyYAML>=6.0.1; extra == "tools-with-ax"
Requires-Dist: optuna>=3.0.0; extra == "tools-with-ax"
Requires-Dist: ax-platform; extra == "tools-with-ax"
Provides-Extra: web
Requires-Dist: fastapi<0.103.0,>=0.89.1; extra == "web"
Requires-Dist: websocket-client>=1.8.0; extra == "web"
Requires-Dist: platformdirs<4.2.0,>=3.11.0; extra == "web"
Provides-Extra: web-dev
Requires-Dist: fastapi<0.104.0,>=0.103.0; extra == "web-dev"
Requires-Dist: pydantic<3.0.0,>=2.6.4; extra == "web-dev"
Requires-Dist: platformdirs<5.0.0,>=4.2.0; extra == "web-dev"
Provides-Extra: external-tools
Requires-Dist: spyder>=5.0.0; extra == "external-tools"
Requires-Dist: selenium<4.25.0,>=4.15.0; extra == "external-tools"
Requires-Dist: spotdl<4.2.5,>=4.2.0; extra == "external-tools"
Provides-Extra: typing
Requires-Dist: typing-extensions>=4.9.0; extra == "typing"
Provides-Extra: r-integration
Requires-Dist: rpy2>=3.6.0; extra == "r-integration"
Provides-Extra: dev-tools
Requires-Dist: jedi>=0.18.2; extra == "dev-tools"
Provides-Extra: special
Requires-Dist: pykan; extra == "special"
Provides-Extra: ml-full
Requires-Dist: torch>=2.1.0; extra == "ml-full"
Requires-Dist: torchvision>=0.16.0; extra == "ml-full"
Requires-Dist: torch-geometric; extra == "ml-full"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "ml-full"
Requires-Dist: typing-extensions>=4.9.0; extra == "ml-full"
Requires-Dist: numpy<2.3,>=1.24; extra == "ml-full"
Requires-Dist: six>=1.16.0; extra == "ml-full"
Requires-Dist: tensorboardX; extra == "ml-full"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "ml-full"
Requires-Dist: sqlalchemy>=2.0.0; extra == "ml-full"
Requires-Dist: urllib3>=1.26.7; extra == "ml-full"
Provides-Extra: analysis
Requires-Dist: notebook>=7.0.0; extra == "analysis"
Requires-Dist: ipywidgets>=8.0.0; extra == "analysis"
Requires-Dist: jupyterlab>=4.0.0; extra == "analysis"
Requires-Dist: packaging>=21.0; extra == "analysis"
Requires-Dist: python-dateutil>=2.8.2; extra == "analysis"
Requires-Dist: PyYAML>=6.0.1; extra == "analysis"
Requires-Dist: optuna>=3.0.0; extra == "analysis"
Requires-Dist: pykan; extra == "analysis"
Provides-Extra: analysis-with-ax
Requires-Dist: notebook>=7.0.0; extra == "analysis-with-ax"
Requires-Dist: ipywidgets>=8.0.0; extra == "analysis-with-ax"
Requires-Dist: jupyterlab>=4.0.0; extra == "analysis-with-ax"
Requires-Dist: packaging>=21.0; extra == "analysis-with-ax"
Requires-Dist: python-dateutil>=2.8.2; extra == "analysis-with-ax"
Requires-Dist: PyYAML>=6.0.1; extra == "analysis-with-ax"
Requires-Dist: optuna>=3.0.0; extra == "analysis-with-ax"
Requires-Dist: ax-platform; extra == "analysis-with-ax"
Requires-Dist: pykan; extra == "analysis-with-ax"
Provides-Extra: development
Requires-Dist: shap; extra == "development"
Requires-Dist: pytest; extra == "development"
Requires-Dist: pytest-cov; extra == "development"
Requires-Dist: cython>=0.29.21; extra == "development"
Requires-Dist: jedi>=0.18.2; extra == "development"
Requires-Dist: fastapi<0.103.0,>=0.89.1; extra == "development"
Requires-Dist: websocket-client>=1.8.0; extra == "development"
Requires-Dist: platformdirs<4.2.0,>=3.11.0; extra == "development"
Provides-Extra: modern-web
Requires-Dist: fastapi<0.104.0,>=0.103.0; extra == "modern-web"
Requires-Dist: pydantic<3.0.0,>=2.6.4; extra == "modern-web"
Requires-Dist: platformdirs<5.0.0,>=4.2.0; extra == "modern-web"
Requires-Dist: typing-extensions>=4.9.0; extra == "modern-web"
Provides-Extra: ide-tools
Requires-Dist: spyder>=5.0.0; extra == "ide-tools"
Requires-Dist: selenium<4.25.0,>=4.15.0; extra == "ide-tools"
Requires-Dist: spotdl<4.2.5,>=4.2.0; extra == "ide-tools"
Requires-Dist: typing-extensions>=4.9.0; extra == "ide-tools"
Provides-Extra: python313-ml-minimal
Requires-Dist: torch>=2.1.0; extra == "python313-ml-minimal"
Requires-Dist: torchvision>=0.16.0; extra == "python313-ml-minimal"
Requires-Dist: torch-geometric; extra == "python313-ml-minimal"
Requires-Dist: scikit-learn>=1.6.0; extra == "python313-ml-minimal"
Requires-Dist: typing-extensions>=4.9.0; extra == "python313-ml-minimal"
Requires-Dist: numpy<2.3,>=1.24; extra == "python313-ml-minimal"
Provides-Extra: python313-ml-stable
Requires-Dist: torch>=2.1.0; extra == "python313-ml-stable"
Requires-Dist: torchvision>=0.16.0; extra == "python313-ml-stable"
Requires-Dist: torch-geometric; extra == "python313-ml-stable"
Requires-Dist: tensorflow>=2.15.0; extra == "python313-ml-stable"
Requires-Dist: scikit-learn>=1.6.0; extra == "python313-ml-stable"
Requires-Dist: typing-extensions>=4.9.0; extra == "python313-ml-stable"
Requires-Dist: numpy<2.3,>=1.24; extra == "python313-ml-stable"
Provides-Extra: python313-minimal-safe
Requires-Dist: torch>=2.1.0; extra == "python313-minimal-safe"
Requires-Dist: torchvision>=0.16.0; extra == "python313-minimal-safe"
Requires-Dist: torch-geometric; extra == "python313-minimal-safe"
Requires-Dist: scikit-learn>=1.6.0; extra == "python313-minimal-safe"
Requires-Dist: pandas>=2.2.0; extra == "python313-minimal-safe"
Requires-Dist: matplotlib>=3.7.0; extra == "python313-minimal-safe"
Requires-Dist: seaborn>=0.12.2; extra == "python313-minimal-safe"
Requires-Dist: numpy<2.3,>=1.24; extra == "python313-minimal-safe"
Requires-Dist: scipy>=1.11.0; extra == "python313-minimal-safe"
Requires-Dist: jupyter>=1.0.0; extra == "python313-minimal-safe"
Provides-Extra: python313-safe
Requires-Dist: shap; extra == "python313-safe"
Requires-Dist: pytest; extra == "python313-safe"
Requires-Dist: pytest-cov; extra == "python313-safe"
Requires-Dist: cython>=0.29.21; extra == "python313-safe"
Requires-Dist: FuzzyTM>=0.4.0; extra == "python313-safe"
Requires-Dist: blosc2<3.0.0,>=2.0.0; extra == "python313-safe"
Requires-Dist: llvmlite>=0.40.1; extra == "python313-safe"
Requires-Dist: pycombat; extra == "python313-safe"
Requires-Dist: torch>=2.1.0; extra == "python313-safe"
Requires-Dist: torchvision>=0.16.0; extra == "python313-safe"
Requires-Dist: torch-geometric; extra == "python313-safe"
Requires-Dist: scikit-learn>=1.6.0; extra == "python313-safe"
Requires-Dist: typing-extensions>=4.9.0; extra == "python313-safe"
Requires-Dist: numpy<2.3,>=1.24; extra == "python313-safe"
Requires-Dist: notebook>=7.0.0; extra == "python313-safe"
Requires-Dist: ipywidgets>=8.0.0; extra == "python313-safe"
Requires-Dist: jupyterlab>=4.0.0; extra == "python313-safe"
Requires-Dist: pykan; extra == "python313-safe"
Provides-Extra: full-no-ax
Requires-Dist: shap; extra == "full-no-ax"
Requires-Dist: pytest; extra == "full-no-ax"
Requires-Dist: pytest-cov; extra == "full-no-ax"
Requires-Dist: cython>=0.29.21; extra == "full-no-ax"
Requires-Dist: FuzzyTM>=0.4.0; extra == "full-no-ax"
Requires-Dist: blosc2<3.0.0,>=2.0.0; extra == "full-no-ax"
Requires-Dist: llvmlite>=0.40.1; extra == "full-no-ax"
Requires-Dist: pycombat; extra == "full-no-ax"
Requires-Dist: torch>=2.1.0; extra == "full-no-ax"
Requires-Dist: torchvision>=0.16.0; extra == "full-no-ax"
Requires-Dist: torch-geometric; extra == "full-no-ax"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "full-no-ax"
Requires-Dist: typing-extensions>=4.9.0; extra == "full-no-ax"
Requires-Dist: numpy<2.3,>=1.24; extra == "full-no-ax"
Requires-Dist: six>=1.16.0; extra == "full-no-ax"
Requires-Dist: tensorboardX; extra == "full-no-ax"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "full-no-ax"
Requires-Dist: sqlalchemy>=2.0.0; extra == "full-no-ax"
Requires-Dist: urllib3>=1.26.7; extra == "full-no-ax"
Requires-Dist: notebook>=7.0.0; extra == "full-no-ax"
Requires-Dist: ipywidgets>=8.0.0; extra == "full-no-ax"
Requires-Dist: jupyterlab>=4.0.0; extra == "full-no-ax"
Requires-Dist: packaging>=21.0; extra == "full-no-ax"
Requires-Dist: python-dateutil>=2.8.2; extra == "full-no-ax"
Requires-Dist: PyYAML>=6.0.1; extra == "full-no-ax"
Requires-Dist: optuna>=3.0.0; extra == "full-no-ax"
Requires-Dist: fastapi<0.103.0,>=0.89.1; extra == "full-no-ax"
Requires-Dist: websocket-client>=1.8.0; extra == "full-no-ax"
Requires-Dist: platformdirs<4.2.0,>=3.11.0; extra == "full-no-ax"
Requires-Dist: rpy2>=3.6.0; extra == "full-no-ax"
Requires-Dist: pykan; extra == "full-no-ax"
Provides-Extra: full
Requires-Dist: shap; extra == "full"
Requires-Dist: pytest; extra == "full"
Requires-Dist: pytest-cov; extra == "full"
Requires-Dist: cython>=0.29.21; extra == "full"
Requires-Dist: FuzzyTM>=0.4.0; extra == "full"
Requires-Dist: blosc2<3.0.0,>=2.0.0; extra == "full"
Requires-Dist: llvmlite>=0.40.1; extra == "full"
Requires-Dist: pycombat; extra == "full"
Requires-Dist: torch>=2.1.0; extra == "full"
Requires-Dist: torchvision>=0.16.0; extra == "full"
Requires-Dist: torch-geometric; extra == "full"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "full"
Requires-Dist: typing-extensions>=4.9.0; extra == "full"
Requires-Dist: numpy<2.3,>=1.24; extra == "full"
Requires-Dist: six>=1.16.0; extra == "full"
Requires-Dist: tensorboardX; extra == "full"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "full"
Requires-Dist: sqlalchemy>=2.0.0; extra == "full"
Requires-Dist: urllib3>=1.26.7; extra == "full"
Requires-Dist: notebook>=7.0.0; extra == "full"
Requires-Dist: ipywidgets>=8.0.0; extra == "full"
Requires-Dist: jupyterlab>=4.0.0; extra == "full"
Requires-Dist: packaging>=21.0; extra == "full"
Requires-Dist: python-dateutil>=2.8.2; extra == "full"
Requires-Dist: PyYAML>=6.0.1; extra == "full"
Requires-Dist: optuna>=3.0.0; extra == "full"
Requires-Dist: ax-platform; extra == "full"
Requires-Dist: fastapi<0.103.0,>=0.89.1; extra == "full"
Requires-Dist: websocket-client>=1.8.0; extra == "full"
Requires-Dist: platformdirs<4.2.0,>=3.11.0; extra == "full"
Requires-Dist: rpy2>=3.6.0; extra == "full"
Requires-Dist: pykan; extra == "full"
Provides-Extra: full-safe
Requires-Dist: shap; extra == "full-safe"
Requires-Dist: pytest; extra == "full-safe"
Requires-Dist: pytest-cov; extra == "full-safe"
Requires-Dist: cython>=0.29.21; extra == "full-safe"
Requires-Dist: torch>=2.1.0; extra == "full-safe"
Requires-Dist: torchvision>=0.16.0; extra == "full-safe"
Requires-Dist: torch-geometric; extra == "full-safe"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "full-safe"
Requires-Dist: typing-extensions>=4.9.0; extra == "full-safe"
Requires-Dist: numpy<2.3,>=1.24; extra == "full-safe"
Requires-Dist: six>=1.16.0; extra == "full-safe"
Requires-Dist: tensorboardX; extra == "full-safe"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "full-safe"
Requires-Dist: sqlalchemy>=2.0.0; extra == "full-safe"
Requires-Dist: urllib3>=1.26.7; extra == "full-safe"
Requires-Dist: notebook>=7.0.0; extra == "full-safe"
Requires-Dist: ipywidgets>=8.0.0; extra == "full-safe"
Requires-Dist: jupyterlab>=4.0.0; extra == "full-safe"
Requires-Dist: pykan; extra == "full-safe"
Provides-Extra: py311-plus
Requires-Dist: torch>=2.1.0; extra == "py311-plus"
Requires-Dist: torchvision>=0.16.0; extra == "py311-plus"
Requires-Dist: torch-geometric; extra == "py311-plus"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "py311-plus"
Requires-Dist: typing-extensions>=4.9.0; extra == "py311-plus"
Requires-Dist: numpy<2.3,>=1.24; extra == "py311-plus"
Requires-Dist: six>=1.16.0; extra == "py311-plus"
Requires-Dist: tensorboardX; extra == "py311-plus"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "py311-plus"
Requires-Dist: sqlalchemy>=2.0.0; extra == "py311-plus"
Requires-Dist: urllib3>=1.26.7; extra == "py311-plus"
Provides-Extra: py312-plus
Requires-Dist: torch>=2.1.0; extra == "py312-plus"
Requires-Dist: torchvision>=0.16.0; extra == "py312-plus"
Requires-Dist: torch-geometric; extra == "py312-plus"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "py312-plus"
Requires-Dist: typing-extensions>=4.9.0; extra == "py312-plus"
Requires-Dist: numpy<2.3,>=1.24; extra == "py312-plus"
Requires-Dist: six>=1.16.0; extra == "py312-plus"
Requires-Dist: tensorboardX; extra == "py312-plus"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "py312-plus"
Requires-Dist: sqlalchemy>=2.0.0; extra == "py312-plus"
Requires-Dist: urllib3>=1.26.7; extra == "py312-plus"
Provides-Extra: py313-plus
Requires-Dist: torch>=2.1.0; extra == "py313-plus"
Requires-Dist: torchvision>=0.16.0; extra == "py313-plus"
Requires-Dist: torch-geometric; extra == "py313-plus"
Requires-Dist: tensorflow>=2.20.0rc0; extra == "py313-plus"
Requires-Dist: typing-extensions>=4.9.0; extra == "py313-plus"
Requires-Dist: numpy<2.3,>=1.24; extra == "py313-plus"
Requires-Dist: six>=1.16.0; extra == "py313-plus"
Requires-Dist: tensorboardX; extra == "py313-plus"
Requires-Dist: mlflow[extras]>=2.12.1; extra == "py313-plus"
Requires-Dist: sqlalchemy>=2.0.0; extra == "py313-plus"
Requires-Dist: urllib3>=1.26.7; extra == "py313-plus"
Dynamic: author
Dynamic: classifier
Dynamic: description
Dynamic: description-content-type
Dynamic: home-page
Dynamic: license
Dynamic: provides-extra
Dynamic: requires-dist
Dynamic: requires-python
Dynamic: summary

# BERNN-MSMS

Batch-effect-aware representation learning and classification with an
estimator-style `fit(...); predict(...)` interface.

## Install

```bash
pip install bernn
```

## Basic usage

```python
from bernn import TrainAEClassifierHoldout
from bernn.config.training_config import TrainingConfig

config = TrainingConfig(
    n_epochs=100,
    optimize_hyperparams=False,
    device="cpu",
)
model = TrainAEClassifierHoldout(config=config, log_metrics=False)

# Inductive fit: training data only.
model.fit(X_train, y_train, groups_train=batch_train)
y_pred = model.predict(X_new, groups_test=batch_new)
```

Read the complete [BERNN usage guide](TUTORIAL.md) for input shapes, inductive
and fully transductive fitting, prediction, and hyperparameter guidance.

Important runtime contract:

- `groups_train` is mandatory.
- If `X_valid` or `X_test` is supplied, its matching batch vector is mandatory.
- BERNN hyperparameters are dataset-dependent. The examples are interface
  demonstrations, not universal performance-optimal defaults.

## Important parameters

Focus on these first:

- optimize_hyperparams: enable/disable Ax optimization.
- n_trials: number of optimization trials.
- fixed_hyperparams: force values and remove them from search.
- n_repeats: number of holdout repeats.
- n_layers, layer1: classifier depth and width seed.
- dloss: domain loss mode.
- warmup, n_epochs: core training schedule.
- device: cpu/cuda target.
- scaler, bs: preprocessing and batch size.
- num_workers: PyTorch DataLoader subprocesses; defaults to 0.

## Documentation

- [Usage tutorial](TUTORIAL.md)
