Skip to content

Commit 3e34e8e

Browse files
machine_learning: pin numeric output of k_means_clust with doctests (#15285)
1 parent 0fe748a commit 3e34e8e

1 file changed

Lines changed: 34 additions & 1 deletion

File tree

‎machine_learning/k_means_clust.py‎

Lines changed: 34 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -81,6 +81,13 @@ def centroid_pairwise_dist(x, centroids):
8181

8282

8383
def assign_clusters(data, centroids):
84+
"""Assign each data point to the index of its nearest centroid.
85+
86+
>>> data = np.array([[0.0, 0.0], [0.0, 1.0], [10.0, 10.0], [10.0, 11.0]])
87+
>>> centroids = np.array([[0.0, 0.0], [10.0, 10.0]])
88+
>>> assign_clusters(data, centroids).tolist()
89+
[0, 0, 1, 1]
90+
"""
8491
# Compute distances between each data point and the set of centroids:
8592
# Fill in the blank (RHS only)
8693
distances_from_centroids = centroid_pairwise_dist(data, centroids)
@@ -93,6 +100,13 @@ def assign_clusters(data, centroids):
93100

94101

95102
def revise_centroids(data, k, cluster_assignment):
103+
"""Recompute each centroid as the mean of the points assigned to it.
104+
105+
>>> data = np.array([[0.0, 0.0], [0.0, 1.0], [10.0, 10.0], [10.0, 11.0]])
106+
>>> assignment = np.array([0, 0, 1, 1])
107+
>>> revise_centroids(data, 2, assignment).tolist()
108+
[[0.0, 0.5], [10.0, 10.5]]
109+
"""
96110
new_centroids = []
97111
for i in range(k):
98112
# Select all data points that belong to cluster i. Fill in the blank (RHS only)
@@ -106,6 +120,16 @@ def revise_centroids(data, k, cluster_assignment):
106120

107121

108122
def compute_heterogeneity(data, k, centroids, cluster_assignment):
123+
"""Sum of squared distances from each point to its assigned centroid.
124+
125+
This is the objective k-means minimises; lower is a tighter clustering.
126+
127+
>>> data = np.array([[0.0, 0.0], [0.0, 1.0], [10.0, 10.0], [10.0, 11.0]])
128+
>>> centroids = np.array([[0.0, 0.5], [10.0, 10.5]])
129+
>>> assignment = np.array([0, 0, 1, 1])
130+
>>> float(compute_heterogeneity(data, 2, centroids, assignment))
131+
1.0
132+
"""
109133
heterogeneity = 0.0
110134
for i in range(k):
111135
# Select all data points that belong to cluster i. Fill in the blank (RHS only)
@@ -154,7 +178,16 @@ def kmeans(
154178
as function of iterations
155179
if None, do not store the history.
156180
verbose: if True, print how many data points changed their cluster labels in
157-
each iteration"""
181+
each iteration
182+
183+
>>> data = np.array([[0.0, 0.0], [0.0, 1.0], [10.0, 10.0], [10.0, 11.0]])
184+
>>> initial_centroids = np.array([[0.0, 0.0], [10.0, 10.0]])
185+
>>> centroids, assignment = kmeans(data, 2, initial_centroids, maxiter=10)
186+
>>> centroids.tolist()
187+
[[0.0, 0.5], [10.0, 10.5]]
188+
>>> assignment.tolist()
189+
[0, 0, 1, 1]
190+
"""
158191
centroids = initial_centroids[:]
159192
prev_cluster_assignment = None
160193

0 commit comments

Comments
 (0)