Supervisor: Thomas Christie
Recent work has shown that, for a data distribution generated by a Markov chain, the quantities required for generating samples from the corresponding masked discrete diffusion process are available exactly, cheaply, and in closed-form via the forward-backward algorithm, and forward-filtering backward-sampling. Given that we know what an optimal algorithm looks like, it would be interesting to train a transformer on this task and try to reverse-engineer what it has learned, drawing on techniques from the field of mechanistic interpretability. We can also monitor the training dynamics of the model - to what extent does the model just memorise the examples it has seen during training, and how does its ability to generate novel (and valid) samples improve over the course of training? Does the Grokking phenomenon occur and, if so, can we tie it to learning the forward-backward algorithm via the mechanistic interpretability aspect of this project?
Prerequisites:
▶ Proficiency with PyTorch or JAX
▶ Strong maths skills.
▶ The ability to work independently.