defdistance(data, dots): rs = [] for d in data: rs.append([]) for dot in dots: rs[-1].append(((d[0]-dot[0])**2+(d[1]-dot[1])**2)**0.5) return rs
3、分组。根据计算得到的各个点距离聚类中心的距离,可以判断这个点是属于哪一个聚类的。
1 2
defclassify(dis): return [i.index(min(i)) for i in dis]
返回一个分类索引数组。这个时候也可以用plt将此时的分类打印出来,和最开始的正确聚类对比一下:
upload successful
upload successful
考虑一种特殊情况:只分出两个聚类。这时候将无法继续计算下去,应当重新随机生成中心。
4、移动聚类中心。根据随机得到的初始聚类,重新计算每个聚类的质心作为其新的聚类中心。
1 2 3 4 5
defcalCenters(data, clazz): clz = [[] for i inrange(3)] for index,dot inenumerate(data): clz[clazz[index]].append(dot) return [sum(x)/len(x) for x in clz]