西瓜书9.4 k均值算法实现

一.需求:

​ 编程实现k均值算法,在西瓜数据集4.0上进行实验比较

二.伪代码:

image-20201208231608873

三.代码实现:

​ 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);

    }
posted on 2020-12-08 23:53  计网好难  阅读(529)  评论(0)    收藏  举报