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
8 changes: 8 additions & 0 deletions RELEASES.md
Original file line number Diff line number Diff line change
Expand Up @@ -5,14 +5,22 @@
### Added

### Changed
- Changed `torml.model_selection.GridSearchCV` docs: added `examples/model_selection/plot_grid_search.py`.

### Deprecated

### Removed

### Fixed
- Fixed Bandit scan scope to `torml/` (test asserts no longer reported as B101).
- Fixed `check_estimator` name check (class name instead of full repr) and params round-trip comparison.
- Fixed `get_tags` mixin detection (was crashing on every estimator).
- Fixed `RandomState` attribute bookkeeping (reset/pickle/seed handling).
- Fixed `MultiOutputRegressor`/`MultiOutputClassifier` nested `estimator__param` support.

### Security
- Bumped `setuptools` 81.0.0 to 84.0.0 (CVE-2026-59890).
- Pinned GitHub Actions to commit SHAs; added Dependabot config and `SECURITY.md`.

### Contributors

Expand Down
29 changes: 29 additions & 0 deletions examples/model_selection/plot_grid_search.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
"""Tune k with GridSearchCV and report the best params.

Run with ``python examples/model_selection/plot_grid_search.py``.
"""

from __future__ import annotations

import torch

from torml.model_selection import GridSearchCV
from torml.neighbors import KNeighborsClassifier


def main() -> None:
"""Run the grid search example."""
torch.manual_seed(0)
x0 = torch.randn(30, 2) + torch.tensor([-2.0, 0.0])
x1 = torch.randn(30, 2) + torch.tensor([2.0, 0.0])
X = torch.cat([x0, x1])
y = torch.cat([torch.zeros(30), torch.ones(30)]).long()

gs = GridSearchCV(KNeighborsClassifier(), {"n_neighbors": [1, 3, 7]}, cv=3)
gs.fit(X, y)
print(f"best params: {gs.best_params_}")
print(f"best score: {gs.best_score_:.4f}")


if __name__ == "__main__":
main()
Loading