如何基于__dict__属性泛化决策树的准确率预测函数?
Got it, let's work through this. You're using __dict__ to dynamically build decision tree nodes (your __init__ converts input dicts into DecisionNode instances recursively), and you want to generalize the predicts function so it relies entirely on the node's own __dict__ instead of hardcoded attributes like attribute or children.
I'll cover two common scenarios based on how your tree is structured, both using __dict__ for full flexibility:
Scenario 1: Your tree uses attribute/children keys (original structure)
If your input dicts follow the pattern {"attribute": "feature_name", "children": {"feature_value": sub_node_or_leaf}}, you can adjust predicts to access these values via __dict__ instead of direct attribute calls. This keeps your original structure but makes the function rely on the dynamic __dict__ rather than fixed property names:
def predicts(self, x): # Check if we're at a leaf: the children node's __dict__ is empty (since leaf children are empty dicts) if not self.__dict__["children"].__dict__: # Return the leaf's predicted value from the attribute field return self.__dict__["attribute"] # Get the feature name we're splitting on from __dict__ split_feature = self.__dict__["attribute"] # Get the value of this feature from the input x feature_value = x[split_feature] # Grab the corresponding subtree from the children node's __dict__ subtree = self.__dict__["children"].__dict__[feature_value] # Recurse if it's a DecisionNode, otherwise return the leaf value directly if isinstance(subtree, DecisionNode): return subtree.predicts(x) else: return subtree
Scenario 2: Fully dynamic tree (no fixed key names)
If your tree is structured with feature names as direct keys (e.g., {"color": {"red": "apple", "green": {"shape": ...}}}), you can make predicts entirely dynamic—no hardcoded keys at all. The function will automatically detect which key corresponds to the split feature (by looking for DecisionNode values) and handle leaves by returning the non-node value:
def predicts(self, x): # Filter __dict__ to find keys that point to DecisionNode instances (our split feature) child_nodes = [(key, val) for key, val in self.__dict__.items() if isinstance(val, DecisionNode)] if not child_nodes: # We're at a leaf: return the predicted value (assuming leaf has one key-value pair like {"value": "apple"}) # Adjust this if your leaves store values differently (e.g., direct value as a key) return next(iter(self.__dict__.values())) # Grab the split feature and its corresponding subtree root split_feature, subtree_root = child_nodes[0] # Get the input's value for this feature feature_value = x[split_feature] # Get the next node/leaf from the subtree's __dict__ next_step = subtree_root.__dict__[feature_value] # Recurse if it's a node, return if it's a leaf value return next_step.predicts(x) if isinstance(next_step, DecisionNode) else next_step
How this works
- Leaf detection: We check if the node's
__dict__contains anyDecisionNodevalues. If not, it's a leaf, and we return the stored prediction. - Dynamic split feature: For non-leaf nodes, we automatically find which key maps to a subtree (a
DecisionNode), which tells us the feature we're splitting on. - Recursive traversal: We use the input's feature value to jump to the correct subtree, then repeat the process until we hit a leaf.
Quick test example
Let's use the fully dynamic scenario to see it in action:
class DecisionNode: def __init__(self, _d): self.__dict__ = {a: DecisionNode(b) if isinstance(b, dict) else b for a, b in _d.items()} def predicts(self, x): child_nodes = [(key, val) for key, val in self.__dict__.items() if isinstance(val, DecisionNode)] if not child_nodes: return next(iter(self.__dict__.values())) split_feature, subtree_root = child_nodes[0] feature_value = x[split_feature] next_step = subtree_root.__dict__[feature_value] return next_step.predicts(x) if isinstance(next_step, DecisionNode) else next_step # Build a sample tree tree_dict = { "color": { "red": {"value": "apple"}, "green": { "shape": { "round": {"value": "pear"}, "long": {"value": "banana"} } } } } my_tree = DecisionNode(tree_dict) # Test predictions print(my_tree.predicts({"color": "red"})) # Output: apple print(my_tree.predicts({"color": "green", "shape": "long"})) # Output: banana
This version of predicts is fully generalized to work with any tree structure you build via __dict__, no hardcoded property names required.
内容的提问来源于stack exchange,提问作者user9755172

