KMeans11--微博营销案例

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()
  }
}

  

posted @ 2018-07-08 21:19  uuhh  阅读(210)  评论(0)    收藏  举报