mlx is apples machine learning library, its what i used to finetune lilchat on my macbook. i found an open issue on it, #3710: mx.jvp crashed if you used it on logcumsumexp. i fixed it, and on oct 4 it got merged into main. so theres a bit of my code in mlx now which is kinda crazy.
whats a jvp
most of ml uses gradients backwards. you run the model, get the loss, then go back through it working out how much each number was to blame. thats a vjp, and mlx already had it for logcumsumexp.
a jvp goes forwards instead. you nudge the inputs a tiny bit in some direction and ask how much the output moves. its used for forward mode autodiff and stuff like hessian vector products. mlx just didnt have it for this one function, it threw a RuntimeError.
whats logcumsumexp
its a running total, but in log space. at every position it does this:
out_i = log(exp(x_1) + exp(x_2) + ... + exp(x_i))
you need it whenever adding up the plain exps would overflow, which happens a lot with probabilities.
the fix
inside mlx, logcumsumexp is a Scan with a LogAddExp reduction, and Scan::jvp only knew how to do Sum. so i added the LogAddExp case.
the maths ends up pretty nice. if you nudge input j, output i moves by the softmax weight of j (softmax over everything up to i) times the nudge. so the jvp is just a softmax weighted running sum of the tangents:
dout_i = sum over j <= i of exp(x_j - out_i) * t_j
the annoying part is doing that without overflowing. the existing vjp works in log space, so i did the same. but you cant take the log of a negative number, and tangents can be negative. so i split them into the positive part and the negative part, ran each one through logcumsumexp in log space, and subtracted at the end.
two edge cases. with inclusive=False the first position has nothing added into it, so its tangent is zero now instead of NaN. and complex tangents raise a ValueError for now, because the vjp assumes real numbers too. i saw that come up in another pr so i added it before anyone asked.
testing
the test, test_logcumsumexp_jvp, checks inclusive and exclusive, both directions, really big inputs, and compares the jvp against finite differences (nudge the input by a tiny number and measure the change by hand). all 58 autograd tests passed.
getting it merged
opened it on oct 2. it got a "low priority" label, which is fair. main kept moving so i had to merge it back into my branch a few times. zcbenz, one of the mlx maintainers, approved it and merged it on oct 4. 6 commits, 30 checks passed.
i asked claude for a quick rundown of how jvps are done in mlx before i started, and said so in the pr. the code and tests are mine.
heres the pr if you want to see the actual diff.
next
find another issue. one merged pr is cool but i want it to be a normal thing.