diff --git a/docs/source/design/transfer-engine/ascend_direct_transport.md b/docs/source/design/transfer-engine/ascend_direct_transport.md index d96cfa90..13a14bb8 100644 --- a/docs/source/design/transfer-engine/ascend_direct_transport.md +++ b/docs/source/design/transfer-engine/ascend_direct_transport.md @@ -72,3 +72,5 @@ Complete command format is shown below: 10. **Async transfer**: The asynchronous transfer mode can be enabled by configuring the ASCEND_USE_ASYNC_TRANSFER environment variable. 12. **Fabric Memory mode**: On the A3, with the latest drivers and CANN installed, when using Mooncake store, the ASCEND_ENABLE_USE_FABRIC_MEM environment variable can be set to enable fabric memory transfer mode (which allows direct access remote HOST memory). + +13. **Auto Connect**: The auto connect feature can be enabled by configuring the `ASCEND_AUTO_CONNECT` environment variable. The default value is 0 (disabled). diff --git a/mooncake-transfer-engine/include/transport/ascend_transport/ascend_direct_transport/ascend_direct_transport.h b/mooncake-transfer-engine/include/transport/ascend_transport/ascend_direct_transport/ascend_direct_transport.h index 4ca2fc6b..79bb1100 100644 --- a/mooncake-transfer-engine/include/transport/ascend_transport/ascend_direct_transport/ascend_direct_transport.h +++ b/mooncake-transfer-engine/include/transport/ascend_transport/ascend_direct_transport/ascend_direct_transport.h @@ -138,6 +138,7 @@ class AscendDirectTransport : public Transport { std::string local_adxl_engine_name_{}; aclrtStream stream_{}; bool use_buffer_pool_{false}; + bool auto_connect_{false}; int32_t base_port_ = 20000; std::unordered_set need_update_metadata_segs_; bool use_short_connection_{false}; diff --git a/mooncake-transfer-engine/src/transport/ascend_transport/ascend_direct_transport/ascend_direct_transport.cpp b/mooncake-transfer-engine/src/transport/ascend_transport/ascend_direct_transport/ascend_direct_transport.cpp index 08db1e85..27d319de 100644 --- a/mooncake-transfer-engine/src/transport/ascend_transport/ascend_direct_transport/ascend_direct_transport.cpp +++ b/mooncake-transfer-engine/src/transport/ascend_transport/ascend_direct_transport/ascend_direct_transport.cpp @@ -44,6 +44,9 @@ constexpr size_t kAsyncTaskLimit = 100U; constexpr int32_t kPortRange = 100; constexpr int32_t kDefaultDisconnectTime = 1000; constexpr int32_t kMaxGenPortAttempts = 500; +constexpr const char *kAutoConnect = "AutoConnect"; +constexpr const char *kEnabled = "1"; +constexpr const char *kDisabled = "0"; } // namespace AscendDirectTransport::AscendDirectTransport() : running_(false) {} @@ -241,6 +244,15 @@ int AscendDirectTransport::InitAdxlEngine() { LOG(INFO) << "Use async transfer"; use_async_transfer_ = true; } + char *auto_connect = std::getenv("ASCEND_AUTO_CONNECT"); + if (auto_connect) { + auto auto_connect_opt = parseFromString(auto_connect); + if (auto_connect_opt.has_value()) { + auto_connect_ = (*auto_connect_opt == 1); + options[kAutoConnect] = auto_connect_ ? kEnabled : kDisabled; + LOG(INFO) << "Set AutoConnect to: " << auto_connect; + } + } // set default buffer pool options["adxl.BufferPool"] = "0:0"; use_buffer_pool_ = false; @@ -804,12 +816,14 @@ void AscendDirectTransport::processSliceList( void AscendDirectTransport::connectAndTransfer( const std::string &target_adxl_engine_name, adxl::TransferOp operation, const std::vector &slice_list, int32_t times) { - int ret = checkAndConnect(target_adxl_engine_name); - if (ret != 0) { - for (auto &slice : slice_list) { - slice->markFailed(); + if (!auto_connect_) { + int ret = checkAndConnect(target_adxl_engine_name); + if (ret != 0) { + for (auto &slice : slice_list) { + slice->markFailed(); + } + return; } - return; } auto start = std::chrono::steady_clock::now(); std::vector op_descs; @@ -1139,6 +1153,16 @@ int AscendDirectTransport::checkAndConnect( int AscendDirectTransport::disconnect( const std::string &target_adxl_engine_name, int32_t timeout_in_millis, bool force) { + if (auto_connect_) { + auto status = adxl_->Disconnect(target_adxl_engine_name.c_str(), + timeout_in_millis); + if (status != adxl::SUCCESS) { + LOG(ERROR) << "Failed to disconnect to: " << target_adxl_engine_name + << ", status: " << status; + return -1; + } + return 0; + } std::lock_guard lock(connection_mutex_); auto it = connected_segments_.find(target_adxl_engine_name); if (it == connected_segments_.end()) { diff --git a/mooncake-transfer-engine/tent/include/tent/transport/ascend/ascend_direct_transport.h b/mooncake-transfer-engine/tent/include/tent/transport/ascend/ascend_direct_transport.h index 6492e257..1c83ecf3 100644 --- a/mooncake-transfer-engine/tent/include/tent/transport/ascend/ascend_direct_transport.h +++ b/mooncake-transfer-engine/tent/include/tent/transport/ascend/ascend_direct_transport.h @@ -115,6 +115,7 @@ class AscendDirectTransport : public Transport { aclrtContext rt_context_{nullptr}; std::unique_ptr hixl_; bool use_buffer_pool_{false}; + bool auto_connect_{false}; uint64_t connect_timeout_; uint64_t transfer_timeout_; std::string local_hixl_name_{}; diff --git a/mooncake-transfer-engine/tent/src/transport/ascend/ascend_direct_transport.cpp b/mooncake-transfer-engine/tent/src/transport/ascend/ascend_direct_transport.cpp index aa8aeaac..004e1de9 100644 --- a/mooncake-transfer-engine/tent/src/transport/ascend/ascend_direct_transport.cpp +++ b/mooncake-transfer-engine/tent/src/transport/ascend/ascend_direct_transport.cpp @@ -208,6 +208,14 @@ Status AscendDirectTransport::initHixl(const std::shared_ptr &conf) { LOG(INFO) << "Set RdmaServiceLevel to:" << rdma_sl_env; } } + std::string auto_connect = + conf->get("transports/ascend_direct/auto_connect", ""); + auto auto_connect_opt = parseFromString(auto_connect); + if (auto_connect_opt.has_value()) { + auto_connect_ = (*auto_connect_opt == 1); + options["AutoConnect"] = auto_connect_ ? "1" : "0"; + LOG(INFO) << "Set AutoConnect to: " << auto_connect; + } std::string buffer_pool = conf->get("transports/ascend_direct/buffer_pool", ""); if (!buffer_pool.empty()) { @@ -349,9 +357,11 @@ void AscendDirectTransport::startTransfer(SegmentID target_id, << " us."; return; } else { - auto ret = checkAndConnect(remote_hixl); - if (!ret.ok()) { - return; + if (!auto_connect_) { + auto ret = checkAndConnect(remote_hixl); + if (!ret.ok()) { + return; + } } } auto op = (opcode == Request::WRITE) ? hixl::WRITE : hixl::READ; @@ -453,6 +463,14 @@ Status AscendDirectTransport::getTransferStatus(SubBatchRef batch, int task_id, void AscendDirectTransport::disconnect(const std::string &remote_hixl, int32_t timeout_in_millis) { + if (auto_connect_) { + auto status = hixl_->Disconnect(remote_hixl.c_str(), timeout_in_millis); + if (status != hixl::SUCCESS) { + LOG(ERROR) << "Failed to disconnect to: " << remote_hixl + << ", status: " << status; + } + return; + } std::lock_guard lock(connection_mutex_); auto it = connected_segments_.find(remote_hixl); if (it == connected_segments_.end()) {