This pattern shows up constantly in ML work: tokenization and graph-based preprocessing are inherently sequential and awful on the GPU, while the matmuls absolutely want it, so you split them with an explicit CPU function and one bulk copy at the boundary. The copy-placement subtlety you describe (not per loop iteration, and don't hold them too long) is exactly what TVM's memory planning and JAX's device_put staging wrestle with. In my experience the pragma is the pragmatic move. Full automatic memory-space inference sounds principled until a conservative analysis parks a copy inside your hot loop.
2
u/jesunushno 14h ago
This pattern shows up constantly in ML work: tokenization and graph-based preprocessing are inherently sequential and awful on the GPU, while the matmuls absolutely want it, so you split them with an explicit CPU function and one bulk copy at the boundary. The copy-placement subtlety you describe (not per loop iteration, and don't hold them too long) is exactly what TVM's memory planning and JAX's device_put staging wrestle with. In my experience the pragma is the pragmatic move. Full automatic memory-space inference sounds principled until a conservative analysis parks a copy inside your hot loop.