diff --git a/doc/releases/changelog-dev.md b/doc/releases/changelog-dev.md index 1b271f2e16..d13084e3fc 100644 --- a/doc/releases/changelog-dev.md +++ b/doc/releases/changelog-dev.md @@ -19,6 +19,10 @@ lifecycle at the IR level. [(#3047)](https://github.com/PennyLaneAI/catalyst/pull/3047) +* A `convert-transport-to-llvm` pass is added, lowering the `Transport` dialect ops to the + transport runtime CAPI. + [(#3048)](https://github.com/PennyLaneAI/catalyst/pull/3048) + * A `BufferizableOpInterface` implementation is now added for `catalyst.launch_kernel` operation and it is now bufferizable. [(#3024)](https://github.com/PennyLaneAI/catalyst/pull/3024) diff --git a/mlir/include/RegisterAllPasses.h b/mlir/include/RegisterAllPasses.h index 791fe4a784..a5195536c5 100644 --- a/mlir/include/RegisterAllPasses.h +++ b/mlir/include/RegisterAllPasses.h @@ -28,6 +28,8 @@ #include "Test/Transforms/Passes.h" #include "hlo-extensions/Transforms/Passes.h" +#include "Transport/Transforms/Passes.h" + namespace catalyst { inline void registerAllPasses() @@ -44,6 +46,7 @@ inline void registerAllPasses() qref::registerQRefPasses(); quantum::registerQuantumPasses(); rtio::registerRTIOPasses(); + transport::registerTransportPasses(); test::registerTestPasses(); } diff --git a/mlir/include/Transport/CMakeLists.txt b/mlir/include/Transport/CMakeLists.txt index f33061b2d8..9f57627c32 100644 --- a/mlir/include/Transport/CMakeLists.txt +++ b/mlir/include/Transport/CMakeLists.txt @@ -1 +1,2 @@ add_subdirectory(IR) +add_subdirectory(Transforms) diff --git a/mlir/include/Transport/Transforms/CMakeLists.txt b/mlir/include/Transport/Transforms/CMakeLists.txt new file mode 100644 index 0000000000..fb92bac94b --- /dev/null +++ b/mlir/include/Transport/Transforms/CMakeLists.txt @@ -0,0 +1,4 @@ +set(LLVM_TARGET_DEFINITIONS Passes.td) +mlir_tablegen(Passes.h.inc -gen-pass-decls -name Transport) +add_public_tablegen_target(MLIRTransportPassIncGen) +add_mlir_doc(Passes TransportPasses ./ -gen-pass-doc) diff --git a/mlir/include/Transport/Transforms/Passes.h b/mlir/include/Transport/Transforms/Passes.h new file mode 100644 index 0000000000..78d1d1c464 --- /dev/null +++ b/mlir/include/Transport/Transforms/Passes.h @@ -0,0 +1,30 @@ +// Copyright 2026 Xanadu Quantum Technologies Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#pragma once + +#include "mlir/Pass/Pass.h" + +#include "Transport/IR/TransportDialect.h" +#include "Transport/IR/TransportOps.h" + +namespace catalyst { +namespace transport { + +#define GEN_PASS_DECL +#define GEN_PASS_REGISTRATION +#include "Transport/Transforms/Passes.h.inc" + +} // namespace transport +} // namespace catalyst diff --git a/mlir/include/Transport/Transforms/Passes.td b/mlir/include/Transport/Transforms/Passes.td new file mode 100644 index 0000000000..a33e9bd80e --- /dev/null +++ b/mlir/include/Transport/Transforms/Passes.td @@ -0,0 +1,32 @@ +// Copyright 2026 Xanadu Quantum Technologies Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#ifndef TRANSPORT_PASSES +#define TRANSPORT_PASSES + +include "mlir/Pass/PassBase.td" + +def ConvertTransportToLLVMPass : Pass<"convert-transport-to-llvm", "mlir::ModuleOp"> { + let summary = "Lower the transport dialect to LLVM dialect with runtime calls."; + let description = [{ + Lowers each `transport` op to an `llvm.call` on the matching + `__catalyst__transport__*` symbol. + }]; + + let dependentDialects = [ + "mlir::LLVM::LLVMDialect" + ]; +} + +#endif // TRANSPORT_PASSES diff --git a/mlir/lib/Driver/CMakeLists.txt b/mlir/lib/Driver/CMakeLists.txt index cd55702b17..c4f6318cba 100644 --- a/mlir/lib/Driver/CMakeLists.txt +++ b/mlir/lib/Driver/CMakeLists.txt @@ -77,6 +77,7 @@ set(LIBS MLIRRTIO rtio-transforms MLIRTransport + transport-transforms MLIRCatalystTest ${ENZYME_LIB} fmt::fmt diff --git a/mlir/lib/Transport/CMakeLists.txt b/mlir/lib/Transport/CMakeLists.txt index f33061b2d8..9f57627c32 100644 --- a/mlir/lib/Transport/CMakeLists.txt +++ b/mlir/lib/Transport/CMakeLists.txt @@ -1 +1,2 @@ add_subdirectory(IR) +add_subdirectory(Transforms) diff --git a/mlir/lib/Transport/Transforms/CMakeLists.txt b/mlir/lib/Transport/Transforms/CMakeLists.txt new file mode 100644 index 0000000000..efd3bd5462 --- /dev/null +++ b/mlir/lib/Transport/Transforms/CMakeLists.txt @@ -0,0 +1,25 @@ +set(LIBRARY_NAME transport-transforms) + +file(GLOB SRC + TransportToLLVM.cpp +) + +get_property(dialect_libs GLOBAL PROPERTY MLIR_DIALECT_LIBS) +get_property(conversion_libs GLOBAL PROPERTY MLIR_CONVERSION_LIBS) +set(LIBS + ${dialect_libs} + ${conversion_libs} + MLIRTransport +) + +set(DEPENDS + MLIRTransportPassIncGen + MLIRTransportEnumsIncGen +) + +add_mlir_library(${LIBRARY_NAME} STATIC ${SRC} LINK_LIBS PRIVATE ${LIBS} DEPENDS ${DEPENDS}) +target_compile_features(${LIBRARY_NAME} PUBLIC cxx_std_20) +target_include_directories(${LIBRARY_NAME} PUBLIC + . + ${PROJECT_SOURCE_DIR}/include + ${CMAKE_BINARY_DIR}/include) diff --git a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp new file mode 100644 index 0000000000..ce5e1698a2 --- /dev/null +++ b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp @@ -0,0 +1,345 @@ +// Copyright 2026 Xanadu Quantum Technologies Inc. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Lower the `transport` dialect to `llvm.call`s on the __catalyst__transport__* +// CAPI (runtime/include/TransportCAPI.h). + +#include "llvm/ADT/Twine.h" +#include "mlir/Conversion/LLVMCommon/TypeConverter.h" +#include "mlir/Dialect/LLVMIR/FunctionCallUtils.h" +#include "mlir/Dialect/LLVMIR/LLVMDialect.h" +#include "mlir/IR/BuiltinOps.h" +#include "mlir/Pass/Pass.h" +#include "mlir/Transforms/DialectConversion.h" + +#include "Transport/IR/TransportOps.h" +#include "Transport/Transforms/Passes.h" + +using namespace mlir; +using namespace catalyst::transport; + +namespace catalyst { +namespace transport { + +#define GEN_PASS_DEF_CONVERTTRANSPORTTOLLVMPASS +#include "Transport/Transforms/Passes.h.inc" + +namespace { + +LLVM::LLVMPointerType ptrTy(MLIRContext *ctx) { return LLVM::LLVMPointerType::get(ctx); } +IntegerType i32Ty(MLIRContext *ctx) { return IntegerType::get(ctx, 32); } +IntegerType i64Ty(MLIRContext *ctx) { return IntegerType::get(ctx, 64); } + +ModuleOp moduleOf(Operation *op) { return op->getParentOfType(); } + +Value emitCall(ConversionPatternRewriter &rewriter, Location loc, ModuleOp mod, StringRef name, + ArrayRef paramTys, Type resultTy, ValueRange args) +{ + Type rty = resultTy ? resultTy : LLVM::LLVMVoidType::get(rewriter.getContext()); + auto fn = LLVM::lookupOrCreateFn(rewriter, mod, name, paramTys, rty); + assert(succeeded(fn) && "failed to declare transport CAPI function"); + auto call = LLVM::CallOp::create(rewriter, loc, *fn, args); + return call.getNumResults() ? call.getResult() : Value(); +} + +Value globalStr(ConversionPatternRewriter &rewriter, Location loc, ModuleOp mod, StringRef prefix, + StringRef value) +{ + static int counter = 0; + std::string symName = (prefix + Twine(counter++)).str(); + return LLVM::createGlobalString(loc, rewriter, symName, Twine(value).concat(Twine('\0')).str(), + LLVM::Linkage::Internal); +} + +Value constInt(ConversionPatternRewriter &rewriter, Location loc, Type ty, int64_t v) +{ + return LLVM::ConstantOp::create(rewriter, loc, ty, rewriter.getIntegerAttr(ty, v)); +} + +// From a lowered 1-D memref descriptor (an LLVM struct), extract the aligned data +// pointer and the buffer's size in bytes (num elements * element byte width). +std::pair memrefPtrAndBytes(ConversionPatternRewriter &rewriter, Location loc, + Value descriptor, MemRefType memTy) +{ + Value ptr = LLVM::ExtractValueOp::create(rewriter, loc, descriptor, ArrayRef{1}); + Value nelem = LLVM::ExtractValueOp::create(rewriter, loc, descriptor, ArrayRef{3, 0}); + Type elemTy = memTy.getElementType(); + int64_t elemBytes = isa(elemTy) ? 8 : (elemTy.getIntOrFloatBitWidth() + 7) / 8; + Value ebytes = constInt(rewriter, loc, i64Ty(rewriter.getContext()), elemBytes); + Value bytes = LLVM::MulOp::create(rewriter, loc, nelem, ebytes); + return {ptr, bytes}; +} + +//===----------------------------------------------------------------------===// +// Patterns +//===----------------------------------------------------------------------===// + +struct CreateLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(CreateOp op, OpAdaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + auto sessTy = cast(op.getSession().getType()); + Value lib = globalStr(rewriter, op.getLoc(), mod, "transport_backend_", op.getBackendLib()); + Value cfg = globalStr(rewriter, op.getLoc(), mod, "transport_config_", op.getConfig()); + Value key = globalStr(rewriter, op.getLoc(), mod, "transport_key_", op.getKey()); + Value role = + constInt(rewriter, op.getLoc(), i32Ty(ctx), static_cast(sessTy.getRole())); + Value s = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__create", + {ptrTy(ctx), ptrTy(ctx), i32Ty(ctx), ptrTy(ctx)}, ptrTy(ctx), + {lib, cfg, role, key}); + rewriter.replaceOp(op, s); + return success(); + } +}; + +template struct ConnectLoweringBase : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(OpT op, typename OpT::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = op->template getParentOfType(); + Value peer = globalStr(rewriter, op.getLoc(), mod, "transport_peer_", op.getPeer()); + Value port = constInt(rewriter, op.getLoc(), IntegerType::get(ctx, 16), op.getOobPort()); + if (Async) { + Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__connect_async", + {ptrTy(ctx), ptrTy(ctx), IntegerType::get(ctx, 16)}, i64Ty(ctx), + {adaptor.getSession(), peer, port}); + rewriter.replaceOp(op, r); + } + else { + emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__connect", + {ptrTy(ctx), ptrTy(ctx), IntegerType::get(ctx, 16)}, i32Ty(ctx), + {adaptor.getSession(), peer, port}); + rewriter.eraseOp(op); + } + return success(); + } +}; +using ConnectLowering = ConnectLoweringBase; +using ConnectAsyncLowering = ConnectLoweringBase; + +template +struct ExchangeKeysLoweringBase : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(OpT op, typename OpT::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = op->template getParentOfType(); + if (Async) { + Value r = + emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__exchange_keys_async", + {ptrTy(ctx)}, i64Ty(ctx), {adaptor.getSession()}); + rewriter.replaceOp(op, r); + } + else { + emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__exchange_keys", + {ptrTy(ctx)}, i32Ty(ctx), {adaptor.getSession()}); + rewriter.eraseOp(op); + } + return success(); + } +}; +using ExchangeKeysLowering = ExchangeKeysLoweringBase; +using ExchangeKeysAsyncLowering = ExchangeKeysLoweringBase; + +struct BarrierLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(BarrierOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__barrier", + {i64Ty(ctx)}, i32Ty(ctx), {adaptor.getToken()}); + rewriter.eraseOp(op); + return success(); + } +}; + +struct EstablishChannelLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(EstablishChannelOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + Value dp = globalStr(rewriter, op.getLoc(), moduleOf(op), "transport_data_path_", + op.getDataPath()); + emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__establish_channel", + {ptrTy(ctx), ptrTy(ctx)}, i32Ty(ctx), {adaptor.getSession(), dp}); + rewriter.eraseOp(op); + return success(); + } +}; + +struct SetCoprocessorFnLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(SetCoprocessorFnOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + Value sym = globalStr(rewriter, op.getLoc(), mod, "transport_coproc_fn_", op.getSymbol()); + emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__set_coprocessor_fn", + {ptrTy(ctx), ptrTy(ctx)}, i32Ty(ctx), {adaptor.getSession(), sym}); + rewriter.eraseOp(op); + return success(); + } +}; + +struct CommitWorkItemLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(CommitWorkItemOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + Value idx = constInt(rewriter, op.getLoc(), i32Ty(ctx), op.getWorkItemIdx()); + Value inB = constInt(rewriter, op.getLoc(), i64Ty(ctx), op.getInBytes()); + Value outB = constInt(rewriter, op.getLoc(), i64Ty(ctx), op.getOutBytes()); + emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__commit_work_item", + {ptrTy(ctx), i32Ty(ctx), i64Ty(ctx), i64Ty(ctx)}, i32Ty(ctx), + {adaptor.getSession(), idx, inB, outB}); + rewriter.eraseOp(op); + return success(); + } +}; + +struct KickLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(KickOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + auto memTy = dyn_cast(op.getPayload().getType()); + if (!memTy) + return rewriter.notifyMatchFailure(op, "kick payload must be bufferized (memref)"); + auto [srcPtr, bytes] = + memrefPtrAndBytes(rewriter, op.getLoc(), adaptor.getPayload(), memTy); + Value slot = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__data_slot", + {ptrTy(ctx)}, ptrTy(ctx), {adaptor.getSession()}); + LLVM::MemcpyOp::create(rewriter, op.getLoc(), slot, srcPtr, bytes, /*isVolatile=*/false); + Value idx = constInt(rewriter, op.getLoc(), i32Ty(ctx), op.getWorkItemIdx()); + emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__kick", + {ptrTy(ctx), i32Ty(ctx)}, i32Ty(ctx), {adaptor.getSession(), idx}); + rewriter.eraseOp(op); + return success(); + } +}; + +struct CollectLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(CollectOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + if (!op.getDest()) + return rewriter.notifyMatchFailure(op, + "collect must be bufferized (dest-passing form)"); + auto memTy = cast(op.getDest().getType()); + auto [dstPtr, bytes] = memrefPtrAndBytes(rewriter, op.getLoc(), adaptor.getDest(), memTy); + emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__collect", + {ptrTy(ctx), ptrTy(ctx), i64Ty(ctx)}, i32Ty(ctx), + {adaptor.getSession(), dstPtr, bytes}); + rewriter.eraseOp(op); + return success(); + } +}; + +struct LastRttLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(LastRttNsOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + Value r = + emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__last_rtt_ns", + {ptrTy(op.getContext())}, i64Ty(op.getContext()), {adaptor.getSession()}); + rewriter.replaceOp(op, r); + return success(); + } +}; + +// Void-returning single-session ops: start / stop / close / destroy. +template struct VoidSessionLowering : public OpConversionPattern { + VoidSessionLowering(const TypeConverter &tc, MLIRContext *ctx, StringRef sym) + : OpConversionPattern(tc, ctx), symbol(sym) + { + } + LogicalResult matchAndRewrite(OpT op, typename OpT::Adaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + emitCall(rewriter, op.getLoc(), op->template getParentOfType(), symbol, + {ptrTy(op.getContext())}, Type(), {adaptor.getSession()}); + rewriter.eraseOp(op); + return success(); + } + std::string symbol; +}; + +// get_session: look the session up by role from the runtime registry (populated at create). +struct GetSessionLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(GetSessionOp op, OpAdaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + auto sessTy = cast(op.getSession().getType()); + Value role = + constInt(rewriter, op.getLoc(), i32Ty(ctx), static_cast(sessTy.getRole())); + Value key = globalStr(rewriter, op.getLoc(), mod, "transport_key_", op.getKey()); + Value s = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__get_session", + {i32Ty(ctx), ptrTy(ctx)}, ptrTy(ctx), {role, key}); + rewriter.replaceOp(op, s); + return success(); + } +}; + +} // namespace + +struct ConvertTransportToLLVMPass + : public impl::ConvertTransportToLLVMPassBase { + using ConvertTransportToLLVMPassBase::ConvertTransportToLLVMPassBase; + + void runOnOperation() override + { + MLIRContext *ctx = &getContext(); + LLVMTypeConverter tc(ctx); + tc.addConversion([ctx](SessionType) -> Type { return LLVM::LLVMPointerType::get(ctx); }); + tc.addConversion([ctx](TokenType) -> Type { return IntegerType::get(ctx, 64); }); + + RewritePatternSet patterns(ctx); + patterns.add(tc, ctx); + patterns.add>(tc, ctx, "__catalyst__transport__start"); + patterns.add>(tc, ctx, "__catalyst__transport__stop"); + patterns.add>(tc, ctx, "__catalyst__transport__destroy"); + + ConversionTarget target(*ctx); + target.addLegalDialect(); + target.addIllegalDialect(); + + if (failed(applyPartialConversion(getOperation(), target, std::move(patterns)))) + signalPassFailure(); + } +}; + +} // namespace transport +} // namespace catalyst diff --git a/mlir/test/Transport/ConvertTransportToLLVM.mlir b/mlir/test/Transport/ConvertTransportToLLVM.mlir new file mode 100644 index 0000000000..e2590bb88d --- /dev/null +++ b/mlir/test/Transport/ConvertTransportToLLVM.mlir @@ -0,0 +1,89 @@ +// Copyright 2026 Xanadu Quantum Technologies Inc. + +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at + +// http://www.apache.org/licenses/LICENSE-2.0 + +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// RUN: quantum-opt %s --convert-transport-to-llvm --split-input-file | FileCheck %s + +// CHECK-DAG: llvm.func @__catalyst__transport__create(!llvm.ptr, !llvm.ptr, i32, !llvm.ptr) -> !llvm.ptr +// CHECK-DAG: llvm.func @__catalyst__transport__connect(!llvm.ptr, !llvm.ptr, i16) -> i32 +// CHECK-DAG: llvm.func @__catalyst__transport__exchange_keys(!llvm.ptr) -> i32 +// CHECK-DAG: llvm.func @__catalyst__transport__establish_channel(!llvm.ptr, !llvm.ptr) -> i32 +// CHECK-DAG: llvm.func @__catalyst__transport__commit_work_item(!llvm.ptr, i32, i64, i64) -> i32 +// CHECK-DAG: llvm.func @__catalyst__transport__data_slot(!llvm.ptr) -> !llvm.ptr +// CHECK-DAG: llvm.func @__catalyst__transport__kick(!llvm.ptr, i32) -> i32 +// CHECK-DAG: llvm.func @__catalyst__transport__collect(!llvm.ptr, !llvm.ptr, i64) -> i32 +// CHECK-DAG: llvm.func @__catalyst__transport__start(!llvm.ptr) +// CHECK-DAG: llvm.func @__catalyst__transport__stop(!llvm.ptr) +// CHECK-DAG: llvm.func @__catalyst__transport__destroy(!llvm.ptr) + +// Controller: create (role in the result type) -> bring-up -> kick/collect over buffers. +// CHECK-LABEL: func.func @controller +func.func @controller(%syndrome: memref, %correction: memref) { + // CHECK: %[[S:.*]] = llvm.call @__catalyst__transport__create({{.*}}) : (!llvm.ptr, !llvm.ptr, i32, !llvm.ptr) -> !llvm.ptr + %s = transport.create {backend_lib = "libbackend.so", config = "cfg"} -> !transport.session + // CHECK: llvm.call @__catalyst__transport__connect(%[[S]] + transport.connect %s {peer = "127.0.0.1", oob_port = 18560 : i16} : !transport.session + // CHECK: llvm.call @__catalyst__transport__exchange_keys(%[[S]]) + transport.exchange_keys %s : !transport.session + // CHECK: llvm.call @__catalyst__transport__establish_channel(%[[S]] + transport.establish_channel %s "cpu_verbs" : !transport.session + // CHECK: llvm.call @__catalyst__transport__commit_work_item(%[[S]] + transport.commit_work_item %s {work_item_idx = 0 : i32, in_bytes = 8 : i64, out_bytes = 8 : i64} : !transport.session + // CHECK: llvm.call @__catalyst__transport__start(%[[S]]) + transport.start %s : !transport.session + // CHECK: %[[SLOT:.*]] = llvm.call @__catalyst__transport__data_slot(%[[S]]) + // CHECK: "llvm.intr.memcpy"(%[[SLOT]] + // CHECK: llvm.call @__catalyst__transport__kick(%[[S]] + transport.kick %s, %syndrome {work_item_idx = 0 : i32} : !transport.session, memref + // CHECK: llvm.call @__catalyst__transport__collect(%[[S]] + transport.collect %s, %correction : !transport.session, memref + // CHECK: llvm.call @__catalyst__transport__stop(%[[S]]) + transport.stop %s : !transport.session + // CHECK: llvm.call @__catalyst__transport__destroy(%[[S]]) + transport.destroy %s : !transport.session + return +} + +// ----- + +// Coprocessor: create + bind the coprocessor function symbol + async bring-up. +// CHECK-LABEL: func.func @coprocessor +func.func @coprocessor() { + // CHECK: %[[C:.*]] = llvm.call @__catalyst__transport__create({{.*}}) : (!llvm.ptr, !llvm.ptr, i32, !llvm.ptr) -> !llvm.ptr + %c = transport.create {backend_lib = "libbackend.so", config = "cfg"} -> !transport.session + // CHECK: llvm.call @__catalyst__transport__connect_async(%[[C]] + %t = transport.connect_async %c {peer = "127.0.0.1", oob_port = 18560 : i16} : !transport.session -> !transport.token + // CHECK: llvm.call @__catalyst__transport__barrier + transport.barrier %t : !transport.token + // CHECK: llvm.call @__catalyst__transport__set_coprocessor_fn(%[[C]], {{.*}}) : (!llvm.ptr, !llvm.ptr) -> i32 + transport.set_coprocessor_fn %c {symbol = "foo"} : !transport.session + // CHECK: llvm.call @__catalyst__transport__destroy(%[[C]]) + transport.destroy %c : !transport.session + return +} + +// ----- + +// get_session resolves a session by (role from result type, key) via the runtime registry. +// CHECK-DAG: llvm.func @__catalyst__transport__get_session(i32, !llvm.ptr) -> !llvm.ptr +// CHECK-LABEL: func.func @resolve +func.func @resolve(%syndrome: memref, %correction: memref) { + // CHECK: %[[R:.*]] = llvm.mlir.constant(0 : i32) : i32 + // CHECK: %[[S:.*]] = llvm.call @__catalyst__transport__get_session(%[[R]], {{.*}}) : (i32, !llvm.ptr) -> !llvm.ptr + %s = transport.get_session {key = "cop0"} : !transport.session + // CHECK: llvm.call @__catalyst__transport__kick(%[[S]] + transport.kick %s, %syndrome {work_item_idx = 0 : i32} : !transport.session, memref + // CHECK: llvm.call @__catalyst__transport__collect(%[[S]] + transport.collect %s, %correction : !transport.session, memref + return +} diff --git a/mlir/tools/quantum-opt/CMakeLists.txt b/mlir/tools/quantum-opt/CMakeLists.txt index cf7a3b2053..81847edd38 100644 --- a/mlir/tools/quantum-opt/CMakeLists.txt +++ b/mlir/tools/quantum-opt/CMakeLists.txt @@ -34,6 +34,7 @@ set(LIBS MLIRRTIO rtio-transforms MLIRTransport + transport-transforms MLIRQecLogical MLIRQecPhysical MLIRCatalystTest