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

Python实现MNIST数据集K-means算法结果不佳,请求排查Bug

Troubleshooting Your MNIST K-Means Implementation

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:12:47