@@ -523,6 +523,57 @@ def test_partial_dependence_easy_target(est, power):
523523 assert r2 > 0.99
524524
525525
526+ def test_partial_dependence_recursion_decision_tree_categorical_target_feature ():
527+ # Non-ordinal categorical signal: categories {1, 2, 5} -> 1 and
528+ # {0, 3, 4} -> 0. Recursion must route with the categorical bitset.
529+ category_values = np .arange (6 , dtype = np .float64 )
530+ X = np .repeat (category_values , 4 ).reshape (- 1 , 1 )
531+ y = np .isin (X .ravel (), [1 , 2 , 5 ]).astype (np .float64 )
532+
533+ est = DecisionTreeRegressor (
534+ categorical_features = [0 ], max_depth = 1 , random_state = 0
535+ ).fit (X , y )
536+
537+ recursion = partial_dependence (
538+ est , X , features = [0 ], method = "recursion" , categorical_features = [0 ]
539+ )
540+ brute = partial_dependence (
541+ est , X , features = [0 ], method = "brute" , categorical_features = [0 ]
542+ )
543+
544+ grid = recursion ["grid_values" ][0 ].reshape (- 1 , 1 )
545+ expected = est .predict (grid )
546+
547+ assert_array_equal (recursion ["grid_values" ][0 ], category_values )
548+ assert_allclose (recursion ["average" ][0 ], brute ["average" ][0 ])
549+ assert_allclose (recursion ["average" ][0 ], expected )
550+
551+
552+ def test_partial_dependence_recursion_decision_tree_missing_target_feature ():
553+ # Missing values are isolated by the split. Recursion must route np.nan
554+ # according to the fitted node's missing_go_to_left flag.
555+ X = np .array ([0.0 , 0.0 , np .nan , np .nan , 1.0 , 1.0 , 1.0 ]).reshape (- 1 , 1 )
556+ y = np .array ([0.0 , 0.0 , 0.0 , 0.0 , 1.0 , 1.0 , 1.0 ])
557+
558+ est = DecisionTreeRegressor (max_depth = 1 , random_state = 0 ).fit (X , y )
559+ assert est .tree_ .missing_go_to_left [0 ]
560+ custom_values = {0 : [0.0 , np .nan , 1.0 ]}
561+
562+ recursion = partial_dependence (
563+ est , X , features = [0 ], method = "recursion" , custom_values = custom_values
564+ )
565+ brute = partial_dependence (
566+ est , X , features = [0 ], method = "brute" , custom_values = custom_values
567+ )
568+
569+ grid = recursion ["grid_values" ][0 ].reshape (- 1 , 1 )
570+ expected = est .predict (grid )
571+
572+ assert_array_equal (recursion ["grid_values" ][0 ], custom_values [0 ])
573+ assert_allclose (recursion ["average" ][0 ], brute ["average" ][0 ])
574+ assert_allclose (recursion ["average" ][0 ], expected )
575+
576+
526577@pytest .mark .parametrize (
527578 "Estimator" ,
528579 (
0 commit comments