@@ -69,22 +69,25 @@ def _get_rules(data: pd.DataFrame, outcome: str) -> list:
6969 """
7070 # Discover 5 times and get the one with more confidence
7171 best_confidence = 0
72+ best_support = 0
7273 best_rules = []
73- for i in range (5 ):
74+ for i in range (3 ):
7475 # Train new model to extract 1 rule
7576 new_model = DecisionTreeClassifier ()
7677 new_model .fit (data [[column for column in data .columns if column is not outcome ]], data [outcome ])
77- best_rules = _tree_to_best_rules (new_model , [column for column in data .columns if column is not outcome ])
78+ rules = _tree_to_best_rules (new_model , [column for column in data .columns if column is not outcome ])
7879 # If any rule has been discovered
79- if len (best_rules ) > 0 :
80+ if len (rules ) > 0 :
8081 # Measure confidence
81- predictions = _predict (best_rules , data .drop ([outcome ], axis = 1 ))
82+ predictions = _predict (rules , data .drop ([outcome ], axis = 1 ))
8283 true_positives = [p and a for (p , a ) in zip (predictions , data [outcome ])]
8384 confidence = sum (true_positives ) / sum (predictions )
85+ support = sum (true_positives ) / len (data )
8486 # Retain if it's better than the previous one
85- if confidence > best_confidence :
87+ if confidence > best_confidence or ( confidence == best_confidence and support > best_support ) :
8688 best_confidence = confidence
87- best_rules = best_rules
89+ best_support = support
90+ best_rules = rules
8891 # Return the best one, or None if no rules found in any iteration
8992 return best_rules
9093
@@ -114,7 +117,7 @@ def _tree_to_best_rules(tree, feature_names) -> list:
114117 else :
115118 # Leaf node
116119 current_impurity = tree_ .impurity [current_node ]
117- current_sample_sizes = tree_ .value [current_node ][0 ] # Number of positive samples
120+ current_sample_sizes = tree_ .value [current_node ][0 ] * tree_ . n_node_samples [ current_node ] # #PositiveSamples
118121 # If it is the best leaf node, save it
119122 if current_sample_sizes [0 ] < current_sample_sizes [1 ] and ( # Less samples with negative outcome
120123 current_impurity < best_rule ["impurity" ]
@@ -139,8 +142,14 @@ def _summarize_rules(rules: list) -> list:
139142 'attribute' : attribute ,
140143 'comparison' : 'in' ,
141144 'value' : "({},{}]" .format (
142- max ([rule ['value' ] for rule in rules if rule ['attribute' ] == attribute and rule ['comparison' ] == ">" ]),
143- min ([rule ['value' ] for rule in rules if rule ['attribute' ] == attribute and rule ['comparison' ] == "<=" ])
145+ max ([
146+ rule ['value' ] for rule in rules
147+ if rule ['attribute' ] == attribute and rule ['comparison' ] == ">"
148+ ]),
149+ min ([
150+ rule ['value' ] for rule in rules
151+ if rule ['attribute' ] == attribute and rule ['comparison' ] == "<="
152+ ])
144153 )
145154 }]
146155 else :
0 commit comments