Skip to content

Commit 1856e0e

Browse files
committed
align 16
1 parent cf73485 commit 1856e0e

2 files changed

Lines changed: 100 additions & 44 deletions

File tree

include/mori/core/transport/p2p/device_primitives.hpp

Lines changed: 95 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -1407,48 +1407,104 @@ __device__ __forceinline__ void WarpAccumBf16Fp8Mixed(OutT* __restrict__ outTok,
14071407
const uint8_t* const* __restrict__ fp8Ptrs,
14081408
int nNodes, int myNode, int hiddenDimSize) {
14091409
const int laneId = threadIdx.x & (warpSize - 1);
1410-
const int vecEnd4 = (hiddenDimSize / 4) * 4;
1411-
for (int idx = laneId * 4; idx < vecEnd4; idx += warpSize * 4) {
1412-
float2 sum01 = float2{0.0f, 0.0f};
1413-
float2 sum23 = float2{0.0f, 0.0f};
1414-
if (localTok != nullptr) {
1415-
const uint64_t packed8 = load<8>(localTok + idx);
1416-
sum01.x = Bf16BitsToF32(static_cast<uint16_t>(packed8 & 0xFFFF));
1417-
sum01.y = Bf16BitsToF32(static_cast<uint16_t>((packed8 >> 16) & 0xFFFF));
1418-
sum23.x = Bf16BitsToF32(static_cast<uint16_t>((packed8 >> 32) & 0xFFFF));
1419-
sum23.y = Bf16BitsToF32(static_cast<uint16_t>((packed8 >> 48) & 0xFFFF));
1420-
}
14211410

1422-
for (int n = 0; n < nNodes; n++) {
1423-
if (n == myNode) continue;
1424-
const uint8_t* fp8TokBytes = fp8Ptrs[n];
1425-
if (fp8TokBytes == nullptr) continue;
1426-
const uint32_t packed4 = static_cast<uint32_t>(load<4>(fp8TokBytes + idx));
1427-
const __hip_fp8x2_storage_t p01 =
1428-
static_cast<__hip_fp8x2_storage_t>(static_cast<uint16_t>(packed4 & 0xFFFF));
1429-
const __hip_fp8x2_storage_t p23 =
1430-
static_cast<__hip_fp8x2_storage_t>(static_cast<uint16_t>(packed4 >> 16));
1431-
const float2 v01 = CvtFp8x2ToFloat2<Fp8T>(p01);
1432-
const float2 v23 = CvtFp8x2ToFloat2<Fp8T>(p23);
1433-
sum01.x += v01.x;
1434-
sum01.y += v01.y;
1435-
sum23.x += v23.x;
1436-
sum23.y += v23.y;
1411+
// Use 8-element path when hiddenDimSize is large enough for good lane utilization
1412+
if (hiddenDimSize >= warpSize * 8) {
1413+
const int vecEnd8 = (hiddenDimSize / 8) * 8;
1414+
for (int idx = laneId * 8; idx < vecEnd8; idx += warpSize * 8) {
1415+
float2 sum01{0.0f, 0.0f}, sum23{0.0f, 0.0f}, sum45{0.0f, 0.0f}, sum67{0.0f, 0.0f};
1416+
if (localTok != nullptr) {
1417+
const ulong2 packed16 = load<16>(localTok + idx);
1418+
const uint64_t p0 = packed16.x, p1 = packed16.y;
1419+
sum01.x = Bf16BitsToF32(static_cast<uint16_t>(p0 & 0xFFFF));
1420+
sum01.y = Bf16BitsToF32(static_cast<uint16_t>((p0 >> 16) & 0xFFFF));
1421+
sum23.x = Bf16BitsToF32(static_cast<uint16_t>((p0 >> 32) & 0xFFFF));
1422+
sum23.y = Bf16BitsToF32(static_cast<uint16_t>((p0 >> 48) & 0xFFFF));
1423+
sum45.x = Bf16BitsToF32(static_cast<uint16_t>(p1 & 0xFFFF));
1424+
sum45.y = Bf16BitsToF32(static_cast<uint16_t>((p1 >> 16) & 0xFFFF));
1425+
sum67.x = Bf16BitsToF32(static_cast<uint16_t>((p1 >> 32) & 0xFFFF));
1426+
sum67.y = Bf16BitsToF32(static_cast<uint16_t>((p1 >> 48) & 0xFFFF));
1427+
}
1428+
for (int n = 0; n < nNodes; n++) {
1429+
if (n == myNode) continue;
1430+
const uint8_t* fp8TokBytes = fp8Ptrs[n];
1431+
if (fp8TokBytes == nullptr) continue;
1432+
const uint64_t packed8 = static_cast<uint64_t>(load<8>(fp8TokBytes + idx));
1433+
const float2 v01 = CvtFp8x2ToFloat2<Fp8T>(
1434+
static_cast<__hip_fp8x2_storage_t>(static_cast<uint16_t>(packed8 & 0xFFFF)));
1435+
const float2 v23 = CvtFp8x2ToFloat2<Fp8T>(
1436+
static_cast<__hip_fp8x2_storage_t>(static_cast<uint16_t>((packed8 >> 16) & 0xFFFF)));
1437+
const float2 v45 = CvtFp8x2ToFloat2<Fp8T>(
1438+
static_cast<__hip_fp8x2_storage_t>(static_cast<uint16_t>((packed8 >> 32) & 0xFFFF)));
1439+
const float2 v67 = CvtFp8x2ToFloat2<Fp8T>(
1440+
static_cast<__hip_fp8x2_storage_t>(static_cast<uint16_t>((packed8 >> 48) & 0xFFFF)));
1441+
sum01.x += v01.x;
1442+
sum01.y += v01.y;
1443+
sum23.x += v23.x;
1444+
sum23.y += v23.y;
1445+
sum45.x += v45.x;
1446+
sum45.y += v45.y;
1447+
sum67.x += v67.x;
1448+
sum67.y += v67.y;
1449+
}
1450+
StoreOutPair<OutT>(outTok, idx, sum01);
1451+
StoreOutPair<OutT>(outTok, idx + 2, sum23);
1452+
StoreOutPair<OutT>(outTok, idx + 4, sum45);
1453+
StoreOutPair<OutT>(outTok, idx + 6, sum67);
14371454
}
1438-
StoreOutPair<OutT>(outTok, idx, sum01);
1439-
StoreOutPair<OutT>(outTok, idx + 2, sum23);
1440-
}
1441-
for (int idx = vecEnd4 + laneId; idx < hiddenDimSize; idx += warpSize) {
1442-
float sum = 0.0f;
1443-
if (localTok != nullptr) sum += static_cast<float>(localTok[idx]);
1444-
for (int n = 0; n < nNodes; n++) {
1445-
if (n == myNode) continue;
1446-
const uint8_t* fp8TokBytes = fp8Ptrs[n];
1447-
if (fp8TokBytes == nullptr) continue;
1448-
const auto* fp8Tok = reinterpret_cast<const Fp8T*>(fp8TokBytes);
1449-
sum += static_cast<float>(fp8Tok[idx]);
1455+
for (int idx = vecEnd8 + laneId; idx < hiddenDimSize; idx += warpSize) {
1456+
float sum = 0.0f;
1457+
if (localTok != nullptr) sum += static_cast<float>(localTok[idx]);
1458+
for (int n = 0; n < nNodes; n++) {
1459+
if (n == myNode) continue;
1460+
const uint8_t* fp8TokBytes = fp8Ptrs[n];
1461+
if (fp8TokBytes == nullptr) continue;
1462+
sum += static_cast<float>(reinterpret_cast<const Fp8T*>(fp8TokBytes)[idx]);
1463+
}
1464+
outTok[idx] = OutT(sum);
1465+
}
1466+
} else {
1467+
const int vecEnd4 = (hiddenDimSize / 4) * 4;
1468+
for (int idx = laneId * 4; idx < vecEnd4; idx += warpSize * 4) {
1469+
float2 sum01 = float2{0.0f, 0.0f};
1470+
float2 sum23 = float2{0.0f, 0.0f};
1471+
if (localTok != nullptr) {
1472+
const uint64_t packed8 = load<8>(localTok + idx);
1473+
sum01.x = Bf16BitsToF32(static_cast<uint16_t>(packed8 & 0xFFFF));
1474+
sum01.y = Bf16BitsToF32(static_cast<uint16_t>((packed8 >> 16) & 0xFFFF));
1475+
sum23.x = Bf16BitsToF32(static_cast<uint16_t>((packed8 >> 32) & 0xFFFF));
1476+
sum23.y = Bf16BitsToF32(static_cast<uint16_t>((packed8 >> 48) & 0xFFFF));
1477+
}
1478+
for (int n = 0; n < nNodes; n++) {
1479+
if (n == myNode) continue;
1480+
const uint8_t* fp8TokBytes = fp8Ptrs[n];
1481+
if (fp8TokBytes == nullptr) continue;
1482+
const uint32_t packed4 = static_cast<uint32_t>(load<4>(fp8TokBytes + idx));
1483+
const __hip_fp8x2_storage_t p01 =
1484+
static_cast<__hip_fp8x2_storage_t>(static_cast<uint16_t>(packed4 & 0xFFFF));
1485+
const __hip_fp8x2_storage_t p23 =
1486+
static_cast<__hip_fp8x2_storage_t>(static_cast<uint16_t>(packed4 >> 16));
1487+
const float2 v01 = CvtFp8x2ToFloat2<Fp8T>(p01);
1488+
const float2 v23 = CvtFp8x2ToFloat2<Fp8T>(p23);
1489+
sum01.x += v01.x;
1490+
sum01.y += v01.y;
1491+
sum23.x += v23.x;
1492+
sum23.y += v23.y;
1493+
}
1494+
StoreOutPair<OutT>(outTok, idx, sum01);
1495+
StoreOutPair<OutT>(outTok, idx + 2, sum23);
1496+
}
1497+
for (int idx = vecEnd4 + laneId; idx < hiddenDimSize; idx += warpSize) {
1498+
float sum = 0.0f;
1499+
if (localTok != nullptr) sum += static_cast<float>(localTok[idx]);
1500+
for (int n = 0; n < nNodes; n++) {
1501+
if (n == myNode) continue;
1502+
const uint8_t* fp8TokBytes = fp8Ptrs[n];
1503+
if (fp8TokBytes == nullptr) continue;
1504+
sum += static_cast<float>(reinterpret_cast<const Fp8T*>(fp8TokBytes)[idx]);
1505+
}
1506+
outTok[idx] = OutT(sum);
14501507
}
1451-
outTok[idx] = OutT(sum);
14521508
}
14531509
}
14541510

src/ops/dispatch_combine/internode_v1.cpp

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -630,7 +630,7 @@ __global__ void EpDispatchCopyToStaging(EpDispatchCombineArgs<T> args) {
630630
if (args.curRankNumToken == 0) return;
631631

632632
index_t warpsPerToken = (globalWarpNum + args.curRankNumToken - 1) / args.curRankNumToken;
633-
index_t hiddenDimPerWarp = (config.hiddenDim + warpsPerToken - 1) / warpsPerToken;
633+
index_t hiddenDimPerWarp = ((config.hiddenDim + warpsPerToken - 1) / warpsPerToken + 15) & ~15;
634634

635635
// First copy to staging buffer
636636
for (int i = globalWarpId; i < (args.curRankNumToken * warpsPerToken); i += globalWarpNum) {
@@ -767,7 +767,7 @@ inline __device__ void CombineIntraNodeLL(EpDispatchCombineArgs<T>& args) {
767767
(nNodes + myNode) * config.MaxNumTokensToRecvPerRank() * combXferBytes;
768768

769769
index_t warpsPerToken = (xgmiWarpNum + args.curRankNumToken - 1) / args.curRankNumToken;
770-
index_t hiddenDimPerWarp = (config.hiddenDim + warpsPerToken - 1) / warpsPerToken;
770+
index_t hiddenDimPerWarp = ((config.hiddenDim + warpsPerToken - 1) / warpsPerToken + 15) & ~15;
771771

772772
for (int i = globalWarpId - blockOffset * warpNum; i < (args.curRankNumToken * warpsPerToken);
773773
i += xgmiWarpNum) {
@@ -1046,7 +1046,7 @@ inline __device__ void CombineInterNodeLL(EpDispatchCombineArgs<T>& args) {
10461046
// int warpsPerToken = (rdmaWarpNum + nodeCount - 1) / nodeCount;
10471047
// NOTE: Keep a small fixed value to reduce overhead; use fewer warps for FP8 path.
10481048
int warpsPerToken = useInternalFp8 ? 2 : 4;
1049-
int hiddenDimPerWarp = (config.hiddenDim + warpsPerToken - 1) / warpsPerToken;
1049+
int hiddenDimPerWarp = ((config.hiddenDim + warpsPerToken - 1) / warpsPerToken + 15) & ~15;
10501050

10511051
for (int i = globalWarpId; i < (nodeCount * warpsPerToken); i += rdmaWarpNum) {
10521052
int tokenId = i / warpsPerToken;
@@ -1225,7 +1225,7 @@ inline __device__ void CombineAll(EpDispatchCombineArgs<T>& args) {
12251225
nNodes * config.MaxNumTokensToRecvPerRank() * combXferBytes;
12261226

12271227
index_t warpsPerToken = (globalWarpNum + args.curRankNumToken - 1) / args.curRankNumToken;
1228-
index_t hiddenDimPerWarp = (config.hiddenDim + warpsPerToken - 1) / warpsPerToken;
1228+
index_t hiddenDimPerWarp = ((config.hiddenDim + warpsPerToken - 1) / warpsPerToken + 15) & ~15;
12291229

12301230
for (int i = globalWarpId; i < (args.curRankNumToken * warpsPerToken); i += globalWarpNum) {
12311231
index_t tokenId = i / warpsPerToken;
@@ -1307,7 +1307,7 @@ __global__ void EpCombineAll(EpDispatchCombineArgs<T> args) {
13071307
#endif
13081308

13091309
index_t warpsPerToken = (globalWarpNum + args.curRankNumToken - 1) / args.curRankNumToken;
1310-
index_t hiddenDimPerWarp = (config.hiddenDim + warpsPerToken - 1) / warpsPerToken;
1310+
index_t hiddenDimPerWarp = ((config.hiddenDim + warpsPerToken - 1) / warpsPerToken + 15) & ~15;
13111311

13121312
for (int i = globalWarpId; i < (args.curRankNumToken * warpsPerToken); i += globalWarpNum) {
13131313
index_t tokenId = i / warpsPerToken;

0 commit comments

Comments
 (0)