[TE] Add auto-connect feature controlled by config and environment variable (#1482)

Co-authored-by: zhangcheng <zhangcheng299@huawei.com>
This commit is contained in:
ZhangCheng 2026-02-05 14:07:56 +08:00 committed by GitHub
parent 17b9df8fb8
commit 01d1c6a8b3
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
5 changed files with 54 additions and 8 deletions

View File

@ -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).

View File

@ -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<SegmentID> need_update_metadata_segs_;
bool use_short_connection_{false};

View File

@ -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<int32_t>(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 *> &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<adxl::TransferOpDesc> 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<std::mutex> lock(connection_mutex_);
auto it = connected_segments_.find(target_adxl_engine_name);
if (it == connected_segments_.end()) {

View File

@ -115,6 +115,7 @@ class AscendDirectTransport : public Transport {
aclrtContext rt_context_{nullptr};
std::unique_ptr<hixl::Hixl> hixl_;
bool use_buffer_pool_{false};
bool auto_connect_{false};
uint64_t connect_timeout_;
uint64_t transfer_timeout_;
std::string local_hixl_name_{};

View File

@ -208,6 +208,14 @@ Status AscendDirectTransport::initHixl(const std::shared_ptr<Config> &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<int32_t>(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<std::mutex> lock(connection_mutex_);
auto it = connected_segments_.find(remote_hixl);
if (it == connected_segments_.end()) {