[mlir][SPIR-V] Guard update-vce pass against ops exceeding target max version (#212939) GitOrigin-RevId: c67a0e3a942286b84a9c84780f8fab1ee92392f7
diff --git a/lib/Dialect/SPIRV/Transforms/UpdateVCEPass.cpp b/lib/Dialect/SPIRV/Transforms/UpdateVCEPass.cpp index febfc0f..68e4183 100644 --- a/lib/Dialect/SPIRV/Transforms/UpdateVCEPass.cpp +++ b/lib/Dialect/SPIRV/Transforms/UpdateVCEPass.cpp
@@ -136,6 +136,18 @@ } } + // Op max version requirements + if (auto maxVersionIfx = dyn_cast<spirv::QueryMaxVersionInterface>(op)) { + std::optional<spirv::Version> maxVersion = maxVersionIfx.getMaxVersion(); + if (maxVersion && *maxVersion < allowedVersion) { + return op->emitError("'") + << op->getName() << "' is missing after version " + << spirv::stringifyVersion(*maxVersion) + << " but target environment is " + << spirv::stringifyVersion(allowedVersion); + } + } + // Op extension requirements if (auto extensions = dyn_cast<spirv::QueryExtensionInterface>(op)) if (failed(checkAndUpdateExtensionRequirements( @@ -244,9 +256,6 @@ } } - // TODO: verify that the deduced version is consistent with - // SPIR-V ops' maximal version requirements. - auto triple = spirv::VerCapExtAttr::get( deducedVersion, deducedCapabilities.getArrayRef(), deducedExtensions.getArrayRef(), &getContext());
diff --git a/test/Dialect/SPIRV/Transforms/vce-deduction.mlir b/test/Dialect/SPIRV/Transforms/vce-deduction.mlir index ad9653e..e451483 100644 --- a/test/Dialect/SPIRV/Transforms/vce-deduction.mlir +++ b/test/Dialect/SPIRV/Transforms/vce-deduction.mlir
@@ -1,4 +1,4 @@ -// RUN: mlir-opt -spirv-update-vce %s | FileCheck %s +// RUN: mlir-opt -spirv-update-vce -split-input-file -verify-diagnostics %s | FileCheck %s //===----------------------------------------------------------------------===// // Version @@ -42,6 +42,25 @@ } } +// ----- + +// Test rejecting an op whose max version is below what the target +// environment allows. +// spirv.AtomicCompareExchangeWeak is only available up to v1.3. + +spirv.module Logical GLSL450 attributes { + spirv.target_env = #spirv.target_env< + #spirv.vce<v1.6, [Kernel], []>, #spirv.resource_limits<>> +} { + spirv.func @atomic_compare_exchange_weak(%ptr : !spirv.ptr<i32, Workgroup>, %value : i32, %comparator : i32) -> i32 "None" { + // expected-error @+1 {{'spirv.AtomicCompareExchangeWeak' is missing after version v1.3 but target environment is v1.6}} + %0 = spirv.AtomicCompareExchangeWeak <Workgroup> <Acquire> <None> %ptr, %value, %comparator : !spirv.ptr<i32, Workgroup> + spirv.ReturnValue %0 : i32 + } +} + +// ----- + //===----------------------------------------------------------------------===// // Capability //===----------------------------------------------------------------------===//