Hardware-Aware Features for CUTLASS Kernel Selection
Organizations: ETH Zurich Zurich, Switzerland
Abstract
GPU libraries such as CUTLASS expose tens of thousands of semantically equivalent kernels for a single operation, making exhaustive autotuning expensive and execution-free selection difficult. Existing analytical selectors require hand-designed performance rules, while learned selectors operate on raw configuration parameters and must infer hardware consequences from data. We introduce a hardware-aware representation for CUTLASS kernel selection that augments candidate configurations with statically computable estimates of induced hardware behavior. We construct a dataset of 4.9 million CUTLASS kernels and train gradient-boosted and neural learning-to-rank models to rank candidates within each problem. On held-out exhaustive evaluation problems, hardware-aware representations reduce selection regret by up to 40% relative to structural baselines and 64.2% relative to NVIDIA's matrix-multiply heuristics. We further evaluate data-efficient cross-precision and epilogue-fusion transfer within CUTLASS GEMM, showing that explicitly representing candidate-induced hardware behavior provides a useful inductive bias for learned kernel selection.
Figures & tables
| Knob | Values / domain | Meaning |
| tile_m | Thread-block tile extent | |
| tile_n | Thread-block tile extent | |
| tile_k | Thread-block tile / pipeline atom | |
| stages | Mainloop pipeline stage count | |
| cluster_m , cluster_n | , | CTA cluster shape in and |
| cluster_k | fixed | -clustering not searched |
| Constraint | Reason |
| Stream-K cooperative mainloop | CUTLASS Stream-K scheduler requires the cooperative TMA mainloop. |
| Cooperative mainloop | Two consumer warpgroups split the tile; each half must stay 64-row WGMMA aligned. |
| Hopper shared-memory capacity per CTA; depends on epilogue schedule and tile via estimate_epilogue_smem_bytes . | |
| If is column-major: | TMA smem box extent is multicast-split across the cluster dimension. |
| If is row-major: | Symmetric -box limit along the cluster dimension for the transposed operand. |
| Metric | Global | Mean | Median |
| Throughput and pipeline activity | |||
| tensor_active_pct | |||
| dram_throughput_pct | |||
| compute_throughput_pct | |||
| tma_active_pct | |||
| sm_active_pct | |||
| Feature | Global | Mean | Median |
| tile_m | |||
| tile_n | |||
| tile_k | |||
| k_iters | |||
| reg_pressure_proxy | |||
| cluster_size |
| Regime | Mean | Median |
| Small | ||
| Medium | ||
| Large | ||
| Compute-heavy | ||
| Memory-heavy | ||
| Wide |
| Group | Features and purpose |
| Structural core | ; operand layouts; input, accumulation, and output types; tile and instruction shapes; pipeline stages; mainloop and epilogue schedules; cluster shape; tile scheduler; and architecture identifiers. |
| Work decomposition | Output-tile counts, reduction iterations, edge-tile waste, SM subscription, final-wave efficiency, Stream-K applicability, and cluster fit. |
| Memory behavior | Problem and tile arithmetic intensity, operand re-streaming, working-set and resident-panel size relative to L2, and input, data/epilogue/conversion traffic. |
| Resource pressure | Bytes per pipeline stage, total mainloop and epilogue shared-memory storage, register-pressure proxy, and shared-memory- and register-limited occupancy. |
| Pipeline behavior | Estimated TMA load time, WGMMA compute time, producer–consumer ratio, pipeline-fill fraction, stage amortization, and interactions between tile depth, stage count, and occupancy. |
| Threshold | Mean fraction | Median fraction | Minimum fraction |
| Within 1% | 0.43% | 0.17% | 0.17% |
| Within 5% | 2.41% | 1.00% | 0.33% |
| Within 10% | 10.91% | 4.42% | 1.17% |
| Heuristic-biased fraction | Mean sampling regret | Groups within 5% of oracle |
| 0.00 | 8.5% | 25.6% |
| 0.25 | 1.4% | 87.0% |
| 0.50 | 1.1% | 88.9% |
| 0.75 | 0.9% | 91.3% |
| 1.00 | 0.6% | 97.1% |
| Shape | Kernel count | ||||||
| TN | TT | NN | NT | Total | |||
| 2048 | 2048 | 2048 | 60,795 | 60,740 | 60,740 | 60,696 | 242,971 |
| 4096 | 4096 | 4096 | 60,795 | 60,740 | 60,740 | 60,696 | 242,971 |
| 64 | 64 | 64 | 32,190 | 32,135 | 32,135 | 32,091 | 128,551 |
| 128 | 128 | 128 | 60,795 | 60,740 | 60,740 | 60,696 | 242,971 |
| 256 | 256 | 256 | 60,795 | 60,740 | 60,740 | 60,696 | 242,971 |
| Knob | Range / values | Notes |
| (learning rate) | log-uniform | |
| max_depth | ||
| min_child_weight | log-uniform | |
| subsample | row subsampling | |
| colsample_bytree | column subsampling | |
| (L2) | log-uniform |
| Knob | Range / values | Notes |
| loss | optional categorical | |
| learning rate | log-uniform | Adam |
| hidden topology | preset tuples | depths – , widths – |
| dropout | after each ReLU | |
| weight decay | ||
| epochs |
| Method | Mean | Median | Within 1% | Within 5% | Top-1 |
| MLP (full) | 6.2% | 2.6% | 36.8% | 63.2% | 27.9% |
| MLP (structural) | 7.4% | 3.1% | 32.4% | 55.9% | 26.5% |
| XGBoost (full) | 6.4% | 4.9% | 25.0% | 52.9% | 19.1% |
| XGBoost (structural) | 10.7% | 8.3% | 17.6% | 39.7% | 13.2% |
| nvMMH (best-of- ) | 12.0% | 10.2% | 0.0% | 25.0% | 0.0% |
| nvMMH (top-1) | 17.3% | 14.6% | 0.0% | 10.3% | 0.0% |
| Model | Features | Geo-mean % roof | vs. nvMMH | Win % | |
| MLP MSE | Full | 11.9% | 1.14 | 79.6% | |
| MLP MSE | Structural only | 11.5% | 1.10 | 73.9% | |
| XGBoost MSE | Full | 11.8% | 1.13 | 77.4% | |
| XGBoost MSE | Structural only | 11.4% | 1.09 | 70.1% | |
| nvMMH | — | 37.6 | 10.5% | — | — |
| Dtype | Method | Coverage | Speedup | Win % |
| FP32 | MLP (full) | 89.5% | 1.24 | 74.6% |
| FP32 | XGBoost (full) | 96.7% | 1.24 | 78.4% |
| FP32 | MLP (structural) | 96.1% | 1.18 | 71.8% |
| FP32 | XGBoost (structural) | 91.7% | 1.20 | 68.6% |
| FP8 E4M3 | MLP (full) | 65.4% | 1.19 | 47.4% |
| FP8 E4M3 | XGBoost (full) | 73.0% | 1.17 | 51.2% |
| Model | Features | vs. nvMMH | Win % |
| MLP MSE | Full | 1.12 | 96.0% |
| MLP MSE | Structural only | 1.13 | 92.7% |
| XGBoost MSE | Full | 1.12 | 84.8% |
| XGBoost MSE | Structural only | 1.05 | 59.6% |
| Ridge | Full | 0.83 | 40.4% |
| Ridge | Structural only | 0.27 | 0.0% |
| Split | MLP (full) | MLP (struct.) | XGBoost (full) | |
| Training | 107 | 1.09 | 1.09 | 1.07 |
| Inference server | 39 | 1.25 | 1.28 | 1.29 |
| Inference device | 5 | 1.08 | 1.00 | 1.04 |