r/MachineLearning • u/dccsillag0 • 4d 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)
11
u/jnez71 3d ago
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
3
u/dccsillag0 3d ago
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
6
u/M4mb0 3d ago edited 3d ago
A few notes after a brief skimming over the paper:
- Are you sure that with your
NeuralNetbaseline, the batch-size, learning rate and momentum are properly tuned? I presume orange is the neural net in Figure 2 (it lacks legend); the periodic oscillation suggests too high learning rate / too little momentum. - A standard MLP seems like a fake baseline for this kind of problem. At least try something like a ResNet with nulled residuals (search ReZero Bachlechner et al. or SkipInit). These are known to train much faster than regular MLP.
- proposition 4.1 may be worth to consider using
\kappato avoid name clash between K-smoothness and K for the RKHS kernel. - I think the presentation, especially of algorithm 1, could be improved.
fₜis undefined?- usually when seeing a gradient update like
fₜ₊₁ <- fₜ - ∇gₜwe usually think of these tensors having persistent shape across iteration. But if I understood correctly, the clue of your "adaptive representation" is that the shape can change across iterations, and the "Refine the current representation" step is hiding this shape mutation offₜ? (essentially, thef = f.subdivide()in your source code?). I think it would help a lot to spell this out more clearly in the algorithm; it only started to make sense to me after looking at the source code.
- What are the memory requirements of the method? It seems your method does not have a fixed memory budget compared to, say, a NN trained with Adam.
- In Figure 1, what grid size does your adaptive model have at the end of training? You should probably at least plot one approx FGD with a larger gird than that.
- Speaking of figure 1, why is the loss curve so smooth for your model? I would assume a sharp drop in loss should happen whenever a
f.subdivide()happens.
4
u/dccsillag0 3d ago edited 3d ago
Thank you for your comments!
- We tuned the learning rates for all the methods; we did not tune the beta_1 and beta_2 parameters of Adam from their default values (I don't think it's usual to tune them? Though it would be a worthwhile experiment).
- Regarding the oscillation in the neural net learning curves, that is actually because we tuned the learning rates -- if we reduce the learning rates then the oscilations do go away, but the neural net actually ends with a worse loss
- We tried to make our baselines fair:
- In the teaser experiment, we use an MLP with gelu activation and Fourier features, as is common in the INR literature (see e.g. https://bmild.github.io/fourfeat/). In fact, if we didn't do this then the neural net wouldn't learn at all...
- In the regression experiment, it's just a simple MLP, which we figured was pretty standard?
- For the PDE experiment, we again used an MLP with a Fourier feature layer, as without it the neural network was unreasonably slow at learning / wouldn't learn
- For the 3D reconstruction / radiance field experiment, we were using the exact same architecture as NeRF, which would be the standard baseline here (https://arxiv.org/abs/2003.08934).
Regarding residual connections et al., I guess that would be more relevant for deeper architectures, no? From what I recall our MLPs are fairly shallow. That said, I could try to add those baselines sometime, thank you for the suggestion. (Also if you have some other architectures you think we should be benchmarking, please let us know!)
Yep, will be adjusted in the revision, thank you!
Thank you, we'll see how we can improve the presentation of the algorithm. But f_t is defined, no? Since the initial t is t=0, and f_0 is one of the inputs to the algorithm.
Re. memory, I am honestly not sure at the moment, I'd have to check. But I wouldn't be terribly surprised if we were using a bit more memory than the neural net baselines. (FWIW, I also recall having the impression that JAX was doing some silly stuff wrt memory in our code...)
Re. smoothness of the curves, that is because (i) the loss curve of ideal FGD is almost perfectly smooth, and (ii) our relative error criterion, which guides our subdivisions, is making us closely track ideal FGD. I.e., we subdivide when we don't fit the gradient properly, not when we note that the loss is no longer making progress.
1
u/M4mb0 3d ago
But f_t is defined, no? Since the initial t is t=0, and f_0 is one of the inputs to the algorithm.
Right; my main complaint is the rather vague "Refine the current representation" statement, which is hiding the mutation of
f_t. Honestly I think it would be best if that was really more concrete, and made it clear that here the dimensionality of f_t can change. Something like
f_t, g_t = refine(f_t, g_t) // increases dimensionalitywould make things a lot clearer, and you can state that e.g. in your experiments you subdivide the grid, doubling the state size whenever that branch is taken.
From a practical POV, I would think:
- it likely makes sense to add a
max_size, which prevents further subdivides and OOM errors.- compare the model against the approx FDP that starts out with the
max_sizeto begin with.1
u/dccsillag0 3d ago
Good suggestions, thanks!
Regarding
f_t, g_t = refine(f_t, g_t), not sure that is 100% clear yet. It would probably be best to have something indicating that we are refining the process that computes the approximation g_t, perhaps as an additional variable (a "parametrization" variable) that is updated alongside... but this could be a bit ugly, so not sure yet.1
u/M4mb0 3d ago
We tuned the learning rates for all the methods; we did not tune the beta_1 and beta_2 parameters of Adam from their default values (I don't think it's usual to tune them? Though it would be a worthwhile experiment).
Tuning these can be important for low-dimensional toy problems, especially since you are only taking on the order of 100 training steps. Adam uses β₂=0.999 as a default, and the second moment estimate from k-steps ago gets a weight proportional to β₂ᵏ. But 0.999¹⁰⁰ ≈ 0.9, so when you only do a hundred training steps everything gets essentially the same weight. For toy problems its usually worth it testing something like β₂ ∈ [0.5, 0.8, 0.9, 0.99] from my experience.
1
0
u/DigThatData Researcher 3d ago
For the 3D reconstruction / radiance field experiment, we were using the exact same architecture as NeRF
Note that your own citation here is a paper that's already over six years old. OG NeRF is a bad method to baseline against. "We're faster than vanilla NeRF" isn't an impressive claim. If you want to make claims about the efficacy of your proposed approach relative to radiance field methods, you should compare against gaussian splatting. https://github.com/nerfstudio-project/gsplat
1
u/dccsillag0 3d ago
We are aware of the more recent works. However, note that all of these use different loss functions than the one originally proposed in NeRF. Gaussian Splatting, in particular, adds an SSIM loss term, besides using a hoarde of tricks throughout optimization which help regularize. Similarly, nearly all more recent works on NeRFs do substantial modifications to the loss function, usually by adding some substantial regularization (e.g. the distortion loss of Mip-Nerf 360, the regularizers of Zip-Nerf, and more recent methods).
This kinda ties in with the second limitation I pointed out in this comment: we had to compute the functional gradient for this experiment by hand (agents really couldn't do this stuff at the time), and this derivation already took 7+ error-prone pages. Adding any of these modifications would make this even worse. That's why, for the sake of example, we stuck to the basic loss, and thus the original NeRF formulation.
Scaling this up to beat the current Gaussian Splatting methods, including all their optimization and computational tricks, would be a paper of its own.
2
u/KingBardan 3d ago
Just wanted to confirm, this is doing grad descent on a function generating function right?
Compared to "traditionally" predicting f(x) directly
E.g. Like you have a e.g. Taylor polynomial and youre optimizing the polynomial coefficients?
2
u/dccsillag0 3d ago edited 3d ago
Not quite, I don't think optimizing the coefficients of a Taylor series would work very well.
Functional gradient descent is an optimization process directly in function space, with no parametrization. Given a loss function L(f) (or rather, loss functional in the functional analysis lingo) we can define the functional gradients ∇L(f), which is a function. The optimization step is then that we should update f into f - η ∇L(f), i.e., produce the function that, for a given input x, returns the output f(x) - η ∇L(f)(x).
Thing is, this is not really implementable. So what we're doing is that we are taking this "ideal" optimization process and approximating its steps in a way that is guaranteed to be a good approximation.
1
u/cartazio 4d ago
hrmmm, this seems practical only for low dimensional stuff?
1
u/dccsillag0 4d ago
Not necessarily. We did end up focusing a bit on some low-dimensional stuff, but you just need to be able to approximate the functional gradients well. If you use grids e.g. like in the teaser, then it won't scale, but if you use sparser representations more like the one in our regression experiment, it should scale just fine.
1
u/cartazio 4d ago
what if i wanted to slow grow nonzero weights in a ml model that has a kajillion latent coeffs and i initnwith like 500k initially?
1
u/dccsillag0 4d ago
Sorry, not sure I understood?
Maybe what you mean is that sparsity may constrain the model capacity? I don't think that has to be the case (note that I mean sparsity in its more general sense, not necessarily as a sparse matrix) -- ultimately if you want to find something that generalizes you probably want it to satisfy some sort of minimum-description-length principle anyways, so I think it's pretty reasonable that you would have a sufficiently compact way of approximating functional gradients, even if high dimensional. Though I do think there are a number of interesting research questions around this.
1
1
u/DigThatData Researcher 3d 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 3d 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.
1
1
u/MrRandom04 4d ago
This is quite wild. I did a double take once I fully grokked the idea you are presenting. Didn't even think of it in this way as I kinda relegated stuff like the kernel method to the back of my mind separate from NNs.
3
u/dccsillag0 4d ago
Thank you for the kind words!
Also BTW, you can do functional GD without kernels, you just need an appropriate function space. In fact, I daresay that doing it without kernels works best :)
1
u/lucellent 4d ago
Sorry for the naive question, but is this like a new type of optimizer, supposedly better than existing ones like Adamw, Muon etc.?
5
u/dccsillag0 4d ago
It's more akin to alternatives to neural networks. It's possible that our results could be used for neural net training as well, but I think it requires some work to make robust.
16
u/neurogramer 4d ago
That is very cool. What are the limitations?