diff --git a/mlir/include/mlir/Dialect/QIR/Transforms/Passes.td b/mlir/include/mlir/Dialect/QIR/Transforms/Passes.td index d88f5d79fa..74881e5bc1 100644 --- a/mlir/include/mlir/Dialect/QIR/Transforms/Passes.td +++ b/mlir/include/mlir/Dialect/QIR/Transforms/Passes.td @@ -16,6 +16,11 @@ def QIRSetAttributesAndMetadata let dependentDialects = ["mlir::LLVM::LLVMDialect"]; let summary = "Sets the required attributes to the entry point function and " "adds the required module flags, compliant with QIR 2.1"; + let description = [{ + Static resource counts specify the capacity through the highest qubit ID + used by a quantum instruction and the highest recorded result ID. Rejects + IDs whose required capacity cannot fit in an unsigned 64-bit integer. + }]; let options = [Option<"useAdaptive", "use-adaptive", "bool", /*default= */ "true", "Specifies the profile.">]; } diff --git a/mlir/lib/Dialect/QIR/Transforms/AttachQIRAttributes.cpp b/mlir/lib/Dialect/QIR/Transforms/AttachQIRAttributes.cpp index 058ee5a8ff..2bf9e09ecf 100644 --- a/mlir/lib/Dialect/QIR/Transforms/AttachQIRAttributes.cpp +++ b/mlir/lib/Dialect/QIR/Transforms/AttachQIRAttributes.cpp @@ -27,10 +27,13 @@ #include #include #include +#include +#include #include #include #include +#include #include #include #include @@ -43,10 +46,10 @@ namespace { /// State object for tracking QIR metadata during conversion struct Metadata { - /// Number of qubits used in the module - size_t numQubits{0}; - /// Number of measurement results stored in the module - size_t numResults{0}; + /// Required capacity for static qubit IDs + uint64_t numQubits{0}; + /// Required capacity for static result IDs + uint64_t numResults{0}; /// Whether the module uses dynamic qubit management bool useDynamicQubit{false}; /// Whether the module uses dynamic result management @@ -81,7 +84,23 @@ struct QIRSetAttributesAndMetadata final return; } - Metadata metadata = useAdaptive ? getAdaptive(main) : getBase(main); + Metadata metadata = useAdaptive ? getAdaptive(main) : Metadata{}; + if (!metadata.useDynamicQubit) { + const auto numQubits = getNumQubits(main); + if (failed(numQubits)) { + signalPassFailure(); + return; + } + metadata.numQubits = *numQubits; + } + if (!metadata.useDynamicResult) { + const auto numResults = getNumResults(main); + if (failed(numResults)) { + signalPassFailure(); + return; + } + metadata.numResults = *numResults; + } if (useAdaptive) { collectOptionalFeatures(getOperation(), main, metadata); } @@ -96,8 +115,8 @@ struct QIRSetAttributesAndMetadata final /// - `entry_point`: Marks the main entry point function /// - `output_labeling_schema`: labeled /// - `qir_profiles`: base_profile - /// - `required_num_qubits`: Number of qubits used - /// - `required_num_results`: Number of measurement results + /// - `required_num_qubits`: Capacity through the highest static qubit ID + /// - `required_num_results`: Capacity through the highest recorded result ID /// - `qir_major_version`: 2 /// - `qir_minor_version`: 1 /// - `dynamic_qubit_management`: true/false @@ -228,15 +247,28 @@ struct QIRSetAttributesAndMetadata final return preserved; } - /// Count the number of uniquely indexed qubit pointers. + /// Extend the resource capacity to include an ID without overflow. + static FailureOr includeResourceId(IntegerAttr index, + uint64_t capacity, + Operation* operation) { + const auto& value = index.getValue(); + if (value.getActiveBits() > 64 || + value.getZExtValue() == std::numeric_limits::max()) { + return operation->emitError("static QIR resource ID requires a capacity " + "that does not fit in 64 bits"); + } + return std::max(capacity, value.getZExtValue() + 1); + } + + /// Return the capacity required to address every static qubit ID. /// Assumes that qubits are constant integers that are converted to /// an integer pointer and then used in (at least) one quantum instruction. - static size_t getNumQubits(LLVM::LLVMFuncOp& main) { + static FailureOr getNumQubits(LLVM::LLVMFuncOp& main) { static constexpr StringRef QIS_PREFIX = "__quantum__qis"; - DenseSet seen; + FailureOr numQubits = uint64_t{0}; main->walk([&](LLVM::ConstantOp constOp) { - if (constOp.use_empty()) { + if (failed(numQubits) || constOp.use_empty()) { return; } @@ -285,18 +317,17 @@ struct QIRSetAttributesAndMetadata final return; } - // The set ensures that we don't insert the same index multiple times. - seen.insert(intAttr.getValue()); + numQubits = includeResourceId(intAttr, *numQubits, constOp); }); - return seen.size(); + return numQubits; } - /// Count the number of uniquely indexed result_record_output statements. - static size_t getNumResults(LLVM::LLVMFuncOp& main) { - DenseSet seen; + /// Return the capacity required to address every recorded static result ID. + static FailureOr getNumResults(LLVM::LLVMFuncOp& main) { + FailureOr numResults = uint64_t{0}; main->walk([&](LLVM::CallOp callOp) { - if (!callOp.getCallee()) { + if (failed(numResults) || !callOp.getCallee()) { return; } @@ -321,11 +352,10 @@ struct QIRSetAttributesAndMetadata final return; } - // The set ensures that we don't insert the same index multiple times. - seen.insert(intAttr.getValue()); + numResults = includeResourceId(intAttr, *numResults, constOp); }); - return seen.size(); + return numResults; } /// Determine whether a loop (as a set of blocks) is an iterative loop (true) @@ -487,19 +517,7 @@ struct QIRSetAttributesAndMetadata final }); } - /// Return the metadata for a QIR base profile compliant program. - static Metadata getBase(LLVM::LLVMFuncOp& main) { - return { - .numQubits = getNumQubits(main), - .numResults = getNumResults(main), - .useDynamicQubit = false, - .useDynamicResult = false, - .useArrays = false, - .backwardsBranching = 0, - }; - } - - /// Return the metadata for a QIR adaptive profile compliant program. + /// Return the dynamic resource and control flow metadata for QIR adaptive. Metadata getAdaptive(LLVM::LLVMFuncOp& main) { const auto& domInfo = getAnalysis(); const auto [useIteration, useCondTerm] = @@ -512,14 +530,6 @@ struct QIRSetAttributesAndMetadata final md.useDynamicResult = useDynamicResult; md.useArrays = useArrays; - if (!useDynamicQubit) { - md.numQubits = getNumQubits(main); - } - - if (!useDynamicResult) { - md.numResults = getNumResults(main); - } - if (useIteration) { md.backwardsBranching = useCondTerm ? 3 : 1; } else if (useCondTerm) { diff --git a/mlir/unittests/Compiler/test_compiler_pipeline.cpp b/mlir/unittests/Compiler/test_compiler_pipeline.cpp index 25c2a3454a..b4d32f2035 100644 --- a/mlir/unittests/Compiler/test_compiler_pipeline.cpp +++ b/mlir/unittests/Compiler/test_compiler_pipeline.cpp @@ -50,6 +50,7 @@ #include #include #include +#include #include #include #include @@ -124,6 +125,8 @@ class CompilerPipelineTest protected: std::unique_ptr context; + // GoogleTest requires this override name. + // NOLINTNEXTLINE(readability-identifier-naming) void SetUp() override { DialectRegistry registry; registry.insertmodule()))); + + for (const auto format : + {ProgramFormat::QIRBase, ProgramFormat::QIRAdaptive}) { + auto output = runDefaultPipeline(CompilerInput{qc->copy()}, format); + ASSERT_TRUE(output); + auto moduleOp = std::get(*output).module(); + ASSERT_TRUE(succeeded(verify(moduleOp))); + auto main = getMainFunction(moduleOp); + ASSERT_TRUE(main); + OpBuilder builder(moduleOp.getContext()); + EXPECT_TRUE(llvm::is_contained( + main.getPassthroughAttr(), + builder.getStrArrayAttr({"required_num_qubits", "8"}))); + + LLVM::CallOp gate; + main.walk([&](LLVM::CallOp call) { + if (call.getCallee() == QIR_X) { + gate = call; + } + }); + ASSERT_TRUE(gate); + auto pointer = gate.getOperand(0).getDefiningOp(); + ASSERT_TRUE(pointer); + auto index = pointer.getArg().getDefiningOp(); + ASSERT_TRUE(index); + EXPECT_EQ(cast(index.getValue()).getInt(), 7); + } +} + TEST_F(CompilerPipelineTest, JeffRejectsMutableClassicalHelperArguments) { auto qco = QCOProgram::fromMLIRString(R"mlir(module { func.func private @helper(%bits: !cbit.reg<1>) { diff --git a/mlir/unittests/Dialect/QIR/IR/test_qir_ir.cpp b/mlir/unittests/Dialect/QIR/IR/test_qir_ir.cpp index f83aed8f45..5775e20a0d 100644 --- a/mlir/unittests/Dialect/QIR/IR/test_qir_ir.cpp +++ b/mlir/unittests/Dialect/QIR/IR/test_qir_ir.cpp @@ -19,6 +19,7 @@ #include #include #include +#include #include #include #include @@ -36,12 +37,14 @@ #include #include #include +#include #include #include #include #include #include +#include #include #include #include @@ -94,6 +97,8 @@ class QIRTest : public testing::TestWithParam { protected: std::unique_ptr context; + // GoogleTest requires this override name. + // NOLINTNEXTLINE(readability-identifier-naming) void SetUp() override { DialectRegistry registry; registry.insert(); @@ -309,6 +314,76 @@ TEST_F(QIRTest, PreservesUnrelatedMetadataIdempotently) { OperationEquivalence::Flags::None)); } +TEST_F(QIRTest, MetadataDeclaresCapacityForStaticResourceIds) { + struct CapacityCase { + SmallVector indices; + StringRef requiredCapacity; + }; + const std::array cases{ + CapacityCase{.indices = {}, .requiredCapacity = "0"}, + CapacityCase{.indices = {0}, .requiredCapacity = "1"}, + CapacityCase{.indices = {0, 1, 0}, .requiredCapacity = "2"}, + CapacityCase{.indices = {7, 2, 7}, .requiredCapacity = "8"}, + }; + + for (const auto profile : { + QIRProgramBuilder::Profile::Base, + QIRProgramBuilder::Profile::Adaptive, + }) { + for (const auto& testCase : cases) { + SCOPED_TRACE(testing::Message() + << "profile=" << static_cast(profile) + << ", requiredCapacity=" << testCase.requiredCapacity.str()); + auto moduleOp = QIRProgramBuilder::build( + context.get(), + [&](QIRProgramBuilder& builder) { + for (const auto index : testCase.indices) { + auto qubit = builder.staticQubit(index); + builder.x(qubit); + builder.measure(qubit, index); + } + return builder.intConstant(0); + }, + profile); + + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + auto main = getMainFunction(moduleOp.get()); + ASSERT_TRUE(main); + const auto passthrough = main.getPassthroughAttr(); + ASSERT_TRUE(passthrough); + OpBuilder builder(context.get()); + for (const StringRef attribute : + {"required_num_qubits", "required_num_results"}) { + EXPECT_TRUE(llvm::is_contained( + passthrough, + builder.getStrArrayAttr({attribute, testCase.requiredCapacity}))); + } + } + } +} + +TEST_F(QIRTest, MetadataRejectsUnrepresentableStaticResourceCapacity) { + auto moduleOp = parseSourceString(R"mlir(module { + llvm.func @__quantum__qis__x__body(!llvm.ptr) + llvm.func @main() attributes {passthrough = ["entry_point"]} { + %index = llvm.mlir.constant(-1 : i64) : i64 + %qubit = llvm.inttoptr %index : i64 to !llvm.ptr + llvm.call @__quantum__qis__x__body(%qubit) : (!llvm.ptr) -> () + llvm.return + } + })mlir", + context.get()); + ASSERT_TRUE(moduleOp); + ASSERT_TRUE(succeeded(verify(*moduleOp))); + OwningOpRef before = moduleOp->clone(); + + EXPECT_TRUE(failed(attachQIRMetadata(moduleOp.get()))); + EXPECT_TRUE(OperationEquivalence::isEquivalentTo( + moduleOp.get(), before->getOperation(), + OperationEquivalence::Flags::None)); +} + TEST_F(QIRTest, AdaptiveBuilderSelectsControlledSpecializationsByArity) { auto module = QIRProgramBuilder::build( context.get(),