In plain words: Shaping how the model takes in data and gives answers, plus copying a teacher, yields programs that work on any list size; then practice cuts wasted steps. It sorted perfectly at every tested size, using fewer operations than quicksort even on lists larger than training.
Abstract · Strong Generalization and Efficiency in Neural Programs
We study the problem of learning efficient algorithms that strongly generalize in the framework of neural program induction. By carefully designing the input / output interfaces of the neural model and through imitation, we are able to learn models that produce correct results for arbitrary input sizes, achieving strong generalization. Moreover, by using reinforcement learning, we optimize for program efficiency metrics, and discover new algorithms that surpass the teacher used in imitation. With this, our approach can learn to outperform custom-written solutions for a variety of problems, as we tested it on sorting, searching in ordered lists and the NP-complete 0/1 knapsack problem, which sets a notable milestone in the field of Neural Program Induction. As highlights, our learned model can perform sorting perfectly on any input data size we tested on, with $O(n log n)$ complexity, whilst outperforming hand-coded algorithms, including quick sort, in number of operations even for list sizes far beyond those seen during training.
Yujia Li, Felix Gimeno, Pushmeet Kohli, Oriol Vinyals
arXiv:2007.03629 · cs.LG, cs.AI, cs.NE, stat.ML · submitted Jul 7, 2020 · updated Jul 8, 2020
abstract · pdf · html
Sorting was broken in java and nobody noticed for a long time: http://envisage-project.eu/wp-content/uploads/2015/02/sortin...
The same was true for java's binary search: https://ai.googleblog.com/2006/06/extra-extra-read-all-about...
So I am not sure I will ever trust an ML algorithm trained on inputs/outputs only (which is what I think "neural program induction" means). The above bugs are only hit because the programmers who wrote it didn't think about overflow. What is this "neural programmer" assuming about their inputs? We'll never know.
OTOH the standard library sort can be improved if you know the distribution of the numbers you're sorting (e.g., small numbers can bucket sort, small lists can use a sorting network, etc). If this thing can efficiently almost-correctly-sort because it's better at these types of pattern matching, we can just run a final pass of insertion sort to make it useful!