Skip to content

Commit 96c621d

Browse files
Hoist data pointers out of sdpa_bitwise_mask_gen packing loops (pytorch#23129)
Summary: Both packing loops in the generic Cadence `sdpa_bitwise_mask_gen_out` kernel called the ExecuTorch tensor accessors from inside the innermost loop, so `mask.mutable_data_ptr<T>()` ran once per element and `out.mutable_data_ptr<uint8_t>()` once per output byte. This hoists both pointers out of the loops as `const T* __restrict` / `uint8_t* __restrict`. Nothing else changes. Bit polarity and layout are untouched (bit 1 = masked/blocked, LSB-first, groups of 8 along the last dimension, bool inputs still inverted because PyTorch bool masks use True = keep), and the `ET_KERNEL_CHECK` validation is left intact because the HiFi TIE kernel at `on_device_ai/Assistant/Jarvis/min_runtime/operators/HiFi/tie/op_sdpa_bitwise_mask_gen_tie.cpp` falls back to this function for unsupported inputs and relies on it doing the full check. Why this kernel is hot: on Jazz (`DLA_V130_1p0v2`) no vision/V130 kernel is registered for `cadence::sdpa_bitwise_mask_gen.out`, so `operator_fallback.bzl` resolves vision -> V130 -> xtensa -> generic and lands here. Measured on the Jazz DSP core `DLA_V130_1p0v2` under `xt-run` (Xtensa ISS), mask shape SEQ=256 x KV=256 -> 8192 packed bytes: accessor per element: 1,302,566 cycles = 159.00 cyc/byte hoisted pointer: 163,875 cycles = 20.00 cyc/byte -> 7.95x At the production shape KV_SEQ=8192 (256 KB mask) this is 41,680,933 -> 5,242,915 cycles, about 36.4M cycles saved. Scaling is linear (159.00 cyc/byte at both sizes). Caveat on those numbers: they come from a standalone model of the loop in which the accessor was simulated with a `noinline` + `volatile` stub, so 7.95x is an upper bound, not this kernel's speedup. The real kernel has since been measured on the same core in D120855219, which adds the ISS benchmark: **1.66x for bool and 1.41x for float**. The standalone model overstated the gain by roughly 5x, as expected once `mutable_data_ptr<T>()` is a real partially-inlinable call rather than a stub. Treat 1.66x / 1.41x as this diff's result and the figures above as the proxy they came from. For end-to-end context, an ETDump captured on the Jazz virtual platform running `test_llama3_8b_tce_softmax_1_layer_small` attributes 22.6% of `Method::execute` to `native_call_sdpa_bitwise_mask_gen.out`, making it the single most expensive operation in the graph -- larger than any individual Zion delegate. The VP is functional rather than cycle-accurate, so only the op's share is quoted here, not absolute wall-clock times. Reviewed By: ThomasJannaud Differential Revision: D120841844 Pull Request resolved: pytorch#23129
1 parent f55681f commit 96c621d

1 file changed

Lines changed: 8 additions & 4 deletions

File tree

‎backends/cadence/generic/operators/op_sdpa_bitwise_mask_gen.cpp‎

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -55,22 +55,26 @@ ::executorch::aten::Tensor& sdpa_bitwise_mask_gen_out(
5555
const auto dtype = mask.dtype();
5656
const int64_t numel = mask.numel();
5757
if (dtype == ::executorch::aten::ScalarType::Bool) {
58+
const bool* __restrict mask_data = mask.mutable_data_ptr<bool>();
59+
uint8_t* __restrict out_data = out.mutable_data_ptr<uint8_t>();
5860
// Generate bitwise mask by iterating boolean tensor elements and inverting
5961
// each
6062
for (int64_t i = 0, out_index = 0; i < numel; i += 8, out_index++) {
6163
uint8_t packed_mask = 0;
6264
for (int64_t j = 0; j < 8; j++) {
63-
packed_mask |= (!mask.mutable_data_ptr<bool>()[i + j]) << j;
65+
packed_mask |= (!mask_data[i + j]) << j;
6466
}
65-
out.mutable_data_ptr<uint8_t>()[out_index] = packed_mask;
67+
out_data[out_index] = packed_mask;
6668
}
6769
} else if (dtype == ::executorch::aten::ScalarType::Float) {
70+
const float* __restrict mask_data = mask.mutable_data_ptr<float>();
71+
uint8_t* __restrict out_data = out.mutable_data_ptr<uint8_t>();
6872
for (int64_t i = 0, out_index = 0; i < numel; i += 8, out_index++) {
6973
uint8_t packed_mask = 0;
7074
for (int64_t j = 0; j < 8; j++) {
71-
packed_mask |= (mask.mutable_data_ptr<float>()[i + j] < threshold) << j;
75+
packed_mask |= (mask_data[i + j] < threshold) << j;
7276
}
73-
out.mutable_data_ptr<uint8_t>()[out_index] = packed_mask;
77+
out_data[out_index] = packed_mask;
7478
}
7579
} else {
7680
ET_KERNEL_CHECK(ctx, false, InvalidArgument, out);

0 commit comments

Comments
 (0)