diff --git a/src/qxmx_gpu.cpp b/src/qxmx_gpu.cpp index dadb77d..6c308ba 100644 --- a/src/qxmx_gpu.cpp +++ b/src/qxmx_gpu.cpp @@ -131,7 +131,15 @@ static void gemv_dpas8(sycl::queue& q, const uint8_t* A_dpas, const float* A_sca constexpr int LWS = WGM8; int NGM = (M8 + WGM8 - 1) / WGM8; int S = 1; - while (NGM * S < 64 && S < 8 && ngrp % (GPP * S * 2) == 0) S *= 2; + /* Raise split-K to hide latency: the decode GEMVs are 0.27-0.31 WGs/EU at + * the old `NGM*S<64` cap, starving the EUs (50% idle @ 100% occupancy). + * Target ~1 WG/EU (256). Only double while ngrp stays divisible by 2*S, + * which guarantees gps=ngrp/S is integral and [0,ngrp) is fully covered + * (no dropped groups -> qxmx_diff stable). Partial last GPP-tile is + * handled by the gmax guard below. See latency-hunt-plan.md. */ + constexpr int GV_SMAX = 16; + constexpr int GV_TGTWGS = 256; /* 1 WG/EU on B70 (256 EU) */ + while (S < GV_SMAX && ngrp % (S * 2) == 0 && NGM * S * 2 <= GV_TGTWGS) S *= 2; int G = NGM * S * LWS; int gps = ngrp / S; auto props = esimd::properties{esimd::cache_hint_L1, @@ -229,7 +237,10 @@ static void gemv_dpas8_dual(sycl::queue& q, const uint8_t* A0_dpas, constexpr int LWS = WGM8; int NGM = (M8 + WGM8 - 1) / WGM8; int S = 1; - while (NGM * S < 64 && S < 8 && ngrp % (GPP * S * 2) == 0) S *= 2; + /* Same raised split-K heuristic as gemv_dpas8 (see comment there). */ + constexpr int GV_SMAX = 16; + constexpr int GV_TGTWGS = 256; /* 1 WG/EU on B70 (256 EU) */ + while (S < GV_SMAX && ngrp % (S * 2) == 0 && NGM * S * 2 <= GV_TGTWGS) S *= 2; int G = NGM * S * LWS; int gps = ngrp / S; auto props = esimd::properties{esimd::cache_hint_L1,