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

基于Python的K-Means聚类:寻求结果可视化实现方法

Alright, nice work putting together your own K-Means implementation! To visualize your clustering results clearly, matplotlib is the perfect tool—super straightforward and widely used for this kind of task. Let’s walk through adding visualization to your code step by step.

1. First, Import Matplotlib

You’ll need to add this import at the top of your script to access plotting functionality:

import matplotlib.pyplot as plt

2. Create a Visualization Function

Add this function to your code—it’ll take your final centroids (mu) and clusters, then plot everything neatly with distinct colors and clear markers:

def plot_clusters(mu, clusters):
    # Generate a unique color for each cluster using a built-in color map
    colors = plt.cm.tab10(np.linspace(0, 1, len(clusters)))
    
    # Plot each cluster's data points
    for cluster_idx, (color, points) in enumerate(zip(colors, clusters.values())):
        points = np.array(points)
        plt.scatter(points[:, 0], points[:, 1], color=color, label=f"Cluster {cluster_idx+1}", alpha=0.7)
    
    # Plot centroids with a distinct marker to stand out
    mu_array = np.array(mu)
    plt.scatter(mu_array[:, 0], mu_array[:, 1], color='red', marker='*', s=200, label="Centroids")
    
    # Add labels, title, and legend for readability
    plt.xlabel("X-axis")
    plt.ylabel("Y-axis")
    plt.title("K-Means Clustering Results")
    plt.legend()
    plt.grid(True, alpha=0.3)
    plt.show()

A quick breakdown of what this does:

  • Uses the tab10 color map to automatically assign unique, easy-to-distinguish colors to each cluster
  • Adds subtle transparency (alpha=0.7) to data points so overlapping points are easier to see
  • Marks centroids with a large red star (*) so you can immediately spot where each cluster’s center lies
  • Includes a grid, labels, and legend to make the plot intuitive to read

3. Full Working Code with Visualization

Here’s your complete code with the visualization added, plus an example of how to run it (I also fixed a tiny numpy array compatibility issue in find_centers):

import numpy as np
import random
import matplotlib.pyplot as plt

def cluster_points(X, mu):
    clusters = {}
    for x in X:
        bestmukey = min([(i[0], np.linalg.norm(x - mu[i[0]])) \
                        for i in enumerate(mu)], key=lambda t: t[1])[0]
        try:
            clusters[bestmukey].append(x)
        except KeyError:
            clusters[bestmukey] = [x]
    return clusters

def reevaluate_centers(mu, clusters):
    newmu = []
    keys = sorted(clusters.keys())
    for k in keys:
        newmu.append(np.mean(clusters[k], axis=0))
    return newmu

def find_centers(x, k):
    # Convert numpy array to list for random.sample compatibility
    oldmu = random.sample(x.tolist(), k)
    mu = random.sample(x.tolist(), k)
    while not (set([tuple(a) for a in mu]) == set([tuple(a) for a in oldmu])):
        oldmu = mu
        # Assign all points in X to clusters
        clusters = cluster_points(x, mu)
        # Reevaluate centers
        mu = reevaluate_centers(oldmu, clusters)
    return (mu, clusters)

def init_board(N):
    X = np.array([(random.uniform(-1, 1), random.uniform(-1, 1)) for i in range(N)])
    return X

def init_board_gauss(N, k):
    n = float(N)/k
    X = []
    for i in range(k):
        c = (random.uniform(-1, 1), random.uniform(-1, 1))
        s = random.uniform(0.05,0.5)
        x = []
        while len(x) < n:
            a, b = np.array([np.random.normal(c[0], s), np.random.normal(c[1], s)])
            # Continue drawing points from the distribution in the range [-1,1]
            if abs(a) < 1 and abs(b) < 1:
                x.append([a,b])
        X.extend(x)
    X = np.array(X)[:N]
    return X

def plot_clusters(mu, clusters):
    # Generate a unique color for each cluster using a built-in color map
    colors = plt.cm.tab10(np.linspace(0, 1, len(clusters)))
    
    # Plot each cluster's data points
    for cluster_idx, (color, points) in enumerate(zip(colors, clusters.values())):
        points = np.array(points)
        plt.scatter(points[:, 0], points[:, 1], color=color, label=f"Cluster {cluster_idx+1}", alpha=0.7)
    
    # Plot centroids with a distinct marker to stand out
    mu_array = np.array(mu)
    plt.scatter(mu_array[:, 0], mu_array[:, 1], color='red', marker='*', s=200, label="Centroids")
    
    # Add labels, title, and legend for readability
    plt.xlabel("X-axis")
    plt.ylabel("Y-axis")
    plt.title("K-Means Clustering Results")
    plt.legend()
    plt.grid(True, alpha=0.3)
    plt.show()

# Example usage to test the code
if __name__ == "__main__":
    # Generate 300 points grouped into 3 Gaussian clusters
    X = init_board_gauss(300, 3)
    # Run your K-Means implementation
    centroids, clusters = find_centers(X, 3)
    # Visualize the final clusters and centroids
    plot_clusters(centroids, clusters)

What You’ll See

When you run the script, a plot will pop up showing:

  • 3 distinct clusters, each with a unique color
  • Red star markers at the exact center of each cluster (your converged centroids)
  • Clean labels and a legend to make the output easy to interpret

That’s all you need to visualize your K-Means results effectively!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.13 08:04:16