4 ms·
There is no argument for why the LSH would work well, especially at the beginning of training. As the weights are initially random, bucket assignment would be
by lapink 7y ago
There is no argument for why the LSH would work well, especially at the beginning of training. As the weights are initially random, bucket assignment would be random as well. If predicting at position A requires info from position B, but they are not in the same bucket, there will be no gradient to get the query embedding of A closer to the key embedding of B. The reversible layer trick is neat though.
- gwern 7y agoWhy is that any worse than, say, starting with randomly initialized weights in general?
- MiroF 7y agoI haven't read this paper yet, but to answer your question: because bucket choice is a discrete decision - and discrete decisions are hard to pass gradients through
- gwern 7y agoBut you aren't 'making decision's to pass gradients through them at all. It's just fixed random projections AFAICT (https://openreview.net/pdf?id=rkgNKkHtvB#page=3 https://openreview.net/pdf?id=rkgNKkHtvB#page=3).
- MiroF 7y agoOkay, I've now done a brief perusal of the paper. 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]
- octbash 7y agoYes, the Reformer is basically trading off noisier for faster training / memory savings.
- MiroF 7y agoYep, and I'm not saying its a bad approach! Just trying to answer "why is that any worse than, say, starting with randomly initialized weights in general?" wrt gradient passing 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.
- lapink 7y agoI agree that eventually this could work because on some training examples, two related entries will be in the same bucket. However, I'm not sure this would really scale to parsing an entire book all at once like the author suggest. While the algorithmic complexity might scale, the odds for two related items that could be chapters appart to end up in the same bucket seems so close to zero that training time would explode. 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.
- lucidrains 7y agoI think what they did in the paper is to increase the chances by multiple rounds of hashing, up to 8 times. They show experimentally that 8 times was good enough to be equivalent to a regular transformer.
- marcinzm 7y agoWouldn't the important embedding be of the token at position A and the token at position B rather than the positions themselves? Since there's a lot fewer tokens than positions you're fairly likely to get the two tokens to hash together at least a few times. edit: I'd image the position itself (ie: word number in text rather than token) could be embedded using sin-cosine or by breaking it up into chapter/paragraph/word. Seems more meaningful and efficient than word number in text. That would prevent this issue on that side of things.