Bias–Variance decomposition toolkit for regression (MSE) and classification (0–1 loss)
Project description
Bias–Variance Decomposition Toolkit
A lightweight Python toolkit for estimating bias and variance components of machine learning models.
Supports both regression (via Mean Squared Error) and classification (via 0–1 loss).
Works seamlessly with PyTorch models and scikit-learn style data workflows.
✨ Features
- 📊 Bias–variance decomposition for:
- Regression using Mean Squared Error (MSE)
- Classification using 0–1 loss
- 🔄 Bootstrap resampling for reliable estimates
- 🧠 Works with PyTorch models (custom architectures supported)
- ⏱️ Early stopping with patience-based validation
- ⚡ Device support: CPU and CUDA (GPU)
📦 Installation
pip install biasvariance-toolkit
🚀 Quick Start
1. Define your PyTorch model
import torch.nn as nn
class SimpleNet(nn.Module):
def __init__(self, input_dim, hidden_dim, output_dim):
super().__init__()
self.fc1 = nn.Linear(input_dim, hidden_dim)
self.relu = nn.ReLU()
self.fc2 = nn.Linear(hidden_dim, output_dim)
def forward(self, x):
return self.fc2(self.relu(self.fc1(x)))
2. Bias–Variance Decomposition for Regression
from biasvariance_toolkit import estimate_bias_variance_mse
import torch.nn as nn
bias_sq, variance, total_error, bias_plus_var, avg_train_loss, test_loss = estimate_bias_variance_mse(
model_class=SimpleNet,
model_kwargs={"input_dim": 10, "hidden_dim": 32, "output_dim": 1},
X_train=X_train,
y_train=y_train,
X_test=X_test,
y_test=y_test,
loss_fn=nn.MSELoss(),
num_models=10,
max_epochs=100,
lr=0.001,
device="cpu"
)
3. Bias–Variance Decomposition for Classification
from biasvariance_toolkit import estimate_bias_variance_0_1
import torch.nn as nn
avg_bias, avg_var, expected_loss, empirical_loss, avg_train_loss, test_loss = estimate_bias_variance_0_1(
model_class=SimpleNet,
model_kwargs={"input_dim": 20, "hidden_dim": 64, "output_dim": 3},
X_train=X_train,
y_train=y_train,
X_test=X_test,
y_test=y_test,
loss_fn=nn.CrossEntropyLoss(),
num_models=10,
max_epochs=100,
lr=0.001,
device="cpu"
)
📘 API Overview
estimate_bias_variance_mse
- Task: Regression (MSE)
- Returns:
bias_sq– Squared biasvariance– Variancetotal_error– Average test error (MSE)bias_plus_variance– Bias² + Varianceavg_train_loss– Average training loss across modelstest_loss– Loss of the ensemble’s mean prediction
estimate_bias_variance_0_1
- Task: Classification (0–1 Loss)
- Returns:
avg_bias– Average biasavg_var– Average varianceexpected_loss– Expected 0–1 lossempirical_loss– Empirical 0–1 lossavg_train_loss– Average training loss across modelstest_loss– Loss of the ensemble’s mean prediction
🛠 Requirements
- Python ≥ 3.8
- PyTorch ≥ 1.9
- NumPy
- SciPy
- scikit-learn
📊 Example Use Cases
- Analyzing model stability under resampling
- Comparing architectures (e.g., shallow vs deep networks)
- Studying underfitting vs overfitting tradeoffs
- Teaching / demonstrating bias–variance decomposition concepts
📜 License
MIT License © 2025
Project details
Release history Release notifications | RSS feed
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file biasvariance_toolkit-1.0.0.tar.gz.
File metadata
- Download URL: biasvariance_toolkit-1.0.0.tar.gz
- Upload date:
- Size: 6.3 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.10.18
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
bc150618ec43bbaa2b459a2713710f52971e60d1c303be796f53e856e45b209b
|
|
| MD5 |
76241e7e0fdf5ef68351fc3a8f734e18
|
|
| BLAKE2b-256 |
fae2151a371fb92f8be3958265c3177695a53c85dd86d32844757975d9ae764e
|
File details
Details for the file biasvariance_toolkit-1.0.0-py3-none-any.whl.
File metadata
- Download URL: biasvariance_toolkit-1.0.0-py3-none-any.whl
- Upload date:
- Size: 6.5 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/6.2.0 CPython/3.10.18
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
be05ad167152b3c71aadb57fb30ddf9cf6ceee3cfbd4b5d5e30cadcda0c2e0e5
|
|
| MD5 |
551e9f3eeddc130db110eb5f6be428f3
|
|
| BLAKE2b-256 |
be4de4a3a436b39cc0fad0bc76963ed6556110fa0d8f5a9c3ca8238eee9d35f6
|