Differentiable image compression operations in PyTorch

Viewed 1087

During a CNN classification model training while calculating the loss I am applying the encoding jpeg compression on the image in PyTorch. While I call loss.backward() it must also backpropagate through encoding and compression operation performed on the images.

Are those compression algorithms (e.g. encoding and JPEG compression) are differentiable otherwise how to backpropagate the loss gradient through those operations?

If those operations are not differentiable is there any differentiable compression algorithm that exists in PyTorch which performs H.264 encoding and JPEG compression?

Any suggestions will be highly helpful.

1 Answers

To start with, carefully consider whether you need to differentiate across the JPEG compression step. The vast majority of projects do not differentiate across this step, and if you're unsure if you need to, you probably don't.


If you really need to differentiate across an image compressor, you might consider a codec that is easier to implement than JPEG. Wavelett-based compression (the technology behind the ill-fated JPEG 2000 format) is mathematically elegant and easy to differentiate across. In a recent application of this technique, Thies et al. 2019 represent an image as a laplacian pyramid, with a loss component that serves to force sparsity in the higher resolution levels.


Now, as a thought experiment, we can look at the different steps within JPEG compression and determine if they could be implemented in a differentiable way.

  • Color transform (RBG to YCbCr): We can represent this as a point-wise convolution.

  • Chroma downsampling: Easy enough with torch.nn.functional.interpolate on the chroma channels.

  • Discrete Cosine Transform (DCT): Now things are getting interesting. Here is a Pytorch implementation of DCT that might work: https://github.com/zh217/torch-dct.

  • Quantization table: Easy again. This should just be multiplying output of the DCT with the values in the table.

  • Huffman encoding: Hard; I'm not sure this is possible. The number of output elements is going to vary based on the image entropy, which rules out many differentiable building blocks. Depending on your application, you might be able to skip this step (this step is lossless compression; so if you're trying to differentiate across the compression artifacts introduced by JPEG, the previous steps should be sufficient).

For an interesting related work on inputting JPEG DCT components directly into a neural net, see Faster Neural Networks Straight from JPEG.

Related