Start fixing ring
This commit is contained in:
@@ -54,7 +54,7 @@ class RingImpl {
|
||||
ReduceOp reduce_op) {
|
||||
constexpr int PIPELINE = 2;
|
||||
constexpr int WC_NUM = PIPELINE * RING_MAX_CONNS * 2 * MAX_DIR;
|
||||
int64_t chunk_size = (size + size_ - 1) / size_;
|
||||
int64_t chunk_size = size / size_;
|
||||
int64_t size_per_wire =
|
||||
(chunk_size + (MAX_DIR * n_wires) - 1) / (MAX_DIR * n_wires);
|
||||
auto [sz, N] = buffer_size_from_message(size_per_wire * sizeof(T));
|
||||
@@ -63,18 +63,17 @@ class RingImpl {
|
||||
|
||||
// Counters to maintain the state of transfers
|
||||
int in_flight = 0;
|
||||
int64_t chunk_multiple_size = size_ * chunk_size;
|
||||
int64_t send_offset[MAX_DIR];
|
||||
int64_t recv_offset[MAX_DIR];
|
||||
int64_t send_limits[MAX_DIR];
|
||||
int64_t recv_limits[MAX_DIR];
|
||||
int send_count[MAX_DIR * RING_MAX_CONNS] = {0};
|
||||
int recv_count[MAX_DIR * RING_MAX_CONNS] = {0};
|
||||
send_offset[0] = rank_ * chunk_size;
|
||||
recv_offset[0] = ((rank_ + size_ - 1) % size_) * chunk_size;
|
||||
send_offset[0] = ((rank_ + size_ - 1) % size_) * chunk_size;
|
||||
recv_offset[0] = ((rank_ + size_ - 2) % size_) * chunk_size;
|
||||
if constexpr (MAX_DIR == 2) {
|
||||
send_offset[1] = rank_ * chunk_size;
|
||||
recv_offset[1] = ((rank_ + 1) % size_) * chunk_size;
|
||||
send_offset[1] = ((rank_ + 1) % size_) * chunk_size;
|
||||
recv_offset[1] = ((rank_ + 2) % size_) * chunk_size;
|
||||
send_limits[0] = std::min(
|
||||
n_wires * size_per_wire, std::max<int64_t>(0, size - send_offset[0]));
|
||||
send_limits[1] =
|
||||
@@ -91,30 +90,49 @@ class RingImpl {
|
||||
}
|
||||
|
||||
for (int k = 0; k < size_ - 1; k++) {
|
||||
const T* read_ptr = (k == 0) ? in_ptr : out_ptr;
|
||||
int64_t read_offset[MAX_DIR];
|
||||
read_offset[0] = (k == 0) ? send_offset[0] : 0;
|
||||
if constexpr (MAX_DIR == 2) {
|
||||
read_offset[1] = (k == 0) ? send_offset[1] : 0;
|
||||
}
|
||||
|
||||
// Prefill the pipeline
|
||||
int buff = 0;
|
||||
while (buff < n_steps && buff < PIPELINE) {
|
||||
// Post receives
|
||||
post_recv_all<MAX_DIR>(sz, buff, n_wires);
|
||||
in_flight += MAX_DIR * n_wires;
|
||||
|
||||
// Copy to send buffers and also to the output buffer
|
||||
for (int lr = 0; lr < MAX_DIR; lr++) {
|
||||
for (int lw = 0; lw < n_wires; lw++) {
|
||||
// send
|
||||
int64_t offset = lw * N +
|
||||
send_count[lr * RING_MAX_CONNS + lw] * n_wires * N +
|
||||
lr * n_wires * size_per_wire;
|
||||
std::copy(
|
||||
in_ptr + send_offset[lr] + offset,
|
||||
in_ptr + send_offset[lr] +
|
||||
read_ptr + read_offset[lr] + offset,
|
||||
read_ptr + read_offset[lr] +
|
||||
std::max(offset, std::min(offset + N, send_limits[lr])),
|
||||
send_buffer(sz, buff, lr, lw).begin<T>());
|
||||
send_count[lr * RING_MAX_CONNS + lw]++;
|
||||
in_flight++;
|
||||
send_to(sz, buff, lr, lw);
|
||||
|
||||
// copy for in-place recv reduce
|
||||
offset -= n_wires * N;
|
||||
std::copy(
|
||||
in_ptr + recv_offset[lr] + offset,
|
||||
in_ptr + recv_offset[lr] +
|
||||
std::max(offset, std::min(offset + N, recv_limits[lr])),
|
||||
out_ptr + offset);
|
||||
}
|
||||
}
|
||||
post_send_all<MAX_DIR>(sz, buff, n_wires);
|
||||
|
||||
buff++;
|
||||
in_flight += 2 * MAX_DIR * n_wires;
|
||||
}
|
||||
|
||||
// Main loop
|
||||
while (in_flight > 0) {
|
||||
ibv_wc wc[WC_NUM];
|
||||
int n = poll(left_, right_, WC_NUM, wc);
|
||||
@@ -127,17 +145,29 @@ class RingImpl {
|
||||
|
||||
in_flight--;
|
||||
|
||||
if (work_type == SEND_WR && send_count[wire] < n_steps) {
|
||||
if (work_type == SEND_WR) {
|
||||
int64_t offset = lw * N + send_count[wire] * n_wires * N +
|
||||
lr * n_wires * size_per_wire;
|
||||
|
||||
// More stuff to send
|
||||
if (send_count[wire] < n_steps) {
|
||||
std::copy(
|
||||
read_ptr + read_offset[lr] + offset,
|
||||
read_ptr + read_offset[lr] +
|
||||
std::max(offset, std::min(offset + N, send_limits[lr])),
|
||||
send_buffer(sz, buff, lr, lw).begin<T>());
|
||||
send_to(sz, buff, lr, lw);
|
||||
in_flight++;
|
||||
send_count[wire]++;
|
||||
}
|
||||
|
||||
// Copy the input chunk into the output
|
||||
offset -= n_wires * N;
|
||||
std::copy(
|
||||
in_ptr + send_offset[lr] + offset,
|
||||
in_ptr + send_offset[lr] +
|
||||
std::max(offset, std::min(offset + N, send_limits[lr])),
|
||||
send_buffer(sz, buff, lr, lw).begin<T>());
|
||||
send_to(sz, buff, lr, lw);
|
||||
in_flight++;
|
||||
send_count[wire]++;
|
||||
in_ptr + recv_offset[lr] + offset,
|
||||
in_ptr + recv_offset[lr] +
|
||||
std::max(offset, std::min(offset + N, recv_limits[lr])),
|
||||
out_ptr + offset);
|
||||
}
|
||||
|
||||
else if (work_type == RECV_WR) {
|
||||
@@ -156,13 +186,11 @@ class RingImpl {
|
||||
}
|
||||
}
|
||||
|
||||
send_offset[0] = (send_offset[0] + chunk_multiple_size - chunk_size) %
|
||||
chunk_multiple_size;
|
||||
recv_offset[0] = (recv_offset[0] + chunk_multiple_size - chunk_size) %
|
||||
chunk_multiple_size;
|
||||
send_offset[0] = (send_offset[0] + size - chunk_size) % size;
|
||||
recv_offset[0] = (recv_offset[0] + size - chunk_size) % size;
|
||||
if constexpr (MAX_DIR == 2) {
|
||||
send_offset[1] = (send_offset[1] + chunk_size) % chunk_multiple_size;
|
||||
recv_offset[1] = (recv_offset[1] + chunk_size) % chunk_multiple_size;
|
||||
send_offset[1] = (send_offset[1] + chunk_size) % size;
|
||||
recv_offset[1] = (recv_offset[1] + chunk_size) % size;
|
||||
send_limits[0] = std::min(
|
||||
n_wires * size_per_wire,
|
||||
std::max<int64_t>(0, size - send_offset[0]));
|
||||
|
||||
Reference in New Issue
Block a user