標籤:nis 實現 utf-8 var arc 樣本 isl spl 機器學習
一、kdTree 資料結構節點
- left: 左子樹
- right:右子樹
- fea:所選軸(特徵)
- dataNode:所選軸中點的樣本
二、kdTree實現主要包括兩部分:
- 1、建樹 :計算軸方差,選出方差最大的軸,進行遞迴二分
- 2、查詢:根據當前kdTree節點軸的值與要查詢節點軸的值比較,選擇向左子樹(或右子樹)遞迴查詢,得到兩點間左子樹(或右子樹)的最小距離dis;根據當前kdTree節點軸的值與要查詢節點軸的差值作比較,若差值較大,則說明(超球面是否與超矩形交割)要對右子樹(或左子樹)回溯
三、代碼實現
1 # -*- coding: utf-8 -*- 2 """ 3 Created on Sun Sep 30 12:44:51 2018 4 5 @author: Administrator 6 """ 7 import pandas as pd 8 import numpy as np 9 import math10 #定義treeNode11 class Node:12 def __init__(self,lTree,rTree,fea,dataNode): #fea表示選擇的軸,dataNode 以該節點進行分割左右子樹13 self.left=lTree;14 self.right=rTree;15 self.fea=fea;16 self.dataNode=dataNode #標籤包含在其中、17 18 19 ##直接用 DataFrame 作為資料結構20 def getInfo():21 data=[[2,3,‘羊‘],[5,4,‘猴‘],[9,6,‘雞‘],[4,7,‘狗‘],[8,1,‘豬‘],[7,2,‘猴‘]]; 22 data=pd.DataFrame(data,columns=[‘fea1‘,‘fea2‘,‘label‘])23 return data;24 25 # 計算方差,選擇軸 根據軸方差26 def calSq(data):27 sq=data.var(); 28 pos=data.columns[0];29 val=sq[0];30 for i in data.columns[1:-1]: #選擇方差最大的31 if(val<sq[i]):32 val=sq[i];33 pos=i;34 return pos;35 36 #按軸將資料拆分37 def splitAxis(data): 38 fea=calSq(data);39 sortData=data.sort_values(by=fea); #按軸排序40 sortData=(np.array(sortData)).tolist(); #轉list41 dataNode=pd.DataFrame( [ sortData[len(sortData)//2] ], columns=list(data.columns)); #資料節點42 leftSet=pd.DataFrame( sortData[0:len(sortData)//2] , columns=list(data.columns) ); #左子樹43 rightSet=pd.DataFrame(sortData[len(sortData)//2+1:] , columns=list(data.columns) ); #右子樹44 return fea,dataNode,leftSet,rightSet;45 46 #建樹47 def createTree(data): #遞迴建樹48 if(len(data)>0): #如果有資料49 fea,dataNode,leftSet,rightSet=splitAxis(data)50 treeNode=Node(None,None,fea,dataNode);51 if(len(leftSet)>0): #左邊是否可分52 treeNode.left=createTree(leftSet);53 if(len(rightSet)>0): #右邊是否可分54 treeNode.right=createTree(rightSet);55 return treeNode;56 57 #遞迴搜尋 58 def search(tree,preNode): #perNode 表示要查詢一個樣本;59 dis=0;60 for i in tree.dataNode.columns[:-1]: #計算距離61 dis=dis+( tree.dataNode[i][0]-preNode[i][0] )**2;62 dis=math.sqrt(dis);63 label=tree.dataNode[tree.dataNode.columns[-1]][0]; #當前節點標記64 labelL=‘‘;65 labelR=‘‘;66 if(tree.left!=None and preNode[tree.fea][0] < tree.dataNode[tree.fea][0] ): #左邊搜尋67 disL,labelL = search( tree.left, preNode );68 if(disL<dis): #取距離最小的69 dis=disL70 label=labelL;71 if( dis > abs(preNode[tree.fea][0] - tree.dataNode[tree.fea][0])): #超球面是否與超矩形交割 判斷是否要回溯72 disHR,labelHR=search(tree.right,preNode); #回溯右子樹73 if(disHR<dis):74 return disHR,labelHR75 else:76 return dis,label77 78 if(tree.right!=None and preNode[tree.fea][0] >= tree.dataNode[tree.fea][0] ): #右邊搜尋79 disR,labelR=search(tree.right,preNode);80 if(disR < dis): #取距離最小的81 dis=disR;82 label=labelR;83 if( dis > abs(preNode[tree.fea][0] - tree.dataNode[tree.fea][0])): #超球面是否與超矩形交割 判斷是否要回溯84 disHL,labelHL=search(tree.left,preNode); #回溯左子樹85 if(disHL<dis):86 return disHL,labelHL87 else:88 return dis,label89 return dis,label;90 91 data=getInfo();92 root=createTree(data);93 test=pd.DataFrame( [ [7.1,1] ], columns=list(data.columns[:-1]));94 dis,label=search(root,test)
機器學習——kdTree實踐