Python金融数据挖掘 第11章 第2节 K均值聚类代码
·
1、库
import pandas as pd
import numpy as np
import matplotlib as mpl
import matplotlib.pyplot as plt
2、随机生成聚类中心点
def initCentroids(dataSet,k):
numSamples,dim=dataSet.shape
centroids=np.zeros((k,dim))
for i in range(k):
index=int(np.random.uniform(0,numSamples))
centroids[i,:]=dataSet[index,:]
return centroids
3、欧氏距离
def euclDistance(vector1,vector2):
return np.sqrt(np.sum(np.power(vector2-vector1,2)))
4、K-均值聚类
def kmeans(dataSet,k):
numSamples=dataSet.shape[0]
clusterAssment=np.mat(np.zeros((numSamples,2)))
clusterChanged=True
centroids=initCentroids(dataSet,k)
while clusterChanged:
clusterChanged=False
for i in range(numSamples):
minDist=100000.0
minIndex=0
for j in range(k):
distance=euclDistance(centroids[j,:],dataSet[i,:])
if distance<minDist:
minDist=distance
minIndex=j
if clusterAssment[i,0]!=minIndex:
clusterChanged=True
clusterAssment[i,:]=minIndex,minDist**2
for j in range(k):
pointsInCluster=dataSet[np.nonzero(clusterAssment[:,0].A==j)[0]]
centroids[j,:]=np.mean(pointsInCluster,axis=0)
print('KNN聚类完成!')
return centroids,clusterAssment
5、二维平面显示聚类结果
def showCluster(dataSet,k,centroids,clusterAssment):
fig_2d_clustered=plt.figure()
ax2d_clustered=fig_2d_clustered.add_subplot(111)
numSamples,dim=dataSet.shape
if dim!=2:
print("只能绘制二维图形:")
return 1
mark=['.r','+b','*g','lk','^r','vr','sr','dr','<r','pr']
if k>len(mark):
print("K值过大!")
return 1
for i in range(numSamples):
markIndex=int(clusterAssment[i,0])
ax2d_clustered.plot(dataSet[i,0],dataSet[i,1],mark[markIndex])
for i in range(k):
ax2d_clustered.plot(centroids[i,0],centroids[i,1],mark[i],markersize=20)
fig_2d_clustered.savefig('clusterRes.png',dpi=300,bbox_inches='tight')
fig_2d_clustered.show()
6、结果
#调用以上函数,对读入数据进行聚类
print("step1:读入数据:")
dataSetKNN=[]
fileIn=open('testSet.txt')
for line in fileIn.readlines():
lineArr=line.strip().split(' ')
dataSetKNN.append([float(lineArr[0]),float(lineArr[1])])
dataSetKNNSize=len(dataSetKNN)
dataSetKNN1=np.mat(dataSetKNN)
for i in range(dataSetKNNSize):
plt.plot(dataSetKNN1[i,0],dataSetKNN1[i,1],'b*')
print("原始数据分布:")
plt.savefig('ch12_knn_orig.png',dpi=300,bbox_inches='tight')
plt.show()
#K 取值4,调用K均值算法聚类
print("step2:聚类")
k=4
centroids,clusterAssment=kmeans(dataSetKNN1,k)
print('数据类型:',dataSetKNN1.dtype)
print("step3:结果输出:")
showCluster(dataSetKNN1,k,centroids,clusterAssment)
step1:读入数据:
原始数据分布:
step2:聚类
KNN聚类完成!
数据类型: float64
step3:结果输出:
C:\Users\86186\Desktop\2021-2022(2)课程\2021-2022(2)Python金融数据挖掘\课件\第11章\例11-2.py:94: UserWarning: Matplotlib is currently using module://matplotlib_inline.backend_inline, which is a non-GUI backend, so cannot show the figure.
fig_2d_clustered.show()

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐




所有评论(0)