Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 35 additions & 13 deletions csrc/cpu/comm/ccl.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,15 @@
#include <iostream>
#include <oneapi/ccl.hpp>

// states for collectives
enum coll_state {
coll_begin = 0,
// coll states for naive allreduce
coll_allreduce_naive__copy_in_done, // this state is for rank != 0
coll_allreduce_naive__reduce_done, // this state is for rank == 0
coll_allreduce_naive__copy_out_done, // this state is for rank != 0
};

// SHM building blocks
struct SharedData {
const char* name;
Expand Down Expand Up @@ -63,19 +72,27 @@ void shared_close(SharedData* data)
#define SHM_BUFFER_NAME "deepspeed_allreduce_buffer"
SharedData allreduce_buffer;
struct allreduce_workspace {
int state;
enum coll_state state;
char buffer[MAX_BUF_SIZE];
};
struct allreduce_workspace* workspace;

void wait_buffer_state_until(int index, int state)
void wait_buffer_state_until(int index, enum coll_state state)
{
volatile int* state_ptr = &(workspace[index].state);
volatile enum coll_state* state_ptr = &(workspace[index].state);

while (*state_ptr != state)
;
}

void wait_buffer_state_until_not(int index, enum coll_state state)
{
volatile enum coll_state* state_ptr = &(workspace[index].state);

while (*state_ptr == state)
;
}

__m512 cvt_bf16_to_fp32(const __m256i src) __attribute__((target("avx512bw")));
inline __m512 cvt_bf16_to_fp32(const __m256i src)
{
Expand Down Expand Up @@ -313,7 +330,7 @@ void initialize(int size, int rank, torch::Tensor& kvs_data)
workspace,
size * sizeof(struct allreduce_workspace));
workspace = (struct allreduce_workspace*)allreduce_buffer.bytes;
for (int i = 0; i < size; i++) { workspace[i].state = 0; }
for (int i = 0; i < size; i++) { workspace[i].state = coll_begin; }
}
CCLCHECK(ccl::barrier(_get_comm_from_group()).wait());
if (rank != 0) {
Expand Down Expand Up @@ -506,33 +523,38 @@ void inference_all_reduce(torch::Tensor& data, py::object op, py::object group,

memcpy(workspace[world_rank].buffer, data_ptr, data_size);
std::atomic_thread_fence(std::memory_order_release);
workspace[world_rank].state = 1;
workspace[world_rank].state = coll_allreduce_naive__copy_in_done;

if (world_rank == 0) {
// compute allreduce result on rank 0
for (int i = 1; i < world_size; i++) {
// wait until the other rank copy the buffer
wait_buffer_state_until(i, 1);
wait_buffer_state_until(i, coll_allreduce_naive__copy_in_done);
}
reduce_all_buffers(workspace, numel, data.scalar_type(), world_size);
std::atomic_thread_fence(std::memory_order_release);
workspace[world_rank].state = 2;
workspace[world_rank].state = coll_allreduce_naive__reduce_done;
memcpy(data_ptr, workspace[0].buffer, data_size);
}
if (world_rank != 0) {
wait_buffer_state_until(0, 2);
wait_buffer_state_until(0, coll_allreduce_naive__reduce_done);
memcpy(data_ptr, workspace[0].buffer, data_size);
std::atomic_thread_fence(std::memory_order_release);
workspace[world_rank].state = 2;
workspace[world_rank].state = coll_allreduce_naive__copy_out_done;
}
if (world_rank == 0) {
for (int i = 1; i < world_size; i++) { wait_buffer_state_until(i, 2); }
for (int i = 1; i < world_size; i++) {
wait_buffer_state_until(i, coll_allreduce_naive__copy_out_done);
}
std::atomic_thread_fence(std::memory_order_release);
workspace[world_rank].state = 0;
workspace[world_rank].state = coll_begin;
}
if (world_rank != 0) {
wait_buffer_state_until(0, 0);
workspace[world_rank].state = 0;
// if rank 0 spin too fast it could be in state 1 of next allreduce
// in this case wait_buffer_state_until(0, 0) may cause deadlock
// what we are certain is when rank 0 finishes the state won't be 2
wait_buffer_state_until_not(0, coll_allreduce_naive__reduce_done);
workspace[world_rank].state = coll_begin;
}
}

Expand Down