基于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
tab10color 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

