轉載請註明出處:http://blog.csdn.net/xiaojimanman/article/details/51064307
http://www.llwjy.com/blogdetail/f74b497c2ad6261b0ea651454b97a390.html
個人部落格站已經上線了,網址 www.llwjy.com ~歡迎各位吐槽~
-------------------------------------------------------------------------------------------------
在開始之前先打一個小小的廣告,自己建立一個QQ群:321903218,點選連結加入群【Lucene案例開發】,主要用於交流如何使用Lucene來建立站內搜尋後台,同時還會不週期性在群內開相關的公開課,感興趣的童鞋可以加入交流。
KNN演算法又叫近鄰演算法,是資料採礦中一種常用的分類演算法,接單的介紹KNN演算法的核心思想就是:尋找與目標最近的K個個體,這些樣本屬於類別最多的那個類別就是目標的類別。比如K為7,那麼我們就從資料中找到和目標最近(或者相似性最高)的7個樣本,加入這7個樣本對應的類別分別為A、B、C、A、A、A、B,那麼目標屬於的分類就是A(因為這7個樣本中屬於A類別的樣本個數最多)。
演算法實現
一、訓練資料格式定義
下面就簡單的介紹下如何用JAVA來實現KNN分類,首先我們需要儲存訓練集(包括屬性以及對應的類別),這裡我們對未知的屬性使用泛型,類別我們使用字串儲存。
/** *@Description: KNN分類模型中一條記錄的儲存格式 */ package com.lulei.datamining.knn.bean; public class KnnValueBean<T>{private T value;//記錄值private String typeId;//分類IDpublic KnnValueBean(T value, String typeId) {this.value = value;this.typeId = typeId;}public T getValue() {return value;}public void setValue(T value) {this.value = value;}public String getTypeId() {return typeId;}public void setTypeId(String typeId) {this.typeId = typeId;}}
二、K個最近鄰類別資料格式定義
在統計得到K個最近鄰中,我們需要記錄前K個樣本的分類以及對應的相似性,我們這裡使用如下資料格式:
/** *@Description: K個最近鄰的類別得分 */ package com.lulei.datamining.knn.bean; public class KnnValueSort {private String typeId;//分類IDprivate double score;//該分類得分public KnnValueSort(String typeId, double score) {this.typeId = typeId;this.score = score;}public String getTypeId() {return typeId;}public void setTypeId(String typeId) {this.typeId = typeId;}public double getScore() {return score;}public void setScore(double score) {this.score = score;}}
三、KNN演算法基本屬性
在KNN演算法中,最重要的一個指標就是K的取值,因此我們在基類中需要設定一個屬性K以及設定一個數組用於儲存已知分類的資料。
private List<KnnValueBean> dataArray;private int K = 3;
四、添加已知分類資料
在使用KNN分類之前,我們需要先向其中添加我們已知分類的資料,我們後面就是使用這些資料來預測未知資料的分類。
/** * @param value * @param typeId * @Author:lulei * @Description: 向模型中添加記錄 */public void addRecord(T value, String typeId) {if (dataArray == null) {dataArray = new ArrayList<KnnValueBean>();}dataArray.add(new KnnValueBean<T>(value, typeId));}
五、兩個樣本之間的相似性(或者距離)
在KNN演算法中,最重要的一個方法就是如何確定兩個樣本之間的相似性(或者距離),由於這裡我們使用的是泛型,並沒有辦法確定兩個對象之間的相似性,一次這裡我們把它設定為抽象方法,讓子類來實現。這裡我們方法定義為相似性,也就是傳回值越大,兩者越相似,之間的距離越短。
/** * @param o1 * @param o2 * @return * @Author:lulei * @Description: o1 o2之間的相似性 */public abstract double similarScore(T o1, T o2);
六、擷取最近的K個樣本的分類
KNN演算法的核心思想就是找到最近的K個近鄰,因此這一步也是整個演算法的核心部分。這裡我們使用數組來儲存相似性最大的K個樣本的分類和相似性,在計算的過程中通過迴圈遍曆所有的樣本,數組儲存截至當前計算點最相似的K個樣本對應的類別和相似性,具體實現如下:
/** * @param value * @return * @Author:lulei * @Description: 擷取距離最近的K個分類 */private KnnValueSort[] getKType(T value) {int k = 0;KnnValueSort[] topK = new KnnValueSort[K];for (KnnValueBean<T> bean : dataArray) {double score = similarScore(bean.getValue(), value);if (k == 0) {//數組中的記錄個數為0是直接添加topK[k] = new KnnValueSort(bean.getTypeId(), score);k++;} else {if (!(k == K && score < topK[k -1].getScore())){int i = 0;//找到要插入的點for (; i < k && score < topK[i].getScore(); i++);int j = k - 1;if (k < K) {j = k;k++;}for (; j > i; j--) {topK[j] = topK[j - 1];}topK[i] = new KnnValueSort(bean.getTypeId(), score);}}}return topK;}
七、統計K個樣本出現次數最多的類別
這一步就是一個簡單的計數,統計K個樣本中出現次數最多的分類,該分類就是我們要預測的目標資料的分類。
/** * @param value * @return * @Author:lulei * @Description: KNN分類判斷value的類別 */public String getTypeId(T value) {KnnValueSort[] array = getKType(value);HashMap<String, Integer> map = new HashMap<String, Integer>(K);for (KnnValueSort bean : array) {if (bean != null) {if (map.containsKey(bean.getTypeId())) {map.put(bean.getTypeId(), map.get(bean.getTypeId()) + 1);} else {map.put(bean.getTypeId(), 1);}}}String maxTypeId = null;int maxCount = 0;Iterator<Entry<String, Integer>> iter = map.entrySet().iterator();while (iter.hasNext()) {Entry<String, Integer> entry = iter.next();if (maxCount < entry.getValue()) {maxCount = entry.getValue();maxTypeId = entry.getKey();}}return maxTypeId;}
到現在為止KNN分類的抽象基類已經編寫完成,在測試之前我們先多說幾句,KNN分類是統計K個樣本中出現次數最多的分類,這種在有些情況下並不是特別合理,比如K=5,前5個樣本對應的分類分別為A、A、B、B、B,對應的相似性得分分別為10、9、2、2、1,如果使用上面的方法,那預測的分類就是B,但是看這些資料,預測的分類是A感覺更合理。基於這種情況,自己對KNN演算法提出如下最佳化(這裡並不提供代碼,只提供簡單的思路):在擷取最相似K個樣本和相似性後,可以對相似性和出現次數K做一種函數運算,比如加權,得到的函數值最大的分類就是目標的預測分類。
基類源碼
/** *@Description: KNN分類 */ package com.lulei.datamining.knn; import java.util.ArrayList;import java.util.HashMap;import java.util.Iterator;import java.util.List;import java.util.Map.Entry;import com.lulei.datamining.knn.bean.KnnValueBean;import com.lulei.datamining.knn.bean.KnnValueSort;import com.lulei.util.JsonUtil; @SuppressWarnings({"rawtypes"})public abstract class KnnClassification<T> {private List<KnnValueBean> dataArray;private int K = 3;public int getK() {return K;}public void setK(int K) {if (K < 1) {throw new IllegalArgumentException("K must greater than 0");}this.K = K;}/** * @param value * @param typeId * @Author:lulei * @Description: 向模型中添加記錄 */public void addRecord(T value, String typeId) {if (dataArray == null) {dataArray = new ArrayList<KnnValueBean>();}dataArray.add(new KnnValueBean<T>(value, typeId));}/** * @param value * @return * @Author:lulei * @Description: KNN分類判斷value的類別 */public String getTypeId(T value) {KnnValueSort[] array = getKType(value);System.out.println(JsonUtil.parseJson(array));HashMap<String, Integer> map = new HashMap<String, Integer>(K);for (KnnValueSort bean : array) {if (bean != null) {if (map.containsKey(bean.getTypeId())) {map.put(bean.getTypeId(), map.get(bean.getTypeId()) + 1);} else {map.put(bean.getTypeId(), 1);}}}String maxTypeId = null;int maxCount = 0;Iterator<Entry<String, Integer>> iter = map.entrySet().iterator();while (iter.hasNext()) {Entry<String, Integer> entry = iter.next();if (maxCount < entry.getValue()) {maxCount = entry.getValue();maxTypeId = entry.getKey();}}return maxTypeId;}/** * @param value * @return * @Author:lulei * @Description: 擷取距離最近的K個分類 */private KnnValueSort[] getKType(T value) {int k = 0;KnnValueSort[] topK = new KnnValueSort[K];for (KnnValueBean<T> bean : dataArray) {double score = similarScore(bean.getValue(), value);if (k == 0) {//數組中的記錄個數為0是直接添加topK[k] = new KnnValueSort(bean.getTypeId(), score);k++;} else {if (!(k == K && score < topK[k -1].getScore())){int i = 0;//找到要插入的點for (; i < k && score < topK[i].getScore(); i++);int j = k - 1;if (k < K) {j = k;k++;}for (; j > i; j--) {topK[j] = topK[j - 1];}topK[i] = new KnnValueSort(bean.getTypeId(), score);}}}return topK;}/** * @param o1 * @param o2 * @return * @Author:lulei * @Description: o1 o2之間的相似性 */public abstract double similarScore(T o1, T o2);}
具體子類實現
對於上面介紹的都在KNN分類的抽象基類中,對於實際的問題我們需要繼承基類並實現基類中的相似性抽象方法,這裡我們做一個簡單的實現。
/** *@Description: */ package com.lulei.datamining.knn.test; import com.lulei.datamining.knn.KnnClassification;import com.lulei.util.JsonUtil; public class Test extends KnnClassification<Integer>{@Overridepublic double similarScore(Integer o1, Integer o2) {return -1 * Math.abs(o1 - o2);}/** * @param args * @Author:lulei * @Description: */public static void main(String[] args) {Test test = new Test();for (int i = 1; i < 10; i++) {test.addRecord(i, i > 5 ? "0" : "1");}System.out.println(JsonUtil.parseJson(test.getTypeId(0)));}}
這裡我們一共添加了1、2、3、4、5、6、7、8、9這9組資料,前5組的類別為1,後4組的類別為0,兩個資料之間的相似性為兩者之間的差值的絕對值的相反數,下面預測0應該屬於的分類,這裡K的預設值為3,因此最近的K個樣本分別為1、2、3,對應的分類分別為"1"、"1"、"1",因為最後預測的分類為"1"。
-------------------------------------------------------------------------------------------------
小福利
-------------------------------------------------------------------------------------------------
個人在極客學院上《Lucene案例開發》課程已經上線了,歡迎大家吐槽~
第一課:Lucene概述
第二課:Lucene 常用功能介紹
第三課:網路爬蟲
第四課:資料庫連接池
第五課:小說網站的採集
第六課:小說網站資料庫操作
第七課:小說網站分布式爬蟲的實現
第八課:Lucene即時搜尋