Tensor program bugs found faster by checking outputs one location at a time
The Output-Space Hypothesis: Enumerative Equivalence Checking for Tensor Programs
Programming LanguagesMachine LearningSoftware Engineering
Summary
Tensor programs used in AI models are tricky to optimize because small bugs can cause big problems. Traditional tests check many outputs for one input, but this misses rare bugs hidden in complex data. The authors flipped this by checking one output location for all possible inputs, catching subtle mistakes more reliably. Their system, Dirigo, found hundreds of bugs that older tests missed, most within minutes.
What this means in practice
- •For deep learning engineers: Detect hidden errors in AI model tensor computations that traditional tests miss, improving model reliability.
- •For gpu kernel developers: Find subtle bugs in CUDA kernels for tensor processing efficiently without relying on rare test inputs.
Authors
Paul Biberstein, Joseph Devietti, Mayur Naik
Abstract
Tensor programs, as used in deep learning models, are a prime target for optimization, as small performance improvements can have a large impact across training or inference workloads. However, such optimizations are complicated and can produce subtle bugs. Traditionally, correctness is assumed when differential testing against a reference on random inputs fails to reveal bugs. However, the inputs to these programs are massive tensors, and finding bugs can require generating extremely low likelihood inputs with precise relationships among their values. We propose a novel way to find bugs more consistently by flipping the quantifiers. Rather than generating a single input and checking all output tensor locations for equivalence, what if you could check a single output tensor location's equivalence for all inputs? We implement this idea in a system, \dirigo, by using a novel symbolic execution strategy. We demonstrate that \dirigo can find bugs effectively in a public dataset of 6,988 AI-written CUDA kernels that are all marked correct by differential testing. Of these, \dirigo finds 600 kernels that are actually buggy, and finds 97.3\% of those bugs within two minutes.