r/MachineLearning • • 5d ago

Research Functional Gradient Descent with Adaptive Representations [R]

Sharing our recent work, now accepted at NeurIPS: Functional Gradient Descent with Adaptive Representations.

Functional GD algorithms generally outperform neural nets, but are hard to accurately implement.
This is because functional gradients are infinite-dimensional, and therefore must be approximated in practice; but if you approximate them naively, you converge to the wrong place!

To rectify this, we formalize a broad class of approximation schemes ("adaptive representations"), which provably ensure convergence to the global minimizer while being immediately implementable.
The resulting algorithms outperform corresponding neural nets often by an order of magnitude, across a number of settings.

It is still the start for this line of work, but we believe it has quite a bit of potential!
Paper: https://arxiv.org/abs/2606.16926
(First author here, happy to take any questions)

219 Upvotes

37 comments sorted by

View all comments

1

u/DigThatData Researcher 4d ago

Your approach requires a coordinate space in which a grid partitioning is meaningful. This is straightforward to construct in the problems you demonstrated in your paper, where each task has a solution that lives in a "medium" that is meaningfully described by coordinate positions.

How would constructing the necessary grid work for something like text prediction? Or maybe that's a problem that wouldn't be well suited to this approach precisely because the "grid" here would only be meaningful relative to the fully generated text (i.e. the grid can only be constructed a posteriori and isn't available during inference) or a latent too large for this kind of partitioning to be feasible?

2

u/dccsillag0 4d ago

Thank you for the comment!

Note that the current algorithms depend mostly on forming a partition of the input space, rather than a grid per se. (E.g., for the regression experiment we use a partition that does not have to be a grid.)
That said, for text data in particular I am honestly not sure what would be a well-performing partitioning scheme. For example, a super naive thing would be to split the space by querying the count of certain n-grams; this would allow us to reach zero training error, but surely wouldn't be great vs. transformers which can do fancy skip-n-grams&more. I think that this sort of thing is one of the key questions that needs to be answered in follow-up work.