Skip to content

Median branch assignment - #41

Merged
upadhyan merged 4 commits into
mainfrom
median-branch-assignment
Aug 25, 2026
Merged

Median branch assignment#41
upadhyan merged 4 commits into
mainfrom
median-branch-assignment

Conversation

@josh-lee-MIE2t5

Copy link
Copy Markdown
Collaborator

Summary

Speeds up AbsoluteError (median / MAE) branch assignment during coordinate descent by replacing full partition re-sorts with dynamic structures.

Previously, AbsoluteError branch assignment either skipped CD or re-sorted each partition’s samples on every bin move (O(n log n) per trial). This PR keeps sorted structure across leave/join so moves cost closer to O(n) or O(k log n).

Algorithmic approaches

  1. Merge/filter (default, production)

    • Each bin stores a pre-sorted (y, w) array.
    • Each partition stores a sorted (y, w) array plus source-bin ids.
    • Join: mergesort-style merge of the bin into the partition — O(n + k).
    • Leave: scan/filter the partition by bin id — O(n).
    • MAE uses a presorted pinball/median pass (no re-sort).
  2. BST (deprecated, A/B)

    • Augmented AVL (WeightedMAETree) per partition/output with subtree (Σw, Σw·y).
    • Batch insert/erase of a bin — O(k log N); median/MAE — O(log N).
  3. Sort (deprecated, reference)

    • Original path: collect partition samples and std::sort on every add/remove — O(n log n).

Backends are hot-swappable via SGTLEARN_MAE_BACKEND=merge|bst|sort. MAE CD remains gated by SGTLEARN_MAE_CD (off by default for sklearn CART parity). Implementations are split into separate files: AbsoluteErrorBranchAssignment (merge), …Bst, …Sort, plus shared helpers.

Benchmarks

Timing scripts and CSVs are kept out of this PR and live on the fork:

josh-lee-MIE2t5/sgtlearn-MAE-CD-SpeedUp @ median-branch-assignment

CD-only (branch assignment)

case sort bst merge bst× merge×
small_64×20 1119 ms 77 ms 364 ms 14.6× 3.1×
medium_128×50 10356 ms 382 ms 2201 ms 27.1× 4.7×
large_256×100 57983 ms 2859 ms 13313 ms 20.3× 4.4×
multiout_128×40×3 13717 ms 747 ms 3840 ms 18.4× 3.6×

Objectives matched across backends.

End-to-end shape-tree fit (median predictor)

case sort bst merge bst× merge×
2k × 8 15.7 s 7.0 s 5.2 s 2.24× 2.99×
5k × 12 81.3 s 39.3 s 31.2 s 2.07× 2.61×
8k × 16 259 s 149 s 101 s 1.74× 2.57×

Leaf counts identical across backends. Merge wins end-to-end fit; BST wins the CD-only microbench.

@upadhyan upadhyan left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

looks good!

@upadhyan
upadhyan merged commit 8782ed1 into main Aug 25, 2026
24 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[Exploration] More efficient ways of branch routing a discretization function for MAE

2 participants