[Mlir-commits] [mlir] e2679f3 - [mlir][SPIR-V] Fix ArrayType::getSizeInBytes double counting the stride (#203821)

llvmlistbot at llvm.org llvmlistbot at llvm.org
Wed Aug 19 08:46:01 PDT 2026


Author: Arseniy Obolenskiy
Date: 2026-08-19T17:45:57+02:00
New Revision: e2679f3e50a8004dd68b15ca2e5b125c09cedfcc

URL: https://github.com/llvm/llvm-project/commit/e2679f3e50a8004dd68b15ca2e5b125c09cedfcc
DIFF: https://github.com/llvm/llvm-project/commit/e2679f3e50a8004dd68b15ca2e5b125c09cedfcc.diff

LOG: [mlir][SPIR-V] Fix ArrayType::getSizeInBytes double counting the stride (#203821)

The array stride is the per-element byte distance and already includes
the element size, so adding it to the element size overcounted strided
arrays

Added: 
    mlir/unittests/Dialect/SPIRV/SPIRVTypeTest.cpp

Modified: 
    mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
    mlir/test/Conversion/SPIRVToLLVM/spirv-types-to-llvm.mlir
    mlir/unittests/Dialect/SPIRV/CMakeLists.txt

Removed: 
    


################################################################################
diff  --git a/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp b/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
index f260b51902e9a..efb6cd0b33613 100644
--- a/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
+++ b/mlir/lib/Dialect/SPIRV/IR/SPIRVTypes.cpp
@@ -795,11 +795,13 @@ std::optional<int64_t> SPIRVType::getSizeInBytes() {
         return bitWidth / 8;
       })
       .Case([](ArrayType type) -> std::optional<int64_t> {
-        // Since array type may have an explicit stride declaration (in bytes),
-        // we also include it in the calculation.
+        // The stride, if set, is the per-element byte distance and already
+        // includes the element size; otherwise the array is tightly packed.
         auto elementType = cast<SPIRVType>(type.getElementType());
+        if (unsigned stride = type.getArrayStride())
+          return stride * type.getNumElements();
         if (std::optional<int64_t> size = elementType.getSizeInBytes())
-          return (*size + type.getArrayStride()) * type.getNumElements();
+          return *size * type.getNumElements();
         return std::nullopt;
       })
       .Case<VectorType, TensorArmType>([](auto type) -> std::optional<int64_t> {

diff  --git a/mlir/test/Conversion/SPIRVToLLVM/spirv-types-to-llvm.mlir b/mlir/test/Conversion/SPIRVToLLVM/spirv-types-to-llvm.mlir
index 0f2dbf8ef1155..04510227a9bd9 100644
--- a/mlir/test/Conversion/SPIRVToLLVM/spirv-types-to-llvm.mlir
+++ b/mlir/test/Conversion/SPIRVToLLVM/spirv-types-to-llvm.mlir
@@ -10,6 +10,9 @@ spirv.func @array(!spirv.array<16 x f32>, !spirv.array< 32 x vector<4xf32> >) "N
 // CHECK-LABEL: @array_with_natural_stride(!llvm.array<16 x f32>)
 spirv.func @array_with_natural_stride(!spirv.array<16 x f32, stride=4>) "None"
 
+// CHECK-LABEL: @array_of_strided_array(!llvm.array<3 x array<4 x f32>>)
+spirv.func @array_of_strided_array(!spirv.array<3 x !spirv.array<4 x f32, stride=4>, stride=16>) "None"
+
 //===----------------------------------------------------------------------===//
 // Pointer type
 //===----------------------------------------------------------------------===//

diff  --git a/mlir/unittests/Dialect/SPIRV/CMakeLists.txt b/mlir/unittests/Dialect/SPIRV/CMakeLists.txt
index 3aa0512459f21..967fb6d402ff1 100644
--- a/mlir/unittests/Dialect/SPIRV/CMakeLists.txt
+++ b/mlir/unittests/Dialect/SPIRV/CMakeLists.txt
@@ -1,6 +1,7 @@
 add_mlir_unittest(MLIRSPIRVImportExportTests
   DeserializationTest.cpp
   SerializationTest.cpp
+  SPIRVTypeTest.cpp
 )
 mlir_target_link_libraries(MLIRSPIRVImportExportTests
   PRIVATE

diff  --git a/mlir/unittests/Dialect/SPIRV/SPIRVTypeTest.cpp b/mlir/unittests/Dialect/SPIRV/SPIRVTypeTest.cpp
new file mode 100644
index 0000000000000..d74d7db06408d
--- /dev/null
+++ b/mlir/unittests/Dialect/SPIRV/SPIRVTypeTest.cpp
@@ -0,0 +1,66 @@
+//===- SPIRVTypeTest.cpp - SPIR-V Type Tests ------------------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#include "mlir/Dialect/SPIRV/IR/SPIRVDialect.h"
+#include "mlir/Dialect/SPIRV/IR/SPIRVTypes.h"
+#include "mlir/IR/Builders.h"
+#include "mlir/IR/MLIRContext.h"
+#include "gtest/gtest.h"
+
+using namespace mlir;
+
+TEST(SPIRVTypeTest, ArraySizeInBytes) {
+  MLIRContext context;
+  context.loadDialect<spirv::SPIRVDialect>();
+  Builder b(&context);
+  Type f32 = b.getF32Type();
+
+  // Tightly packed array: size is element size times element count.
+  auto array = spirv::ArrayType::get(f32, 16);
+  EXPECT_EQ(array.getSizeInBytes(), std::optional<int64_t>(64));
+
+  // Explicitly strided array: stride is the per-element byte distance, so it
+  // already accounts for the element size and must not be added to it.
+  auto stridedArray = spirv::ArrayType::get(f32, 16, /*stride=*/4);
+  EXPECT_EQ(stridedArray.getSizeInBytes(), std::optional<int64_t>(64));
+
+  // Padded array: stride larger than the element size includes the padding.
+  auto paddedArray = spirv::ArrayType::get(f32, 16, /*stride=*/8);
+  EXPECT_EQ(paddedArray.getSizeInBytes(), std::optional<int64_t>(128));
+}
+
+TEST(SPIRVTypeTest, ArrayOfVectorSizeInBytes) {
+  MLIRContext context;
+  context.loadDialect<spirv::SPIRVDialect>();
+  Builder b(&context);
+  auto vec4f32 = VectorType::get({4}, b.getF32Type());
+
+  // Tightly packed array of vectors: 4 * 4 bytes per element, 8 elements.
+  auto array = spirv::ArrayType::get(vec4f32, 8);
+  EXPECT_EQ(array.getSizeInBytes(), std::optional<int64_t>(128));
+
+  // Strided array of vectors: stride already covers the whole vector element.
+  auto stridedArray = spirv::ArrayType::get(vec4f32, 8, /*stride=*/16);
+  EXPECT_EQ(stridedArray.getSizeInBytes(), std::optional<int64_t>(128));
+}
+
+TEST(SPIRVTypeTest, NestedArraySizeInBytes) {
+  MLIRContext context;
+  context.loadDialect<spirv::SPIRVDialect>();
+  Builder b(&context);
+  Type f32 = b.getF32Type();
+
+  // Array of tightly packed arrays: inner is 4 * 4 = 16 bytes, outer has 3.
+  auto inner = spirv::ArrayType::get(f32, 4);
+  auto outer = spirv::ArrayType::get(inner, 3);
+  EXPECT_EQ(outer.getSizeInBytes(), std::optional<int64_t>(48));
+
+  // Outer stride dominates and includes any inner padding.
+  auto stridedOuter = spirv::ArrayType::get(inner, 3, /*stride=*/32);
+  EXPECT_EQ(stridedOuter.getSizeInBytes(), std::optional<int64_t>(96));
+}


        


More information about the Mlir-commits mailing list