From 6984afaf17a764d060aaea8e97be5732ca554801 Mon Sep 17 00:00:00 2001 From: Mehrdad Malekmohammadi Date: Tue, 21 Jul 2026 07:10:23 -0400 Subject: [PATCH 1/7] Add transport-to-llvm pass Assisted-by: Claude Opus 4.8 --- .../Transport/Transforms/CMakeLists.txt | 4 + mlir/include/Transport/Transforms/Passes.h | 30 ++ mlir/include/Transport/Transforms/Passes.td | 32 +++ mlir/lib/Transport/Transforms/CMakeLists.txt | 24 ++ .../Transport/Transforms/TransportToLLVM.cpp | 266 ++++++++++++++++++ .../Transport/ConvertTransportToLLVM.mlir | 57 ++++ 6 files changed, 413 insertions(+) create mode 100644 mlir/include/Transport/Transforms/CMakeLists.txt create mode 100644 mlir/include/Transport/Transforms/Passes.h create mode 100644 mlir/include/Transport/Transforms/Passes.td create mode 100644 mlir/lib/Transport/Transforms/CMakeLists.txt create mode 100644 mlir/lib/Transport/Transforms/TransportToLLVM.cpp create mode 100644 mlir/test/Transport/ConvertTransportToLLVM.mlir 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/Transport/Transforms/CMakeLists.txt b/mlir/lib/Transport/Transforms/CMakeLists.txt new file mode 100644 index 0000000000..17ee33278d --- /dev/null +++ b/mlir/lib/Transport/Transforms/CMakeLists.txt @@ -0,0 +1,24 @@ +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 +) + +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..e96f68d1f7 --- /dev/null +++ b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp @@ -0,0 +1,266 @@ +// 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). Controller-side only. + +#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 { + +// sizeof(CatalystTransportPeerRef) = {u32, u64, u64}; over-allocate for alignment. +constexpr int64_t kPeerRefBytes = 32; + +LLVM::LLVMPointerType ptrTy(MLIRContext *ctx) { return LLVM::LLVMPointerType::get(ctx); } + +ModuleOp moduleOf(Operation *op) { return op->getParentOfType(); } + +// Declare-or-reuse a CAPI function and emit a call to it. A null resultTy means +// the function returns void. +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(); +} + +// Materialize a null-terminated global string and return a ptr to its data. +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)); +} + +//===----------------------------------------------------------------------===// +// Patterns +//===----------------------------------------------------------------------===// + +struct ControllerCreateLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(ControllerCreateOp op, OpAdaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + Value lib = globalStr(rewriter, op.getLoc(), mod, "transport_backend_", op.getBackendLib()); + Value cfg = globalStr(rewriter, op.getLoc(), mod, "transport_config_", op.getConfig()); + Value s = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__controller_create", + {ptrTy(ctx), ptrTy(ctx)}, ptrTy(ctx), {lib, cfg}); + rewriter.replaceOp(op, s); + return success(); + } +}; + +struct ConnectLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(ConnectOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + Value peer = globalStr(rewriter, op.getLoc(), mod, "transport_peer_", op.getPeer()); + Value port = constInt(rewriter, op.getLoc(), rewriter.getI16Type(), op.getOobPort()); + Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__connect", + {ptrTy(ctx), ptrTy(ctx), rewriter.getI16Type()}, rewriter.getI32Type(), + {adaptor.getSession(), peer, port}); + rewriter.replaceOp(op, r); + return success(); + } +}; + +struct ExchangeKeysLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(ExchangeKeysOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + Value peerBuf = LLVM::AllocaOp::create( + rewriter, op.getLoc(), ptrTy(ctx), rewriter.getI8Type(), + constInt(rewriter, op.getLoc(), rewriter.getI64Type(), kPeerRefBytes), /*alignment=*/8); + Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__exchange_keys", + {ptrTy(ctx), ptrTy(ctx)}, rewriter.getI32Type(), + {adaptor.getSession(), peerBuf}); + rewriter.replaceOp(op, {r, peerBuf}); + return success(); + } +}; + +struct EstablishChannelLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(EstablishChannelOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + Value dp = constInt(rewriter, op.getLoc(), rewriter.getI32Type(), op.getDataPath()); + Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__establish_channel", + {ptrTy(ctx), rewriter.getI32Type(), ptrTy(ctx)}, rewriter.getI32Type(), + {adaptor.getSession(), dp, adaptor.getPeer()}); + rewriter.replaceOp(op, r); + return success(); + } +}; + +struct CommitWorkItemLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(CommitWorkItemOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override + { + auto *ctx = op.getContext(); + ModuleOp mod = moduleOf(op); + Value idx = constInt(rewriter, op.getLoc(), rewriter.getI32Type(), op.getWorkItemIdx()); + Value inB = constInt(rewriter, op.getLoc(), rewriter.getI64Type(), op.getInBytes()); + Value outB = constInt(rewriter, op.getLoc(), rewriter.getI64Type(), op.getOutBytes()); + Value r = emitCall( + rewriter, op.getLoc(), mod, "__catalyst__transport__commit_work_item", + {ptrTy(ctx), rewriter.getI32Type(), rewriter.getI64Type(), rewriter.getI64Type()}, + rewriter.getI32Type(), {adaptor.getSession(), idx, inB, outB}); + 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; +}; + +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); + // slot = data_slot(s); store payload -> slot; kick(s, idx) + Value slot = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__data_slot", + {ptrTy(ctx)}, ptrTy(ctx), {adaptor.getSession()}); + LLVM::StoreOp::create(rewriter, op.getLoc(), adaptor.getPayload(), slot); + Value idx = constInt(rewriter, op.getLoc(), rewriter.getI32Type(), op.getWorkItemIdx()); + Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__kick", + {ptrTy(ctx), rewriter.getI32Type()}, rewriter.getI32Type(), + {adaptor.getSession(), idx}); + rewriter.replaceOp(op, r); + 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); + Value one = constInt(rewriter, op.getLoc(), rewriter.getI64Type(), 1); + Value buf = LLVM::AllocaOp::create(rewriter, op.getLoc(), ptrTy(ctx), rewriter.getI64Type(), + one, /*alignment=*/8); + Value bytes = constInt(rewriter, op.getLoc(), rewriter.getI64Type(), op.getBytes()); + emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__collect", + {ptrTy(ctx), ptrTy(ctx), rewriter.getI64Type()}, rewriter.getI32Type(), + {adaptor.getSession(), buf, bytes}); + Value loaded = LLVM::LoadOp::create(rewriter, op.getLoc(), rewriter.getI64Type(), buf); + rewriter.replaceOp(op, loaded); + 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())}, rewriter.getI64Type(), {adaptor.getSession()}); + rewriter.replaceOp(op, r); + 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](PeerType) -> Type { return LLVM::LLVMPointerType::get(ctx); }); + + 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__close"); + 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..d9cdb1fd06 --- /dev/null +++ b/mlir/test/Transport/ConvertTransportToLLVM.mlir @@ -0,0 +1,57 @@ +// 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__controller_create(!llvm.ptr, !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, !llvm.ptr) -> i32 +// CHECK-DAG: llvm.func @__catalyst__transport__establish_channel(!llvm.ptr, i32, !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) + +// CHECK-LABEL: func.func @controller_roundtrip +func.func @controller_roundtrip() -> i64 { + // CHECK: %[[S:.*]] = llvm.call @__catalyst__transport__controller_create + %s = transport.controller_create {backend_lib = "libtransport_backend.so", config = "key=value"} -> !transport.session + // CHECK: llvm.call @__catalyst__transport__connect(%[[S]] + %c = transport.connect %s {peer = "127.0.0.1", oob_port = 18560 : i16} : (!transport.session) -> i32 + // CHECK: %[[PEER:.*]] = llvm.alloca + // CHECK: llvm.call @__catalyst__transport__exchange_keys(%[[S]], %[[PEER]]) + %cs, %peer = transport.exchange_keys %s : !transport.session -> !transport.peer + // CHECK: llvm.call @__catalyst__transport__establish_channel(%[[S]], {{.*}}, %[[PEER]]) + %e = transport.establish_channel %s, %peer {data_path = 0 : i32} : !transport.session, !transport.peer + // CHECK: llvm.call @__catalyst__transport__commit_work_item(%[[S]] + %w = 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 + %payload = arith.constant 81985529216486895 : i64 + // CHECK: %[[SLOT:.*]] = llvm.call @__catalyst__transport__data_slot(%[[S]]) + // CHECK: llvm.store %{{.*}}, %[[SLOT]] + // CHECK: llvm.call @__catalyst__transport__kick(%[[S]] + %k = transport.kick %s, %payload {work_item_idx = 0 : i32} : !transport.session, i64 + // CHECK: llvm.call @__catalyst__transport__collect(%[[S]] + // CHECK: %[[RESULT:.*]] = llvm.load + %result = transport.collect %s {bytes = 8 : i64} : !transport.session -> i64 + // 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 %result : i64 +} From 9cd39b52556d4721a1be0d926d0f3a4974507b14 Mon Sep 17 00:00:00 2001 From: Mehrdad Malekmohammadi Date: Wed, 22 Jul 2026 11:15:03 -0400 Subject: [PATCH 2/7] update --- mlir/include/Transport/IR/CMakeLists.txt | 5 + mlir/include/Transport/IR/TransportDialect.h | 1 + mlir/include/Transport/IR/TransportDialect.td | 63 +++++++++--- mlir/include/Transport/IR/TransportOps.td | 97 +++++++++++++------ mlir/lib/Transport/IR/CMakeLists.txt | 1 + mlir/lib/Transport/IR/TransportDialect.cpp | 1 + 6 files changed, 126 insertions(+), 42 deletions(-) diff --git a/mlir/include/Transport/IR/CMakeLists.txt b/mlir/include/Transport/IR/CMakeLists.txt index d1bc88dbed..d2b9081dc0 100644 --- a/mlir/include/Transport/IR/CMakeLists.txt +++ b/mlir/include/Transport/IR/CMakeLists.txt @@ -1,3 +1,8 @@ add_mlir_dialect(TransportOps transport) add_mlir_doc(TransportDialect TransportDialect Transport/ -gen-dialect-doc) add_mlir_doc(TransportOps TransportOps Transport/ -gen-op-doc) + +set(LLVM_TARGET_DEFINITIONS TransportOps.td) +mlir_tablegen(TransportEnums.h.inc -gen-enum-decls) +mlir_tablegen(TransportEnums.cpp.inc -gen-enum-defs) +add_public_tablegen_target(MLIRTransportEnumsIncGen) diff --git a/mlir/include/Transport/IR/TransportDialect.h b/mlir/include/Transport/IR/TransportDialect.h index ebfdd9ba47..57ab9b3dcb 100644 --- a/mlir/include/Transport/IR/TransportDialect.h +++ b/mlir/include/Transport/IR/TransportDialect.h @@ -19,6 +19,7 @@ #include "mlir/IR/OpDefinition.h" #include "Transport/IR/TransportOpsDialect.h.inc" +#include "Transport/IR/TransportEnums.h.inc" #define GET_TYPEDEF_CLASSES #include "Transport/IR/TransportOpsTypes.h.inc" diff --git a/mlir/include/Transport/IR/TransportDialect.td b/mlir/include/Transport/IR/TransportDialect.td index 78ed32fb70..118aae1734 100644 --- a/mlir/include/Transport/IR/TransportDialect.td +++ b/mlir/include/Transport/IR/TransportDialect.td @@ -18,20 +18,20 @@ include "mlir/IR/OpBase.td" include "mlir/IR/DialectBase.td" include "mlir/IR/AttrTypeBase.td" +include "mlir/IR/EnumAttr.td" +include "mlir/Interfaces/SideEffectInterfaces.td" //===----------------------------------------------------------------------===// // Transport dialect definition. //===----------------------------------------------------------------------===// def Transport_Dialect : Dialect { - let summary = "Runtime transport-layer ops that lower to the __catalyst__transport__* CAPI."; + let summary = "Typed ops for setting up and driving a transport session."; let description = [{ - A thin, typed MLIR representation of the Catalyst runtime transport CAPI - (runtime/include/TransportCAPI.h). One op per controller-side CAPI entry - point; `convert-transport-to-llvm` lowers each to an `llvm.call` on the - matching `__catalyst__transport__*` symbol. The coprocessor role is an - out-of-tree backend loaded by the runtime, so only the controller side is - emitted here. + The transport dialect models a connection-oriented data-movement session + between two endpoints: creating a session, bringing up the connection, + exchanging memory handles, establishing a data path, and running rounds + of request/reply traffic until teardown. }]; let name = "transport"; @@ -41,7 +41,26 @@ def Transport_Dialect : Dialect { } //===----------------------------------------------------------------------===// -// Transport dialect types. +// Enums. +//===----------------------------------------------------------------------===// + +def Transport_Role : I32EnumAttr<"Role", "transport session role", [ + I32EnumAttrCase<"Controller", 0, "controller">, + I32EnumAttrCase<"Coprocessor", 1, "coprocessor"> + ]> { + let cppNamespace = "::catalyst::transport"; +} + +def Transport_DataPath : I32EnumAttr<"DataPath", "transport data-movement path", [ + I32EnumAttrCase<"CpuVerbs", 0, "cpu_verbs">, + I32EnumAttrCase<"GpuEngine", 1, "gpu_engine">, + I32EnumAttrCase<"Other", 2, "other"> + ]> { + let cppNamespace = "::catalyst::transport"; +} + +//===----------------------------------------------------------------------===// +// Types. //===----------------------------------------------------------------------===// class Transport_Type traits = []> @@ -49,19 +68,37 @@ class Transport_Type traits = []> let mnemonic = typeMnemonic; } +// Opaque session handle, parameterized by role. Lowers to !llvm.ptr; the role is +// compile-time only and drives op verification + the create factory selection. def Transport_SessionType : Transport_Type<"Session", "session"> { - let summary = "Opaque transport controller session handle (CatalystTransportSession*)."; + let summary = "Opaque transport session handle (CatalystTransportSession*), tagged with its role."; + let parameters = (ins EnumParameter:$role); + let assemblyFormat = "`<` $role `>`"; } -def Transport_PeerType : Transport_Type<"Peer", "peer"> { - let summary = "Opaque peer-region descriptor handle (CatalystTransportPeerRef*)."; +def Transport_TokenType : Transport_Type<"Token", "token"> { + let summary = "Handle to an in-flight async transport step, awaited with transport.barrier."; } +// Role-constrained session-type constraints for role-specific ops. +def Transport_ControllerSession : Type< + CPred<"::llvm::isa<::catalyst::transport::SessionType>($_self) && " + "::llvm::cast<::catalyst::transport::SessionType>($_self).getRole() == " + "::catalyst::transport::Role::Controller">, + "controller transport session">; + +def Transport_CoprocessorSession : Type< + CPred<"::llvm::isa<::catalyst::transport::SessionType>($_self) && " + "::llvm::cast<::catalyst::transport::SessionType>($_self).getRole() == " + "::catalyst::transport::Role::Coprocessor">, + "coprocessor transport session">; + //===----------------------------------------------------------------------===// -// Transport operation base. +// Operation base. //===----------------------------------------------------------------------===// +// All transport ops perform side-effecting I/O class Transport_Op traits = []> : - Op; + Op])>; #endif // TRANSPORT_DIALECT diff --git a/mlir/include/Transport/IR/TransportOps.td b/mlir/include/Transport/IR/TransportOps.td index ccd7b2308e..cc5406f8be 100644 --- a/mlir/include/Transport/IR/TransportOps.td +++ b/mlir/include/Transport/IR/TransportOps.td @@ -21,61 +21,104 @@ include "mlir/Interfaces/SideEffectInterfaces.td" include "Transport/IR/TransportDialect.td" //===----------------------------------------------------------------------===// -// Bring-up ops +// Session creation //===----------------------------------------------------------------------===// -def Transport_ControllerCreateOp : Transport_Op<"controller_create"> { - let summary = "Create a controller session from a named backend plugin .so."; +def Transport_CreateOp : Transport_Op<"create"> { + let summary = "Create a session from a backend plugin .so with a certain role"; + let description = [{ + Loads the backend `.so` and builds a session. The result type's role + (`!transport.session`) selects which factory the + runtime looks up and constrains the role-specific ops downstream. + }]; let arguments = (ins StrAttr:$backend_lib, StrAttr:$config); let results = (outs Transport_SessionType:$session); let assemblyFormat = "attr-dict `->` type($session)"; } +//===----------------------------------------------------------------------===// +// Bring-up ops +//===----------------------------------------------------------------------===// + def Transport_ConnectOp : Transport_Op<"connect"> { - let summary = "Bring up the connection to the peer."; + let summary = "Bring up the connection to the peer (blocking)."; let arguments = (ins Transport_SessionType:$session, StrAttr:$peer, I16Attr:$oob_port); - let results = (outs I32:$status); - let assemblyFormat = "$session attr-dict `:` functional-type($session, $status)"; + let assemblyFormat = "$session attr-dict `:` type($session)"; +} + +def Transport_ConnectAsyncOp : Transport_Op<"connect_async"> { + let summary = "connect() on a worker; await with transport.barrier."; + let arguments = (ins Transport_SessionType:$session, StrAttr:$peer, I16Attr:$oob_port); + let results = (outs Transport_TokenType:$token); + let assemblyFormat = "$session attr-dict `:` type($session) `->` type($token)"; } def Transport_ExchangeKeysOp : Transport_Op<"exchange_keys"> { - let summary = "Exchange local and peer region handles; yields the peer descriptor."; + let summary = "Exchange region handles with the peer (blocking); result kept in the session."; + let arguments = (ins Transport_SessionType:$session); + let assemblyFormat = "$session attr-dict `:` type($session)"; +} + +def Transport_ExchangeKeysAsyncOp : Transport_Op<"exchange_keys_async"> { + let summary = "exchange_keys() on a worker; await with transport.barrier."; let arguments = (ins Transport_SessionType:$session); - let results = (outs I32:$status, Transport_PeerType:$peer); - let assemblyFormat = "$session attr-dict `:` type($session) `->` type($peer)"; + let results = (outs Transport_TokenType:$token); + let assemblyFormat = "$session attr-dict `:` type($session) `->` type($token)"; +} + +def Transport_BarrierOp : Transport_Op<"barrier"> { + let summary = "Await an async step (connect_async / exchange_keys_async)."; + let arguments = (ins Transport_TokenType:$token); + let assemblyFormat = "$token attr-dict `:` type($token)"; } def Transport_EstablishChannelOp : Transport_Op<"establish_channel"> { - let summary = "Program the data-movement channel from the local + peer regions."; - let arguments = (ins Transport_SessionType:$session, Transport_PeerType:$peer, - I32Attr:$data_path); - let results = (outs I32:$status); - let assemblyFormat = "$session `,` $peer attr-dict `:` type($session) `,` type($peer)"; + let summary = "Set up the data channel used to transfer payloads each round."; + let description = [{ + Arms the data channel for the given `data_path`, using this side's + registered memory region together with the peer's region that + `exchange_keys` learned earlier (both stored in the session). After this + the channel is ready for `kick`/`collect`. + }]; + let arguments = (ins Transport_SessionType:$session, Transport_DataPath:$data_path); + let assemblyFormat = "$session $data_path attr-dict `:` type($session)"; } +//===----------------------------------------------------------------------===// +// Controller-only ops +//===----------------------------------------------------------------------===// + def Transport_CommitWorkItemOp : Transport_Op<"commit_work_item"> { let summary = "Build a work item (I/O sizes) in a slot before kicking rounds."; - let arguments = (ins Transport_SessionType:$session, I32Attr:$work_item_idx, + let arguments = (ins Transport_ControllerSession:$session, I32Attr:$work_item_idx, I64Attr:$in_bytes, I64Attr:$out_bytes); - let results = (outs I32:$status); let assemblyFormat = "$session attr-dict `:` type($session)"; } -def Transport_StartOp : Transport_Op<"start"> { - let summary = "Start the session (non-blocking; runs until stop())."; - let arguments = (ins Transport_SessionType:$session); +def Transport_KickOp : Transport_Op<"kick"> { + let summary = "Write the payload into the outbound slot and fire one round."; + let arguments = (ins Transport_ControllerSession:$session, I64:$payload, I32Attr:$work_item_idx); + let assemblyFormat = "$session `,` $payload attr-dict `:` type($session) `,` type($payload)"; +} + +//===----------------------------------------------------------------------===// +// Coprocessor-only: bind the coprocessor function (requires a coprocessor session) +//===----------------------------------------------------------------------===// + +def Transport_SetCoprocessorFnOp : Transport_Op<"set_coprocessor_fn"> { + let summary = "Bind the built-in coprocessor function (echo / on-device kernel)."; + let arguments = (ins Transport_CoprocessorSession:$session); let assemblyFormat = "$session attr-dict `:` type($session)"; } //===----------------------------------------------------------------------===// -// Per-round ops +// Run / collect / teardown //===----------------------------------------------------------------------===// -def Transport_KickOp : Transport_Op<"kick"> { - let summary = "Write the payload into the outbound slot and fire one round."; - let arguments = (ins Transport_SessionType:$session, I64:$payload, I32Attr:$work_item_idx); - let results = (outs I32:$status); - let assemblyFormat = "$session `,` $payload attr-dict `:` type($session) `,` type($payload)"; +def Transport_StartOp : Transport_Op<"start"> { + let summary = "Start the session (non-blocking; runs until stop())."; + let arguments = (ins Transport_SessionType:$session); + let assemblyFormat = "$session attr-dict `:` type($session)"; } def Transport_CollectOp : Transport_Op<"collect"> { @@ -92,10 +135,6 @@ def Transport_LastRttNsOp : Transport_Op<"last_rtt_ns"> { let assemblyFormat = "$session attr-dict `:` type($session) `->` type($rtt_ns)"; } -//===----------------------------------------------------------------------===// -// Teardown ops -//===----------------------------------------------------------------------===// - def Transport_StopOp : Transport_Op<"stop"> { let summary = "Stop the session. Idempotent."; let arguments = (ins Transport_SessionType:$session); diff --git a/mlir/lib/Transport/IR/CMakeLists.txt b/mlir/lib/Transport/IR/CMakeLists.txt index 9554c24fe3..1dbeed05e2 100644 --- a/mlir/lib/Transport/IR/CMakeLists.txt +++ b/mlir/lib/Transport/IR/CMakeLists.txt @@ -7,6 +7,7 @@ add_mlir_library(MLIRTransport DEPENDS MLIRTransportOpsIncGen + MLIRTransportEnumsIncGen LINK_LIBS PRIVATE MLIRLLVMDialect diff --git a/mlir/lib/Transport/IR/TransportDialect.cpp b/mlir/lib/Transport/IR/TransportDialect.cpp index 96bb105661..aea2c90196 100644 --- a/mlir/lib/Transport/IR/TransportDialect.cpp +++ b/mlir/lib/Transport/IR/TransportDialect.cpp @@ -28,6 +28,7 @@ using namespace catalyst::transport; //===----------------------------------------------------------------------===// #include "Transport/IR/TransportOpsDialect.cpp.inc" +#include "Transport/IR/TransportEnums.cpp.inc" //===----------------------------------------------------------------------===// // Transport type definitions. From 7b3ff71cd922d0d03c26703979726fe62a44f9c5 Mon Sep 17 00:00:00 2001 From: Mehrdad Malekmohammadi Date: Wed, 22 Jul 2026 11:29:19 -0400 Subject: [PATCH 3/7] update --- mlir/include/RegisterAllPasses.h | 2 + mlir/include/Transport/CMakeLists.txt | 1 + mlir/lib/Driver/CMakeLists.txt | 1 + mlir/lib/Transport/CMakeLists.txt | 1 + mlir/lib/Transport/Transforms/CMakeLists.txt | 1 + .../Transport/Transforms/TransportToLLVM.cpp | 193 +++++++++++------- mlir/tools/quantum-opt/CMakeLists.txt | 1 + 7 files changed, 122 insertions(+), 78 deletions(-) diff --git a/mlir/include/RegisterAllPasses.h b/mlir/include/RegisterAllPasses.h index 791fe4a784..e92a0850d4 100644 --- a/mlir/include/RegisterAllPasses.h +++ b/mlir/include/RegisterAllPasses.h @@ -25,6 +25,7 @@ #include "QecPhysical/Transforms/Passes.h" #include "Quantum/Transforms/Passes.h" #include "RTIO/Transforms/Passes.h" +#include "Transport/Transforms/Passes.h" #include "Test/Transforms/Passes.h" #include "hlo-extensions/Transforms/Passes.h" @@ -44,6 +45,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/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 index 17ee33278d..efd3bd5462 100644 --- a/mlir/lib/Transport/Transforms/CMakeLists.txt +++ b/mlir/lib/Transport/Transforms/CMakeLists.txt @@ -14,6 +14,7 @@ set(LIBS set(DEPENDS MLIRTransportPassIncGen + MLIRTransportEnumsIncGen ) add_mlir_library(${LIBRARY_NAME} STATIC ${SRC} LINK_LIBS PRIVATE ${LIBS} DEPENDS ${DEPENDS}) diff --git a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp index e96f68d1f7..4fbc1ae56c 100644 --- a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp +++ b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp @@ -13,15 +13,15 @@ // limitations under the License. // Lower the `transport` dialect to `llvm.call`s on the __catalyst__transport__* -// CAPI (runtime/include/TransportCAPI.h). Controller-side only. +// 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 "llvm/ADT/Twine.h" #include "Transport/IR/TransportOps.h" #include "Transport/Transforms/Passes.h" @@ -37,15 +37,12 @@ namespace transport { namespace { -// sizeof(CatalystTransportPeerRef) = {u32, u64, u64}; over-allocate for alignment. -constexpr int64_t kPeerRefBytes = 32; - 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(); } -// Declare-or-reuse a CAPI function and emit a call to it. A null resultTy means -// the function returns void. Value emitCall(ConversionPatternRewriter &rewriter, Location loc, ModuleOp mod, StringRef name, ArrayRef paramTys, Type resultTy, ValueRange args) { @@ -56,7 +53,6 @@ Value emitCall(ConversionPatternRewriter &rewriter, Location loc, ModuleOp mod, return call.getNumResults() ? call.getResult() : Value(); } -// Materialize a null-terminated global string and return a ptr to its data. Value globalStr(ConversionPatternRewriter &rewriter, Location loc, ModuleOp mod, StringRef prefix, StringRef value) { @@ -75,107 +71,132 @@ Value constInt(ConversionPatternRewriter &rewriter, Location loc, Type ty, int64 // Patterns //===----------------------------------------------------------------------===// -struct ControllerCreateLowering : public OpConversionPattern { +struct CreateLowering : public OpConversionPattern { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite(ControllerCreateOp op, OpAdaptor, + 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 s = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__controller_create", - {ptrTy(ctx), ptrTy(ctx)}, ptrTy(ctx), {lib, cfg}); + 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), {lib, cfg, role}); rewriter.replaceOp(op, s); return success(); } }; -struct ConnectLowering : public OpConversionPattern { - using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite(ConnectOp op, OpAdaptor adaptor, +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 = moduleOf(op); + ModuleOp mod = op->template getParentOfType(); Value peer = globalStr(rewriter, op.getLoc(), mod, "transport_peer_", op.getPeer()); - Value port = constInt(rewriter, op.getLoc(), rewriter.getI16Type(), op.getOobPort()); - Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__connect", - {ptrTy(ctx), ptrTy(ctx), rewriter.getI16Type()}, rewriter.getI32Type(), - {adaptor.getSession(), peer, port}); - rewriter.replaceOp(op, r); + 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; -struct ExchangeKeysLowering : public OpConversionPattern { - using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite(ExchangeKeysOp op, OpAdaptor adaptor, +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 = moduleOf(op); - Value peerBuf = LLVM::AllocaOp::create( - rewriter, op.getLoc(), ptrTy(ctx), rewriter.getI8Type(), - constInt(rewriter, op.getLoc(), rewriter.getI64Type(), kPeerRefBytes), /*alignment=*/8); - Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__exchange_keys", - {ptrTy(ctx), ptrTy(ctx)}, rewriter.getI32Type(), - {adaptor.getSession(), peerBuf}); - rewriter.replaceOp(op, {r, peerBuf}); + 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 EstablishChannelLowering : public OpConversionPattern { +struct BarrierLowering : public OpConversionPattern { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite(EstablishChannelOp op, OpAdaptor adaptor, + LogicalResult matchAndRewrite(BarrierOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { auto *ctx = op.getContext(); - ModuleOp mod = moduleOf(op); - Value dp = constInt(rewriter, op.getLoc(), rewriter.getI32Type(), op.getDataPath()); - Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__establish_channel", - {ptrTy(ctx), rewriter.getI32Type(), ptrTy(ctx)}, rewriter.getI32Type(), - {adaptor.getSession(), dp, adaptor.getPeer()}); - rewriter.replaceOp(op, r); + emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__barrier", + {i64Ty(ctx)}, i32Ty(ctx), {adaptor.getToken()}); + rewriter.eraseOp(op); return success(); } }; -struct CommitWorkItemLowering : public OpConversionPattern { +struct EstablishChannelLowering : public OpConversionPattern { using OpConversionPattern::OpConversionPattern; - LogicalResult matchAndRewrite(CommitWorkItemOp op, OpAdaptor adaptor, + LogicalResult matchAndRewrite(EstablishChannelOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { auto *ctx = op.getContext(); - ModuleOp mod = moduleOf(op); - Value idx = constInt(rewriter, op.getLoc(), rewriter.getI32Type(), op.getWorkItemIdx()); - Value inB = constInt(rewriter, op.getLoc(), rewriter.getI64Type(), op.getInBytes()); - Value outB = constInt(rewriter, op.getLoc(), rewriter.getI64Type(), op.getOutBytes()); - Value r = emitCall( - rewriter, op.getLoc(), mod, "__catalyst__transport__commit_work_item", - {ptrTy(ctx), rewriter.getI32Type(), rewriter.getI64Type(), rewriter.getI64Type()}, - rewriter.getI32Type(), {adaptor.getSession(), idx, inB, outB}); - rewriter.replaceOp(op, r); + Value dp = + constInt(rewriter, op.getLoc(), i32Ty(ctx), static_cast(op.getDataPath())); + emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__establish_channel", + {ptrTy(ctx), i32Ty(ctx)}, i32Ty(ctx), {adaptor.getSession(), dp}); + rewriter.eraseOp(op); 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) +struct SetCoprocessorFnLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(SetCoprocessorFnOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__set_coprocessor_fn", + {ptrTy(op.getContext())}, i32Ty(op.getContext()), {adaptor.getSession()}); + rewriter.eraseOp(op); + return success(); } - LogicalResult matchAndRewrite(OpT op, typename OpT::Adaptor adaptor, +}; + +struct CommitWorkItemLowering : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + LogicalResult matchAndRewrite(CommitWorkItemOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { - emitCall(rewriter, op.getLoc(), op->template getParentOfType(), symbol, - {ptrTy(op.getContext())}, Type(), {adaptor.getSession()}); + 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(); } - std::string symbol; }; struct KickLowering : public OpConversionPattern { @@ -185,15 +206,13 @@ struct KickLowering : public OpConversionPattern { { auto *ctx = op.getContext(); ModuleOp mod = moduleOf(op); - // slot = data_slot(s); store payload -> slot; kick(s, idx) Value slot = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__data_slot", {ptrTy(ctx)}, ptrTy(ctx), {adaptor.getSession()}); LLVM::StoreOp::create(rewriter, op.getLoc(), adaptor.getPayload(), slot); - Value idx = constInt(rewriter, op.getLoc(), rewriter.getI32Type(), op.getWorkItemIdx()); - Value r = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__kick", - {ptrTy(ctx), rewriter.getI32Type()}, rewriter.getI32Type(), - {adaptor.getSession(), idx}); - rewriter.replaceOp(op, r); + 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(); } }; @@ -205,14 +224,14 @@ struct CollectLowering : public OpConversionPattern { { auto *ctx = op.getContext(); ModuleOp mod = moduleOf(op); - Value one = constInt(rewriter, op.getLoc(), rewriter.getI64Type(), 1); - Value buf = LLVM::AllocaOp::create(rewriter, op.getLoc(), ptrTy(ctx), rewriter.getI64Type(), - one, /*alignment=*/8); - Value bytes = constInt(rewriter, op.getLoc(), rewriter.getI64Type(), op.getBytes()); + Value one = constInt(rewriter, op.getLoc(), i64Ty(ctx), 1); + Value buf = LLVM::AllocaOp::create(rewriter, op.getLoc(), ptrTy(ctx), i64Ty(ctx), one, + /*alignment=*/8); + Value bytes = constInt(rewriter, op.getLoc(), i64Ty(ctx), op.getBytes()); emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__collect", - {ptrTy(ctx), ptrTy(ctx), rewriter.getI64Type()}, rewriter.getI32Type(), + {ptrTy(ctx), ptrTy(ctx), i64Ty(ctx)}, i32Ty(ctx), {adaptor.getSession(), buf, bytes}); - Value loaded = LLVM::LoadOp::create(rewriter, op.getLoc(), rewriter.getI64Type(), buf); + Value loaded = LLVM::LoadOp::create(rewriter, op.getLoc(), i64Ty(ctx), buf); rewriter.replaceOp(op, loaded); return success(); } @@ -223,14 +242,31 @@ struct LastRttLowering : public 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())}, rewriter.getI64Type(), {adaptor.getSession()}); + 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; +}; + } // namespace struct ConvertTransportToLLVMPass @@ -242,12 +278,13 @@ struct ConvertTransportToLLVMPass MLIRContext *ctx = &getContext(); LLVMTypeConverter tc(ctx); tc.addConversion([ctx](SessionType) -> Type { return LLVM::LLVMPointerType::get(ctx); }); - tc.addConversion([ctx](PeerType) -> 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); patterns.add>(tc, ctx, "__catalyst__transport__start"); patterns.add>(tc, ctx, "__catalyst__transport__stop"); patterns.add>(tc, ctx, "__catalyst__transport__close"); 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 From a5150205e70ef88e53283c1d818ceb8a320b9c82 Mon Sep 17 00:00:00 2001 From: Mehrdad Malekmohammadi Date: Wed, 22 Jul 2026 12:00:44 -0400 Subject: [PATCH 4/7] format --- mlir/include/RegisterAllPasses.h | 3 ++- mlir/lib/Transport/Transforms/TransportToLLVM.cpp | 12 ++++++------ 2 files changed, 8 insertions(+), 7 deletions(-) diff --git a/mlir/include/RegisterAllPasses.h b/mlir/include/RegisterAllPasses.h index e92a0850d4..a5195536c5 100644 --- a/mlir/include/RegisterAllPasses.h +++ b/mlir/include/RegisterAllPasses.h @@ -25,10 +25,11 @@ #include "QecPhysical/Transforms/Passes.h" #include "Quantum/Transforms/Passes.h" #include "RTIO/Transforms/Passes.h" -#include "Transport/Transforms/Passes.h" #include "Test/Transforms/Passes.h" #include "hlo-extensions/Transforms/Passes.h" +#include "Transport/Transforms/Passes.h" + namespace catalyst { inline void registerAllPasses() diff --git a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp index 4fbc1ae56c..0a3b7cab90 100644 --- a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp +++ b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp @@ -15,13 +15,13 @@ // 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 "llvm/ADT/Twine.h" #include "Transport/IR/TransportOps.h" #include "Transport/Transforms/Passes.h" @@ -242,9 +242,9 @@ struct LastRttLowering : public 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()}); + 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(); } @@ -283,8 +283,8 @@ struct ConvertTransportToLLVMPass RewritePatternSet patterns(ctx); patterns.add(tc, ctx); + SetCoprocessorFnLowering, CommitWorkItemLowering, KickLowering, + CollectLowering, LastRttLowering>(tc, ctx); patterns.add>(tc, ctx, "__catalyst__transport__start"); patterns.add>(tc, ctx, "__catalyst__transport__stop"); patterns.add>(tc, ctx, "__catalyst__transport__close"); From 1c2ea7a8271a21f0f69d95e7baf22f1499c668e0 Mon Sep 17 00:00:00 2001 From: Mehrdad Malekmohammadi Date: Fri, 24 Jul 2026 10:47:27 -0400 Subject: [PATCH 5/7] update get_session --- .../Transport/Transforms/TransportToLLVM.cpp | 85 ++++++++++++++----- .../Transport/ConvertTransportToLLVM.mlir | 78 ++++++++++++----- 2 files changed, 119 insertions(+), 44 deletions(-) diff --git a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp index 0a3b7cab90..9cc7bf609e 100644 --- a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp +++ b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp @@ -15,13 +15,13 @@ // 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 "llvm/ADT/Twine.h" #include "Transport/IR/TransportOps.h" #include "Transport/Transforms/Passes.h" @@ -67,6 +67,22 @@ Value constInt(ConversionPatternRewriter &rewriter, Location loc, Type ty, int64 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 //===----------------------------------------------------------------------===// @@ -81,10 +97,12 @@ struct CreateLowering : public OpConversionPattern { 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), {lib, cfg, role}); + {ptrTy(ctx), ptrTy(ctx), i32Ty(ctx), ptrTy(ctx)}, ptrTy(ctx), + {lib, cfg, role, key}); rewriter.replaceOp(op, s); return success(); } @@ -161,10 +179,10 @@ struct EstablishChannelLowering : public OpConversionPattern ConversionPatternRewriter &rewriter) const override { auto *ctx = op.getContext(); - Value dp = - constInt(rewriter, op.getLoc(), i32Ty(ctx), static_cast(op.getDataPath())); + 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), i32Ty(ctx)}, i32Ty(ctx), {adaptor.getSession(), dp}); + {ptrTy(ctx), ptrTy(ctx)}, i32Ty(ctx), {adaptor.getSession(), dp}); rewriter.eraseOp(op); return success(); } @@ -175,8 +193,11 @@ struct SetCoprocessorFnLowering : public OpConversionPattern LogicalResult matchAndRewrite(SetCoprocessorFnOp op, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override { - emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__set_coprocessor_fn", - {ptrTy(op.getContext())}, i32Ty(op.getContext()), {adaptor.getSession()}); + 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(); } @@ -206,9 +227,14 @@ struct KickLowering : public OpConversionPattern { { 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::StoreOp::create(rewriter, op.getLoc(), adaptor.getPayload(), slot); + 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}); @@ -224,15 +250,14 @@ struct CollectLowering : public OpConversionPattern { { auto *ctx = op.getContext(); ModuleOp mod = moduleOf(op); - Value one = constInt(rewriter, op.getLoc(), i64Ty(ctx), 1); - Value buf = LLVM::AllocaOp::create(rewriter, op.getLoc(), ptrTy(ctx), i64Ty(ctx), one, - /*alignment=*/8); - Value bytes = constInt(rewriter, op.getLoc(), i64Ty(ctx), op.getBytes()); + 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(), buf, bytes}); - Value loaded = LLVM::LoadOp::create(rewriter, op.getLoc(), i64Ty(ctx), buf); - rewriter.replaceOp(op, loaded); + {adaptor.getSession(), dstPtr, bytes}); + rewriter.eraseOp(op); return success(); } }; @@ -242,9 +267,9 @@ struct LastRttLowering : public 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()}); + 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(); } @@ -267,6 +292,25 @@ template struct VoidSessionLowering : public OpConversionPattern< 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 @@ -283,11 +327,10 @@ struct ConvertTransportToLLVMPass RewritePatternSet patterns(ctx); patterns.add(tc, ctx); + SetCoprocessorFnLowering, CommitWorkItemLowering, KickLowering, CollectLowering, + LastRttLowering, GetSessionLowering>(tc, ctx); patterns.add>(tc, ctx, "__catalyst__transport__start"); patterns.add>(tc, ctx, "__catalyst__transport__stop"); - patterns.add>(tc, ctx, "__catalyst__transport__close"); patterns.add>(tc, ctx, "__catalyst__transport__destroy"); ConversionTarget target(*ctx); diff --git a/mlir/test/Transport/ConvertTransportToLLVM.mlir b/mlir/test/Transport/ConvertTransportToLLVM.mlir index d9cdb1fd06..e2590bb88d 100644 --- a/mlir/test/Transport/ConvertTransportToLLVM.mlir +++ b/mlir/test/Transport/ConvertTransportToLLVM.mlir @@ -14,10 +14,10 @@ // RUN: quantum-opt %s --convert-transport-to-llvm --split-input-file | FileCheck %s -// CHECK-DAG: llvm.func @__catalyst__transport__controller_create(!llvm.ptr, !llvm.ptr) -> !llvm.ptr +// 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, !llvm.ptr) -> i32 -// CHECK-DAG: llvm.func @__catalyst__transport__establish_channel(!llvm.ptr, i32, !llvm.ptr) -> 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 @@ -26,32 +26,64 @@ // CHECK-DAG: llvm.func @__catalyst__transport__stop(!llvm.ptr) // CHECK-DAG: llvm.func @__catalyst__transport__destroy(!llvm.ptr) -// CHECK-LABEL: func.func @controller_roundtrip -func.func @controller_roundtrip() -> i64 { - // CHECK: %[[S:.*]] = llvm.call @__catalyst__transport__controller_create - %s = transport.controller_create {backend_lib = "libtransport_backend.so", config = "key=value"} -> !transport.session +// 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]] - %c = transport.connect %s {peer = "127.0.0.1", oob_port = 18560 : i16} : (!transport.session) -> i32 - // CHECK: %[[PEER:.*]] = llvm.alloca - // CHECK: llvm.call @__catalyst__transport__exchange_keys(%[[S]], %[[PEER]]) - %cs, %peer = transport.exchange_keys %s : !transport.session -> !transport.peer - // CHECK: llvm.call @__catalyst__transport__establish_channel(%[[S]], {{.*}}, %[[PEER]]) - %e = transport.establish_channel %s, %peer {data_path = 0 : i32} : !transport.session, !transport.peer + 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]] - %w = transport.commit_work_item %s {work_item_idx = 0 : i32, in_bytes = 8 : i64, out_bytes = 8 : i64} : !transport.session + 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 - %payload = arith.constant 81985529216486895 : i64 + transport.start %s : !transport.session // CHECK: %[[SLOT:.*]] = llvm.call @__catalyst__transport__data_slot(%[[S]]) - // CHECK: llvm.store %{{.*}}, %[[SLOT]] + // CHECK: "llvm.intr.memcpy"(%[[SLOT]] // CHECK: llvm.call @__catalyst__transport__kick(%[[S]] - %k = transport.kick %s, %payload {work_item_idx = 0 : i32} : !transport.session, i64 + transport.kick %s, %syndrome {work_item_idx = 0 : i32} : !transport.session, memref // CHECK: llvm.call @__catalyst__transport__collect(%[[S]] - // CHECK: %[[RESULT:.*]] = llvm.load - %result = transport.collect %s {bytes = 8 : i64} : !transport.session -> i64 + transport.collect %s, %correction : !transport.session, memref // CHECK: llvm.call @__catalyst__transport__stop(%[[S]]) - transport.stop %s : !transport.session + transport.stop %s : !transport.session // CHECK: llvm.call @__catalyst__transport__destroy(%[[S]]) - transport.destroy %s : !transport.session - return %result : i64 + 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 } From 88874665e1abb2454e0c977fe8373d56807a60c2 Mon Sep 17 00:00:00 2001 From: Mehrdad Malekmohammadi Date: Fri, 24 Jul 2026 10:48:01 -0400 Subject: [PATCH 6/7] add changelog --- doc/releases/changelog-dev.md | 4 ++++ 1 file changed, 4 insertions(+) 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) From 747c0ed20ae7966576b1d410b924189b7e3bfadf Mon Sep 17 00:00:00 2001 From: Mehrdad Malekmohammadi Date: Fri, 24 Jul 2026 12:10:02 -0400 Subject: [PATCH 7/7] format --- .../Transport/Transforms/TransportToLLVM.cpp | 19 +++++++++---------- 1 file changed, 9 insertions(+), 10 deletions(-) diff --git a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp index 9cc7bf609e..ce5e1698a2 100644 --- a/mlir/lib/Transport/Transforms/TransportToLLVM.cpp +++ b/mlir/lib/Transport/Transforms/TransportToLLVM.cpp @@ -15,13 +15,13 @@ // 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 "llvm/ADT/Twine.h" #include "Transport/IR/TransportOps.h" #include "Transport/Transforms/Passes.h" @@ -67,7 +67,6 @@ Value constInt(ConversionPatternRewriter &rewriter, Location loc, Type ty, int64 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, @@ -76,8 +75,7 @@ std::pair memrefPtrAndBytes(ConversionPatternRewriter &rewriter, L 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; + 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}; @@ -251,7 +249,8 @@ struct CollectLowering : public OpConversionPattern { auto *ctx = op.getContext(); ModuleOp mod = moduleOf(op); if (!op.getDest()) - return rewriter.notifyMatchFailure(op, "collect must be bufferized (dest-passing form)"); + 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", @@ -267,9 +266,9 @@ struct LastRttLowering : public 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()}); + 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(); } @@ -327,8 +326,8 @@ struct ConvertTransportToLLVMPass RewritePatternSet patterns(ctx); patterns.add(tc, ctx); + SetCoprocessorFnLowering, CommitWorkItemLowering, KickLowering, + CollectLowering, LastRttLowering, GetSessionLowering>(tc, ctx); patterns.add>(tc, ctx, "__catalyst__transport__start"); patterns.add>(tc, ctx, "__catalyst__transport__stop"); patterns.add>(tc, ctx, "__catalyst__transport__destroy");