jax
a4f0e402 - [Mosaic GPU] Resolve different tile transforms using the largest common divisor.

Commit
1 year ago
[Mosaic GPU] Resolve different tile transforms using the largest common divisor. Before this change, transform inference would fail if different tilings needed to be used together (e.g. `(4 ,64)` and `(8, 32)`). After this change such tilings are resolved to use a tiling compatible with both, by taking the largest common divisor of each dimension. In the case above, the final tiling would be `(4, 32)`. This change also adds two additional traversals in the transform inference, as this was necessary in a real-world example of tile resolution (taken from a ragged dot kernel). The new test exercises both changes. This CL also moves `is_known_divisible` from pallas to mgpu and extends it. PiperOrigin-RevId: 771032093
Parents
Loading