package com.bjsxt.kmeans
import org.apache.spark.mllib.clustering.KMeans
import org.apache.spark.{ SparkContext, SparkConf }
import org.apache.spark.mllib.feature.HashingTF
import org.wltea.analyzer.lucene.IKAnalyzer
import java.io.StringReader
import org.apache.lucene.analysis.tokenattributes.CharTermAttribute
import scala.collection.mutable.ListBuffer
import scala.collection.mutable.StringBuilder
import scala.collection.mutable.ArrayBuffer
import org.apache.spark.mllib.feature.IDF
import org.apache.spark.sql.SQLContext
import org.apache.spark.mllib.feature.IDFModel
import org.apache.spark.rdd.RDD
/**
* Created by zfg on 2017/8/14.
*/
object KMeans11 {
def main(args: Array[String]) {
val conf = new SparkConf().setAppName("KMeans1").setMaster("local[*]")
val sc = new SparkContext(conf)
val rdd = sc.textFile("e:/original.txt")
/**
* wordRDD 是一个KV格式的RDD
* K:微博ID
* V:微博内容分词后的结果 ArrayBuffer
*/
var wordRDD: RDD[(String, ArrayBuffer[String])] = rdd.mapPartitions(iterator => {
val list = new ListBuffer[(String, ArrayBuffer[String])]
while (iterator.hasNext) {
val analyzer = new IKAnalyzer(true)
val line = iterator.next()
val textArr = line.split("\t")
val text = textArr(1)
val id = textArr(0)
val ts = analyzer.tokenStream("", text)
val term = ts.getAttribute(classOf[CharTermAttribute])
ts.reset()
val arr = new ArrayBuffer[String]
while (ts.incrementToken()) {
arr.+=(term.toString())
}
list.append((id, arr))
analyzer.close()
}
list.iterator
})
wordRDD = wordRDD.cache()
//1000:只是计算每篇微博中1000个单词的词频 最大似然估计思想
val hashingTF = new HashingTF(1000)
/**
* tfRDD
* K:微博ID
* V:Vector(tf,tf,tf.....)
*/
val tfRDD = wordRDD.map(x => {
(x._1, hashingTF.transform(x._2))
})
/**
* idf对象中就会保存每一个单词的IDF值
*/
val idf: IDFModel = new IDF().fit(tfRDD.map(_._2))
/**
* K:weibo ID
* V:每一个单词的TF-IDF值
* tfIdfs这个RDD就是训练模型的训练集
* 不
*/
val tfIdfs = tfRDD.mapValues(idf.transform(_))
val kcluster = 5
val kmeans = new KMeans()
kmeans.setK(kcluster)
//使用的是kemans++算法来训练模型
kmeans.setInitializationMode("k-means||")
kmeans.setMaxIterations(100)
/**
* 每一类的中心点坐标
*/
val kmeansModel = kmeans.run(tfIdfs.map(_._2))
/**
* kmeansModel.save(sc, "d:/model001")
* kmeansModel:5个中心点的 坐标
*/
println(kmeansModel.clusterCenters)
/**
* 模型预测
*/
val modelBroadcast = sc.broadcast(kmeansModel)
/**
* predicetionRDD KV格式的RDD
* K:微博ID
* V:分类号
*/
val predicetionRDD = tfIdfs.mapValues(sample => {
val model = modelBroadcast.value
model.predict(sample)
})
// predicetionRDD.saveAsTextFile("d:/resultttt")
/**
* 总结预测结果
* tfIdfs2wordsRDD:kv格式的RDD
* K:微博ID
* V:二元组(Vector(tfidf1,tfidf2....),ArrayBuffer(word,word,word....))
*/
val tfIdfs2wordsRDD = tfIdfs.join(wordRDD)
/**
* result:KV
* K:微博ID
* V:(类别号,(Vector(tfidf1,tfidf2....),ArrayBuffer(word,word,word....)))
*/
val result = predicetionRDD.join(tfIdfs2wordsRDD)
result
.filter(x => x._2._1 == 1)
.flatMap(line => {
val tfIdfV = line._2._2._1
val words = line._2._2._2
val tfIdfA = tfIdfV.toArray
val wordL = new ListBuffer[String]()
val tfIdfL = new ListBuffer[Double]()
var index = 0
for(i <- 0 until tfIdfA.length;if tfIdfV(i) != 0){
wordL.+=(words(index))
tfIdfL.+=(tfIdfA(index))
index += 1
}
println(wordL.length + "===" + tfIdfL.length)
val list = new ListBuffer[(Double, String)]
for (i <- 0 until wordL.length) {
list.append((tfIdfV(i), words(i)))
}
list
})
.sortBy(x => x._1, false)
.map(_._2)
.distinct()
.take(30).foreach(println)
/* val str1 = new StringBuilder
val str2 = new StringBuilder
val str3 = new StringBuilder
val str4 = new StringBuilder
val str5 = new StringBuilder
result
.filter(x=> x._2._1 == 0)
.flatMap(x=>x._2._2._1.toArray)
.sortBy(x=>x,false)
.distinct
.take(20)
.foreach { x => {
str1.append("," + tfIdf2Words.get(x).get)
} }
result
.filter(x=> x._2._1 == 1)
.flatMap(x=>x._2._2._1.toArray)
.sortBy(x=>x,false)
.distinct
.take(20)
.foreach { x => {
str2.append("," + tfIdf2Words.get(x).get)
} }
result
.filter(x=> x._2._1 == 2)
.flatMap(x=>x._2._2._1.toArray)
.sortBy(x=>x,false)
.distinct
.take(20)
.foreach { x => {
str3.append("," + tfIdf2Words.get(x).get)
} }
result
.filter(x=> x._2._1 == 3)
.flatMap(x=>x._2._2._1.toArray)
.sortBy(x=>x,false)
.distinct
.take(20)
.foreach { x => {
str4.append("," + tfIdf2Words.get(x).get)
} }
result
.filter(x=> x._2._1 == 4)
.flatMap(x=>x._2._2._1.toArray)
.sortBy(x=>x,false)
.distinct
.take(20)
.foreach { x => {
str5.append("," + tfIdf2Words.get(x).get)
} }
println(str1)
println(str2)
println(str3)
println(str4)
println(str5)*/
while(true){}
// sc.stop()
}
}