[SDPA][hipDNN] generate_stats=true not supported
- Dominant language
- C++
- Stars
- 25
- Forks
- 16
- PR merge metrics
- No merged PRs in 30d
Description
SDPA forward graphs with `generate_stats=true` are rejected by the fusilli plugin.
We're integrating hipDNN SDPA into PyTorch. To reach feature parity with cuDNN and pass the upstream `test_cudnn_attention_*` tests, the fusilli plugin would need to support this. PyTorch sets `generate_stats=true` whenever `requires_grad=True`, so any training workload that uses hipDNN SDPA hits this. For example, [`test_cudnn_attention_d192_heuristic`](https://github.com/pytorch/pytorch/blob/6035abb8d168a7e2c9524a86efa2b976d178ccca/test/test_transformers.py#L2851-L2882) creates tensors with `requires_grad=True` and calls `backward()`.
The same failure affects all `test_cudnn_attention_*` tests that use `requires_grad=True`:
- [`test_cudnn_attention_d256_heuristic`](https://github.com/pytorch/pytorch/blob/6035abb8d168a7e2c9524a86efa2b976d178ccca/test/test_transformers.py#L2819)
- [`test_cudnn_attention_different_dk_dv`](https://github.com/pytorch/pytorch/blob/6035abb8d168a7e2c9524a86efa2b976d178ccca/test/test_transformers.py#L2773)
- [`test_cudnn_attention_trivial_output_transpose`](https://github.com/pytorch/pytorch/blob/6035abb8d168a7e2c9524a86efa2b976d178ccca/test/test_transformers.py#L2930)
- [`test_cudnn_attention_preserves_query_layout`](https://github.com/pytorch/pytorch/blob/6035abb8d168a7e2c9524a86efa2b976d178ccca/test/test_transformers.py#L2984)
- [`test_cudnn_attention_nonmodulo64seqlen`](https://github.com/pytorch/pytorch/blob/6035abb8d168a7e2c9524a86efa2b976d178ccca/test/test_transformers.py#L2945)
- [`test_cudnn_attention_broken_166211`](https://github.com/pytorch/pytorch/blob/6035abb8d168a7e2c9524a86efa2b976d178ccca/test/test_transformers.py#L3124)
**Repro:**
Graph serialized from `test_cudnn_attention_d192_heuristic` via `mha_graph->toJson()`. Apply to `dnn-providers/fusilli-provider/test/integration/sdpa/simple_sdpa.cpp`:
```diff
--- a/dnn-providers/fusilli-provider/test/integration/sdpa/simple_sdpa.cpp
+++ b/dnn-providers/fusilli-provider/test/integration/sdpa/simple_sdpa.cpp
@@ -18,0 +19,2 @@
+#include
+
@@ -137,0 +140,68 @@ TEST(SdpaIntegrationTest, SimpleSdpa) {
+// Graph from PyTorch test_cudnn_attention_d192_heuristic (bf16, B=32 H=16 S=640 D=192).
+// Fails at create_execution_plans because generate_stats=true is not supported.
+TEST(SdpaIntegrationTest, StatsOutput) {
+ SKIP_IF_NO_DEVICES();
+ ASSERT_EQ(hipInit(0), hipSuccess);
+
+ auto pluginPath = std::filesystem::canonical(getCurrentExecutableDirectory() /
+ FUSILLI_PLUGIN_PATH);
+ const std::array paths = {pluginPath.c_str()};
+ ASSERT_EQ(hipdnnSetEnginePluginPaths_ext(paths.size(), paths.data(),
+ HIPDNN_PLUGIN_LOADING_ABSOLUTE),
+ HIPDNN_STATUS_SUCCESS);
+ hipdnnHandle_t handle;
+ ASSERT_EQ(hipdnnCreate(&handle), HIPDNN_STATUS_SUCCESS);
+
+ // clang-format off
+ auto graph = std::make_shared();
+ auto result = graph->deserialize(nlohmann::json::parse(R"({
+ "compute_data_type": "float",
+ "intermediate_data_type": "float",
+ "io_data_type": "bfloat16",
+ "name": "",
+ "tensors": [
+ { "uid": 0, "name": "Q", "data_type": "unset", "dims": [32, 16, 640, 192], "strides": [1966080, 122880, 192, 1], "virtual": false },
+ { "uid": 1, "name": "K", "data_type": "unset", "dims": [32, 16, 640, 192], "strides": [1966080, 122880, 192, 1], "virtual": false },
+ { "uid": 2, "name": "V", "data_type": "unset", "dims": [32, 16, 640, 128], "strides": [1310720, 81920, 128, 1], "virtual": false },
+ { "uid": 3, "name": "CUDNN_SDPA::O", "data_type": "unset", "dims": [32, 16, 640, 128], "strides": [1310720, 81920, 128, 1], "virtual": false },
+ { "uid": 5, "name": "Attn_scale", "data_type": "float", "dims": [1], "strides": [1], "value": 0.07216878235340118, "value_type": "Float32Value", "virtual": false },
+ { "uid": 8, "name": "CUDNN_SDPA::STATS","data_type": "float", "dims": [32, 16, 640, 1], "strides": [10240, 640, 1, 1], "virtual": false }
+ ],
+ "nodes": [{
+ "type": "SdpaAttributes", "name": "CUDNN_SDPA", "compute_data_type": "unset",
+ "attributes": {
+ "generate_stats": true,
+ "causal_mask": false, "dropout_probability": null,
+ "alibi_mask": false, "attn_scale_value": null, "causal_mask_bottom_right": false,
+ "diagonal_alignment": "TOP_LEFT", "implementation": "AUTO",
+ "left_bound": null, "max_seq_len_kv": null, "mma_core_mode": "unset",
+ "padding_mask": false, "right_bound": null
+ },
+ "inputs": {
+ "q_tensor_uid": 0, "k_tensor_uid": 1, "v_tensor_uid": 2, "scale_tensor_uid": 5,
+ "attn_mask_tensor_uid": null, "block_mask_tensor_uid": null,
+ "descale_k_tensor_uid": null, "descale_q_tensor_uid": null,
+ "descale_s_tensor_uid": null, "descale_v_tensor_uid": null,
+ "dropout_mask_tensor_uid": null, "dropout_scale_tensor_uid": null,
+ "offset_tensor_uid": null, "page_table_k_tensor_uid": null,
+ "page_table_v_tensor_uid": null, "scale_o_tensor_uid": null,
+ "scale_s_tensor_uid": null, "seed_tensor_uid": null,
+ "seq_len_kv_tensor_uid": null, "seq_len_q_tensor_uid": null,
+ "sink_token_tensor_uid": null
+ },
+ "outputs": {
+ "o_tensor_uid": 3, "stats_tensor_uid": 8,
+ "amax_o_tensor_uid": null, "amax_s_tensor_uid": null,
+ "max_tensor_uid": null, "rng_dump_tensor_uid": null,
+ "sum_exp_tensor_uid": null
+ }
+ }]
+ })"));
+ // clang-format on
+ ASSERT_EQ(result.code, error_code_t::OK);
+ ASSERT_EQ(graph->validate().code, error_code_t::OK);
+ ASSERT_EQ(graph->build_operation_graph(handle).code, error_code_t::OK);
+ // Fails: "No engine configurations available for the graph."
+ ASSERT_EQ(graph->create_execution_plans({HeurMode_t::FALLBACK}).code, error_code_t::OK);
+}
+
```
**Error output:**
`create_execution_plans` returns `HIPDNN_BACKEND_ERROR`:
```
No engine configurations available for the graph.
```
With `HIPDNN_LOG_LEVEL=info`:
```
[FUSILLI_ENGINE] Graph not supported: SDPA with stats output not supported.
```
**Environment:**
- fusilli-provider from `rocm-libraries` at `8511c3d320`
- fusilli at `6165d53`
- ROCm 7.1.0, MI300X (gfx942)
Contributor guide
No contributing guide indexed for this repository
Assessment
This issue has not been assessed yet.