已合并
[shmem] fix with aclshmem #29045
[shmem] fix with aclshmem #29045
已合并
王超创建于 1月5日
7 个文件变更+261-122
Rthird_party/shmem/include/shmem_types.hthird_party/shmem/include/shmem_common_types.h+5-58
@@ -13,71 +13,18 @@
13#ifdef __cplusplus13#ifdef __cplusplus
14extern "C" {14extern "C" {
15#endif15#endif
16-/**
17- * @private
18-*/
19-#define SHMEM_GLOBAL __global__ __aicore__
20- 
21-/// \def SHMEM_DEVICE
22-/// \brief A macro that identifies a function on the device side.
23-#define SHMEM_DEVICE __attribute__((always_inline)) __aicore__ __inline__
24 16 
25/**17/**
26- * @addtogroup group_enums18+ * @brief Data operation engine type.
27- * @{
28-*/
29- 
30-/**
31- * @brief Team's index.
32-*/
33-enum shmem_team_index_t{
34- SHMEM_TEAM_INVALID = -1,
35- SHMEM_TEAM_WORLD = 0
36-};
37- 
38-/**
39- * @brief Data op engine type.
40*/19*/
41enum data_op_engine_type_t {20enum data_op_engine_type_t {
42- SHMEM_DATA_OP_MTE = 0x01,21+ ACLSHMEM_DATA_OP_MTE = 0x01,
22+ ACLSHMEM_DATA_OP_SDMA = 0x02,
23+ ACLSHMEM_DATA_OP_ROCE = 0x04,
43};24};
44 25 
45-/**
46- * @brief signal ops, used by signaler in p2p synchronization
47- */
48-enum {
49- SHMEM_SIGNAL_SET,
50- SHMEM_SIGNAL_ADD
51-};
52- 
53-/**
54- * @brief signal compare ops, used by signalee in p2p synchronization
55- */
56-enum {
57- SHMEM_CMP_EQ,
58- SHMEM_CMP_NE,
59- SHMEM_CMP_GT,
60- SHMEM_CMP_GE,
61- SHMEM_CMP_LT,
62- SHMEM_CMP_LE
63-};
64- 
65-/**@} */ // end of group_enums
66- 
67-/**
68- * @defgroup group_typedef Typedef
69- * @{
70- 
71-*/
72-/**
73- * @brief A typedef of int
74-*/
75-typedef int shmem_team_t;
76- 
77-/**@} */ // end of group_typedef
78- 
79#ifdef __cplusplus26#ifdef __cplusplus
80}27}
81#endif28#endif
82 29 
83-#endif /*SHMEM_TYPES_H*/30+#endif // SHMEM_TYPES_H
Mthird_party/shmem/include/shmem_host_def.h+92-1
@@ -10,7 +10,7 @@
10#ifndef SHMEM_HOST_DEF_H10#ifndef SHMEM_HOST_DEF_H
11#define SHMEM_HOST_DEF_H11#define SHMEM_HOST_DEF_H
12#include <climits>12#include <climits>
13-#include "third_party/shmem/include/shmem_types.h"13+#include "third_party/shmem/include/shmem_common_types.h"
14 14 
15#ifdef __cplusplus15#ifdef __cplusplus
16extern "C" {16extern "C" {
@@ -42,6 +42,16 @@ enum shmem_error_code_t : int {
42 SHMEM_NOT_INITED = -5, ///< This is a problem caused by an uninitialization.42 SHMEM_NOT_INITED = -5, ///< This is a problem caused by an uninitialization.
43};43};
44 44 
45+/**
46+ * @brief init flags
47+*/
48+enum aclshmemx_bootstrap_t : int {
49+ ACLSHMEMX_INIT_WITH_DEFAULT = 1 << 0,
50+ ACLSHMEMX_INIT_WITH_MPI = 1 << 1,
51+ ACLSHMEMX_INIT_WITH_UNIQUEID = 1 << 3,
52+ ACLSHMEMX_INIT_MAX = 1 << 31
53+};
54+ 
45/**55/**
46 * @brief The state of the SHMEM library initialization.56 * @brief The state of the SHMEM library initialization.
47*/57*/
@@ -53,6 +63,14 @@ enum shmem_init_status_t{
53};63};
54 64 
55constexpr uint16_t SHMEM_UNIQUE_ID_INNER_LEN = 60;65constexpr uint16_t SHMEM_UNIQUE_ID_INNER_LEN = 60;
66+/// \brief Inner length of the unique ID buffer for ACLSHMEM
67+constexpr uint16_t ACLSHMEM_UNIQUE_ID_INNER_LEN = 124;
68+/// \brief Default timeout value (in seconds) for ACLSHMEM operations
69+constexpr int DEFAULT_TIMEOUT = 120;
70+ 
71+/// \def ACLSHMEM_MAX_IP_PORT_LEN
72+/// \brief Maximum length of the IP and port string in ACLSHMEM (including null terminator)
73+#define ACLSHMEM_MAX_IP_PORT_LEN 64
AtlasAccount
AtlasAccountAtlasAccount1月5日

代码一致性: 宏'ACLSHMEM_MAX_IP_PORT_LEN'使用#define定义,而其他类似的长度常量如'SHMEM_UNIQUE_ID_INNER_LEN'和'ACLSHMEM_UNIQUE_ID_INNER_LEN'使用constexpr定义。这种不一致的常量定义方式降低了代码的一致性。

问题类型: 代码一致性 文件路径: third_party/shmem/include/shmem_host_def.h 行号: 74 问题代码:

#define ACLSHMEM_MAX_IP_PORT_LEN 64

修改建议:

统一常量定义风格。由于这是C/C++混合头文件,建议:
1. 如果支持C++11及以上,使用constexpr:
   constexpr uint16_t ACLSHMEM_MAX_IP_PORT_LEN = 64;
2. 如果需要C兼容,使用enum或static const:
   enum { ACLSHMEM_MAX_IP_PORT_LEN = 64 };

此评论由代码审查工具自动生成

likedislike
56 74 
57typedef struct {75typedef struct {
58 int32_t version;76 int32_t version;
@@ -69,6 +87,40 @@ constexpr int32_t SHMEM_UNIQUEID_VERSION = (1 << 16) + sizeof(shmem_uniqueid_t);
69 } \87 } \
70 } \88 } \
71 89 
90+/**
91+ * @struct aclshmemx_uniqueid_t
92+ * @brief Structure required for SHMEM unique ID (uid) initialization
93+ *
94+ * - int32_t version: version.
95+ * - int my_pe: The pe of the current process.
96+ * - int n_pes: The total pe number of all processes.
97+ * - char internal[ACLSHMEM_UNIQUE_ID_INNER_LEN]: Internal information of uid.
98+*/
99+typedef struct {
100+ int32_t version;
101+ int my_pe;
102+ int n_pes;
103+ char internal[ACLSHMEM_UNIQUE_ID_INNER_LEN];
104+} aclshmemx_uniqueid_t;
105+ 
106+/// \brief Version number of the ACLSHMEM unique ID structure
107+constexpr int32_t ACLSHMEM_UNIQUEID_VERSION = (1 << 16) + sizeof(aclshmemx_uniqueid_t);
108+/**@} */ // end of group_constants
109+ 
110+/**
111+ * @addtogroup group_macros
112+ * @{
113+*/
114+/// \def ACLSHMEM_UNIQUEID_INITIALIZER
115+/// \brief Initializer macro for the ACLSHMEM unique ID structure
116+#define ACLSHMEM_UNIQUEID_INITIALIZER \
AtlasAccount
AtlasAccountAtlasAccount1月5日

代码结构与可维护性: 宏ACLSHMEM_UNIQUEID_INITIALIZER的初始化列表不完整。结构体aclshmemx_uniqueid_t有四个成员:versionmy_pen_pesinternal。但该宏只初始化了versioninternal(通过一个嵌套的{0}),遗漏了my_pen_pes。这会导致使用该宏初始化的结构体中my_pen_pes成员处于未初始化状态,可能包含任意值,从而引发未定义行为。

问题类型: 代码结构与可维护性 文件路径: third_party/shmem/include/shmem_host_def.h 行号: 117 问题代码:

#define ACLSHMEM_UNIQUEID_INITIALIZER                   \
    {                                                   \
        ACLSHMEM_UNIQUEID_VERSION,                      \
        {                                               \
            0                                           \
        }                                               \
    }

修改建议:

修正初始化宏,显式初始化所有成员。根据结构体定义,应初始化`version`、`my_pe`、`n_pes`和`internal`。例如:
```c
#define ACLSHMEM_UNIQUEID_INITIALIZER                   \
    {                                                   \
        ACLSHMEM_UNIQUEID_VERSION,                      \
        0, /* my_pe */                                  \
        0, /* n_pes */                                  \
        { 0 } /* internal */                            \
    }

这样确保结构体被完全初始化,避免未定义行为。


---
*此评论由代码审查工具自动生成*
likedislike
117+ { \
118+ ACLSHMEM_UNIQUEID_VERSION, \
119+ { \
120+ 0 \
121+ } \
122+ }
123+ 
72/**@} */ // end of group_enums124/**@} */ // end of group_enums
73 125 
74/**126/**
@@ -112,6 +164,45 @@ typedef struct {
112 shmem_init_optional_attr_t option_attr;164 shmem_init_optional_attr_t option_attr;
113} shmem_init_attr_t;165} shmem_init_attr_t;
114 166 
167+/**
168+ * @struct aclshmem_init_optional_attr_t
169+ * @brief Optional parameter for the attributes used for initialization.
170+ *
171+ * - int version: version
172+ * - data_op_engine_type_t data_op_engine_type: data_op_engine_type
173+ * - uint32_t shm_init_timeout: shm_init_timeout
174+ * - uint32_t shm_create_timeout: shm_create_timeout
175+ * - uint32_t control_operation_timeout: control_operation_timeout
176+ * - int32_t sockFd: sock_fd for apply port in advance
177+*/
178+typedef struct {
179+ int version;
180+ data_op_engine_type_t data_op_engine_type;
181+ uint32_t shm_init_timeout;
182+ uint32_t shm_create_timeout;
183+ uint32_t control_operation_timeout;
184+ int32_t sockFd;
185+} aclshmem_init_optional_attr_t;
186+/**
187+ * @struct aclshmemx_init_attr_t
188+ * @brief Mandatory parameter for attributes used for initialization.
189+ *
190+ * - int my_pe: The pe of the current process.
191+ * - int n_pes: The total pe number of all processes.
192+ * - char ip_port[ACLSHMEM_MAX_IP_PORT_LEN]: The ip and port of the communication server. The port must not conflict
193+ * with other modules and processes.
194+ * - uint64_t local_mem_size: The size of shared memory currently occupied by current pe.
195+ * - aclshmem_init_optional_attr_t option_attr: Optional Parameters.
196+*/
197+typedef struct {
198+ int my_pe;
199+ int n_pes;
200+ char ip_port[ACLSHMEM_MAX_IP_PORT_LEN];
201+ uint64_t local_mem_size;
202+ aclshmem_init_optional_attr_t option_attr = {(1 << 16) + sizeof(aclshmem_init_optional_attr_t), ACLSHMEM_DATA_OP_MTE, DEFAULT_TIMEOUT, DEFAULT_TIMEOUT, DEFAULT_TIMEOUT};
AtlasAccount
AtlasAccountAtlasAccount1月5日

代码结构与可维护性: 在C语言头文件中使用了C++风格的默认成员初始化。结构体aclshmemx_init_attr_t的成员option_attr使用了C++的默认初始化语法(= {...})。虽然文件使用了#ifdef __cplusplus来兼容C++,但该结构体定义在extern "C"块之外(第219-221行),这意味着在纯C编译环境中(如C编译器或C++编译器编译C代码时),这种语法是无效的,会导致编译错误。

问题类型: 代码结构与可维护性 文件路径: third_party/shmem/include/shmem_host_def.h 行号: 203 问题代码:

    aclshmem_init_optional_attr_t option_attr = {(1 << 16) + sizeof(aclshmem_init_optional_attr_t), ACLSHMEM_DATA_OP_MTE, DEFAULT_TIMEOUT, DEFAULT_TIMEOUT, DEFAULT_TIMEOUT};

修改建议:

将默认初始化移除,改为在文档中说明用户必须显式初始化该字段,或者提供一个单独的初始化函数/宏(如`ACLSHMEMX_INIT_ATTR_INITIALIZER`)来设置默认值。例如:
1. 删除`= {...}`部分。
2. 添加一个初始化宏:
```c
#define ACLSHMEMX_INIT_ATTR_INITIALIZER \
    { \
        0, 0, "", 0, \
        { (1 << 16) + sizeof(aclshmem_init_optional_attr_t), ACLSHMEM_DATA_OP_MTE, DEFAULT_TIMEOUT, DEFAULT_TIMEOUT, DEFAULT_TIMEOUT, -1 }, \
        NULL \
    }

这样既保证了C/C++兼容性,又提供了便捷的初始化方式。


---
*此评论由代码审查工具自动生成*
likedislike
203+ void *comm_args;
204+} aclshmemx_init_attr_t;
205+ 
115/**206/**
116 * @brief Callback function of private key password decryptor, see shmem_register_decrypt_handler207 * @brief Callback function of private key password decryptor, see shmem_register_decrypt_handler
117 *208 *
Mtorch_npu/csrc/core/npu/sys_ctrl/npu_sys_ctrl.cpp+2-2
@@ -283,8 +283,8 @@ NpuSysCtrl::SysStatus NpuSysCtrl::Finalize()
283 NPU_CHECK_WARN(c10_npu::DestroyUsedStreams());283 NPU_CHECK_WARN(c10_npu::DestroyUsedStreams());
284 NPU_CHECK_WARN(c10_npu::ResetUsedDevices());284 NPU_CHECK_WARN(c10_npu::ResetUsedDevices());
285#ifndef BUILD_LIBTORCH285#ifndef BUILD_LIBTORCH
286- if (c10d::symmetric_memory::Shmem_finalize_exist()) {286+ if (c10d::symmetric_memory::Aclshmem_finalize_exist()) {
287- auto ret = c10d::symmetric_memory::Shmem_finalize();287+ auto ret = c10d::symmetric_memory::Aclshmem_finalize();
288 ASCEND_LOGI("shmem_finalize emd, ret is %d", ret);288 ASCEND_LOGI("shmem_finalize emd, ret is %d", ret);
289 }289 }
290#endif290#endif
Mtorch_npu/csrc/distributed/symm_mem/NPUSHMEMExtension.cpp+47-22
@@ -25,32 +25,57 @@ void initialize_npushmem_with_store(
25 25 
26 logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, rank is %d, world_size is %d.", rank, world_size);26 logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, rank is %d, world_size is %d.", rank, world_size);
27 27 
28- uint32_t status = c10d::symmetric_memory::Shmem_set_conf_store_tls(false, nullptr, 0);28+ uint32_t status = c10d::symmetric_memory::Aclshmemx_set_conf_store_tls(false, nullptr, 0);
29 TORCH_CHECK(status == 0, "shmem_set_conf_store_tls failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));29 TORCH_CHECK(status == 0, "shmem_set_conf_store_tls failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));
30 30 
31- shmem_uniqueid_t unique_id;
32- if (rank == 0) {
33- status = c10d::symmetric_memory::Shmem_get_uniqueid(&unique_id);
34- TORCH_CHECK(status == 0, "shmem_get_uniqueid failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));
35- logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, Shmem_get_uniqueid rank is %d, version %d, internal is %s.",
36- rank, unique_id.version, unique_id.internal);
37- }
38- auto unique_ids = storeExchange.all_gather(store, rank, world_size, unique_id);
39- logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, unique_id rank is %d, version %d, internal is %s.",
40- rank, unique_ids[0].version, unique_ids[0].internal);
41- 
42 int64_t init_size = c10_npu::option::OptionsManager::GetShmemSymmetricSize();31 int64_t init_size = c10_npu::option::OptionsManager::GetShmemSymmetricSize();
43- shmem_init_attr_t* attributes;32+ if (c10d::symmetric_memory::Aclshmemx_get_uniqueid_exist()) {
44- logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, start shmem_set_attr rank is %d, world_size is %d, size is %llu.",33+ // gitcode version
45- rank, world_size, init_size);34+ aclshmemx_uniqueid_t unique_id;
46- status = c10d::symmetric_memory::Shmem_set_attr(rank, world_size, init_size, nullptr, &attributes);35+ if (rank == 0) {
47- TORCH_CHECK(status == 0, "shmem_set_attr failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));36+ status = c10d::symmetric_memory::Aclshmemx_get_uniqueid(&unique_id);
48- logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, end shmem_set_attr rank is %d, world_size is %d, size is %llu.",37+ TORCH_CHECK(status == 0, "aclshmemx_get_uniqueid failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));
49- rank, world_size, init_size);38+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, aclshmemx_get_uniqueid rank is %d, version %d, internal is %s.",
39+ rank, unique_id.version, unique_id.internal);
40+ }
41+ auto unique_ids = storeExchange.all_gather(store, rank, world_size, unique_id);
42+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, unique_id rank is %d, version %d, internal is %s.",
43+ rank, unique_ids[0].version, unique_ids[0].internal);
50 44 
51- status = c10d::symmetric_memory::Shmem_set_attr_uniqueid_args(rank, world_size, &unique_ids[0], attributes);45+ aclshmemx_init_attr_t attr;
52- TORCH_CHECK(status == 0, "Shmem_set_attr_uniqueid_args failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));46+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, start aclshmemx_set_attr_uniqueid_args rank is %d, world_size is %d, size is %llu.",
53- logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store success, rank is %d, world_size is %d.", rank, world_size);47+ rank, world_size, init_size);
48+ 
49+ status = c10d::symmetric_memory::Aclshmemx_set_attr_uniqueid_args(rank, world_size, init_size, &unique_ids[0], &attr);
50+ TORCH_CHECK(status == 0, "aclshmemx_set_attr_uniqueid_args failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));
51+ 
52+ status = c10d::symmetric_memory::Aclshmemx_init_attr(ACLSHMEMX_INIT_WITH_DEFAULT, &attr);
53+ TORCH_CHECK(status == 0, "aclshmemx_init_attr failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));
54+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store success, rank is %d, world_size is %d.", rank, world_size);
55+ } else {
56+ shmem_uniqueid_t unique_id;
57+ if (rank == 0) {
58+ status = c10d::symmetric_memory::Shmemx_get_uniqueid(&unique_id);
59+ TORCH_CHECK(status == 0, "shmem_get_uniqueid failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));
60+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, Shmem_get_uniqueid rank is %d, version %d, internal is %s.",
61+ rank, unique_id.version, unique_id.internal);
62+ }
63+ auto unique_ids = storeExchange.all_gather(store, rank, world_size, unique_id);
64+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, unique_id rank is %d, version %d, internal is %s.",
65+ rank, unique_ids[0].version, unique_ids[0].internal);
66+ 
67+ shmem_init_attr_t* attributes;
68+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, start shmem_set_attr rank is %d, world_size is %d, size is %llu.",
69+ rank, world_size, init_size);
70+ status = c10d::symmetric_memory::Shmem_set_attr(rank, world_size, init_size, nullptr, &attributes);
71+ TORCH_CHECK(status == 0, "shmem_set_attr failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));
72+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store, end shmem_set_attr rank is %d, world_size is %d, size is %llu.",
73+ rank, world_size, init_size);
74+ 
75+ status = c10d::symmetric_memory::Shmem_set_attr_uniqueid_args(rank, world_size, &unique_ids[0], attributes);
76+ TORCH_CHECK(status == 0, "shmem_set_attr_uniqueid_args failed, status is ", status, DIST_ERROR(ErrCode::INTERNAL));
77+ logger->debug("NPUSHMEMSymmetricMemoryAllocator initialize_npushmem_with_store success, rank is %d, world_size is %d.", rank, world_size);
78+ }
54 is_initialized = true;79 is_initialized = true;
55}80}
56 81 
Mtorch_npu/csrc/distributed/symm_mem/NPUSHMEMInterface.cpp+98-28
@@ -13,9 +13,16 @@ namespace symmetric_memory {
13 GET_FUNCTION(libshmem, funcName)13 GET_FUNCTION(libshmem, funcName)
14 14 
15REGISTER_LIBRARY(libshmem)15REGISTER_LIBRARY(libshmem)
16+LOAD_FUNCTION(aclshmemx_set_conf_store_tls)
17+LOAD_FUNCTION(aclshmemx_get_uniqueid)
18+LOAD_FUNCTION(aclshmemx_set_attr_uniqueid_args)
19+LOAD_FUNCTION(aclshmemx_init_attr)
20+LOAD_FUNCTION(aclshmem_malloc)
21+LOAD_FUNCTION(aclshmem_free)
22+LOAD_FUNCTION(aclshmem_ptr)
23+LOAD_FUNCTION(aclshmem_finalize)
16LOAD_FUNCTION(shmem_set_conf_store_tls)24LOAD_FUNCTION(shmem_set_conf_store_tls)
17LOAD_FUNCTION(shmem_set_attr)25LOAD_FUNCTION(shmem_set_attr)
18-LOAD_FUNCTION(shmem_init_attr)
19LOAD_FUNCTION(shmem_get_uniqueid)26LOAD_FUNCTION(shmem_get_uniqueid)
20LOAD_FUNCTION(shmem_set_attr_uniqueid_args)27LOAD_FUNCTION(shmem_set_attr_uniqueid_args)
21LOAD_FUNCTION(shmem_malloc)28LOAD_FUNCTION(shmem_malloc)
@@ -23,14 +30,18 @@ LOAD_FUNCTION(shmem_free)
23LOAD_FUNCTION(shmem_ptr)30LOAD_FUNCTION(shmem_ptr)
24LOAD_FUNCTION(shmem_finalize)31LOAD_FUNCTION(shmem_finalize)
25 32 
26-int32_t Shmem_set_conf_store_tls(bool enable, const char *tls_info, const uint32_t tls_info_len)33+int32_t Aclshmemx_set_conf_store_tls(bool enable, const char *tls_info, const uint32_t tls_info_len)
27{34{
28 typedef int32_t (*ShmemApiFunc)(bool, const char *, const uint32_t);35 typedef int32_t (*ShmemApiFunc)(bool, const char *, const uint32_t);
29 static ShmemApiFunc shmem_set_conf_store_tls_func = nullptr;36 static ShmemApiFunc shmem_set_conf_store_tls_func = nullptr;
37+ if (shmem_set_conf_store_tls_func == nullptr) {
38+ shmem_set_conf_store_tls_func = (ShmemApiFunc)GET_FUNC(aclshmemx_set_conf_store_tls);
39+ }
30 if (shmem_set_conf_store_tls_func == nullptr) {40 if (shmem_set_conf_store_tls_func == nullptr) {
31 shmem_set_conf_store_tls_func = (ShmemApiFunc)GET_FUNC(shmem_set_conf_store_tls);41 shmem_set_conf_store_tls_func = (ShmemApiFunc)GET_FUNC(shmem_set_conf_store_tls);
32 }42 }
33- TORCH_CHECK(shmem_set_conf_store_tls_func, "Failed to find function ", "shmem_set_conf_store_tls", PTA_ERROR(ErrCode::NOT_FOUND));43+ TORCH_CHECK(shmem_set_conf_store_tls_func, "Failed to find function ",
44+ "aclshmemx_set_conf_store_tls or shmem_set_conf_store_tls", PTA_ERROR(ErrCode::NOT_FOUND));
34 return shmem_set_conf_store_tls_func(enable, tls_info, tls_info_len);45 return shmem_set_conf_store_tls_func(enable, tls_info, tls_info_len);
35}46}
36 47 
@@ -46,25 +57,41 @@ int32_t Shmem_set_attr(int32_t my_rank, int32_t n_ranks, uint64_t local_mem_size
46 return shmem_set_attr_func(my_rank, n_ranks, local_mem_size, ip_port, attributes);57 return shmem_set_attr_func(my_rank, n_ranks, local_mem_size, ip_port, attributes);
47}58}
48 59 
49-int32_t Shmem_init_attr(shmem_init_attr_t *attributes)60+int Shmemx_get_uniqueid(shmem_uniqueid_t *uid)
50-{
51- typedef int32_t (*ShmemApiFunc)(shmem_init_attr_t *);
52- static ShmemApiFunc shmem_init_attr_func = nullptr;
53- if (shmem_init_attr_func == nullptr) {
54- shmem_init_attr_func = (ShmemApiFunc)GET_FUNC(shmem_init_attr);
55- }
56- TORCH_CHECK(shmem_init_attr_func, "Failed to find function ", "shmem_init_attr", PTA_ERROR(ErrCode::NOT_FOUND));
57- return shmem_init_attr_func(attributes);
58-}
59- 
60-int32_t Shmem_get_uniqueid(shmem_uniqueid_t *uid)
61{61{
62 typedef int32_t (*ShmemApiFunc)(shmem_uniqueid_t *);62 typedef int32_t (*ShmemApiFunc)(shmem_uniqueid_t *);
63 static ShmemApiFunc shmem_get_uniqueid_func = nullptr;63 static ShmemApiFunc shmem_get_uniqueid_func = nullptr;
64 if (shmem_get_uniqueid_func == nullptr) {64 if (shmem_get_uniqueid_func == nullptr) {
65 shmem_get_uniqueid_func = (ShmemApiFunc)GET_FUNC(shmem_get_uniqueid);65 shmem_get_uniqueid_func = (ShmemApiFunc)GET_FUNC(shmem_get_uniqueid);
66 }66 }
67- TORCH_CHECK(shmem_get_uniqueid_func, "Failed to find function ", "shmem_get_uniqueid", PTA_ERROR(ErrCode::NOT_FOUND));67+ TORCH_CHECK(shmem_get_uniqueid_func, "Failed to find function ",
68+ "shmem_get_uniqueid", PTA_ERROR(ErrCode::NOT_FOUND));
69+ return shmem_get_uniqueid_func(uid);
70+}
71+ 
72+bool Aclshmemx_get_uniqueid_exist()
73+{
74+ const static bool shmemApiFuncExist = []() -> bool {
75+ try {
76+ auto func = GET_FUNC(aclshmemx_get_uniqueid);
77+ return func != nullptr;
78+ } catch (...) {
79+ // libshmem.so not exist
80+ return false;
81+ }
82+ }();
83+ return shmemApiFuncExist;
84+}
85+ 
86+int32_t Aclshmemx_get_uniqueid(aclshmemx_uniqueid_t *uid)
87+{
88+ typedef int32_t (*ShmemApiFunc)(aclshmemx_uniqueid_t *);
89+ static ShmemApiFunc shmem_get_uniqueid_func = nullptr;
90+ if (shmem_get_uniqueid_func == nullptr) {
91+ shmem_get_uniqueid_func = (ShmemApiFunc)GET_FUNC(aclshmemx_get_uniqueid);
92+ }
93+ TORCH_CHECK(shmem_get_uniqueid_func, "Failed to find function ",
94+ "aclshmemx_get_uniqueid or shmem_get_uniqueid", PTA_ERROR(ErrCode::NOT_FOUND));
68 return shmem_get_uniqueid_func(uid);95 return shmem_get_uniqueid_func(uid);
69}96}
70 97 
@@ -75,49 +102,88 @@ int Shmem_set_attr_uniqueid_args(int rank_id, int nranks, const shmem_uniqueid_t
75 if (shmem_set_attr_uniqueid_args_func == nullptr) {102 if (shmem_set_attr_uniqueid_args_func == nullptr) {
76 shmem_set_attr_uniqueid_args_func = (ShmemApiFunc)GET_FUNC(shmem_set_attr_uniqueid_args);103 shmem_set_attr_uniqueid_args_func = (ShmemApiFunc)GET_FUNC(shmem_set_attr_uniqueid_args);
77 }104 }
78- TORCH_CHECK(shmem_set_attr_uniqueid_args_func, "Failed to find function ", "shmem_set_attr_uniqueid_args", PTA_ERROR(ErrCode::NOT_FOUND));105+ TORCH_CHECK(shmem_set_attr_uniqueid_args_func, "Failed to find function ",
106+ "shmem_set_attr_uniqueid_args", PTA_ERROR(ErrCode::NOT_FOUND));
79 return shmem_set_attr_uniqueid_args_func(rank_id, nranks, uid, attr);107 return shmem_set_attr_uniqueid_args_func(rank_id, nranks, uid, attr);
80}108}
81 109 
82-void *Shmem_malloc(size_t size)110+int Aclshmemx_set_attr_uniqueid_args(int rank_id, int nranks, int64_t local_mem_size,
111+ aclshmemx_uniqueid_t *uid, aclshmemx_init_attr_t *aclshmem_attr)
112+{
113+ typedef int32_t (*ShmemApiFunc)(int, int, int64_t, aclshmemx_uniqueid_t *, aclshmemx_init_attr_t *);
114+ static ShmemApiFunc aclshmemx_set_attr_uniqueid_args_func = nullptr;
115+ if (aclshmemx_set_attr_uniqueid_args_func == nullptr) {
116+ aclshmemx_set_attr_uniqueid_args_func = (ShmemApiFunc)GET_FUNC(aclshmemx_set_attr_uniqueid_args);
117+ }
118+ TORCH_CHECK(aclshmemx_set_attr_uniqueid_args_func, "Failed to find function ",
119+ "aclshmemx_set_attr_uniqueid_args", PTA_ERROR(ErrCode::NOT_FOUND));
120+ return aclshmemx_set_attr_uniqueid_args_func(rank_id, nranks, local_mem_size, uid, aclshmem_attr);
121+}
122+ 
123+int Aclshmemx_init_attr(aclshmemx_bootstrap_t bootstrap_flags, aclshmemx_init_attr_t *attributes)
124+{
125+ typedef int32_t (*ShmemApiFunc)(aclshmemx_bootstrap_t, aclshmemx_init_attr_t *);
126+ static ShmemApiFunc aclshmemx_init_attr_func = nullptr;
127+ if (aclshmemx_init_attr_func == nullptr) {
128+ aclshmemx_init_attr_func = (ShmemApiFunc)GET_FUNC(aclshmemx_init_attr);
129+ }
130+ TORCH_CHECK(aclshmemx_init_attr_func, "Failed to find function ",
131+ "aclshmemx_init_attr", PTA_ERROR(ErrCode::NOT_FOUND));
132+ return aclshmemx_init_attr_func(bootstrap_flags, attributes);
133+}
134+ 
135+void *Aclshmem_malloc(size_t size)
83{136{
84 typedef void* (*ShmemApiFunc)(size_t);137 typedef void* (*ShmemApiFunc)(size_t);
85 static ShmemApiFunc shmem_malloc_func = nullptr;138 static ShmemApiFunc shmem_malloc_func = nullptr;
139+ if (shmem_malloc_func == nullptr) {
140+ shmem_malloc_func = (ShmemApiFunc)GET_FUNC(aclshmem_malloc);
141+ }
86 if (shmem_malloc_func == nullptr) {142 if (shmem_malloc_func == nullptr) {
87 shmem_malloc_func = (ShmemApiFunc)GET_FUNC(shmem_malloc);143 shmem_malloc_func = (ShmemApiFunc)GET_FUNC(shmem_malloc);
88 }144 }
89- TORCH_CHECK(shmem_malloc_func, "Failed to find function ", "shmem_malloc", PTA_ERROR(ErrCode::NOT_FOUND));145+ TORCH_CHECK(shmem_malloc_func, "Failed to find function ",
146+ "aclshmem_malloc or shmem_malloc", PTA_ERROR(ErrCode::NOT_FOUND));
90 return shmem_malloc_func(size);147 return shmem_malloc_func(size);
91}148}
92 149 
93-void Shmem_free(void *ptr)150+void Aclshmem_free(void *ptr)
94{151{
95 typedef void (*ShmemApiFunc)(void *);152 typedef void (*ShmemApiFunc)(void *);
96 static ShmemApiFunc shmem_free_func = nullptr;153 static ShmemApiFunc shmem_free_func = nullptr;
154+ if (shmem_free_func == nullptr) {
155+ shmem_free_func = (ShmemApiFunc)GET_FUNC(aclshmem_free);
156+ }
97 if (shmem_free_func == nullptr) {157 if (shmem_free_func == nullptr) {
98 shmem_free_func = (ShmemApiFunc)GET_FUNC(shmem_free);158 shmem_free_func = (ShmemApiFunc)GET_FUNC(shmem_free);
99 }159 }
100- TORCH_CHECK(shmem_free_func, "Failed to find function ", "shmem_free", PTA_ERROR(ErrCode::NOT_FOUND));160+ TORCH_CHECK(shmem_free_func, "Failed to find function ",
161+ "aclshmem_free or shmem_free", PTA_ERROR(ErrCode::NOT_FOUND));
101 return shmem_free_func(ptr);162 return shmem_free_func(ptr);
102}163}
103 164 
104-void *Shmem_ptr(void *ptr, int pe)165+void *Aclshmem_ptr(void *ptr, int pe)
105{166{
106 typedef void* (*ShmemApiFunc)(void *, int);167 typedef void* (*ShmemApiFunc)(void *, int);
107 static ShmemApiFunc shmem_ptr_func = nullptr;168 static ShmemApiFunc shmem_ptr_func = nullptr;
169+ if (shmem_ptr_func == nullptr) {
170+ shmem_ptr_func = (ShmemApiFunc)GET_FUNC(aclshmem_ptr);
171+ }
108 if (shmem_ptr_func == nullptr) {172 if (shmem_ptr_func == nullptr) {
109 shmem_ptr_func = (ShmemApiFunc)GET_FUNC(shmem_ptr);173 shmem_ptr_func = (ShmemApiFunc)GET_FUNC(shmem_ptr);
110 }174 }
111- TORCH_CHECK(shmem_ptr_func, "Failed to find function ", "shmem_ptr", PTA_ERROR(ErrCode::NOT_FOUND));175+ TORCH_CHECK(shmem_ptr_func, "Failed to find function ",
176+ "aclshmem_ptr or shmem_ptr", PTA_ERROR(ErrCode::NOT_FOUND));
112 return shmem_ptr_func(ptr, pe);177 return shmem_ptr_func(ptr, pe);
113}178}
114 179 
115-bool Shmem_finalize_exist()180+bool Aclshmem_finalize_exist()
116{181{
117 const static bool shmemApiFuncExist = []() -> bool {182 const static bool shmemApiFuncExist = []() -> bool {
118 try {183 try {
119- auto func = GET_FUNC(shmem_finalize)184+ auto func1 = GET_FUNC(aclshmem_finalize);
120- return func != nullptr;185+ auto func2 = GET_FUNC(shmem_finalize);
186+ return func1 != nullptr || func2 != nullptr;
121 } catch (...) {187 } catch (...) {
122 // libshmem.so not exist188 // libshmem.so not exist
123 return false;189 return false;
@@ -126,14 +192,18 @@ bool Shmem_finalize_exist()
126 return shmemApiFuncExist;192 return shmemApiFuncExist;
127}193}
128 194 
129-int Shmem_finalize(void)195+int Aclshmem_finalize(void)
130{196{
131 typedef int (*ShmemApiFunc)(void);197 typedef int (*ShmemApiFunc)(void);
132 static ShmemApiFunc shmem_finalize_func = nullptr;198 static ShmemApiFunc shmem_finalize_func = nullptr;
199+ if (shmem_finalize_func == nullptr) {
200+ shmem_finalize_func = (ShmemApiFunc)GET_FUNC(aclshmem_finalize);
201+ }
133 if (shmem_finalize_func == nullptr) {202 if (shmem_finalize_func == nullptr) {
134 shmem_finalize_func = (ShmemApiFunc)GET_FUNC(shmem_finalize);203 shmem_finalize_func = (ShmemApiFunc)GET_FUNC(shmem_finalize);
135 }204 }
136- TORCH_CHECK(shmem_finalize_func, "Failed to find function ", "shmem_finalize", PTA_ERROR(ErrCode::NOT_FOUND));205+ TORCH_CHECK(shmem_finalize_func, "Failed to find function ",
206+ "aclshmem_finalize or shmem_finalize", PTA_ERROR(ErrCode::NOT_FOUND));
137 return shmem_finalize_func();207 return shmem_finalize_func();
138}208}
139 209 
Mtorch_npu/csrc/distributed/symm_mem/NPUSHMEMInterface.h+14-8
@@ -7,26 +7,32 @@
7namespace c10d {7namespace c10d {
8namespace symmetric_memory {8namespace symmetric_memory {
9 9 
10-int32_t Shmem_set_conf_store_tls(bool enable, const char *tls_info, const uint32_t tls_info_len);10+int32_t Aclshmemx_set_conf_store_tls(bool enable, const char *tls_info, const uint32_t tls_info_len);
11 11 
12int32_t Shmem_set_attr(int32_t my_rank, int32_t n_ranks, uint64_t local_mem_size, const char *ip_port,12int32_t Shmem_set_attr(int32_t my_rank, int32_t n_ranks, uint64_t local_mem_size, const char *ip_port,
13 shmem_init_attr_t **attributes);13 shmem_init_attr_t **attributes);
14 14 
15-int32_t Shmem_init_attr(shmem_init_attr_t *attributes);15+int Shmemx_get_uniqueid(shmem_uniqueid_t *uid);
16 16 
17-int Shmem_get_uniqueid(shmem_uniqueid_t *uid);17+bool Aclshmemx_get_uniqueid_exist();
18+ 
19+int Aclshmemx_get_uniqueid(aclshmemx_uniqueid_t *uid);
18 20 
19int Shmem_set_attr_uniqueid_args(int rank_id, int nranks, const shmem_uniqueid_t *uid, shmem_init_attr_t *attr);21int Shmem_set_attr_uniqueid_args(int rank_id, int nranks, const shmem_uniqueid_t *uid, shmem_init_attr_t *attr);
20 22 
21-void *Shmem_malloc(size_t size);23+int Aclshmemx_set_attr_uniqueid_args(int rank_id, int nranks, int64_t local_mem_size, aclshmemx_uniqueid_t *uid, aclshmemx_init_attr_t *aclshmem_attr);
22 24 
23-void Shmem_free(void *ptr);25+int Aclshmemx_init_attr(aclshmemx_bootstrap_t bootstrap_flags, aclshmemx_init_attr_t *attributes);
24 26 
25-void *Shmem_ptr(void *ptr, int pe);27+void *Aclshmem_malloc(size_t size);
26 28 
27-bool Shmem_finalize_exist();29+void Aclshmem_free(void *ptr);
28 30 
29-int Shmem_finalize(void);31+void *Aclshmem_ptr(void *ptr, int pe);
32+ 
33+bool Aclshmem_finalize_exist();
34+ 
35+int Aclshmem_finalize(void);
30 36 
31} // namespace symmetric_memory37} // namespace symmetric_memory
32} // namespace c10d38} // namespace c10d
Mtorch_npu/csrc/distributed/symm_mem/NPUSHMEMSymmetricMemory.cpp+3-3
@@ -23,7 +23,7 @@ NPUSHMEMAllocation::~NPUSHMEMAllocation()
23 auto device = c10::Device(at::DeviceType::PrivateUse1, device_idx);23 auto device = c10::Device(at::DeviceType::PrivateUse1, device_idx);
24 at::DeviceGuard device_guard(device);24 at::DeviceGuard device_guard(device);
25 logger->debug("~NPUSHMEMAllocation, start Shmem_free, ptr is %p.", ptr);25 logger->debug("~NPUSHMEMAllocation, start Shmem_free, ptr is %p.", ptr);
26- Shmem_free(ptr); // shmem_free has no return value26+ Aclshmem_free(ptr); // shmem_free has no return value
27 logger->debug("~NPUSHMEMAllocation, end Shmem_free, ptr is %p.", ptr);27 logger->debug("~NPUSHMEMAllocation, end Shmem_free, ptr is %p.", ptr);
28}28}
29 29 
@@ -66,7 +66,7 @@ NPUSHMEMSymmetricMemory::NPUSHMEMSymmetricMemory(
66 TORCH_INTERNAL_ASSERT(!group_info.rank_to_global_rank.empty());66 TORCH_INTERNAL_ASSERT(!group_info.rank_to_global_rank.empty());
67 rank_to_global_rank_ = group_info.rank_to_global_rank;67 rank_to_global_rank_ = group_info.rank_to_global_rank;
68 for (int r = 0; r < world_size_; ++r) {68 for (int r = 0; r < world_size_; ++r) {
69- auto buffer = Shmem_ptr(allocation->ptr, rank_to_global_rank_[r]);69+ auto buffer = Aclshmem_ptr(allocation->ptr, rank_to_global_rank_[r]);
70 TORCH_CHECK(buffer != nullptr, "shmem_ptr return nullptr with ptr ", allocation->ptr, DIST_ERROR(ErrCode::MEMORY));70 TORCH_CHECK(buffer != nullptr, "shmem_ptr return nullptr with ptr ", allocation->ptr, DIST_ERROR(ErrCode::MEMORY));
71 buffers_.push_back(buffer);71 buffers_.push_back(buffer);
72 logger->debug("[rank %d] NPUSHMEMSymmetricMemory shmem_ptr, r is %d, rank_to_global_rank is %d, ptr is %p, shmem_ptr is %p.",72 logger->debug("[rank %d] NPUSHMEMSymmetricMemory shmem_ptr, r is %d, rank_to_global_rank is %d, ptr is %p, shmem_ptr is %p.",
@@ -205,7 +205,7 @@ void* NPUSHMEMSymmetricMemoryAllocator::alloc(
205 int world_size = group_info.world_size;205 int world_size = group_info.world_size;
206 npushmem_extension::initialize_npushmem_with_store(store, rank, world_size);206 npushmem_extension::initialize_npushmem_with_store(store, rank, world_size);
207 207 
208- auto ptr = Shmem_malloc(size);208+ auto ptr = Aclshmem_malloc(size);
209 TORCH_CHECK(ptr != nullptr, "shmem_malloc return nullptr with size ", size, DIST_ERROR(ErrCode::MEMORY));209 TORCH_CHECK(ptr != nullptr, "shmem_malloc return nullptr with size ", size, DIST_ERROR(ErrCode::MEMORY));
210 auto allocation =210 auto allocation =
211 std::make_shared<NPUSHMEMAllocation>(ptr, size, device_idx);211 std::make_shared<NPUSHMEMAllocation>(ptr, size, device_idx);