how do i plot kmeans and how do i print canvas (im new to kmeans)

Viewed 12
heres the code

import numpy as np
import matplotlib.pyplot as plt
from scipy.spatial.distance import cdist
#gonna change k into a input
np.random.seed(14)
n = 20
p = 3
k = 3 
x = np.random.random((n,p))
plt.scatter(x[:,0], x[:,1])
centers = x[np.random.choice(n, k, replace=False)]
((x[0]-centers[0])**2).sum()**0.5
((x-centers[0])**2).sum(axis=1)
((x.reshape(n,1,p)-centers.reshape(1,k,p))**2).sum(axis=2)**0.5
distances = np.zeros((n,k))
for i in range(k):
    distances[:,i] = ((x-centers[i])**2).sum(axis=1)**0.5
distances
distances = cdist(x, centers)
closest = np.argmin(distances, axis=1)
x[closest == 0].mean(axis=0)
for i in range (k):
    centers[i, :] = x[closest == i].mean(axis=0)
centers
np.random.seed(4160659)
centers = x[np.random.choice(n, k, replace=False)]
closest = np.zeros(n).astype(int)
while True:
    old_closest = closest.copy()
    print(closest)
    distances = cdist(x, centers)
    closest = np.argmin(distances, axis=1)

    for i in range (k):
        centers[i, :] = x[closest == i].mean(axis=0)
    
    if all(closest == old_closest):
        break
plt.scatter(x[:,0],x[:,1],c=closest)
plt.xlabel('age')
plt.ylabel('income ($)')

the purpose of it is to make a means chart also most of it is to find clusters center also, it clusters perfectly warning:try to be as basic talk as possible cuz kmeans new to me

0 Answers
Related