| 50 | } |
| 51 | |
| 52 | Maybe<void> Recv(void* out, size_t elem_cnt, DataType dtype, int64_t src, DeviceType device_type, |
| 53 | ep::Stream* stream) { |
| 54 | if (GlobalProcessCtx::Rank() == src) { |
| 55 | size_t buffer_size = elem_cnt * GetSizeOfDataType(dtype); |
| 56 | auto** src_data_ptr = ThreadLocalSrcDataPtr(); |
| 57 | const void* in = *src_data_ptr; |
| 58 | CHECK_OR_RETURN(*src_data_ptr != nullptr); |
| 59 | std::unique_ptr<ep::primitive::Memcpy> memcpy_primitive = |
| 60 | ep::primitive::NewPrimitive<ep::primitive::MemcpyFactory>(device_type, |
| 61 | ep::primitive::MemcpyKind::kDtoD); |
| 62 | CHECK(memcpy_primitive) << "Can not create Memcpy primitive for device type " << device_type; |
| 63 | memcpy_primitive->Launch(stream, out, in, buffer_size); |
| 64 | *src_data_ptr = nullptr; |
| 65 | } else { |
| 66 | std::unique_ptr<ccl::Recv> recv = |
| 67 | ccl::NewCollectiveCommunication<ccl::Recv>(device_type, dtype); |
| 68 | recv->Launch(stream, out, elem_cnt, src); |
| 69 | } |
| 70 | return Maybe<void>::Ok(); |
| 71 | } |
| 72 | |
| 73 | } // namespace oneflow |