You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

关于实现类sklearn.cluster.MiniBatchKMeans的大规模数据集迷你批K均值的技术问询

Mini Batch K-Means Implementation (Sklearn-Aligned) for Large Datasets

Got it, let's walk through how to implement a Mini Batch K-Means that matches the behavior of sklearn.cluster.MiniBatchKMeans for large datasets. I’ll break down the core steps, key implementation details, and even share a simplified code snippet to get you started:

Core Algorithm Steps

These are the two foundational steps you need to nail to replicate sklearn's behavior:

  • Step 1: Mini-Batch Sampling & Sample Assignment

    • First, randomly draw b samples from your large dataset to form a mini-batch (sklearn uses a default batch size of 100, but you can tune this based on your memory constraints).
    • For every sample in the mini-batch, calculate its distance to each centroid (Euclidean distance is the standard choice here, same as sklearn) and assign the sample to its nearest centroid.
    • Pro tip: Use vectorized operations (like NumPy broadcasting) instead of looping through each sample individually—this will drastically speed up computation, especially on large batches.
  • Step 2: Centroid Update via Moving Average

    • This is where Mini Batch K-Means differs from standard K-Means. Instead of recalculating centroids as the full mean of all assigned samples, you update them incrementally using a moving average.
    • For each sample in the mini-batch:
      1. Let n be the total number of samples that have been assigned to the target centroid so far (including the current sample).
      2. The centroid update can be written in two equivalent ways (both used in sklearn under the hood):
        # Direct moving average calculation
        new_centroid = ((n - 1) * old_centroid + sample) / n
        
        # Learning rate-based approach (same math, different framing)
        learning_rate = 1 / n
        new_centroid = old_centroid + learning_rate * (sample - old_centroid)
        
      • This approach keeps memory usage low because you don’t need to store all past samples assigned to each centroid—you just track the current centroid value and the count of samples used to build it.

Critical Implementation Considerations

To make your implementation robust and aligned with sklearn, don’t overlook these details:

  • Centroid Initialization: Use k-means++ initialization (sklearn's default) instead of random initialization. This picks initial centroids that are well-separated, helping the algorithm converge faster and avoid poor local minima.
  • Convergence Check: Track the maximum shift in centroid positions between iterations. If this shift falls below a small threshold (e.g., 1e-4), you can stop early. Alternatively, run a fixed number of iterations (sklearn uses 100 by default).
  • Empty Cluster Handling: Sometimes a centroid might get no samples assigned in a mini-batch. Sklearn fixes this by reinitializing the centroid to a random sample from the dataset—add this logic to avoid stagnant, unused centroids.
  • Partial Fit Support: For streaming data (where you can’t load the entire dataset into memory), implement a partial_fit method (like sklearn does) that lets you feed mini-batches one at a time without restarting the algorithm.

Simplified Code Snippet

Here’s a stripped-down implementation that covers the core logic:

import numpy as np

class MiniBatchKMeans:
    def __init__(self, n_clusters=8, batch_size=100, max_iter=100, tol=1e-4):
        self.n_clusters = n_clusters
        self.batch_size = batch_size
        self.max_iter = max_iter
        self.tol = tol
        self.centroids = None
        self.counts = None  # Tracks number of samples per centroid

    def fit(self, X):
        # Initialize centroids using k-means++
        self.centroids = self._kmeans_plus_plus(X)
        # Start counts at 1 to avoid division by zero in early updates
        self.counts = np.ones(self.n_clusters, dtype=np.int64)

        for _ in range(self.max_iter):
            prev_centroids = self.centroids.copy()
            
            # Step 1: Sample a mini-batch from the dataset
            batch_indices = np.random.choice(len(X), self.batch_size, replace=False)
            batch = X[batch_indices]

            # Step 2: Assign each sample to the nearest centroid
            # Vectorized distance calculation using broadcasting
            distances = np.sqrt(((batch - self.centroids[:, np.newaxis])**2).sum(axis=2))
            labels = np.argmin(distances, axis=0)

            # Step 3: Update centroids with moving average
            for cluster_idx in range(self.n_clusters):
                # Get all samples in the batch assigned to this cluster
                cluster_samples = batch[labels == cluster_idx]
                for sample in cluster_samples:
                    self.counts[cluster_idx] += 1
                    learning_rate = 1.0 / self.counts[cluster_idx]
                    self.centroids[cluster_idx] += learning_rate * (sample - self.centroids[cluster_idx])

            # Check for convergence
            centroid_shift = np.max(np.linalg.norm(self.centroids - prev_centroids, axis=1))
            if centroid_shift < self.tol:
                break

    def _kmeans_plus_plus(self, X):
        # Simplified k-means++ initialization logic
        centroids = [X[np.random.choice(len(X))]]
        for _ in range(self.n_clusters - 1):
            # Calculate distance from each sample to the nearest existing centroid
            distances = np.min(np.sqrt(((X - np.array(centroids)[:, np.newaxis])**2).sum(axis=2)), axis=0)
            # Probability proportional to distance squared (sklearn uses squared distance)
            probabilities = distances ** 2
            probabilities /= probabilities.sum()
            # Pick next centroid with weighted probability
            next_centroid_idx = np.random.choice(len(X), p=probabilities)
            centroids.append(X[next_centroid_idx])
        return np.array(centroids)

Note that sklearn’s full implementation includes more optimizations (like faster distance calculations using tree-based structures) and additional features (like verbose logging, random state control), but this snippet captures the core behavior you’re targeting.

内容的提问来源于stack exchange,提问作者Logan

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.19 10:37:49