k-means聚类算法hadoop实现源码
·
本篇文章只讲可用的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();
}
}
}
更多推荐
所有评论(0)