@@ -76,11 +76,11 @@ uct_rocm_ipc_map_remote(const uct_rocm_ipc_device_mem_element_t *elem,
7676 elem->mapped_offset );
7777}
7878
79- /* System-wide atomic increment */
79+ /* System-wide atomic add */
8080__device__ static inline void
81- uct_rocm_ipc_atomic_inc (uint64_t *dst, uint64_t inc_value )
81+ uct_rocm_ipc_atomic_add (uint64_t *dst, uint64_t add_value )
8282{
83- atomicAdd_system ((unsigned long long *)dst, (unsigned long long )inc_value );
83+ atomicAdd_system ((unsigned long long *)dst, (unsigned long long )add_value );
8484 __threadfence_system ();
8585}
8686
@@ -101,98 +101,24 @@ __device__ static inline void uct_rocm_ipc_level_sync()
101101 }
102102}
103103
104- /* Copy routines for different parallelism levels */
105- template <ucs_device_level_t level>
106- __device__ void uct_rocm_ipc_copy_level (void *dst, const void *src, size_t len);
107-
108- /* Thread-level copy */
109- template <>
110- __device__ inline void
111- uct_rocm_ipc_copy_level<UCS_DEVICE_LEVEL_THREAD >(void *dst, const void *src,
112- size_t len)
113- {
114- memcpy (dst, src, len);
115- }
116-
117- /* Wavefront-level copy (64 threads) */
118- template <>
119- __device__ inline void
120- uct_rocm_ipc_copy_level<UCS_DEVICE_LEVEL_WARP >(void *dst, const void *src,
121- size_t len)
104+ /* Shared strided copy used by warp- and block-level specializations */
105+ __device__ static inline void
106+ uct_rocm_ipc_copy_strided (void *dst, const void *src, size_t len,
107+ unsigned lane_id, unsigned num_lanes)
122108{
123109 using vec4 = int4;
124110 using vec2 = int2;
125- unsigned int lane_id, num_lanes;
126-
127- uct_rocm_ipc_get_lane<UCS_DEVICE_LEVEL_WARP >(lane_id, num_lanes);
128111 auto s1 = reinterpret_cast <const char *>(src);
129112 auto d1 = reinterpret_cast <char *>(dst);
130113
131114 /* 16B-aligned fast path using vec4 */
132- if (UCS_DEVICE_IS_ALIGNED_POW2 ((intptr_t )s1, sizeof (vec4)) &&
133- UCS_DEVICE_IS_ALIGNED_POW2 ((intptr_t )d1, sizeof (vec4))) {
134- const vec4 *s4 = reinterpret_cast <const vec4*>(s1);
135- vec4 *d4 = reinterpret_cast <vec4*>(d1);
136- size_t n4 = len / sizeof (vec4);
137-
138- for (size_t i = lane_id; i < n4; i += num_lanes) {
139- vec4 v = uct_rocm_ipc_ld_global_cg (s4 + i);
140- uct_rocm_ipc_st_global_cg (d4 + i, v);
141- }
142-
143- len = len - n4 * sizeof (vec4);
144- if (len == 0 ) {
145- return ;
146- }
147-
148- s1 = reinterpret_cast <const char *>(s4 + n4);
149- d1 = reinterpret_cast <char *>(d4 + n4);
150- }
151-
152- /* 8B-aligned fast path using vec2 */
153- if (UCS_DEVICE_IS_ALIGNED_POW2 ((intptr_t )s1, sizeof (vec2)) &&
154- UCS_DEVICE_IS_ALIGNED_POW2 ((intptr_t )d1, sizeof (vec2))) {
155- const vec2 *s2 = reinterpret_cast <const vec2*>(s1);
156- vec2 *d2 = reinterpret_cast <vec2*>(d1);
157- size_t n2 = len / sizeof (vec2);
158-
159- for (size_t i = lane_id; i < n2; i += num_lanes) {
160- vec2 v2 = uct_rocm_ipc_ld_global_cg (s2 + i);
161- uct_rocm_ipc_st_global_cg (d2 + i, v2);
162- }
163-
164- len = len - n2 * sizeof (vec2);
165- if (len == 0 ) {
166- return ;
167- }
168-
169- s1 = reinterpret_cast <const char *>(s2 + n2);
170- d1 = reinterpret_cast <char *>(d2 + n2);
171- }
172-
173- /* Byte tail */
174- for (size_t i = lane_id; i < len; i += num_lanes) {
175- d1[i] = s1[i];
176- }
177- }
178-
179- template <>
180- __device__ inline void
181- uct_rocm_ipc_copy_level<UCS_DEVICE_LEVEL_BLOCK >(void *dst, const void *src,
182- size_t len)
183- {
184- using vec4 = int4;
185- using vec2 = int2;
186- auto s1 = reinterpret_cast <const char *>(src);
187- auto d1 = reinterpret_cast <char *>(dst);
188-
189115 if (UCS_DEVICE_IS_ALIGNED_POW2 ((intptr_t )s1, sizeof (vec4)) &&
190116 UCS_DEVICE_IS_ALIGNED_POW2 ((intptr_t )d1, sizeof (vec4))) {
191117 const vec4 *s4 = reinterpret_cast <const vec4*>(s1);
192118 vec4 *d4 = reinterpret_cast <vec4*>(d1);
193119 size_t num_lines = len / sizeof (vec4);
194120
195- for (size_t line = threadIdx. x ; line < num_lines; line += blockDim. x ) {
121+ for (size_t line = lane_id ; line < num_lines; line += num_lanes ) {
196122 vec4 v = uct_rocm_ipc_ld_global_cg (s4 + line);
197123 uct_rocm_ipc_st_global_cg (d4 + line, v);
198124 }
@@ -213,7 +139,7 @@ uct_rocm_ipc_copy_level<UCS_DEVICE_LEVEL_BLOCK>(void *dst, const void *src,
213139 vec2 *d2 = reinterpret_cast <vec2*>(d1);
214140 size_t num_lines = len / sizeof (vec2);
215141
216- for (size_t line = threadIdx. x ; line < num_lines; line += blockDim. x ) {
142+ for (size_t line = lane_id ; line < num_lines; line += num_lanes ) {
217143 vec2 v2 = uct_rocm_ipc_ld_global_cg (s2 + line);
218144 uct_rocm_ipc_st_global_cg (d2 + line, v2);
219145 }
@@ -228,11 +154,45 @@ uct_rocm_ipc_copy_level<UCS_DEVICE_LEVEL_BLOCK>(void *dst, const void *src,
228154 }
229155
230156 /* Byte tail */
231- for (size_t line = threadIdx. x ; line < len; line += blockDim. x ) {
157+ for (size_t line = lane_id ; line < len; line += num_lanes ) {
232158 d1[line] = s1[line];
233159 }
234160}
235161
162+ /* Copy routines for different parallelism levels */
163+ template <ucs_device_level_t level>
164+ __device__ void uct_rocm_ipc_copy_level (void *dst, const void *src, size_t len);
165+
166+ /* Thread-level copy */
167+ template <>
168+ __device__ inline void
169+ uct_rocm_ipc_copy_level<UCS_DEVICE_LEVEL_THREAD >(void *dst, const void *src,
170+ size_t len)
171+ {
172+ memcpy (dst, src, len);
173+ }
174+
175+ /* Warp- and block-level copy: strided iteration across lanes */
176+ template <>
177+ __device__ inline void
178+ uct_rocm_ipc_copy_level<UCS_DEVICE_LEVEL_WARP >(void *dst, const void *src,
179+ size_t len)
180+ {
181+ unsigned int lane_id, num_lanes;
182+ uct_rocm_ipc_get_lane<UCS_DEVICE_LEVEL_WARP >(lane_id, num_lanes);
183+ uct_rocm_ipc_copy_strided (dst, src, len, lane_id, num_lanes);
184+ }
185+
186+ template <>
187+ __device__ inline void
188+ uct_rocm_ipc_copy_level<UCS_DEVICE_LEVEL_BLOCK >(void *dst, const void *src,
189+ size_t len)
190+ {
191+ unsigned int lane_id, num_lanes;
192+ uct_rocm_ipc_get_lane<UCS_DEVICE_LEVEL_BLOCK >(lane_id, num_lanes);
193+ uct_rocm_ipc_copy_strided (dst, src, len, lane_id, num_lanes);
194+ }
195+
236196/* Grid-level copy - not implemented */
237197template <>
238198__device__ inline void
@@ -278,7 +238,7 @@ __device__ ucs_status_t uct_rocm_ipc_ep_atomic_add(
278238 if (lane_id == 0 ) {
279239 mapped_rem_addr = reinterpret_cast <uint64_t *>(
280240 uct_rocm_ipc_map_remote (rocm_ipc_mem_element, remote_address));
281- uct_rocm_ipc_atomic_inc (mapped_rem_addr, inc_value);
241+ uct_rocm_ipc_atomic_add (mapped_rem_addr, inc_value);
282242 }
283243
284244 uct_rocm_ipc_level_sync<level>();
0 commit comments