代码之家  ›  专栏  ›  技术社区  ›  beesleep

决策树分类器中两片叶子之间的距离

  •  4
  • beesleep  · 技术社区  · 7 年前

    有没有一种方法可以计算 decision tree .

    我指的是从一片叶子到另一片叶子的节点数。

    graph

    例如,在这个示例图中:

    distance(leaf1, leaf2) == 1
    distance(leaf1, leaf3) == 3
    distance(leaf1, leaf4) == 4
    

    谢谢你的帮助!

    1 回复  |  直到 7 年前
        1
  •  5
  •   Kevin    7 年前

    依赖于其他python包的示例,即 networkx pydot . 出于这个原因,我们慷慨地评论了解决方案。这个问题被贴上了标签 scikit-learn 所以这个解决方案是用python给出的。

    一些数据和一个通用的 DecisionTreeClassifier :

    # load example data and classifier
    from sklearn.datasets import load_wine
    from sklearn.tree import DecisionTreeClassifier
    from sklearn.model_selection import train_test_split
    
    # for determining distance
    from sklearn import tree
    import networkx as nx
    import pydot
    
    # load data and fit a DecisionTreeClassifier
    X, y = load_wine(return_X_y=True)
    X_train, X_test, y_train, y_test = train_test_split(X, y, random_state=42)
    clf = DecisionTreeClassifier(max_depth=3, random_state=42)
    clf.fit(X_train, y_train);
    

    此函数转换拟合 决策树分类器 无方向的网络 MultiGraph 使用 tree.export_graphviz , pydot.graph_from_dot_data , nx.drawing.nx_pydot.from_pdyot nx.to_undirected .

    def dt_to_mg(clf):
        """convert a fit DecisionTreeClassifier to a Networkx undirected MultiGraph"""
        # export the classifier to a string DOT format
        dot_data = tree.export_graphviz(clf)
        # Use pydot to convert the dot data to a graph
        dot_graph = pydot.graph_from_dot_data(dot_data)[0]
        # Import the graph data into Networkx 
        MG = nx.drawing.nx_pydot.from_pydot(dot_graph)
        # Convert the tree to an undirected Networkx Graph
        uMG = MG.to_undirected()
        return uMG
    
    uMG = dt_to_mg(clf)
    

    使用 nx.shortest_path_length 找出两者之间的距离 任意两个节点 在树上。

    # get leaves
    leaves = set(str(x) for x in clf.apply(X))
    print(leaves)
    {'10', '7', '9', '5', '3', '4'}
    
    # find the distance for two leaves
    print(nx.shortest_path_length(uMG, source='9', target='5'))
    5
    
    # undirected graph means this should also work
    print(nx.shortest_path_length(uMG, source='5', target='9'))
    5
    

    shortest_path_length 返回介于 source target . 这不是公制OP请求的距离。我 认为 它们之间的节点数量 n_edges - 1 .

    print(nx.shortest_path_length(uMG, source='5', target='9') - 1)
    4
    

    或者找到所有树叶的距离,并将它们存储在字典或其他有用的对象中,以便进行下游计算。

    from itertools import combinations
    leaf_distance_edges = {}
    leaf_distance_nodes = {}
    for leaf1, leaf2 in combinations(leaves, 2):
        d = nx.shortest_path_length(uMG, source=leaf1, target=leaf2)
        leaf_distance_edges[(leaf1, leaf2)] = d
        leaf_distance_nodes[(leaf1, leaf2)] = d - 1 
    
    leaf_distance_nodes
    {('4', '9'): 5,
     ('4', '5'): 2,
     ('4', '10'): 5,
     ('4', '7'): 4,
     ('4', '3'): 1,
     ('9', '5'): 4,
     ('9', '10'): 1,
     ('9', '7'): 2,
     ('9', '3'): 5,
     ('5', '10'): 4,
     ('5', '7'): 3,
     ('5', '3'): 2,
     ('10', '7'): 2,
     ('10', '3'): 5,
     ('7', '3'): 4}