西瓜书9.4 k均值算法实现
一.需求:
编程实现k均值算法,在西瓜数据集4.0上进行实验比较
二.伪代码:

三.代码实现:
1. main方法:
输入为n*2的二维数组,调用KMeans类的train方法训练,输出为每次迭代后的散点图
public class Main {
public static void main(String[] args) throws InterruptedException {
Object[][]objects={{0.697, 0.460},{0.774, 0.376},{0.634, 0.264},{0.608, 0.318},{0.556, 0.215},{0.403, 0.237},{0.481, 0.149},{0.437, 0.211},{0.666, 0.091},{0.243, 0.267},{0.245, 0.057},{0.343, 0.099},{0.639, 0.161},{0.657, 0.198},{0.360, 0.370},{0.593, 0.042},{0.719, 0.103},{0.359, 0.188},{0.339, 0.241},{0.282, 0.257},{0.748, 0.232},{0.714, 0.346},{0.483, 0.312},{0.478, 0.437},{0.525, 0.369},{0.751, 0.489},{0.532, 0.472},{0.473, 0.376},{0.725, 0.445},{0.446, 0.459}};
KMeans kMeans = new KMeans(3, objects);
kMeans.train();
}
}
这里输入改为double二维数组会更好
2.KMeans类实现
package com.fly;
import org.jfree.chart.ChartFactory;
import org.jfree.chart.ChartFrame;
import org.jfree.chart.JFreeChart;
import org.jfree.chart.axis.NumberAxis;
import org.jfree.chart.axis.ValueAxis;
import org.jfree.chart.plot.PlotOrientation;
import org.jfree.chart.plot.XYPlot;
import org.jfree.chart.renderer.xy.XYLineAndShapeRenderer;
import org.jfree.data.xy.DefaultXYDataset;
import javax.jws.Oneway;
import java.awt.*;
import java.util.*;
import java.util.List;
public class KMeans {
//簇的数量
private int k;
//训练集
private Object[][]objects;
//均值向量
private Object[][]meanVectors;
//默认最大循环次数
private int loop=10000;
//属性数
private int valueNum;
//控制训练集循环
boolean flag=true;
public KMeans(int k,Object[][]objects){
this.k=k;
meanVectors=new Object[k][objects[0].length];
this.objects=objects;
this.valueNum=objects[0].length;
Random random = new Random();
//生成k个随机数
HashSet<Integer> integers = new HashSet<Integer>();
while (integers.size()!=k){
integers.add(random.nextInt(objects.length));
}
// integers.add(5);
// integers.add(23);
// integers.add(11);
Iterator<Integer> iterator = integers.iterator();
for(int i=0;i<k;i++){
Integer j = iterator.next();
for(int m=0;m<valueNum;m++){
meanVectors[i][m]=objects[j][m];
}
}
//showMeanVectors();
}
public KMeans(int k,Object[][]objects,int loop){
this(k,objects);
this.loop=loop;
}
//训练
public void train() throws InterruptedException {
List<Set<Integer>>clusters=new ArrayList<Set<Integer>>();
for(int i=0;i<k;i++){
clusters.add(new HashSet<Integer>());
}
//迭代次数
int m=0;
while (flag&&loop>m){
//清空原先分好的簇
for(int i=0;i<k;i++){
clusters.get(i).clear();
}
flag=false;
for(int i=0;i<objects.length;i++){
//System.out.println(i+1+"---->"+getNumOfCluster(objects[i]));
int num = getNumOfCluster(objects[i]);
clusters.get(num).add(i);
}
//System.out.println(clusters);
updateMeanVectors(clusters);
//showMeanVectors();
//System.out.println(clusters);
m++;
showScatterChart(clusters,m);
}
//System.out.println(m);
}
//更新均值向量
public void updateMeanVectors(List<Set<Integer>> clusters){
int x=0;
for(Set<Integer>cluster:clusters){
Iterator<Integer> iterator = cluster.iterator();
double[]temp=new double[valueNum];
while (iterator.hasNext()){
Integer index = iterator.next();
Object[]object=objects[index];
for(int i=0;i<valueNum;i++){
temp[i]+=(Double)object[i];
}
}
for(int i=0;i<valueNum;i++){
temp[i]/=cluster.size();
if((Double) meanVectors[x][i]!=temp[i]){
flag=true;
meanVectors[x][i]=temp[i];
}
}
x++;
}
}
//判断应该放在哪个簇
public int getNumOfCluster(Object[] object){
double t=d(object,meanVectors[0]);
int ans=0;
for(int i=1;i<k;i++){
double x=d(object,meanVectors[i]);
if(t>x){
t=x;
ans=i;
}
}
return ans;
}
//获取x1,x2的距离
public double d(Object[]o1, Object[]o2){
int len=o1.length;
double ans=0;
for(int i=0;i<len;i++){
ans+=( (Double) o1[i] -(Double) o2[i]) *((Double) o1[i] -(Double) o2[i]);
}
ans=Math.sqrt(ans);
return ans;
}
//输出均值向量
private void showMeanVectors(){
System.out.println("-----------mean vectors---------");
for(int i=0;i<k;i++){
for(int j=0;j<valueNum;j++){
System.out.print(meanVectors[i][j]+" ");
}
System.out.println();
}
System.out.println("------------------------------------");
}
//输出散点图
public void showScatterChart(List<Set<Integer>>clusters,int m){
//第j个簇
Integer j=0;
DefaultXYDataset xydataset = new DefaultXYDataset();
for(Set<Integer> cluster:clusters){
double[][] data=new double[2][cluster.size()];
Iterator<Integer> iterator = cluster.iterator();
int i=0;
while (iterator.hasNext()){
Integer index = iterator.next();
Object[] object = objects[index];
//System.out.println(object[0]+" -- "+object[1]);
data[0][i]=(Double) object[0];
data[1][i]=(Double)object[1];
i++;
}
xydataset.addSeries(j,data);
j++;
}
JFreeChart chart = ChartFactory.createScatterPlot("第"+m+"次迭代结果","密度","含糖",xydataset, PlotOrientation.VERTICAL, true, false, false);
ChartFrame frame = new ChartFrame("散点图", chart, true);
chart.setBackgroundPaint(Color.white);
chart.setBorderPaint(Color.GREEN);
chart.setBorderStroke(new BasicStroke(1.5f));
XYPlot xyplot = (XYPlot) chart.getPlot();
xyplot.setBackgroundPaint(new Color(255, 253, 246));
ValueAxis vaaxis = xyplot.getDomainAxis();
vaaxis.setAxisLineStroke(new BasicStroke(1.5f));
ValueAxis va = xyplot.getDomainAxis(0);
va.setAxisLineStroke(new BasicStroke(1.5f));
va.setAxisLineStroke(new BasicStroke(1.5f)); // 坐标轴粗细
va.setAxisLinePaint(new Color(215, 215, 215)); // 坐标轴颜色
xyplot.setOutlineStroke(new BasicStroke(1.5f)); // 边框粗细
va.setLabelPaint(new Color(10, 10, 10)); // 坐标轴标题颜色
va.setTickLabelPaint(new Color(102, 102, 102)); // 坐标轴标尺值颜色
ValueAxis axis = xyplot.getRangeAxis();
axis.setAxisLineStroke(new BasicStroke(1.5f));
XYLineAndShapeRenderer xylineandshaperenderer = (XYLineAndShapeRenderer) xyplot
.getRenderer();
xylineandshaperenderer.setSeriesOutlinePaint(0, Color.WHITE);
xylineandshaperenderer.setUseOutlinePaint(true);
NumberAxis numberaxis = (NumberAxis) xyplot.getDomainAxis();
numberaxis.setAutoRangeIncludesZero(false);
numberaxis.setTickMarkInsideLength(2.0F);
numberaxis.setTickMarkOutsideLength(0.0F);
numberaxis.setAxisLineStroke(new BasicStroke(1.5f));
//输出中文为方框的解决办法
Font font=new Font("黑体",Font.BOLD,18);//测试是可以的
chart.getTitle().setFont(font);
axis.setLabelFont(font);
va.setLabelFont(font);
frame.pack();
frame.setVisible(true);
}
}
簇的数量k,训练集objects,最大循环次数loop都需用户设置,通过构造函数设置这些属性的值,并在构造方法中随机选择k个点作为初始均值向量
public KMeans(int k,Object[][]objects){
this.k=k;
meanVectors=new Object[k][objects[0].length];
this.objects=objects;
this.valueNum=objects[0].length;
Random random = new Random();
//生成k个随机数
HashSet<Integer> integers = new HashSet<Integer>();
while (integers.size()!=k){
integers.add(random.nextInt(objects.length));
}
Iterator<Integer> iterator = integers.iterator();
for(int i=0;i<k;i++){
Integer j = iterator.next();
for(int m=0;m<valueNum;m++){
meanVectors[i][m]=objects[j][m];
}
}
}
public KMeans(int k,Object[][]objects,int loop){
this(k,objects);
this.loop=loop;
}
翻译伪代码,实现train方法,每个簇用Set存储,set中放的是每个元组在训练集的索引,所有的簇放在一个List中
public void train() throws InterruptedException {
List<Set<Integer>>clusters=new ArrayList<Set<Integer>>();
for(int i=0;i<k;i++){
clusters.add(new HashSet<Integer>());
}
//迭代次数
int m=0;
while (flag&&loop>m){
//清空原先分好的簇
for(int i=0;i<k;i++){
clusters.get(i).clear();
}
flag=false;
//更新簇
for(int i=0;i<objects.length;i++){
//获取距离最近的簇
int num = getNumOfCluster(objects[i]);
clusters.get(num).add(i);
}
//更新均值向量
updateMeanVectors(clusters);
m++;
//输出迭代结果
showScatterChart(clusters,m);
}
}
定义输出函数,这里采用的是JFreeChart,一个java的画统计图库。参考 https://blog.csdn.net/C_son/article/details/43954885 ,感谢原作者~
maven依赖
<dependency>
<groupId>org.jfree</groupId>
<artifactId>jfreechart</artifactId>
<version>1.0.19</version>
</dependency>
public void showScatterChart(List<Set<Integer>>clusters,int m){
//第j个簇
Integer j=0;
DefaultXYDataset xydataset = new DefaultXYDataset();
for(Set<Integer> cluster:clusters){
double[][] data=new double[2][cluster.size()];
Iterator<Integer> iterator = cluster.iterator();
int i=0;
while (iterator.hasNext()){
Integer index = iterator.next();
Object[] object = objects[index];
//System.out.println(object[0]+" -- "+object[1]);
data[0][i]=(Double) object[0];
data[1][i]=(Double)object[1];
i++;
}
xydataset.addSeries(j,data);
j++;
}
JFreeChart chart = ChartFactory.createScatterPlot("第"+m+"次迭代结果","密度","含糖",xydataset, PlotOrientation.VERTICAL, true, false, false);
ChartFrame frame = new ChartFrame("散点图", chart, true);
chart.setBackgroundPaint(Color.white);
chart.setBorderPaint(Color.GREEN);
chart.setBorderStroke(new BasicStroke(1.5f));
XYPlot xyplot = (XYPlot) chart.getPlot();
xyplot.setBackgroundPaint(new Color(255, 253, 246));
ValueAxis vaaxis = xyplot.getDomainAxis();
vaaxis.setAxisLineStroke(new BasicStroke(1.5f));
ValueAxis va = xyplot.getDomainAxis(0);
va.setAxisLineStroke(new BasicStroke(1.5f));
va.setAxisLineStroke(new BasicStroke(1.5f)); // 坐标轴粗细
va.setAxisLinePaint(new Color(215, 215, 215)); // 坐标轴颜色
xyplot.setOutlineStroke(new BasicStroke(1.5f)); // 边框粗细
va.setLabelPaint(new Color(10, 10, 10)); // 坐标轴标题颜色
va.setTickLabelPaint(new Color(102, 102, 102)); // 坐标轴标尺值颜色
ValueAxis axis = xyplot.getRangeAxis();
axis.setAxisLineStroke(new BasicStroke(1.5f));
XYLineAndShapeRenderer xylineandshaperenderer = (XYLineAndShapeRenderer) xyplot
.getRenderer();
xylineandshaperenderer.setSeriesOutlinePaint(0, Color.WHITE);
xylineandshaperenderer.setUseOutlinePaint(true);
NumberAxis numberaxis = (NumberAxis) xyplot.getDomainAxis();
numberaxis.setAutoRangeIncludesZero(false);
numberaxis.setTickMarkInsideLength(2.0F);
numberaxis.setTickMarkOutsideLength(0.0F);
numberaxis.setAxisLineStroke(new BasicStroke(1.5f));
//输出中文为方框的解决办法
Font font=new Font("黑体",Font.BOLD,18);//测试是可以的
chart.getTitle().setFont(font);
axis.setLabelFont(font);
va.setLabelFont(font);
frame.pack();
frame.setVisible(true);
}
浙公网安备 33010602011771号