[Mlir-commits] [mlir] [mlir][SPIRV] Fix `StorageBuffer` access conversion for emulated i16 (PR #218693)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Tue Aug 25 07:20:55 PDT 2026


llvmorg-github-actions[bot] wrote:


<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-mlir-spirv

@llvm/pr-subscribers-mlir

Author: Artem Gindinson (AGindinson)

<details>
<summary>Changes</summary>

Follows up on commit 202ece6. In the absence of `Int16` and `StorageBuffer16BitAccess` in the target, `i16` isn't any different from byte & sub-byte types. As exposed by downstream smoke tests of the IREE project, an edge case where this causes issues is a 0/1-rank memref. Semantically:
```
memref<i16>  ->  ptr<struct<array<1 x i32>>>
```
Since the array lengths are the same in the absence of actual packing, just the index bounds check doesn't catch this and `InBoundsAccessChain` still gets chosen. In the end, the memref op fails to lower through the same restriction in `MemRefToSPIRV` that the original change apparently had to work around - only `AccessChain` is expected there.

Long-term, I believe the memref lowering itself could be adjusted to support emulated narrow-int access via `InBoundsAccessChain` through some bit-shifting.

**AI Tool Use disclaimer:** mildly assisted by Codex.

Signed-off-by: Artem Gindinson <gindinson@<!-- -->roofline.ai>

---
Full diff: https://github.com/llvm/llvm-project/pull/218693.diff


3 Files Affected:

- (modified) mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp (+17-19) 
- (modified) mlir/test/Conversion/MemRefToSPIRV/bitwidth-emulation.mlir (+20) 
- (modified) mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir (+6-1) 


``````````diff
diff --git a/mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp b/mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp
index 413ebbbe78548..7f8d30b5fc854 100644
--- a/mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp
+++ b/mlir/lib/Dialect/SPIRV/Transforms/SPIRVConversion.cpp
@@ -1324,43 +1324,41 @@ static std::optional<uint64_t> getMaxLinearizedIndex(ArrayRef<int64_t> shape,
   return maxLinearIndex;
 }
 
-static std::optional<uint64_t> getStorageBufferElementCount(Value basePtr) {
+static spirv::ArrayType getStorageBufferArrayType(Value basePtr) {
   auto pointerType = dyn_cast<spirv::PointerType>(basePtr.getType());
   if (!pointerType ||
       pointerType.getStorageClass() != spirv::StorageClass::StorageBuffer)
-    return std::nullopt;
+    return {};
 
   Type pointeeType = pointerType.getPointeeType();
   if (auto structType = dyn_cast<spirv::StructType>(pointeeType)) {
     if (structType.getNumElements() != 1)
-      return std::nullopt;
+      return {};
     pointeeType = structType.getElementType(0);
   }
-  auto arrayType = dyn_cast<spirv::ArrayType>(pointeeType);
-  if (!arrayType)
-    return std::nullopt;
-  return arrayType.getNumElements();
+  return dyn_cast<spirv::ArrayType>(pointeeType);
 }
 
 static bool shouldEmitInBoundsAccessChain(MemRefType baseType, Value basePtr,
                                           ArrayRef<int64_t> strides,
                                           int64_t offset,
                                           uint64_t accessElementCount) {
-  // Sub-16-bit integer memrefs may be stored using a wider SPIR-V array element
-  // than the source element. Keep a plain access chain so later bitwidth
-  // emulation can adjust the final index in storage-element units.
-  if (auto integerType = dyn_cast<IntegerType>(baseType.getElementType()))
-    if (integerType.getWidth() < 16)
-      return false;
-
   std::optional<uint64_t> maxSourceElementIndex =
       getMaxLinearizedIndex(baseType.getShape(), strides, offset);
-  std::optional<uint64_t> storageElementCount =
-      getStorageBufferElementCount(basePtr);
-  if (!maxSourceElementIndex || !storageElementCount)
+  spirv::ArrayType storageArrayType = getStorageBufferArrayType(basePtr);
+  if (!maxSourceElementIndex || !storageArrayType)
+    return false;
+
+  // Source indices and storage element counts use the same units only when
+  // each source element maps to one SPIR-V array element. An i16 or a narrower
+  // memref source may be stored using a wider SPIR-V array element than that
+  // of the source. Keep a plain access chain so later bitwidth emulation can
+  // adjust the final index in storage-element units.
+  if (baseType.getElementType() != storageArrayType.getElementType())
     return false;
 
-  if (accessElementCount == 0 || accessElementCount > *storageElementCount)
+  uint64_t storageElementCount = storageArrayType.getNumElements();
+  if (accessElementCount == 0 || accessElementCount > storageElementCount)
     return false;
 
   // `InBoundsAccessChain` requires the computed pointer to stay within the
@@ -1369,7 +1367,7 @@ static bool shouldEmitInBoundsAccessChain(MemRefType baseType, Value basePtr,
   // rejects widths that cannot fit in the fixed StorageBuffer object at all.
   // The static proof here is that the memref layout's linear index space maps
   // into that same object.
-  return *maxSourceElementIndex < *storageElementCount;
+  return *maxSourceElementIndex < storageElementCount;
 }
 
 } // namespace
diff --git a/mlir/test/Conversion/MemRefToSPIRV/bitwidth-emulation.mlir b/mlir/test/Conversion/MemRefToSPIRV/bitwidth-emulation.mlir
index a5d51dfa20bbe..a959471d7a6a9 100644
--- a/mlir/test/Conversion/MemRefToSPIRV/bitwidth-emulation.mlir
+++ b/mlir/test/Conversion/MemRefToSPIRV/bitwidth-emulation.mlir
@@ -74,6 +74,26 @@ func.func @load_i16(%arg0: memref<10xi16, #spirv.storage_class<StorageBuffer>>,
   return %0: i16
 }
 
+// The target does not support native i16 storage, so this load requires
+// bitwidth emulation even though the memref has rank zero.
+// CHECK-LABEL: @load_i16_rank0
+//  CHECK-SAME: (%[[ARG0:.+]]: memref<i16, #spirv.storage_class<StorageBuffer>>)
+//       CHECK: %[[BASE:.+]] = builtin.unrealized_conversion_cast %[[ARG0]] : memref<i16, #spirv.storage_class<StorageBuffer>> to !spirv.ptr<!spirv.struct<(!spirv.array<1 x i32, stride=4> [0])>, StorageBuffer>
+//       CHECK: %[[ZERO:.+]] = spirv.Constant 0 : i32
+//       CHECK: %[[PTR:.+]] = spirv.AccessChain %[[BASE]][%[[ZERO]], %[[ZERO]]] : {{.+}} -> !spirv.ptr<i32, StorageBuffer>
+//       CHECK: %[[LOAD:.+]] = spirv.Load "StorageBuffer" %[[PTR]] : i32
+//       CHECK: %[[MASK:.+]] = spirv.Constant 65535 : i32
+//       CHECK: %[[T1:.+]] = spirv.BitwiseAnd %[[LOAD]], %[[MASK]] : i32
+//       CHECK: %[[SIXTEEN:.+]] = spirv.Constant 16 : i32
+//       CHECK: %[[T2:.+]] = spirv.ShiftLeftLogical %[[T1]], %[[SIXTEEN]] : i32, i32
+//       CHECK: %[[T3:.+]] = spirv.ShiftRightArithmetic %[[T2]], %[[SIXTEEN]] : i32, i32
+//       CHECK: %[[RES:.+]] = builtin.unrealized_conversion_cast %[[T3]] : i32 to i16
+//       CHECK: return %[[RES]] : i16
+func.func @load_i16_rank0(%arg0: memref<i16, #spirv.storage_class<StorageBuffer>>) -> i16 {
+  %0 = memref.load %arg0[] : memref<i16, #spirv.storage_class<StorageBuffer>>
+  return %0 : i16
+}
+
 // i64 is a native type with Int64; the access chain index is used as-is without
 // the SDiv/UMod adjustment that emulated sub-32-bit types require.
 // CHECK-LABEL: @load_i64
diff --git a/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir b/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
index 89fc16bbfd07f..1095798483575 100644
--- a/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
+++ b/mlir/test/Conversion/MemRefToSPIRV/memref-to-spirv.mlir
@@ -122,10 +122,15 @@ func.func @store_i1(%dst: memref<4xi1, #spirv.storage_class<StorageBuffer>>, %i:
   return
 }
 
+// COM: Native i16 storage is supported by this test module's target.
 // CHECK-LABEL: @load_i16
+//  CHECK-SAME: (%[[ARG0:.+]]: memref<i16, #spirv.storage_class<StorageBuffer>>)
 func.func @load_i16(%arg0: memref<i16, #spirv.storage_class<StorageBuffer>>) {
   // CHECK-NOT: spirv.SDiv
-  //     CHECK: spirv.Load
+  //     CHECK: %[[BASE:.+]] = builtin.unrealized_conversion_cast %[[ARG0]] : memref<i16, #spirv.storage_class<StorageBuffer>> to !spirv.ptr<!spirv.struct<(!spirv.array<1 x i16, stride=2> [0])>, StorageBuffer>
+  //     CHECK: %[[ZERO:.+]] = spirv.Constant 0 : i32
+  //     CHECK: %[[PTR:.+]] = spirv.InBoundsAccessChain %[[BASE]][%[[ZERO]], %[[ZERO]]] : {{.+}} -> !spirv.ptr<i16, StorageBuffer>
+  //     CHECK: %[[LOAD:.+]] = spirv.Load "StorageBuffer" %[[PTR]] : i16
   // CHECK-NOT: spirv.ShiftRightArithmetic
   %0 = memref.load %arg0[] : memref<i16, #spirv.storage_class<StorageBuffer>>
   return

``````````

</details>


https://github.com/llvm/llvm-project/pull/218693


More information about the Mlir-commits mailing list