Skip to content

Preserve dtypes and propagate CUDA devices; 10k-sample comparison - #14

Merged
ms-kumar merged 2 commits into
mainfrom
feat/dtype-cuda
Sep 25, 2026
Merged

ms-kumar merged 2 commits into
mainfrom
feat/dtype-cuda

Conversation

@ms-kumar

Copy link
Copy Markdown
Owner

Answers: float32/float64 now preserved end to end (ints promote to float32, labels untouched), CUDA inputs stay on CUDA through fit/predict/transform (was: device-mismatch RuntimeError in nearly every estimator). check_array gains dtype=None preservation. Benchmark scaled to 10k/task with honest numbers incl. the knn-predict gap (brute force ~5s vs ball tree ~0.3s). New tests/common/test_dtype_device.py (15 tests, CUDA cases gated). Local: 335 passed + 1 skipped, pre-commit clean, sphinx 0 warnings. Note: no GPU in CI, so CUDA paths are guarded-not-executed there; MPS called out as CPU-only in docs.

@ms-kumar ms-kumar added the enhancement New feature or request label Sep 25, 2026
@ms-kumar ms-kumar self-assigned this Sep 25, 2026
@ms-kumar
ms-kumar merged commit 46135f2 into main Sep 25, 2026
9 of 10 checks passed
@ms-kumar
ms-kumar deleted the feat/dtype-cuda branch September 25, 2026 17:19
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant