These ops are just not needed in PyTorch. while is just a Python while loop. Scan is a for loop, map is a list comprehension that applies modules. No need for anything fancy.
With TF's XLA compiler, they are slowly getting towards kernel fusion, which will then reduce launch overheads.
We have similar things in the works for pytorch: to quickly JIT at runtime the dynamic graph that is getting executed. More news on this will come when time-appropriate.
Also, have you looked at Numba to do the jitting? Probably best not to have yet another separately maintained python JIT.
https://discuss.pytorch.org/t/bayesian-computation-in-pytorc... https://discuss.pytorch.org/t/distribution-implementations/4...