Python实现MNIST数据集K-means算法结果不佳,请求排查Bug
Hey there! Let's dig into why your MNIST k-means implementation isn't hitting the mark. Your code snippet gives us a good starting point, so let's break down the most likely bugs and fixes step by step.
Common Issues & Fixes
1. Broken Random Initialization
Your code cuts off at M = np.rando...—first off, that's a typo (should be np.random). But even if you fix the typo, generating centers from a random normal distribution might not align with MNIST's data distribution, leading to poor convergence.
Instead, sample initial centers directly from your MNIST dataset—this ensures your starting points are in the same data space as your images, which helps the algorithm converge to meaningful clusters:
# Replace your broken initialization with this num_clusters = 10 random_indices = np.random.choice(len(I), num_clusters, replace=False) M = I[random_indices]
2. Missing/Incorrect Distance Calculation
K-means lives or dies by accurate distance calculations. If you skipped this step or messed up the array broadcasting, your cluster assignments will be totally off. Here's how to compute Euclidean distances correctly (efficiently, too):
# Calculate distance from each image to every cluster center distances = np.linalg.norm(I[:, np.newaxis] - M, axis=2) # Assign each image to the closest cluster cluster_assignments = np.argmin(distances, axis=1)
3. Ignoring Empty Clusters
It's common for clusters to become empty during iteration (especially with bad initialization). If you try to compute the mean of an empty cluster, you'll get NaN values that break the rest of the algorithm. Fix this by reinitializing empty clusters with a random sample from your data:
for k in range(num_clusters): # Get all images assigned to cluster k cluster_images = I[cluster_assignments == k] # If cluster is empty, reinitialize with a random image if len(cluster_images) == 0: cluster_images = I[np.random.choice(len(I), 1)] # Update the cluster center to the mean of its images M[k] = cluster_images.mean(axis=0)
4. No Convergence Check
If you're only running a fixed, small number of iterations, the algorithm might not have time to settle into stable clusters. Add a convergence check to stop early when centers stop changing significantly:
max_iterations = 100 tolerance = 1e-4 # How much centers can change before we call it quits for iter in range(max_iterations): old_centers = M.copy() # Step 1: Assign clusters distances = np.linalg.norm(I[:, np.newaxis] - M, axis=2) cluster_assignments = np.argmin(distances, axis=1) # Step 2: Update centers for k in range(num_clusters): cluster_images = I[cluster_assignments == k] if len(cluster_images) == 0: cluster_images = I[np.random.choice(len(I), 1)] M[k] = cluster_images.mean(axis=0) # Check if we've converged if np.linalg.norm(M - old_centers) < tolerance: print(f"Converged after {iter+1} iterations!") break
5. Misjudging Results with Wrong Evaluation
K-means cluster labels don't map directly to MNIST's true digits (e.g., cluster 0 might correspond to digit 5, not 0). To get an accurate accuracy score, you need to match cluster labels to true labels using the Hungarian algorithm:
from scipy.optimize import linear_sum_assignment # Build a confusion matrix: rows = true labels, columns = cluster labels confusion_matrix = np.zeros((10, 10), dtype=int) for true_label, cluster_label in zip(L.flatten(), cluster_assignments): confusion_matrix[true_label, cluster_label] += 1 # Use Hungarian algorithm to find the best label matching row_indices, col_indices = linear_sum_assignment(-confusion_matrix) accuracy = confusion_matrix[row_indices, col_indices].sum() / len(I) print(f"Clustering Accuracy: {accuracy:.2f}")
With these fixes, you should see a big jump in your k-means performance on MNIST—you're aiming for around 70-80% accuracy, which is typical for vanilla k-means on this dataset.
内容的提问来源于stack exchange,提问作者Gilfoyle

