Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 8 additions & 6 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -112,15 +112,17 @@ rerun with `uv run --with scikit-learn python benchmarks/compare_sklearn.py`):

| task | metric | torml | scikit-learn |
| --- | --- | --- | --- |
| linreg | R² | 0.999596 | 0.999596 |
| kmeans | inertia (lower better) | 14235.3 | 14271.8 |
| knn | accuracy | 0.9240 | 0.9240 |
| linreg | R² | 0.999656 | 0.999656 |
| kmeans | inertia (lower better) | 72645.4 | 72611.7 |
| knn | accuracy | 0.9476 | 0.9476 |
| tree | accuracy | 1.0000 | 1.0000 |
| scaler | max \|diff\| | 4.77e-07 | reference |

Correctness matches everywhere; timings are in the same class except tree
fitting (~2x slower — Cython vs Python loops). torml's real edges are
torch-native differentiable I/O, one runtime dependency, and exact
Correctness matches everywhere (10k samples per task); timings are in the
same class except tree fitting (~1.5x slower — Cython vs Python loops) and
10k kNN prediction, where sklearn's ball tree (~0.3s) beats torml's exact
brute force (~5s). torml's real edges are torch-native differentiable I/O
(including CUDA + float64 propagation), one runtime dependency, and exact
statistics. Full honest write-up (including where sklearn wins):
[user guide comparison](doc/user_guide/comparison.rst).

Expand Down
2 changes: 2 additions & 0 deletions RELEASES.md
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,8 @@
## [Unreleased]

### Added
- Added dtype/device propagation: `float32`/`float64` preserved end to end, CUDA inputs stay on CUDA; covered by `tests/common/test_dtype_device.py` (CUDA cases gated on availability).
- Added `doc/user_guide/dtypes_devices.rst` and scaled the sklearn comparison to 10k samples per task.

### Changed
- Changed `torml.model_selection.GridSearchCV` docs: added `examples/model_selection/plot_grid_search.py`.
Expand Down
12 changes: 6 additions & 6 deletions benchmarks/compare_sklearn.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,8 +36,8 @@ def bench_linreg():
from torml.metrics import r2_score

torch.manual_seed(0)
features = torch.randn(5000, 20)
target = features @ torch.randn(20) + 0.1 * torch.randn(5000)
features = torch.randn(10000, 20)
target = features @ torch.randn(20) + 0.1 * torch.randn(10000)
np_x, np_y = features.numpy(), target.numpy()
tm, mine = timed(lambda: ToLin().fit(features, target))
ts, theirs = timed(lambda: SkLin().fit(np_x, np_y))
Expand All @@ -60,7 +60,7 @@ def bench_kmeans():
from torml.cluster import KMeans as ToKM

torch.manual_seed(1)
data = torch.randn(2000, 10)
data = torch.randn(10000, 10)
tm, mine = timed(lambda: ToKM(n_clusters=10, random_state=0, n_init=3).fit(data))
ts, theirs = timed(
lambda: SkKM(n_clusters=10, random_state=0, n_init=3).fit(data.numpy())
Expand All @@ -81,15 +81,15 @@ def bench_knn():
from torml.neighbors import KNeighborsClassifier as ToKNN

torch.manual_seed(2)
features = torch.randn(1000, 10)
features = torch.randn(10000, 10)
labels = (features[:, 0] > 0).long()
np_x, np_y = features.numpy(), labels.numpy()
tm, mine = timed(lambda: ToKNN(5).fit(features, labels))
ts, theirs = timed(lambda: SkKNN(5).fit(np_x, np_y))
print(ROW.format("knn fit", "seconds (best of 3)", f"{tm:.4f}", f"{ts:.4f}"))
tm, my_pred = timed(lambda: mine.predict(features))
ts, their_pred = timed(lambda: theirs.predict(np_x))
print(ROW.format("knn predict-1000", "seconds", f"{tm:.4f}", f"{ts:.4f}"))
print(ROW.format("knn predict-10k", "seconds", f"{tm:.4f}", f"{ts:.4f}"))
print(
ROW.format(
"knn",
Expand Down Expand Up @@ -129,7 +129,7 @@ def bench_scaler():
from torml.preprocessing import StandardScaler as ToSS

torch.manual_seed(0)
data = torch.randn(5000, 20)
data = torch.randn(10000, 20)
tm, mine = timed(lambda: ToSS().fit_transform(data))
ts, theirs = timed(lambda: SkSS().fit_transform(data.numpy()))
print(ROW.format("scaler", "seconds (best of 3)", f"{tm:.4f}", f"{ts:.4f}"))
Expand Down
1 change: 1 addition & 0 deletions doc/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ Start with :doc:`quickstart`, learn the patterns in
user_guide/unsupervised
user_guide/preprocessing
user_guide/model_selection
user_guide/dtypes_devices
user_guide/comparison
api

Expand Down
27 changes: 16 additions & 11 deletions doc/user_guide/comparison.rst
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ Rerun it any time:

uv run --with scikit-learn python benchmarks/compare_sklearn.py

Method: same data, same seeds, CPU, best of 3 timings
Method: same data, same seeds, CPU, best of 3 timings, 10k samples per task
(Apple M-series, 2026-09-25, torch 2.14.0, scikit-learn 1.9.1).

Results
Expand All @@ -18,24 +18,26 @@ Results
=================== ======================= ================== ==================
task metric torml scikit-learn
=================== ======================= ================== ==================
linreg fit seconds (best of 3) 0.0006 0.0009
linreg R² 0.999596 0.999596
kmeans fit seconds (best of 3) 0.0292 0.0278
kmeans inertia (lower better) 14235.3 14271.8
knn fit seconds (best of 3) 0.0000 0.0003
knn predict-1000 seconds 0.0055 0.0057
knn accuracy 0.9240 0.9240
tree fit seconds (best of 3) 0.0010 0.0006
linreg fit seconds (best of 3) 0.0012 0.0014
linreg R² 0.999656 0.999656
kmeans fit seconds (best of 3) 0.1631 0.1032
kmeans inertia (lower better) 72645.4 72611.7
knn fit seconds (best of 3) 0.0003 0.0018
knn predict-10k seconds 5.1368 0.2828
knn accuracy 0.9476 0.9476
tree fit seconds (best of 3) 0.0083 0.0057
tree accuracy 1.0000 1.0000
scaler seconds (best of 3) 0.0005 0.0005
scaler seconds (best of 3) 0.0007 0.0008
scaler max |diff| vs sklearn 4.77e-07 reference
torch-native backward() thru predict True n/a (numpy out)
=================== ======================= ================== ==================

Reading the table honestly: correctness matches everywhere (identical R²,
accuracy, and scaler outputs; kmeans inertia differs only by random
initialization). Timings are in the same class except tree fitting, where
scikit-learn's Cython is ~2x faster than torml's Python CART loops.
scikit-learn's Cython is ~1.5x faster than torml's Python CART loops, and
kNN prediction at 10k, where scikit-learn's ball tree (~0.3s) beats torml's
exact brute-force pairwise distances (~5s).

Where torml is genuinely better
-------------------------------
Expand All @@ -48,6 +50,9 @@ Where torml is genuinely better
p-values are computed exactly, not via approximations.
- **Readable from-scratch code** with strict input validation on every
estimator; useful for teaching and auditing.
- **Dtypes and devices propagate.** ``float32``/``float64`` inputs keep
their precision end to end (integers promote to ``float32``), and CUDA
inputs stay on CUDA — see :doc:`dtypes_devices`.

Where scikit-learn still wins
-----------------------------
Expand Down
40 changes: 40 additions & 0 deletions doc/user_guide/dtypes_devices.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
Dtypes and devices
====================

torml preserves your input dtype and device instead of silently converting
to ``float32``/CPU.

Dtypes
------

- ``float32`` and ``float64`` tensors keep their precision end to end:
statistics (``mean_``, ``coef_``, ``components_``) and predictions come
back in the input dtype.
- Integer feature tensors are promoted to ``float32`` (means and variances
need fractions); integer *label* tensors stay integer.
- Lists default to ``float32``. Pass ``dtype=`` explicitly to
``check_array`` to force a conversion.

.. code-block:: python

X64 = torch.randn(20, 3, dtype=torch.float64)
model = LinearRegression().fit(X64, y64)
model.coef_.dtype # torch.float64

Devices
-------

- CUDA inputs stay on CUDA through ``fit``/``predict``/``transform`` — no
host round-trips, no device-mismatch errors.
- CPU remains the tested path: the suite runs on CPU, and CUDA coverage is
a ``requires CUDA``-gated round-trip test (see
``tests/common/test_dtype_device.py``).

.. code-block:: python

X = torch.randn(40, 3, device="cuda")
StandardScaler().fit_transform(X).device.type # 'cuda'

Caveat: on Apple Silicon, some ``torch.linalg`` ops used internally
(``lstsq``, ``svd``, ``cholesky``) have limited MPS support — CPU is the
supported backend there.
136 changes: 136 additions & 0 deletions tests/common/test_dtype_device.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,136 @@
"""Tests for dtype preservation and device propagation."""

from __future__ import annotations

import pytest
import torch

from torml.cluster import KMeans
from torml.decomposition import PCA
from torml.linear_model import LinearRegression
from torml.neighbors import KNeighborsClassifier
from torml.preprocessing import StandardScaler
from torml.tree import DecisionTreeClassifier
from torml.utils import check_array

cuda_only = pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA")


class TestCheckArrayDtype:
"""Tests for check_array dtype=None preservation."""

def test_preserves_float64(self):
"""Test float64 tensors keep their dtype."""
out = check_array(torch.randn(4, 2, dtype=torch.float64))
assert out.dtype == torch.float64

def test_preserves_float32(self):
"""Test float32 tensors keep their dtype."""
out = check_array(torch.randn(4, 2, dtype=torch.float32))
assert out.dtype == torch.float32

def test_promotes_int(self):
"""Test integer tensors promote to float32."""
out = check_array(torch.ones(4, 2, dtype=torch.int64))
assert out.dtype == torch.float32

def test_list_defaults_float32(self):
"""Test lists default to float32."""
out = check_array([[1.0, 2.0], [3.0, 4.0]])
assert out.dtype == torch.float32

def test_explicit_dtype_still_forces(self):
"""Test explicit dtype overrides preservation."""
out = check_array(torch.randn(4, 2, dtype=torch.float64), dtype=torch.float32)
assert out.dtype == torch.float32


class TestDtypePreservation:
"""Tests that estimators compute in the input dtype."""

def test_scaler_float64(self):
"""Test StandardScaler learns float64 statistics."""
X = torch.randn(20, 3, dtype=torch.float64)
sc = StandardScaler().fit(X)
assert sc.mean_.dtype == torch.float64
assert sc.scale_.dtype == torch.float64
assert sc.transform(X).dtype == torch.float64
assert sc.inverse_transform(sc.transform(X)).dtype == torch.float64

def test_linear_regression_float64(self):
"""Test LinearRegression learns float64 coefficients."""
torch.manual_seed(0)
X = torch.randn(30, 2, dtype=torch.float64)
y = X[:, 0] * 2 - X[:, 1]
model = LinearRegression().fit(X, y)
assert model.coef_.dtype == torch.float64
assert model.predict(X).dtype == torch.float64

def test_numeric_agreement(self):
"""Test float32 and float64 runs agree closely."""
torch.manual_seed(1)
X32 = torch.randn(30, 2)
y32 = X32[:, 0] * 2 - X32[:, 1]
X64, y64 = X32.double(), y32.double()
pred32 = LinearRegression().fit(X32, y32).predict(X32).double()
pred64 = LinearRegression().fit(X64, y64).predict(X64)
torch.testing.assert_close(pred32, pred64, rtol=1e-4, atol=1e-4)

def test_int_features_accepted(self):
"""Test integer features are promoted, not rejected."""
X = torch.randint(0, 5, (20, 2))
y = (X[:, 0] > 2).long()
pred = KNeighborsClassifier(3).fit(X, y).predict(X)
assert pred.shape == (20,)

@pytest.mark.parametrize("dtype", [torch.float32, torch.float64])
def test_predict_accepts_both(self, dtype):
"""Test classifiers predict on both float dtypes."""
torch.manual_seed(2)
X = (torch.randn(20, 2)).to(dtype)
y = (X[:, 0] > 0).long()
for clf in (KNeighborsClassifier(3), DecisionTreeClassifier()):
assert clf.fit(X, y).predict(X).shape == (20,)

def test_pca_float64(self):
"""Test PCA keeps float64 through transform."""
X = torch.randn(20, 4, dtype=torch.float64)
pca = PCA(n_components=2).fit(X)
assert pca.components_.dtype == torch.float64
assert pca.transform(X).dtype == torch.float64

def test_kmeans_float64(self):
"""Test KMeans centers match input dtype."""
X = torch.randn(20, 2, dtype=torch.float64)
km = KMeans(n_clusters=2, random_state=0, n_init=2).fit(X)
assert km.cluster_centers_.dtype == torch.float64


class TestDevice:
"""Tests for device propagation (CPU always, CUDA when present)."""

def test_cpu_round_trip(self):
"""Test CPU tensors stay on CPU."""
X = torch.randn(20, 2)
y = (X[:, 0] > 0).long()
assert StandardScaler().fit_transform(X).device.type == "cpu"
assert LinearRegression().fit(X, y.float()).predict(X).device.type == "cpu"
assert (
KMeans(n_clusters=2, random_state=0, n_init=2).fit(X).labels_.device.type
== "cpu"
)

@cuda_only
def test_cuda_round_trip(self):
"""Test CUDA tensors stay on CUDA end to end."""
torch.manual_seed(0)
X = torch.randn(40, 3, device="cuda")
y = (X[:, 0] > 0).long()
yr = X[:, 0] * 2 - X[:, 1]
assert StandardScaler().fit_transform(X).device.type == "cuda"
assert LinearRegression().fit(X, yr).predict(X).device.type == "cuda"
assert KNeighborsClassifier(3).fit(X, y).predict(X).device.type == "cuda"
assert (
KMeans(n_clusters=2, random_state=0, n_init=2).fit(X).labels_.device.type
== "cuda"
)
2 changes: 0 additions & 2 deletions torml/base/_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,6 @@ def _check_X(

return check_array(
X,
dtype=torch.float32,
ensure_2d=True,
allow_nd=False,
copy=True,
Expand All @@ -204,7 +203,6 @@ def _check_y(

return check_array(
y,
dtype=torch.float32,
ensure_2d=False,
allow_nd=False,
copy=True,
Expand Down
12 changes: 7 additions & 5 deletions torml/cluster/_dbscan.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,7 +78,7 @@ def fit(self, X, y=None):
)
if int(self.min_samples) < 1:
raise ValueError(f"min_samples must be >= 1, got {self.min_samples}.")
Xt = check_array(X, ensure_2d=True, dtype=torch.float32)
Xt = check_array(X, ensure_2d=True)
n_samples = int(Xt.shape[0])
self.n_features_in_ = int(Xt.shape[1])

Expand All @@ -88,10 +88,12 @@ def fit(self, X, y=None):
for i in range(n_samples)
]
is_core = torch.tensor(
[len(nb) >= int(self.min_samples) for nb in neighborhoods], dtype=torch.bool
[len(nb) >= int(self.min_samples) for nb in neighborhoods],
dtype=torch.bool,
device=Xt.device,
)
labels = torch.full((n_samples,), -1, dtype=torch.long)
visited = torch.zeros(n_samples, dtype=torch.bool)
labels = torch.full((n_samples,), -1, dtype=torch.long, device=Xt.device)
visited = torch.zeros(n_samples, dtype=torch.bool, device=Xt.device)
cluster_id = 0
for i in range(n_samples):
if bool(visited[i]) or not bool(is_core[i]):
Expand Down Expand Up @@ -119,5 +121,5 @@ def fit_predict(self, X, y=None):
labels : torch.Tensor of shape (n_samples,)
Cluster labels (-1 for noise).
"""
check_array(X, ensure_2d=True, dtype=torch.float32)
check_array(X, ensure_2d=True)
return self.fit(X, y).labels_
8 changes: 4 additions & 4 deletions torml/cluster/_kmeans.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,7 @@ def fit(self, X, y=None):
"""
self._validate_hyperparams()
generator = check_random_state(self.random_state)
Xt = check_array(X, ensure_2d=True, dtype=torch.float32)
Xt = check_array(X, ensure_2d=True)
n_samples = int(Xt.shape[0])
if n_samples < int(self.n_clusters):
raise ValueError(
Expand All @@ -107,9 +107,9 @@ def fit(self, X, y=None):

best_inertia: float | None = None
for _ in range(int(self.n_init)):
perm = torch.randperm(n_samples, generator=generator)
perm = torch.randperm(n_samples, generator=generator, device=Xt.device)
centers = Xt[perm[: int(self.n_clusters)]].clone()
assign = torch.zeros(n_samples, dtype=torch.long)
assign = torch.zeros(n_samples, dtype=torch.long, device=Xt.device)
n_iter = 0
for n_iter in range(1, int(self.max_iter) + 1):
dist = torch.cdist(Xt, centers, p=2)
Expand Down Expand Up @@ -148,7 +148,7 @@ def predict(self, X):
Closest-center indices.
"""
check_is_fitted(self, attributes=["cluster_centers_"])
Xt = check_array(X, ensure_2d=True, dtype=torch.float32)
Xt = check_array(X, ensure_2d=True)
if int(Xt.shape[1]) != int(self.n_features_in_):
raise ValueError(
f"X has {int(Xt.shape[1])} features, but KMeans was fitted "
Expand Down
Loading
Loading