because bucket choice is a discrete decision - and discrete decisions are hard to pass gradients through
because bucket choice is a discrete decision - and discrete decisions are hard to pass gradients through
Perhaps I'm wrong, but it seems to me that deciding on the bucket is the discrete decision. If you have two "words"/"contexts" in a sequence that ought to attend to each other, but they don't get bucketed together early in training, then there is no gradient pushing those two hidden states to be close to each other, because there is no comparison being done between the two contexts.
In a standard transformer, on the backprop we can see something like "oh, you would have been quite closer to the correct answer on this sentence if you had matched the context for 'dog' with the context for 'treat' about 20 words back." But, here, if 'dog' doesn't get bucketed with 'treat', then there's no such gradient pressure.
Eventually (and with enough hashing+bucketing), the embedding of the more relevant contexts will move closer together, but I'd suspect this might occur more slowly. Here's the authors describing the process:
> We don’t differentiate through the hash bucket assignment procedure, or the choice of what order to sort the items into. Rather, these operations take query/key vectors as input where LSH maps nearby vectors to the same bucket with high probability. Therefore, the sorting re-adjusts any time parameter updates to cause relevant vector pairs to have higher dot product, and “unhelpful” vector pairs to have lower dot products.
e: And here is a reviewer noting what I suspected about number of gradient updates,
> the performance achieved by the proposed method after 140k iterations is achieved by the full attention after ~40k iterations [on imagenet64]
I'm not sure I'd agree with the "noisy" characterization - which to me implies stochasticity-, whereas this is just blocking off the flow of gradient information to save memory.
In particular, this approach removes all kind of domain knowledge. For images, it means ignoring entirely the prior that neighboring pixels are related, which is typically encoded through the use of convolutions. With a Reformer, not only does the locality behavior need to be learnt from scratch, but on top of that it will only happen after a sufficient number of iterations so that neighboring pixels do end up in the same bucket.
For parsing books, I think it would make much more sense to build a hierarchical model with one part parsing only a paragraph at a time and generating an intermediate embedding that could then be used as a representation of the paragraph in a larger scale Transformer working over entire chapter, and then another level going from chapters to the entire book, rather than putting all the words at once together in a giant Reformer with no domain knowledge at all and praying that with enough training data and epochs, the model will learn everything from scratch.