@@ -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
0 commit comments