Repository navigation
a more efficient (log)sofmax #2802
Description
Activity
Just to confirm that Stan does use the latter. See
math/stan/math/prim/fun/log_softmax.hpp
Line 49 in 9207570
a, [](const auto& v) { return v.array() - log_sum_exp(v); }); and
math/stan/math/rev/fun/log_softmax.hpp
Line 106 in 83b3731
(theta.array() - log(theta.exp().sum())).matrix(), I will defer to other developers on commenting whether the proposed approach would be suitable for Stan.
Thanks @bnicenboim. As a general rule, we should do whatever Higham recommends (his blog on numerical analysis is fantastic)! We're using his matrix exponential function algorithm, too.
Our log-sum-exp function uses the shift,
log_sum_exp(v) = max(v) + log(sum(exp(v - max(v)))So the prim version of log_softmax isn't so bad. But the reverse mode goes off course and doesn't even use the log-softmax we implemented in prim. At the very least, line 106 in the reverse implementation should call the prim version. But rather than doing that, I'll just code the approach in the Blanchard et al. paper that Bruno cited.
We're also not computing the derivatives optimally according to that paper. I'll also fix that.
Another question: is it OK if I change the boundary condition? Right now, log-sum-exp applied to an empty container throws an exception. Instead, it should return -infinity, because sum of an empty element is zero, and log of zero is negative infinity.
I would be ok with the change in boundary behavior as you outline. That makes sense and it should have been like that in the first place... which brings me to the question what is the sum of an empty set in Stan?
Sums of empty containers evaluate to 0 and products of empty containers evaluate to 1. These are the usual boundary conditions for empty containers because they generalize inductively properly, like in the accumulators in C++ 11.
then I am inclined to call the current behavior a bug and the fix is what you propose.
I agree with @wds15.
I just wrote to the authors of that paper to ask what they recommend for
log_softmax. They recommend againstsoftmax(x) = exp(x - max(x) - log_sum_exp(x - max(x)))but I couldn't find the recommendations for log softmax in their paper. We implement it in the obvious way as:
log_softmax(x) = x - max(x) - log_sum_exp(x - max(x))@bnicenboim : did you find a mention of log-softmax in their paper? The question is whether we should just implement it as
log_softmax(x) = log(softmax(x))no, I didn't, but given that they recommend against:
softmax(x) = exp(x - max(x) - log_sum_exp(x - max(x)))I assumed that this was a bad idea for the log softmax
(x - max(x) - log_sum_exp(x - max(x))It cannot be that the problem is the removing of the exp(), right? it must be division vs difference....
But I would wait for the answer of the authors, they should have a more thoughtful answer :)
I was going to try to fix this issue, but after spending about 8 hours trying to implement
log_softmax(x)aslog(softmax(x))following the advice of Nick Higham et al., I'm giving up.For reasons I don't understand,
softmaxis coded completely differently thanlog_softmax. Someone tried to makelog_softmaxwork for arbitrary containers, but we only need an implementation for Eigen column vectors---that's all that's exposed in the language. I also don't understand why one throws an exception on size zero input (correct behavior) and one just returns the empty vector (wrong behavior because that's not a simplex). I also don't understand how they're supposed to work with mat_var or even if softmax works for mat_var.I'm hoping this gets easier when @SteveBronder's doc lands, but I fear it's going to be super confusing given that there appear to be a bunch of different ways to code the callback functions.
If you want to start from where I left off, which includes a lot of cleanup on testing, it's on branch
bugfix/2802-softmax-arith. Sadly, the last commit message is, "failed attempt to compile log_softmax". Everything butlog_softmaxseems to be working, but the whole point of doing this is to fix that function.@bob-carpenter I think I managed to fix the compile issue for you. The test under mix for
log_softmaxnow compiles and runs for me ok.Thanks, @wds15. I'm giving up on this issue. I can't keep up with the C++ in the math lib with the limited amount of time I have to code.
To summarize, the minimum fix is to redefine the
doublebased value oflog_softmaxin reverse mode to be implemented aslog(softmax(x))rather than with the unfolded arithmetic. Where I got stuck was in other nice-to-have features:- softmax and log_softmax having identical signatures. They should. Just
Eigen::VectorXdis fine. - softmax and log_softmax throwing with size zero input. Only log_softmax does now.
- having log_softmax delegate to softmax rather than reimplement.
- stop binding the double-based value of a matrix and recompute it in the callback---it's a huge memory sink to save it and I think we should be conservative with memory. This is a "bug" throughout the new reverse mode code.
- both should work for var-mat, but I have no idea how to do that.
- softmax and log_softmax having identical signatures. They should. Just
Old but open issue, I wonder is this still relevant as in are the room for improvement in the current implementation of softmax and/or log-softmax? I am often encountering few seemingly random divergences in logistic normal models which I think could be due to numerical issues with gradients of softmax (just a guess).
If you think you'd be able to fix it, it's welcome!
This should be easy to upgrade in terms of the math, but I couldn't figure out how to unify the calling of the various functions involved, all of which should operate on the same signature.
Thanks, I think I lack the understanding of the internals of Stan to figure that out (especially if Bob couldn't do it either).
If you don't want to refactor the trait guards, you can just use the signatures that are there and fix the math. That shouldn't be too hard as all the unit tests are already in place and the gradients don't change.
My mistake was trying to unify the signatures and then failing to do that. @SteveBronder should be able to help if you want to dive in. Some of this has been refactored since my last effort, so the problem I was facing may already be cleaned up.
I've just recently read https://academic.oup.com/imajna/article/41/4/2311/5893596?login=false
which shows that for log-softmax implemented as
is more accurate than this version (which as far as I understand the hpp files is the one that Stan uses)
(Sorry for the R code, but the point is the same).
I just thought it was worth to point this out. If there are good reasons for the way softmax is implemented, please just close and ignore, otherwise it might be useful...