Close-up of a green circuit board with a smiley face etched into the copper traces

A slow gradient, and the patch we sent upstream

An operation in our document extraction pipeline got slower as the images got larger, and it had no business doing that. The trail ended in PyTorch’s C++ internals, at a memory allocation nobody had asked for. The fix is now in PyTorch itself.

The symptom

Bilinear interpolation estimates the colour value at a point that falls between pixels. In our case the images are static and can be large, and it is the sampling locations we optimise, by gradient descent. The operation only ever looks at a fixed-size neighbourhood around each location, so its cost should barely notice the size of the image.

It noticed. Gradient computation slowed down as the image grew.

Finding it

The core of PyTorch is C++ and CUDA, so a Python profiler sees nothing useful. We used py-spy, which handles Python C++ extensions, and the time turned out to be going somewhere unglamorous: creating a zero-filled array to hold the gradient with respect to the input images. Those gradients were computed every time, whether or not anything downstream wanted them. The interpolation itself is quick, and the allocation grows with the image until it dominates the run time.

The fix

PyTorch has good facilities for adding specialised code paths based on whether an operation’s gradients are needed later. With guidance from the PyTorch maintainers we modified the C++ and CUDA code for bilinear sampling to skip the image gradients when they are not required. For our use case that was a few orders of magnitude faster.

The change went upstream rather than into a fork of our own. Our tools are built on open-source libraries, and the useful thing about building blocks you are allowed to tweak is that the tweak can go back to everyone else using them.

Read the original post