Functional Gradient Descent with Adaptive Representations [R]

Research
media poster

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)

That is very cool. What are the limitations?

Thanks!

IMO, the main limitation (at least for now) is that of inductive bias. If you look at our current experiments, you will see that they are all on settings where we would use MLP-ish neural networks (e.g. plain MLPs, MLPs with positional encodings or Fourier features, MLPs with some feature sharing, etc.). I would not expect our current procedures to work well e.g. on language modelling. I don't think it's a fundamental limitation of the approach, just that it needs some "tweak", analogous to how we opt to use a Transformer instead of an MLP.

Additionally, to use our method you need to compute the form of the functional gradient by hand. And this can get tricky -- see, e.g. the 7-page derivation in the appendix for the functional gradient for our third experiment. Though I think that recent LLM agents can do these derivations on their own now...

still reading through your paper, and a bit of a novice - is there anything limiting using models similar to what you present as a direct swap in replacement for MLP components in larger neural networks?

How does regularization work here, eg in a supervised learning setting where you're fitting to some training data? In some sense, a perfect fit to the data is easy (a lookup table on your training data) but it won't generalize. I think the main point of NN architecture design is to pick one that will generalize well (on some particular domain, when combined with some particular initialization and training procedure, etc).

Was immediately jumping to the same conclusion you mentioned in your last sentence...

Very intersting work guys, kudos!

yes, I don't see limitations clearly laid out in the paper. Other than that it looks like nice work

Neat work. It's remnicient of adaptive refinement in PDE solvers, which also uses error bounds to adjust representation fidelity while minimizing residuals, but nice to see such a mechanism spelled out plainly. Btw you may want to temper your global optimality claim in the intro by mentioning that it requires a functional convexity condition (Polyak-Lojasiewicz). Otherwise you'll bug global nonconvex optimization researchers like myself lol

Thanks you for the comments! Indeed it has some similarities to e.g. some finite element methods.
We'll see if we can tweak the writing of the global minimizer part, haha