Discover how gradient boosting fits entire parameter vectors, not just predictions.

I'm excited to share my latest blog post, where I explore the innovative application of Gradient Boosting to fit entire parameter vectors instead of just predicting a single target value.

3 min readData Science

There's a quiet assumption baked into most gradient boosting work: that a model's job ends at predicting a single value per leaf. This post challenges that assumption head-on, and the result is worth your attention. By fitting entire parameter vectors, like spline coefficients, instead of just a point estimate, this approach opens a door that most practitioners haven't even realized was locked.

Think about what that means in practice. A standard boosted tree can tell you the expected outcome for a given observation, but it can't easily tell you how that outcome changes across a range of inputs, or how uncertainty shifts with different conditions. When the model learns to predict coefficients, it's effectively learning a functional form for each observation. That's not a minor tweak. It's a shift from answering "what's the number?" to answering "what's the shape of the relationship?" For anyone doing survival analysis, causal inference, or probabilistic modeling, that distinction is the difference between a model that gives you a prediction and one that gives you a mechanism.

The use of Jax is also telling. This isn't a theoretical exercise tucked away in a paper, it's a working implementation. That matters because it lowers the barrier for others to experiment. You don't need a custom hardware setup or a team of engineers to explore this. You need a library, a clear explanation, and a willingness to rethink what a leaf node is allowed to hold. The fact that this was built and shared as a practical demonstration rather than a proposal suggests we're closer to this being a usable technique than many would assume.

What we find most compelling is the implication for probabilistic modeling specifically. If a boosted model can output coefficients for a spline basis, it can just as easily output parameters for a distribution, location, scale, shape. That means you could have a single gradient boosting framework that not only predicts a mean but also models heteroscedasticity or tail behavior, all while retaining the interpretability and robustness that trees provide. That's not a promise of some distant future; it's a direct consequence of the approach demonstrated here. The question isn't whether this will be adopted, it's which libraries will integrate it first, and how quickly the broader community recognizes that the leaf node was never the limit. It was just the default.

From Data Science

I’ve always wanted to explore the idea that boosted trees could fit entire coefficients of parameters of a distribution instead of only being able to predict a single value per leaf node. Well using {Jax} I was able to fit a Gradient Boosting Spline model where the model learns to predict the spline coefficients that best fit each individual observation. I think this has an implications for a lot of the advanced modeling techniques available to us; survival modeling, casual inference, and probabilistic modeling. I hope this post is helpful for anyone looking to learn more about gradient boosting.

Read the original at Data Science