Diagram for reference https://x.com/1rreverant/status/2107546198093730287 (In my most recent recent experiment I actually removed the MLP and just used a linear project on the pixel's hidden states)
Can you confirm I'm understanding correctly?
1) Linear unpatchify as usual to go from hidden states to pixel space
2) Attention within a local window (e.g. 3x3, 5x5) to "blend" pixel space data and come up with a better image (as an alternative to MLP or Convolution)
And follow up questions:
1) How do you handle boundaries between your "attention windows"? Do you move the window just like a convolution does or are the "attention windows" all mutually exclusive from one another?
2) How much faster/slower is this operation vs. a linear layer + MLP?
1D grid example for clarity (obviously meant to be done in 2D though)
So embed -> [N dim, N dim, N dim ... , N dim] instead of embed -> [RGB, RGB, RGB ... , RGB]
2) / [followup 1)] Stride of 1, so we place a window at each pixel (so lots of overlapping)
Also I wouldn't phrase it as "come up with a better image", the point is to give less spatial decoding pressure to the patch tokens so that they can almost completely focus on feature learning instead. There is no reason to have the model learn spatial decoding when the structure prior of images is comically strong (especially compared to text), its a waste of training time and parameters
[followup 2)] I haven't measured but it was passable is all I can say (my experiments setup are horrendous right now lol)
(Also don't mind the phrasing, I just wanted to be 100% clear)
If we have extra compute lying around in the coming weeks, we’ll try this out and report back. It’s a good idea :)