Yes, but this uses only two plan-vanilla parallel scans, shown as a* and b* in the pdf.
In the sample code, that's two calls to PyTorch's built-in API.
That looks very efficient to me, almost like it shouldn't be possible. But I tested the sample code[a] and it works.
Not an expert on this, but I don't think I've ever seen that before. Have you?