Skip to content
Closed
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
109 changes: 67 additions & 42 deletions csrc/include/custom_all_reduce.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -908,9 +908,32 @@ template<typename T>
DINLINE T shfl_xor(T var, int mask, int width = opus::get_warp_size())
{
static_assert(sizeof(T) == 4);
int v = __builtin_bit_cast(int, var);
switch(mask)
{
case 1:
return __builtin_bit_cast(T, __builtin_amdgcn_mov_dpp(v, 0xb1, 0xf, 0xf, true));
case 2:
return __builtin_bit_cast(T, __builtin_amdgcn_mov_dpp(v, 0x4e, 0xf, 0xf, true));
case 4:
{
int r = __builtin_amdgcn_update_dpp(v, v, 0x104, 0xf, 0x5, true);
r = __builtin_amdgcn_update_dpp(r, v, 0x114, 0xf, 0xa, true);
Comment on lines +915 to +921
return __builtin_bit_cast(T, r);
}
case 8:
{
int r = __builtin_amdgcn_update_dpp(v, v, 0x108, 0xf, 0x3, true);
r = __builtin_amdgcn_update_dpp(r, v, 0x118, 0xf, 0xc, true);
Comment on lines +926 to +927
return __builtin_bit_cast(T, r);
}
Comment on lines +912 to +929
default:
break;
}
// fallback
int self = opus::lane_id();
int index = (self & ~(width - 1)) + ((self ^ mask) & (width - 1));
return __builtin_bit_cast(T, __builtin_amdgcn_ds_bpermute(index << 2, __builtin_bit_cast(int, var)));
return __builtin_bit_cast(T, __builtin_amdgcn_ds_bpermute(index << 2, v));
}

// shfl_xor support 4bytes dtype only
Expand Down Expand Up @@ -1092,13 +1115,8 @@ __global__ void __launch_bounds__(512, 1) reduce_scatter_cross_device_store(
RankData* _dp, RankSignals sg, Signal* self_sg, int rank, int size)
{
constexpr int pack_size = 16 / sizeof(T);
constexpr int tnum_gpu = THREAD_NUM / ngpus;
using P = typename opus::vector_t<T, pack_size>;
using A = typename opus::vector_t<opus::fp32_t, pack_size>;
__shared__ T tmp_smem[tnum_gpu * ngpus * pack_size];
int warp_id = threadIdx.x / tnum_gpu;
int lane_id = threadIdx.x % tnum_gpu;
int tid = blockIdx.x * tnum_gpu + lane_id;
const P* ptrs[ngpus];
P* tmps[ngpus];
#pragma unroll
Expand All @@ -1109,45 +1127,37 @@ __global__ void __launch_bounds__(512, 1) reduce_scatter_cross_device_store(
}
start_sync<ngpus>(sg, self_sg, rank);

// FLAT reduce-scatter + broadcast. This rank owns the pack range
// [rank*part, (rank+1)*part); each thread reduces ONE pack across all ngpus
// inputs in fp32 and stores the bf16 sum into every rank's tmp at the same
// offset (all-gather via broadcast). Each thread thus touches all ngpus links
// for both the reads and the ngpus broadcast stores; the ngpus stores/thread
// pipeline the write (broadcast) half far better than the previous
// warp-per-rank+LDS scheme's 1 store/thread. Measured faster at every m
// (0.91x at m=64 -> 0.86x at m=512 vs warp-per-rank) with bit-identical output
// (same canonical reduce order, same bf16 sum). No LDS, no intra-block sync.
int part = size / (pack_size * ngpus);
for(int idx = tid; idx < part; idx += gridDim.x * tnum_gpu)
int base = rank * part;
for(int l = blockIdx.x * blockDim.x + threadIdx.x; l < part;
l += gridDim.x * blockDim.x)
{
// cross device read by all warp
P input_reg = ptrs[warp_id][rank * part + idx];
*(reinterpret_cast<P*>(&tmp_smem[0]) + threadIdx.x) = input_reg;
__syncthreads();
// calculate and save in first warp
if(warp_id == 0)
{
A add_reg;
int gp = base + l;
A acc{};
#pragma unroll
for(int i = 0; i < pack_size; ++i)
{
add_reg[i] = upcast_s(tmp_smem[pack_size * threadIdx.x + i]);
}
for(int r = 0; r < ngpus; ++r)
{
P v = ptrs[r][gp];
#pragma unroll
for(int i = 1; i < ngpus; ++i)
{
for(int e = 0; e < pack_size; ++e)
acc[e] += upcast_s(v[e]);
}
P s;
#pragma unroll
for(int j = 0; j < pack_size; ++j)
{
add_reg[j] +=
upcast_s(tmp_smem[i * pack_size * tnum_gpu + pack_size * threadIdx.x + j]);
}
}
P add_rslt;
for(int e = 0; e < pack_size; ++e)
s[e] = downcast_s<T>(acc[e]);
#pragma unroll
for(int i = 0; i < pack_size; ++i)
{
add_rslt[i] = downcast_s<T>(add_reg[i]);
}
*(reinterpret_cast<P*>(&tmp_smem[0]) + lane_id) = add_rslt;
}
__syncthreads();

// cross device store
P rslt = *(reinterpret_cast<P*>(&tmp_smem[0]) + lane_id);
tmps[warp_id][rank * part + idx] = rslt;
for(int w = 0; w < ngpus; ++w)
tmps[w][gp] = s;
}
// NOTE: must use final_sync=false (RELEASE/ACQUIRE) here. Stage 2
// (local_device_load_rmsnorm*) on each rank reads `tmps` on the
Expand Down Expand Up @@ -3130,6 +3140,14 @@ void dispatchFusedAllReduceRMSNorm(hipStream_t stream,

auto pack_size = 16 / sizeof(T);
use_1stage = use_1stage && (n % pack_size == 0) && (n / pack_size <= 1024);
// Win-gate the path by data volume: one-stage's lower overhead (single barrier,
// no HBM round-trip) only wins for small inputs; above the measured crossover the
// FLAT two-stage's halved cross-card traffic (3S -> 1.5S) wins. Crossover ~0.44 MiB
// (m=32 @ hidden 7168); use 0.5 MiB with margin. So even if the caller permits
// one-stage, fall back to two-stage once we're past where it pays (e.g. m=64 decode
// = 0.875 MiB -> two-stage). This only DISABLES one-stage; it never overrides an
// explicit two-stage request (use_1stage already false).
use_1stage = use_1stage && ((size_t)size * sizeof(T) < (size_t)512 * 1024);
#define MAYBE_DISPATCH_1S_KERNEL(NGPUS) \
if(use_1stage) \
{ \
Expand All @@ -3153,22 +3171,29 @@ void dispatchFusedAllReduceRMSNorm(hipStream_t stream,
dim3 block(512);
int block_num = ((size / world_size_) + 512 - 1) / 512;
dim3 grid(std::min(block_num, 80));
// FLAT reduce_scatter_cross_device_store is 1 pack/thread: launch it with block 256
// (not the 512 used by step-2 rmsnorm) so small-m decode engages more CUs -- e.g.
// m=64 -> 56 active blocks vs 28 at block 512. grid covers this rank's pack range,
// capped at kMaxBlocks (start_sync's per-block slot limit).
int rs_packs = size / ((16 / sizeof(T)) * world_size_);
dim3 rs_block(256);
dim3 rs_grid(std::min((rs_packs + 255) / 256, 80));
switch(world_size_)
{
case 8:
MAYBE_DISPATCH_1S_KERNEL(8);
reduce_scatter_cross_device_store<T, 8>
<<<grid, block, 0, stream>>>(ptrs, sg_, self_sg_, rank_, size);
<<<rs_grid, rs_block, 0, stream>>>(ptrs, sg_, self_sg_, rank_, size);
break;
case 4:
MAYBE_DISPATCH_1S_KERNEL(4);
reduce_scatter_cross_device_store<T, 4>
<<<grid, block, 0, stream>>>(ptrs, sg_, self_sg_, rank_, size);
<<<rs_grid, rs_block, 0, stream>>>(ptrs, sg_, self_sg_, rank_, size);
break;
case 2:
MAYBE_DISPATCH_1S_KERNEL(2);
reduce_scatter_cross_device_store<T, 2>
<<<grid, block, 0, stream>>>(ptrs, sg_, self_sg_, rank_, size);
<<<rs_grid, rs_block, 0, stream>>>(ptrs, sg_, self_sg_, rank_, size);
break;
default: throw std::runtime_error("fused allreduce rmsnorm: unsupported world_size=" + std::to_string(world_size_));
}
Expand Down