3 ms·
There is no free lunch:). I remember spending a summer using Template Model Builder (TMB), which is a useful R/C++ automatic differentiation (AD) framework, fo
by marcle 5y ago
There is no free lunch:).
I remember spending a summer using Template Model Builder (TMB), which is a useful R/C++ automatic differentiation (AD) framework, for working with accelerated failure time models. For these models, the survival to time T given covariates X is defined by S(t|X) = P(T>t|X) = S_0(t exp(-beta^T X)) for baseline survival S_0(t). I wanted to use splines for the baseline survival and then use AD for gradients and random effects. Unfortunately, after implementing the splines in template C++, I found a web page entitled "Things you should NOT do in TMB" (https://github.com/kaskr/adcomp/wiki/Things-you-should-NOT-do-in-TMB https://github.com/kaskr/adcomp/wiki/Things-you-should-NOT-d...) - which included using if statements that are based on coefficients. In this case, the splines for S_0 depend on beta, which is this specific excluded case:(. An older framework (ADMB) did not have this constraint, but dissemination of code was more difficult. Finally, PyTorch did not have an implementation of B-splines or an implementation for Laplace's approximation. Returning to my opening comment, there is no free lunch.
- ChrisRackauckas 5y agoThere is definitely no free lunch, it's good to really delineate the engineering trade-offs you're making! A lot of this work actually comes from the fact that some people I work with were building tools that could efficiently handle dynamic control flow without requiring tracing (see the description of Zygote.jl https://arxiv.org/abs/1810.07951 https://arxiv.org/abs/1810.07951). I had to bring up the question: why? It's much harder to build, needs more machinery, and in some cases can make less assumptions/less fusions (a general form of vmap is much harder for example if you cannot trace, see KernelAbstractions.jl for details). This line of inquiry led an example of why you might want to support such dynamic behaviors, so I'll leave it up to someone else to declare whether the maintenance or complexity cost is worth it to them. I wouldn't say that this means Jax or Tensorflow are doomed (far from it: simple ML architectures are quasi-static, so it's building for the correct audience), but it's good to know what exactly you're leaving out when you make a simplifying assumption.
- hyperbovine 5y agoWere you optimizing over the knots as well? Otherwise I can't see why this would be disallowed using either forward or reverse-mode AD. An infinitesimal perturbation of beta will not cause t * exp(-beta^T x) to cross a knot, so the whole thing is smooth. (And, with B-splines the derivatives are continuous from piece to piece anyways.) But in general I agree--a good spline implementation I something I miss the most when moving from scipy.interpolate to jax.scipy. Given that the SciPy implementation is mostly F77 code written before I was born, I do not see this situation resolving itself anytime soon.
- svantana 5y agoIt's not about smoothness, it's about how to JIT the gradient function. ML libs don't generally do interpolation, partly because it's tricky to vectorize (you have to search for which segment to use for each element) and partly because most ML practioners don't need it. What I've done in my code is use all the vertices for all the elements, but with weights that are mostly zero. It's pretty fast on GPU because I don't use that many vertices.
- hyperbovine 5y ago"Interpolation" here I think effectively boils down to np.searchsorted() or its equivalent, which is implemented in all the major ML libs (jittable and backpropable).
- marcle 5y agoA short answer: this requires a basis matrix for the splines rather than interpolation. A longer answer: the splines require a basis matrix B(t), do that g(S_0(t))= B(t) gamma for a vector of parameters gamma and some transformation g of survival. A classical choice would be to use M-splines and I-splines with g(S)=-log(S) and a penalised likelihood, with the constraint that the gammas should be increasing. In R, this would use the splines2 package, while in Julia, one could use the splines2.jl package (disclaimer: which I maintain). The computational challenge is that the basis matrices need to be re-evaluated for changes in the coefficients for the covariates (that is, the betas).
- mdda 5y ago"I wanted to use splines for the baseline survival" - isn't this a modelling step that could have been revisited? It seems a somewhat arbitrary choice (there are other ways to ~interpolate that are much more framework friendly) - and it seems that it forced you down a bit of a rabbit-hole.