perf(model): Replace Sklearn with optimized PyTorch KMeans in CFA - #3760
Conversation
There was a problem hiding this comment.
🟡 Changes recommended
KMeans currently lacks required n_clusters validation and uses CPU-created centroid indices (risking runtime failure / device sync), and touched-file copyright headers need updating to include 2026.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Pull request overview
This PR replaces the CFA model’s use of scikit-learn KMeans (CPU-bound) with anomalib’s PyTorch KMeans implementation and improves KMeans centroid updates using vectorized PyTorch operations to reduce Python-loop overhead and avoid unnecessary device transfers.
Changes:
- Swap CFA centroid initialization from
sklearn.cluster.KMeanstoanomalib.models.components.cluster.kmeans.KMeans. - Optimize KMeans centroid recomputation using
torch.bincount+index_add_instead of per-cluster boolean masking loops.
File summaries
| File | Description |
|---|---|
| src/anomalib/models/image/cfa/torch_model.py | Switch CFA memory-bank clustering from sklearn KMeans (CPU) to anomalib’s PyTorch KMeans. |
| src/anomalib/models/components/cluster/kmeans.py | Vectorize centroid updates and adjust centroid initialization behavior. |
Review details
- Files reviewed: 2/2 changed files
- Comments generated: 3
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
ashwinvaidya17
left a comment
There was a problem hiding this comment.
Thanks for updating this. I have left a very minor comment. Can you also have a look at the copilot comments as well. They are quite minor as well
|
Also, can you update the PR title |
There was a problem hiding this comment.
🟡 Changes recommended
The new KMeans implementation currently lacks convergence-based early exit (critical with max_iter=3000) and CFA’s clustered memory bank dtype now differs from the prior explicit float32 behavior.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
Suppressed comments (1)
Previously missed (1) — in code that hasn't changed since the last review.
src/anomalib/models/components/cluster/kmeans.py:101
- This vectorized centroid-update logic is performance-critical and can materially change clustering behavior (especially around empty clusters and convergence). There are unit tests for
GaussianMixturethat rely on KMeans indirectly, but there are no direct tests forKMeansitself; adding a small deterministic unit test (synthetic blobs, fixed seed, asserts on shapes + stable clustering) would help prevent regressions in this optimized path.
- Files reviewed: 2/2 changed files
- Comments generated: 2
- Review effort level: Lite
There was a problem hiding this comment.
🟡 Changes recommended
The PyTorch KMeans implementation will fail on accelerator devices due to CPU-created index tensors, and the new/updated clustering behavior lacks direct unit test coverage.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
Suppressed comments (1)
src/anomalib/models/components/cluster/kmeans.py:85
centroid_indicesis created on the default device (CPU). Ifinputsis on CUDA/XPU,inputs[centroid_indices]will error because the index tensor must be on the same device as the indexed tensor. This becomes likely now that CFA calls KMeans on-device to avoid CPU transfers.
# Initialize centroids randomly from the data points
centroid_indices = torch.randint(0, batch_size, (self.n_clusters,))
self.cluster_centers_ = inputs[centroid_indices].clone()
- Files reviewed: 2/2 changed files
- Comments generated: 1
- Review effort level: Lite
There was a problem hiding this comment.
🟡 Changes recommended
The new vectorized KMeans update path can produce numerically incorrect centroids under mixed precision (float16 counts) and the CFA migration should keep the prior float32 behavior and add regression coverage.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
Suppressed comments (1)
src/anomalib/models/image/cfa/torch_model.py:248
cluster_centersare now kept in whatever dtypeself.memory_bankhas; previously the sklearn path forcedfloat32. Under mixed precision/autocast this can make the memory bankfloat16, which can reduce numerical stability incompute_distanceand change outputs compared to the prior implementation. Consider running KMeans in float32 and storing the resulting centers as float32 (matching the old behavior).
_, cluster_centers = k_means.fit(self.memory_bank)
self.memory_bank = cluster_centers.detach().to(device)
- Files reviewed: 2/2 changed files
- Comments generated: 2
- Review effort level: Lite
| counts = torch.bincount(self.labels_, minlength=self.n_clusters).to(inputs.dtype) | ||
| new_centers = torch.zeros_like(self.cluster_centers_) | ||
| new_centers.index_add_(0, self.labels_, inputs) | ||
|
|
||
| valid_mask = counts > 0 | ||
| self.cluster_centers_[valid_mask] = new_centers[valid_mask] / counts[valid_mask].unsqueeze(1) |
| if self.gamma_c > 1: | ||
| # TODO(samet-akcay): Create PyTorch KMeans class. | ||
| # CVS-122673 | ||
| k_means = KMeans( |
Added unit tests for KMeans to assert that fit returns expected shapes and dtypes, and that predict is consistent with fitted centers on a synthetic dataset as requested in the PR comments.
|
Hi @ashwinvaidya17 I have added unit tests could you pls let me know if this pr needs any changes or if it is good to merge |
There was a problem hiding this comment.
🟡 Changes recommended
The updated PyTorch KMeans has correctness issues for GPU usage and can return labels inconsistent with the final cluster centers unless adjusted.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
Suppressed comments (1)
src/anomalib/models/components/cluster/kmeans.py:85
centroid_indicesis created on the default device (CPU). Ifinputsis on CUDA (common for CFA/GMM), indexinginputs[centroid_indices]can raise a device-mismatch error. Create the indices oninputs.deviceto ensure KMeans works on GPU.
# Initialize centroids randomly from the data points
centroid_indices = torch.randint(0, batch_size, (self.n_clusters,))
self.cluster_centers_ = inputs[centroid_indices].clone()
- Files reviewed: 3/3 changed files
- Comments generated: 2
- Review effort level: Lite
|
@andersendsa I know continuous copilot comments can be annoying but can you check if they make sense and update the PR if necessary? |
Hi @ashwinvaidya17 sorry for the delay I have implemented the changes required and resolved them pls let me know if this pr needs anymore changes or if it is good to merge |
ashwinvaidya17
left a comment
There was a problem hiding this comment.
Thanks for the efforts. Just one minor comment. I think we should add early exit once the labels stop changing. Otherwise we might keep it running till max_iter
34058e2 to
e3a8dec
Compare
Hi @ashwinvaidya17 I have done the changes pls let me know if pr needs anymore changes or if it is good to merge |
|
@andersendsa can you fix the prek issues? I think we are good to merge after that. |
Hi @ashwinvaidya17 I have fixed the prek issues and This pr is good to merge |
📝 Description
Addresses TODO in
src/anomalib/models/image/cfa/torch_model.pyto use a PyTorch KMeans class instead of Sklearn's version which requires moving data to the CPU. - Usesanomalib.models.components.cluster.kmeans.KMeansand optimizes it using PyTorch's vectorizedbincountandindex_add_operations instead of python loops with boolean indexing, providing significant performance speedups.✨ Changes
Select what type of change your PR is:
✅ Checklist
Before you submit your pull request, please make sure you have completed the following steps:
For more information about code review checklists, see the Code Review Checklist.