本篇文章只讲可用的hadoop实现算法的源码,可直接拷贝用于工程中!

算法实现主要包括了四个类

ParticleModel类:质点特征类,任何业务都必须提取出事务的数字特征才可用程序来处理,例如,一个苹果的质点特征可以是:颜色、重量、口感

在定义质点特征类时可以这样来做,double x,y,z分别代表颜色、重量、口感,例如0代表青色,1代表红色,这样就把具体的事务抽像成了程序可以识别与处理的普通bean

ClusterCenter类:对熟悉k-means算法的人来说,这个类非常容易理解,就是最终要聚类成几个簇

KmeanMRStep1类:这是一个map-reduce类,随机生成k个聚类中心,其中Map函数就是读取需要聚类的原始数据,Reduce过程根据随机原则从原始数据中抽取K个数据作为原始簇中心

KmeanMRStep2类:迭代计算聚类结果,map计算相似度,reduce更新簇中心

KmeansDriver类:这是Map-reduce的驱动函数,所有的工程都是类似的固定的编程风格和方法

/**
 * 质点特征模型,可根据自己的业务需求进行修改
 *
 * @author jianting.zhao
 */
public class ParticleModel {

    //特征x
    public double x;

    //特征y
    public double y;

    public double getX() {
        return x;
    }

    public void setX(double x) {
        this.x = x;
    }

    public double getY() {
        return y;
    }

    public void setY(double y) {
        this.y = y;
    }
}
/**
 * 定义簇中心
 *
 * @author jianting.zhao
 */
public class ClusterCenter {

    //簇的编号
    public int K;

    public ParticleModel particleModel;


    public ParticleModel getParticleModel() {
        return particleModel;
    }

    public void setParticleModel(ParticleModel particleModel) {
        this.particleModel = particleModel;
    }

    public int getK() {
        return K;
    }

    public void setK(int K) {
        this.K = K;
    }


}
import java.io.IOException;
import java.util.ArrayList;
import java.util.HashSet;
import java.util.List;
import java.util.Random;
import java.util.Set;

import org.apache.hadoop.conf.Configuration;
import org.apache.hadoop.fs.FileSystem;
import org.apache.hadoop.fs.Path;
import org.apache.hadoop.io.Text;
import org.apache.hadoop.mapreduce.Job;
import org.apache.hadoop.mapreduce.Mapper;
import org.apache.hadoop.mapreduce.Reducer;
import org.apache.hadoop.mapreduce.lib.input.FileInputFormat;
import org.apache.hadoop.mapreduce.lib.output.FileOutputFormat;
import org.apache.hadoop.util.GenericOptionsParser;

/*
 * kmeans聚类第一步:随机生成k个聚类中心
 */
public class KmeanMRStep1 {

    public static void main(String[] args) throws IOException, ClassNotFoundException, InterruptedException {
        Configuration conf = new Configuration();
        String[] otherArgs = new GenericOptionsParser(conf, args).getRemainingArgs();
        if (otherArgs.length != 3) {
            System.err.println("Usage: Data Deduplication <in> <out> <cluster num>");
            System.exit(2);
        }
        conf.set("ClusterNum", otherArgs[2]);
        FileSystem fs = FileSystem.get(conf);
        Path centerPath = new Path(otherArgs[1]);
        fs.deleteOnExit(centerPath);

        Job job = new Job(conf, "KmeanMRStep1");
        job.setJarByClass(KmeanMRStep1.class);
        job.setMapperClass(KmeanMRStep1Map.class);
        job.setReducerClass(KmeanMRStep1Reduce.class);

        //设置输出类型
        job.setOutputKeyClass(Text.class);
        job.setOutputValueClass(Text.class);
        job.setNumReduceTasks(1);

        //设置输入和输出目录
        FileInputFormat.addInputPath(job, new Path(otherArgs[0]));
        FileOutputFormat.setOutputPath(job, new Path(otherArgs[1]));
        job.waitForCompletion(true);
    }

    public static class KmeanMRStep1Map extends Mapper<Object, Text, Text, Text> {
        @Override
        protected void map(Object key, Text value, Context context)
                throws IOException, InterruptedException {
            if (value.toString() != null) {
                context.write(new Text("kmeans"), value);
            }
        }
    }

    public static class KmeanMRStep1Reduce extends Reducer<Text, Text, Text, Text> {
        int clusterNum;
        private static Set<Integer> indexSet = new HashSet<Integer>();

        @Override
        protected void setup(Context context)
                throws IOException, InterruptedException {
            Configuration conf = context.getConfiguration();
            clusterNum = Integer.parseInt(conf.get("ClusterNum"));
        }

        @Override
        protected void reduce(Text key, Iterable<Text> value, Context context)
                throws IOException, InterruptedException {
            List<String> dataList = new ArrayList<String>();
            for (Text val : value) {
                dataList.add(val.toString());
            }

            for (int k = 1; k <= clusterNum; k++) {
                int index = getIndex(dataList.size());
                String point = dataList.get(index);
                if (point != null) {
                    StringBuffer sb = new StringBuffer();
                    sb.append("cluster" + k).append(":").append(point);
                    context.write(new Text(sb.toString()), new Text());
                }
            }
        }

        /*
         * 生成不重复的随机数
         */
        public static int getIndex(int size) {

            int res;
            Random random = new Random();
            while (true) {
                int index = random.nextInt(size);
                if (!indexSet.contains(index)) {
                    res = index;
                    indexSet.add(index);
                    break;
                }
            }

            return res;
        }
    }

}


import java.io.BufferedReader;
import java.io.IOException;
import java.io.InputStreamReader;
import java.util.ArrayList;
import java.util.Collections;
import java.util.Comparator;
import java.util.HashMap;
import java.util.List;
import java.util.Map.Entry;

import org.apache.hadoop.conf.Configuration;
import org.apache.hadoop.filecache.DistributedCache;
import org.apache.hadoop.fs.FileStatus;
import org.apache.hadoop.fs.FileSystem;
import org.apache.hadoop.fs.Path;
import org.apache.hadoop.io.Text;
import org.apache.hadoop.mapreduce.Counter;
import org.apache.hadoop.mapreduce.Job;
import org.apache.hadoop.mapreduce.Mapper;
import org.apache.hadoop.mapreduce.Reducer;
import org.apache.hadoop.mapreduce.lib.input.FileInputFormat;
import org.apache.hadoop.mapreduce.lib.output.FileOutputFormat;
import org.apache.hadoop.util.GenericOptionsParser;

/**
 * Step2:迭代计算聚类结果,map计算相似度,reduce更新簇中心
 *
 * @author jianting.zhao
 */

public class KmeanMRStep2 {

    /*
     * main函数中加载簇中心信息
     */
    public static void main(String[] args) throws IOException, ClassNotFoundException, InterruptedException {
        Configuration conf = new Configuration();
        String[] otherArgs = new GenericOptionsParser(conf, args).getRemainingArgs();
        if (otherArgs.length != 3) {
            System.err.println("Usage: Data Deduplication <in> <in> <out>");
            System.exit(2);
        }

        //加蒌簇中心
        DistributedCache.createSymlink(conf);
        FileSystem fs = FileSystem.get(conf);
        Path clusterCenter = new Path(otherArgs[0]);
        FileStatus[] user_stat = fs.listStatus(clusterCenter);
        for (FileStatus f : user_stat) {
            if (f.getPath().getName().indexOf("_SUCCESS") == -1) {
                DistributedCache.addCacheFile(f.getPath().toUri(), conf);
            }
        }

        if (fs.exists(new Path(otherArgs[2]))) {
            fs.delete(new Path(otherArgs[2]));
        }

        Job job = new Job(conf, "KmeanMRStep2");
        job.setJarByClass(KmeanMRStep2.class);
        job.setMapperClass(KmeanMRStep2Map.class);
        job.setReducerClass(KmeanMRStep2Reduce.class);

        //设置输出类型
        job.setOutputKeyClass(Text.class);
        job.setOutputValueClass(Text.class);
        job.setNumReduceTasks(1);

        //设置输入和输出目录
        FileInputFormat.addInputPath(job, new Path(otherArgs[1]));
        FileOutputFormat.setOutputPath(job, new Path(otherArgs[2]));
        job.waitForCompletion(true);

    }

    public static class KmeanMRStep2Map extends Mapper<Object, Text, Text, Text> {
        List<ParticleModel> clusterCenter = new ArrayList<ParticleModel>();

        /*
         * 加载簇中心
         */
        protected void setup(Context context)
                throws IOException, InterruptedException {
            Configuration conf = context.getConfiguration();
            Path[] file = DistributedCache.getLocalCacheFiles(conf);
            FileSystem fs = FileSystem.getLocal(conf);
            String line = null;
            for (Path path : file) {
                BufferedReader reader = new BufferedReader(new InputStreamReader(fs.open(path)));
                while ((line = reader.readLine()) != null) {
                    String[] tmp = line.split("\t");

                    String[] array = tmp[0].split(":")[1].split(",");
                    if (array.length != 2) {
                        continue;
                    }

                    ParticleModel pm = new ParticleModel();
                    pm.setX(Double.parseDouble(array[0]));
                    pm.setY(Double.parseDouble(array[1]));
                    clusterCenter.add(pm);
                }
            }
        }

        /*
         * 计算相似度
         */
        @Override
        protected void map(Object key, Text value, Context context)
                throws IOException, InterruptedException {
            String[] line = value.toString().split(",");
            if (line.length == 2) {
                ParticleModel pmSample = new ParticleModel();
                pmSample.setX(Double.parseDouble(line[0]));
                pmSample.setY(Double.parseDouble(line[1]));

                String nearlyCluster = getNearlyCluster(pmSample);

                String[] temp = nearlyCluster.split(":");
                if (temp.length == 2) {
                    context.write(new Text("cluster" + temp[0]), new Text(temp[1]));
                }

            }
        }

        private String getNearlyCluster(ParticleModel pmSample) {

            HashMap<String, Double> nearlyMap = new HashMap<String, Double>();

            for (int k = 0; k < clusterCenter.size(); k++) {
                StringBuffer sb = new StringBuffer();
                double sim = getSimilarity(pmSample, clusterCenter.get(k));
                sb.append(k + 1).append(":").append(pmSample.x).append(",").append(pmSample.y);
                nearlyMap.put(sb.toString(), sim);
            }

            //进行降序
            List<Entry<String, Double>> list_cos = new ArrayList<Entry<String, Double>>(
                    nearlyMap.entrySet());
            Collections.sort(list_cos,
                    new Comparator<Entry<String, Double>>() {
                        // 升序排序
                        public int compare(Entry<String, Double> o1,
                                           Entry<String, Double> o2) {
                            return o1.getValue().compareTo(o2.getValue());
                        }
                    });

            //获取第一个元素,也即最近的一个点
            String nearlyCluster = list_cos.get(0).getKey();

            return nearlyCluster;
        }

        private double getSimilarity(ParticleModel pmSample, ParticleModel pmCluster) {
            double x2 = (pmSample.x - pmCluster.x) * (pmSample.x - pmCluster.x);
            double y2 = (pmSample.y - pmCluster.y) * (pmSample.y - pmCluster.y);

            return Math.sqrt(x2 + y2);
        }


    }

    /*
     * Reduce更新簇中心
     */
    public static class KmeanMRStep2Reduce extends Reducer<Text, Text, Text, Text> {

        Counter counter = null;

        @Override
        protected void reduce(Text key, Iterable<Text> value, Context context)
                throws IOException, InterruptedException {
            //获取簇编号
            String clusterNum = key.toString();

            StringBuilder sbPoint = new StringBuilder();
            List<ParticleModel> clusterInfo = new ArrayList<ParticleModel>();
            for (Text val : value) {
                String[] temp = val.toString().split(",");
                if (temp.length != 2) {
                    continue;
                }
                ParticleModel pm = new ParticleModel();
                pm.setX(Double.parseDouble(temp[0]));
                pm.setY(Double.parseDouble(temp[1]));

                clusterInfo.add(pm);
                sbPoint.append(pm.x).append(",").append(pm.y).append("#");
            }

            ParticleModel newClusterCenter = getNewClusterCenter(clusterInfo);
            StringBuilder sb = new StringBuilder();
            if (newClusterCenter != null) {
                sb.append(clusterNum).append(":").append(newClusterCenter.getX()).append(",").append(newClusterCenter.getY());
                context.write(new Text(sb.toString()), new Text(sbPoint.toString()));
            }

        }

        /*
         * 更新簇中心
         */
        public ParticleModel getNewClusterCenter(List<ParticleModel> clusterInfo) {

            int sumX = 0;
            int sumY = 0;
            for (ParticleModel pm : clusterInfo) {
                sumX += pm.getX();
                sumY += pm.getY();
            }

            ParticleModel pm = new ParticleModel();
            pm.setX(sumX / clusterInfo.size());
            pm.setY(sumY / clusterInfo.size());
            return pm;
        }
    }
}


import org.slf4j.Logger;
import org.slf4j.LoggerFactory;

/**
 * KMeans运行主程序
 *
 * @author jianting.zhao
 */

public class KmeansDriver {
    private static final Logger LOGGER = LoggerFactory.getLogger(KmeansDriver.class);

    //迭代次数
    private static int iteratorNum = 10;

    //样本数据输入路径
    private static String dataInput = "/user/hadoop/tianjungang/kmeans/data/";

    //上一次迭代结果路径,初始为第一次随机初始簇中心路径
    private static String lastInput = "/user/hadoop/tianjungang/kmeans/result/iterator0";

    public static void main(String[] args) {

        //Step1:随机初始化聚类中心
        try {
            LOGGER.info("************************* Run KmeanMRStep1 ******************************");
            String[] parmArgs1 = new String[3];

            //样本数据输入路径
            parmArgs1[0] = dataInput;

            //初始簇中心输出路径
            parmArgs1[1] = lastInput;

            //初始簇的个数
            parmArgs1[2] = "3";

            KmeanMRStep1.main(parmArgs1);

        } catch (Exception e) {
            e.printStackTrace();
        }

        //Step2:迭代计算聚类结果
        //可用org.apache.hadoop.mapreduce.Counter进行优化 在reduce判断簇中心是否变化,退出迭代。
        //如果新的簇中心心跟老的簇中心是一样的,那么相应的计数器加1

        try {
            LOGGER.info("************************* Run KmeanMRStep2 ******************************");

            int k = 1;
            while (k <= iteratorNum) {
                LOGGER.info("************************* 开始第 " + k + " 次迭代计算 ************************");

                String[] parmArgs2 = new String[3];

                parmArgs2[1] = dataInput;
                parmArgs2[0] = lastInput;

                String outPut = "/user/hadoop/tianjungang/kmeans/result/" + "iterator" + k;
                parmArgs2[2] = outPut;

                KmeanMRStep2.main(parmArgs2);

                lastInput = outPut;

                k++;
            }

        } catch (Exception e) {
            e.printStackTrace();
        }

    }

}



Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐