An open-source implementation of the Going deeper with Image Transformers research paper in Google's JAX and Flax.
The paper also notes difficulty in training vision transformers at greater depths and proposes two solutions. First it proposes to do per-channel multiplication of the output of the residual block. Second, it proposes to have the patches attend to one another, and only allow the CLS token to attend to the patches in the last few layers.
CaiT Research Paper: https://arxiv.org/abs/2103.17239
Official Github repository: https://github.com/rwightman/pytorch-image-models
Developer updates can be found on: https://twitter.com/EnricoShippole