How to build a diffusion language model
kuleshov-group.github.io
kuleshov-group.github.io
Once you give names to the larger mathematical structures and understand them a bit better it becomes quite simple. I wish some of the blogs/papers I'd read had named "Importance Sampling".
The probability notation can be pretty confusing too. Sometimes it's hard to understand the "types" of some variables. But I'm inexperienced.
ChatGPT was surprisingly helpful. If you put in the work to truly understand the where the gaps are in your mental model (which parts aren't completely intuitive), it can do an amazing job filling in the gaps.
https://www.youtube.com/watch?v=iv-5mZ_9CPY&pp=ygUVZGlmZnVza...
As in, instead of all the complexities induced by discrete token generation, just generate the image of the text using standard image diffusion methods, then convert it to text.
If you used a single, monospace font, I bet this would be even pretty efficient, because the OCR problem becomes basically just direct template matching.
But I guess probably there is already a paper out there, I haven't searched. I'd be curious to know if it compares on par with token-based methods.
Also, you still need token embeddings (I think you might be confused how that works).
The embeddings are produced in concert with the network, to serve the network, and not created as a separate step.
It’s actually very cool
The look-up table is a matrix. Each row is an embedding and each row number is a token ID.
You get a differentiable transformation from token ID to token embedding using a “one hot vector” and a matrix multiplication
If you take the transpose of this matrix, you can convert an internal representation back to the same token form, but treat it as logits and give it to the sampler.
So token embeddings are produced on demand in service of the model, according to the model’s needs.
I found an example of this strategy in a paper as far back as 1980!
In the other reply I recommend the Bengio paper. But do bite the bullet and try it.
You can actually rig up an embedding variant of a Markov chain with just a few tokens of context, and no position coding, transformers, attention, none of it, and only minutes of training time. As long as you have the embedding lookup table trainable it will do some neat stuff.
I’ve been using diffusion Gemma and it is very fast on GPUs in output token/sec.
In the diffusion Gemma whitepaper, they say they could have done better with more time and compute.
Even with those caveats, it is very uses-able as a local model.
I only learned this the hard way reimplementing diffusiongemma. I had ideas on how to fix it but no cluster to train and experiment, hah.
Diffusion text models are cool, but they're functionally much less reliable than autoregressive transformers... and man that's really saying something. Right now most research on them is trying to figure out what complementary systems they need to be reasonably useful.