mmq.cu 85 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667686970717273747576777879808182838485868788899091929394959697989910010110210310410510610710810911011111211311411511611711811912012112212312412512612712812913013113213313413513613713813914014114214314414514614714814915015115215315415515615715815916016116216316416516616716816917017117217317417517617717817918018118218318418518618718818919019119219319419519619719819920020120220320420520620720820921021121221321421521621721821922022122222322422522622722822923023123223323423523623723823924024124224324424524624724824925025125225325425525625725825926026126226326426526626726826927027127227327427527627727827928028128228328428528628728828929029129229329429529629729829930030130230330430530630730830931031131231331431531631731831932032132232332432532632732832933033133233333433533633733833934034134234334434534634734834935035135235335435535635735835936036136236336436536636736836937037137237337437537637737837938038138238338438538638738838939039139239339439539639739839940040140240340440540640740840941041141241341441541641741841942042142242342442542642742842943043143243343443543643743843944044144244344444544644744844945045145245345445545645745845946046146246346446546646746846947047147247347447547647747847948048148248348448548648748848949049149249349449549649749849950050150250350450550650750850951051151251351451551651751851952052152252352452552652752852953053153253353453553653753853954054154254354454554654754854955055155255355455555655755855956056156256356456556656756856957057157257357457557657757857958058158258358458558658758858959059159259359459559659759859960060160260360460560660760860961061161261361461561661761861962062162262362462562662762862963063163263363463563663763863964064164264364464564664764864965065165265365465565665765865966066166266366466566666766866967067167267367467567667767867968068168268368468568668768868969069169269369469569669769869970070170270370470570670770870971071171271371471571671771871972072172272372472572672772872973073173273373473573673773873974074174274374474574674774874975075175275375475575675775875976076176276376476576676776876977077177277377477577677777877978078178278378478578678778878979079179279379479579679779879980080180280380480580680780880981081181281381481581681781881982082182282382482582682782882983083183283383483583683783883984084184284384484584684784884985085185285385485585685785885986086186286386486586686786886987087187287387487587687787887988088188288388488588688788888989089189289389489589689789889990090190290390490590690790890991091191291391491591691791891992092192292392492592692792892993093193293393493593693793893994094194294394494594694794894995095195295395495595695795895996096196296396496596696796896997097197297397497597697797897998098198298398498598698798898999099199299399499599699799899910001001100210031004100510061007100810091010101110121013101410151016101710181019102010211022102310241025102610271028102910301031103210331034103510361037103810391040104110421043104410451046104710481049105010511052105310541055105610571058105910601061106210631064106510661067106810691070107110721073107410751076107710781079108010811082108310841085108610871088108910901091109210931094109510961097109810991100110111021103110411051106110711081109111011111112111311141115111611171118111911201121112211231124112511261127112811291130113111321133113411351136113711381139114011411142114311441145114611471148114911501151115211531154115511561157115811591160116111621163116411651166116711681169117011711172117311741175117611771178117911801181118211831184118511861187118811891190119111921193119411951196119711981199120012011202120312041205120612071208120912101211121212131214121512161217121812191220122112221223122412251226122712281229123012311232123312341235123612371238123912401241124212431244124512461247124812491250125112521253125412551256125712581259126012611262126312641265126612671268126912701271127212731274127512761277127812791280128112821283128412851286128712881289129012911292129312941295129612971298129913001301130213031304130513061307130813091310131113121313131413151316131713181319132013211322132313241325132613271328132913301331133213331334133513361337133813391340134113421343134413451346134713481349135013511352135313541355135613571358135913601361136213631364136513661367136813691370137113721373137413751376137713781379138013811382138313841385138613871388138913901391139213931394139513961397139813991400140114021403140414051406140714081409141014111412141314141415141614171418141914201421142214231424142514261427142814291430143114321433143414351436143714381439144014411442144314441445144614471448144914501451145214531454145514561457145814591460146114621463146414651466146714681469147014711472147314741475147614771478147914801481148214831484148514861487148814891490149114921493149414951496149714981499150015011502150315041505150615071508150915101511151215131514151515161517151815191520152115221523152415251526152715281529153015311532153315341535153615371538153915401541154215431544154515461547154815491550155115521553155415551556155715581559156015611562156315641565156615671568156915701571157215731574157515761577157815791580158115821583158415851586158715881589159015911592159315941595159615971598159916001601160216031604160516061607160816091610161116121613161416151616161716181619162016211622162316241625162616271628162916301631163216331634163516361637163816391640164116421643164416451646164716481649165016511652165316541655165616571658165916601661166216631664166516661667166816691670167116721673167416751676167716781679168016811682168316841685168616871688168916901691169216931694169516961697169816991700170117021703170417051706170717081709171017111712171317141715171617171718171917201721172217231724172517261727172817291730173117321733173417351736173717381739174017411742174317441745174617471748174917501751175217531754175517561757175817591760176117621763176417651766176717681769177017711772177317741775177617771778177917801781178217831784178517861787178817891790179117921793179417951796179717981799180018011802180318041805180618071808180918101811181218131814181518161817181818191820182118221823182418251826182718281829183018311832183318341835183618371838183918401841184218431844184518461847184818491850185118521853185418551856185718581859186018611862186318641865186618671868186918701871187218731874187518761877187818791880188118821883188418851886188718881889189018911892189318941895189618971898189919001901190219031904190519061907190819091910191119121913191419151916191719181919192019211922192319241925192619271928192919301931193219331934193519361937193819391940194119421943194419451946194719481949195019511952195319541955195619571958195919601961196219631964196519661967196819691970197119721973197419751976197719781979198019811982198319841985198619871988198919901991199219931994199519961997199819992000200120022003200420052006200720082009201020112012201320142015201620172018201920202021202220232024202520262027202820292030203120322033203420352036203720382039204020412042204320442045204620472048204920502051205220532054205520562057205820592060206120622063206420652066206720682069207020712072207320742075207620772078207920802081208220832084208520862087208820892090209120922093209420952096209720982099210021012102210321042105210621072108210921102111211221132114211521162117211821192120212121222123212421252126212721282129213021312132213321342135213621372138213921402141214221432144214521462147214821492150215121522153215421552156215721582159216021612162216321642165216621672168216921702171217221732174217521762177217821792180218121822183218421852186218721882189219021912192219321942195219621972198219922002201220222032204220522062207220822092210221122122213221422152216221722182219222022212222222322242225222622272228222922302231223222332234223522362237223822392240224122422243224422452246224722482249225022512252225322542255
  1. #include "mmq.cuh"
  2. #include "vecdotq.cuh"
  3. typedef void (*allocate_tiles_cuda_t)(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc);
  4. typedef void (*load_tiles_cuda_t)(
  5. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  6. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row);
  7. typedef float (*vec_dot_q_mul_mat_cuda_t)(
  8. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  9. const int * __restrict__ y_qs, const half2 * __restrict__ y_ms, const int & i, const int & j, const int & k);
  10. typedef void (*dot_kernel_k_t)(const void * __restrict__ vx, const int ib, const int iqs, const float * __restrict__ y, float & v);
  11. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q4_0(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  12. GGML_UNUSED(x_qh);
  13. GGML_UNUSED(x_sc);
  14. __shared__ int tile_x_qs[mmq_y * (WARP_SIZE) + mmq_y];
  15. __shared__ float tile_x_d[mmq_y * (WARP_SIZE/QI4_0) + mmq_y/QI4_0];
  16. *x_ql = tile_x_qs;
  17. *x_dm = (half2 *) tile_x_d;
  18. }
  19. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q4_0(
  20. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  21. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  22. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  23. GGML_CUDA_ASSUME(i_offset >= 0);
  24. GGML_CUDA_ASSUME(i_offset < nwarps);
  25. GGML_CUDA_ASSUME(k >= 0);
  26. GGML_CUDA_ASSUME(k < WARP_SIZE);
  27. const int kbx = k / QI4_0;
  28. const int kqsx = k % QI4_0;
  29. const block_q4_0 * bx0 = (const block_q4_0 *) vx;
  30. float * x_dmf = (float *) x_dm;
  31. #pragma unroll
  32. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  33. int i = i0 + i_offset;
  34. if (need_check) {
  35. i = min(i, i_max);
  36. }
  37. const block_q4_0 * bxi = bx0 + i*blocks_per_row + kbx;
  38. x_ql[i * (WARP_SIZE + 1) + k] = get_int_from_uint8(bxi->qs, kqsx);
  39. // x_dmf[i * (WARP_SIZE/QI4_0) + i / QI4_0 + kbx] = bxi->d;
  40. }
  41. const int blocks_per_tile_x_row = WARP_SIZE / QI4_0;
  42. const int kbxd = k % blocks_per_tile_x_row;
  43. #pragma unroll
  44. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI4_0) {
  45. int i = i0 + i_offset * QI4_0 + k / blocks_per_tile_x_row;
  46. if (need_check) {
  47. i = min(i, i_max);
  48. }
  49. const block_q4_0 * bxi = bx0 + i*blocks_per_row + kbxd;
  50. x_dmf[i * (WARP_SIZE/QI4_0) + i / QI4_0 + kbxd] = bxi->d;
  51. }
  52. }
  53. static __device__ __forceinline__ float vec_dot_q4_0_q8_1_mul_mat(
  54. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  55. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  56. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  57. const int kyqs = k % (QI8_1/2) + QI8_1 * (k / (QI8_1/2));
  58. const float * x_dmf = (const float *) x_dm;
  59. int u[2*VDR_Q4_0_Q8_1_MMQ];
  60. #pragma unroll
  61. for (int l = 0; l < VDR_Q4_0_Q8_1_MMQ; ++l) {
  62. u[2*l+0] = y_qs[j * WARP_SIZE + (kyqs + l) % WARP_SIZE];
  63. u[2*l+1] = y_qs[j * WARP_SIZE + (kyqs + l + QI4_0) % WARP_SIZE];
  64. }
  65. return vec_dot_q4_0_q8_1_impl<VDR_Q4_0_Q8_1_MMQ>
  66. (&x_ql[i * (WARP_SIZE + 1) + k], u, x_dmf[i * (WARP_SIZE/QI4_0) + i/QI4_0 + k/QI4_0],
  67. y_ds[j * (WARP_SIZE/QI8_1) + (2*k/QI8_1) % (WARP_SIZE/QI8_1)]);
  68. }
  69. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q4_1(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  70. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  71. __shared__ int tile_x_qs[mmq_y * (WARP_SIZE) + + mmq_y];
  72. __shared__ half2 tile_x_dm[mmq_y * (WARP_SIZE/QI4_1) + mmq_y/QI4_1];
  73. *x_ql = tile_x_qs;
  74. *x_dm = tile_x_dm;
  75. }
  76. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q4_1(
  77. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  78. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  79. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  80. GGML_CUDA_ASSUME(i_offset >= 0);
  81. GGML_CUDA_ASSUME(i_offset < nwarps);
  82. GGML_CUDA_ASSUME(k >= 0);
  83. GGML_CUDA_ASSUME(k < WARP_SIZE);
  84. const int kbx = k / QI4_1;
  85. const int kqsx = k % QI4_1;
  86. const block_q4_1 * bx0 = (const block_q4_1 *) vx;
  87. #pragma unroll
  88. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  89. int i = i0 + i_offset;
  90. if (need_check) {
  91. i = min(i, i_max);
  92. }
  93. const block_q4_1 * bxi = bx0 + i*blocks_per_row + kbx;
  94. x_ql[i * (WARP_SIZE + 1) + k] = get_int_from_uint8_aligned(bxi->qs, kqsx);
  95. }
  96. const int blocks_per_tile_x_row = WARP_SIZE / QI4_1;
  97. const int kbxd = k % blocks_per_tile_x_row;
  98. #pragma unroll
  99. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI4_1) {
  100. int i = i0 + i_offset * QI4_1 + k / blocks_per_tile_x_row;
  101. if (need_check) {
  102. i = min(i, i_max);
  103. }
  104. const block_q4_1 * bxi = bx0 + i*blocks_per_row + kbxd;
  105. x_dm[i * (WARP_SIZE/QI4_1) + i / QI4_1 + kbxd] = bxi->dm;
  106. }
  107. }
  108. static __device__ __forceinline__ float vec_dot_q4_1_q8_1_mul_mat(
  109. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  110. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  111. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  112. const int kyqs = k % (QI8_1/2) + QI8_1 * (k / (QI8_1/2));
  113. int u[2*VDR_Q4_1_Q8_1_MMQ];
  114. #pragma unroll
  115. for (int l = 0; l < VDR_Q4_1_Q8_1_MMQ; ++l) {
  116. u[2*l+0] = y_qs[j * WARP_SIZE + (kyqs + l) % WARP_SIZE];
  117. u[2*l+1] = y_qs[j * WARP_SIZE + (kyqs + l + QI4_1) % WARP_SIZE];
  118. }
  119. return vec_dot_q4_1_q8_1_impl<VDR_Q4_1_Q8_1_MMQ>
  120. (&x_ql[i * (WARP_SIZE + 1) + k], u, x_dm[i * (WARP_SIZE/QI4_1) + i/QI4_1 + k/QI4_1],
  121. y_ds[j * (WARP_SIZE/QI8_1) + (2*k/QI8_1) % (WARP_SIZE/QI8_1)]);
  122. }
  123. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q5_0(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  124. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  125. __shared__ int tile_x_ql[mmq_y * (2*WARP_SIZE) + mmq_y];
  126. __shared__ float tile_x_d[mmq_y * (WARP_SIZE/QI5_0) + mmq_y/QI5_0];
  127. *x_ql = tile_x_ql;
  128. *x_dm = (half2 *) tile_x_d;
  129. }
  130. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q5_0(
  131. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  132. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  133. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  134. GGML_CUDA_ASSUME(i_offset >= 0);
  135. GGML_CUDA_ASSUME(i_offset < nwarps);
  136. GGML_CUDA_ASSUME(k >= 0);
  137. GGML_CUDA_ASSUME(k < WARP_SIZE);
  138. const int kbx = k / QI5_0;
  139. const int kqsx = k % QI5_0;
  140. const block_q5_0 * bx0 = (const block_q5_0 *) vx;
  141. #pragma unroll
  142. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  143. int i = i0 + i_offset;
  144. if (need_check) {
  145. i = min(i, i_max);
  146. }
  147. const block_q5_0 * bxi = bx0 + i*blocks_per_row + kbx;
  148. const int ql = get_int_from_uint8(bxi->qs, kqsx);
  149. const int qh = get_int_from_uint8(bxi->qh, 0) >> (4 * (k % QI5_0));
  150. int qs0 = (ql >> 0) & 0x0F0F0F0F;
  151. qs0 |= (qh << 4) & 0x00000010; // 0 -> 4
  152. qs0 |= (qh << 11) & 0x00001000; // 1 -> 12
  153. qs0 |= (qh << 18) & 0x00100000; // 2 -> 20
  154. qs0 |= (qh << 25) & 0x10000000; // 3 -> 28
  155. qs0 = __vsubss4(qs0, 0x10101010); // subtract 16
  156. x_ql[i * (2*WARP_SIZE + 1) + 2*k+0] = qs0;
  157. int qs1 = (ql >> 4) & 0x0F0F0F0F;
  158. qs1 |= (qh >> 12) & 0x00000010; // 16 -> 4
  159. qs1 |= (qh >> 5) & 0x00001000; // 17 -> 12
  160. qs1 |= (qh << 2) & 0x00100000; // 18 -> 20
  161. qs1 |= (qh << 9) & 0x10000000; // 19 -> 28
  162. qs1 = __vsubss4(qs1, 0x10101010); // subtract 16
  163. x_ql[i * (2*WARP_SIZE + 1) + 2*k+1] = qs1;
  164. }
  165. const int blocks_per_tile_x_row = WARP_SIZE / QI5_0;
  166. const int kbxd = k % blocks_per_tile_x_row;
  167. float * x_dmf = (float *) x_dm;
  168. #pragma unroll
  169. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI5_0) {
  170. int i = i0 + i_offset * QI5_0 + k / blocks_per_tile_x_row;
  171. if (need_check) {
  172. i = min(i, i_max);
  173. }
  174. const block_q5_0 * bxi = bx0 + i*blocks_per_row + kbxd;
  175. x_dmf[i * (WARP_SIZE/QI5_0) + i / QI5_0 + kbxd] = bxi->d;
  176. }
  177. }
  178. static __device__ __forceinline__ float vec_dot_q5_0_q8_1_mul_mat(
  179. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  180. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  181. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  182. const int kyqs = k % (QI8_1/2) + QI8_1 * (k / (QI8_1/2));
  183. const int index_bx = i * (WARP_SIZE/QI5_0) + i/QI5_0 + k/QI5_0;
  184. const float * x_dmf = (const float *) x_dm;
  185. const float * y_df = (const float *) y_ds;
  186. int u[2*VDR_Q5_0_Q8_1_MMQ];
  187. #pragma unroll
  188. for (int l = 0; l < VDR_Q5_0_Q8_1_MMQ; ++l) {
  189. u[2*l+0] = y_qs[j * WARP_SIZE + (kyqs + l) % WARP_SIZE];
  190. u[2*l+1] = y_qs[j * WARP_SIZE + (kyqs + l + QI5_0) % WARP_SIZE];
  191. }
  192. return vec_dot_q8_0_q8_1_impl<QR5_0*VDR_Q5_0_Q8_1_MMQ>
  193. (&x_ql[i * (2*WARP_SIZE + 1) + 2 * k], u, x_dmf[index_bx], y_df[j * (WARP_SIZE/QI8_1) + (2*k/QI8_1) % (WARP_SIZE/QI8_1)]);
  194. }
  195. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q5_1(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  196. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  197. __shared__ int tile_x_ql[mmq_y * (2*WARP_SIZE) + mmq_y];
  198. __shared__ half2 tile_x_dm[mmq_y * (WARP_SIZE/QI5_1) + mmq_y/QI5_1];
  199. *x_ql = tile_x_ql;
  200. *x_dm = tile_x_dm;
  201. }
  202. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q5_1(
  203. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  204. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  205. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  206. GGML_CUDA_ASSUME(i_offset >= 0);
  207. GGML_CUDA_ASSUME(i_offset < nwarps);
  208. GGML_CUDA_ASSUME(k >= 0);
  209. GGML_CUDA_ASSUME(k < WARP_SIZE);
  210. const int kbx = k / QI5_1;
  211. const int kqsx = k % QI5_1;
  212. const block_q5_1 * bx0 = (const block_q5_1 *) vx;
  213. #pragma unroll
  214. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  215. int i = i0 + i_offset;
  216. if (need_check) {
  217. i = min(i, i_max);
  218. }
  219. const block_q5_1 * bxi = bx0 + i*blocks_per_row + kbx;
  220. const int ql = get_int_from_uint8_aligned(bxi->qs, kqsx);
  221. const int qh = get_int_from_uint8_aligned(bxi->qh, 0) >> (4 * (k % QI5_1));
  222. int qs0 = (ql >> 0) & 0x0F0F0F0F;
  223. qs0 |= (qh << 4) & 0x00000010; // 0 -> 4
  224. qs0 |= (qh << 11) & 0x00001000; // 1 -> 12
  225. qs0 |= (qh << 18) & 0x00100000; // 2 -> 20
  226. qs0 |= (qh << 25) & 0x10000000; // 3 -> 28
  227. x_ql[i * (2*WARP_SIZE + 1) + 2*k+0] = qs0;
  228. int qs1 = (ql >> 4) & 0x0F0F0F0F;
  229. qs1 |= (qh >> 12) & 0x00000010; // 16 -> 4
  230. qs1 |= (qh >> 5) & 0x00001000; // 17 -> 12
  231. qs1 |= (qh << 2) & 0x00100000; // 18 -> 20
  232. qs1 |= (qh << 9) & 0x10000000; // 19 -> 28
  233. x_ql[i * (2*WARP_SIZE + 1) + 2*k+1] = qs1;
  234. }
  235. const int blocks_per_tile_x_row = WARP_SIZE / QI5_1;
  236. const int kbxd = k % blocks_per_tile_x_row;
  237. #pragma unroll
  238. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI5_1) {
  239. int i = i0 + i_offset * QI5_1 + k / blocks_per_tile_x_row;
  240. if (need_check) {
  241. i = min(i, i_max);
  242. }
  243. const block_q5_1 * bxi = bx0 + i*blocks_per_row + kbxd;
  244. x_dm[i * (WARP_SIZE/QI5_1) + i / QI5_1 + kbxd] = bxi->dm;
  245. }
  246. }
  247. static __device__ __forceinline__ float vec_dot_q5_1_q8_1_mul_mat(
  248. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  249. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  250. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  251. const int kyqs = k % (QI8_1/2) + QI8_1 * (k / (QI8_1/2));
  252. const int index_bx = i * (WARP_SIZE/QI5_1) + + i/QI5_1 + k/QI5_1;
  253. int u[2*VDR_Q5_1_Q8_1_MMQ];
  254. #pragma unroll
  255. for (int l = 0; l < VDR_Q5_1_Q8_1_MMQ; ++l) {
  256. u[2*l+0] = y_qs[j * WARP_SIZE + (kyqs + l) % WARP_SIZE];
  257. u[2*l+1] = y_qs[j * WARP_SIZE + (kyqs + l + QI5_1) % WARP_SIZE];
  258. }
  259. return vec_dot_q8_1_q8_1_impl<QR5_1*VDR_Q5_1_Q8_1_MMQ>
  260. (&x_ql[i * (2*WARP_SIZE + 1) + 2 * k], u, x_dm[index_bx], y_ds[j * (WARP_SIZE/QI8_1) + (2*k/QI8_1) % (WARP_SIZE/QI8_1)]);
  261. }
  262. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q8_0(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  263. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  264. __shared__ int tile_x_qs[mmq_y * (WARP_SIZE) + mmq_y];
  265. __shared__ float tile_x_d[mmq_y * (WARP_SIZE/QI8_0) + mmq_y/QI8_0];
  266. *x_ql = tile_x_qs;
  267. *x_dm = (half2 *) tile_x_d;
  268. }
  269. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q8_0(
  270. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  271. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  272. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  273. GGML_CUDA_ASSUME(i_offset >= 0);
  274. GGML_CUDA_ASSUME(i_offset < nwarps);
  275. GGML_CUDA_ASSUME(k >= 0);
  276. GGML_CUDA_ASSUME(k < WARP_SIZE);
  277. const int kbx = k / QI8_0;
  278. const int kqsx = k % QI8_0;
  279. float * x_dmf = (float *) x_dm;
  280. const block_q8_0 * bx0 = (const block_q8_0 *) vx;
  281. #pragma unroll
  282. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  283. int i = i0 + i_offset;
  284. if (need_check) {
  285. i = min(i, i_max);
  286. }
  287. const block_q8_0 * bxi = bx0 + i*blocks_per_row + kbx;
  288. x_ql[i * (WARP_SIZE + 1) + k] = get_int_from_int8(bxi->qs, kqsx);
  289. }
  290. const int blocks_per_tile_x_row = WARP_SIZE / QI8_0;
  291. const int kbxd = k % blocks_per_tile_x_row;
  292. #pragma unroll
  293. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI8_0) {
  294. int i = i0 + i_offset * QI8_0 + k / blocks_per_tile_x_row;
  295. if (need_check) {
  296. i = min(i, i_max);
  297. }
  298. const block_q8_0 * bxi = bx0 + i*blocks_per_row + kbxd;
  299. x_dmf[i * (WARP_SIZE/QI8_0) + i / QI8_0 + kbxd] = bxi->d;
  300. }
  301. }
  302. static __device__ __forceinline__ float vec_dot_q8_0_q8_1_mul_mat(
  303. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  304. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  305. GGML_UNUSED(x_qh); GGML_UNUSED(x_sc);
  306. const float * x_dmf = (const float *) x_dm;
  307. const float * y_df = (const float *) y_ds;
  308. return vec_dot_q8_0_q8_1_impl<VDR_Q8_0_Q8_1_MMQ>
  309. (&x_ql[i * (WARP_SIZE + 1) + k], &y_qs[j * WARP_SIZE + k], x_dmf[i * (WARP_SIZE/QI8_0) + i/QI8_0 + k/QI8_0],
  310. y_df[j * (WARP_SIZE/QI8_1) + k/QI8_1]);
  311. }
  312. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q2_K(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  313. GGML_UNUSED(x_qh);
  314. __shared__ int tile_x_ql[mmq_y * (WARP_SIZE) + mmq_y];
  315. __shared__ half2 tile_x_dm[mmq_y * (WARP_SIZE/QI2_K) + mmq_y/QI2_K];
  316. __shared__ int tile_x_sc[mmq_y * (WARP_SIZE/4) + mmq_y/4];
  317. *x_ql = tile_x_ql;
  318. *x_dm = tile_x_dm;
  319. *x_sc = tile_x_sc;
  320. }
  321. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q2_K(
  322. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  323. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  324. GGML_UNUSED(x_qh);
  325. GGML_CUDA_ASSUME(i_offset >= 0);
  326. GGML_CUDA_ASSUME(i_offset < nwarps);
  327. GGML_CUDA_ASSUME(k >= 0);
  328. GGML_CUDA_ASSUME(k < WARP_SIZE);
  329. const int kbx = k / QI2_K;
  330. const int kqsx = k % QI2_K;
  331. const block_q2_K * bx0 = (const block_q2_K *) vx;
  332. #pragma unroll
  333. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  334. int i = i0 + i_offset;
  335. if (need_check) {
  336. i = min(i, i_max);
  337. }
  338. const block_q2_K * bxi = bx0 + i*blocks_per_row + kbx;
  339. x_ql[i * (WARP_SIZE + 1) + k] = get_int_from_uint8_aligned(bxi->qs, kqsx);
  340. }
  341. const int blocks_per_tile_x_row = WARP_SIZE / QI2_K;
  342. const int kbxd = k % blocks_per_tile_x_row;
  343. #pragma unroll
  344. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI2_K) {
  345. int i = (i0 + i_offset * QI2_K + k / blocks_per_tile_x_row) % mmq_y;
  346. if (need_check) {
  347. i = min(i, i_max);
  348. }
  349. const block_q2_K * bxi = bx0 + i*blocks_per_row + kbxd;
  350. x_dm[i * (WARP_SIZE/QI2_K) + i / QI2_K + kbxd] = bxi->dm;
  351. }
  352. #pragma unroll
  353. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * 4) {
  354. int i = i0 + i_offset * 4 + k / (WARP_SIZE/4);
  355. if (need_check) {
  356. i = min(i, i_max);
  357. }
  358. const block_q2_K * bxi = bx0 + i*blocks_per_row + (k % (WARP_SIZE/4)) / (QI2_K/4);
  359. x_sc[i * (WARP_SIZE/4) + i / 4 + k % (WARP_SIZE/4)] = get_int_from_uint8_aligned(bxi->scales, k % (QI2_K/4));
  360. }
  361. }
  362. static __device__ __forceinline__ float vec_dot_q2_K_q8_1_mul_mat(
  363. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  364. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  365. GGML_UNUSED(x_qh);
  366. const int kbx = k / QI2_K;
  367. const int ky = (k % QI2_K) * QR2_K;
  368. const float * y_df = (const float *) y_ds;
  369. int v[QR2_K*VDR_Q2_K_Q8_1_MMQ];
  370. const int kqsx = i * (WARP_SIZE + 1) + kbx*QI2_K + (QI2_K/2) * (ky/(2*QI2_K)) + ky % (QI2_K/2);
  371. const int shift = 2 * ((ky % (2*QI2_K)) / (QI2_K/2));
  372. #pragma unroll
  373. for (int l = 0; l < QR2_K*VDR_Q2_K_Q8_1_MMQ; ++l) {
  374. v[l] = (x_ql[kqsx + l] >> shift) & 0x03030303;
  375. }
  376. const uint8_t * scales = ((const uint8_t *) &x_sc[i * (WARP_SIZE/4) + i/4 + kbx*4]) + ky/4;
  377. const int index_y = j * WARP_SIZE + (QR2_K*k) % WARP_SIZE;
  378. return vec_dot_q2_K_q8_1_impl_mmq(v, &y_qs[index_y], scales, x_dm[i * (WARP_SIZE/QI2_K) + i/QI2_K + kbx], y_df[index_y/QI8_1]);
  379. }
  380. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q3_K(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  381. __shared__ int tile_x_ql[mmq_y * (WARP_SIZE) + mmq_y];
  382. __shared__ half2 tile_x_dm[mmq_y * (WARP_SIZE/QI3_K) + mmq_y/QI3_K];
  383. __shared__ int tile_x_qh[mmq_y * (WARP_SIZE/2) + mmq_y/2];
  384. __shared__ int tile_x_sc[mmq_y * (WARP_SIZE/4) + mmq_y/4];
  385. *x_ql = tile_x_ql;
  386. *x_dm = tile_x_dm;
  387. *x_qh = tile_x_qh;
  388. *x_sc = tile_x_sc;
  389. }
  390. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q3_K(
  391. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  392. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  393. GGML_CUDA_ASSUME(i_offset >= 0);
  394. GGML_CUDA_ASSUME(i_offset < nwarps);
  395. GGML_CUDA_ASSUME(k >= 0);
  396. GGML_CUDA_ASSUME(k < WARP_SIZE);
  397. const int kbx = k / QI3_K;
  398. const int kqsx = k % QI3_K;
  399. const block_q3_K * bx0 = (const block_q3_K *) vx;
  400. #pragma unroll
  401. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  402. int i = i0 + i_offset;
  403. if (need_check) {
  404. i = min(i, i_max);
  405. }
  406. const block_q3_K * bxi = bx0 + i*blocks_per_row + kbx;
  407. x_ql[i * (WARP_SIZE + 1) + k] = get_int_from_uint8(bxi->qs, kqsx);
  408. }
  409. const int blocks_per_tile_x_row = WARP_SIZE / QI3_K;
  410. const int kbxd = k % blocks_per_tile_x_row;
  411. float * x_dmf = (float *) x_dm;
  412. #pragma unroll
  413. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI3_K) {
  414. int i = (i0 + i_offset * QI3_K + k / blocks_per_tile_x_row) % mmq_y;
  415. if (need_check) {
  416. i = min(i, i_max);
  417. }
  418. const block_q3_K * bxi = bx0 + i*blocks_per_row + kbxd;
  419. x_dmf[i * (WARP_SIZE/QI3_K) + i / QI3_K + kbxd] = bxi->d;
  420. }
  421. #pragma unroll
  422. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * 2) {
  423. int i = i0 + i_offset * 2 + k / (WARP_SIZE/2);
  424. if (need_check) {
  425. i = min(i, i_max);
  426. }
  427. const block_q3_K * bxi = bx0 + i*blocks_per_row + (k % (WARP_SIZE/2)) / (QI3_K/2);
  428. // invert the mask with ~ so that a 0/1 results in 4/0 being subtracted
  429. x_qh[i * (WARP_SIZE/2) + i / 2 + k % (WARP_SIZE/2)] = ~get_int_from_uint8(bxi->hmask, k % (QI3_K/2));
  430. }
  431. #pragma unroll
  432. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * 4) {
  433. int i = i0 + i_offset * 4 + k / (WARP_SIZE/4);
  434. if (need_check) {
  435. i = min(i, i_max);
  436. }
  437. const block_q3_K * bxi = bx0 + i*blocks_per_row + (k % (WARP_SIZE/4)) / (QI3_K/4);
  438. const int ksc = k % (QI3_K/4);
  439. const int ksc_low = ksc % (QI3_K/8);
  440. const int shift_low = 4 * (ksc / (QI3_K/8));
  441. const int sc_low = (get_int_from_uint8(bxi->scales, ksc_low) >> shift_low) & 0x0F0F0F0F;
  442. const int ksc_high = QI3_K/8;
  443. const int shift_high = 2 * ksc;
  444. const int sc_high = ((get_int_from_uint8(bxi->scales, ksc_high) >> shift_high) << 4) & 0x30303030;
  445. const int sc = __vsubss4(sc_low | sc_high, 0x20202020);
  446. x_sc[i * (WARP_SIZE/4) + i / 4 + k % (WARP_SIZE/4)] = sc;
  447. }
  448. }
  449. static __device__ __forceinline__ float vec_dot_q3_K_q8_1_mul_mat(
  450. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  451. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  452. const int kbx = k / QI3_K;
  453. const int ky = (k % QI3_K) * QR3_K;
  454. const float * x_dmf = (const float *) x_dm;
  455. const float * y_df = (const float *) y_ds;
  456. const int8_t * scales = ((const int8_t *) (x_sc + i * (WARP_SIZE/4) + i/4 + kbx*4)) + ky/4;
  457. int v[QR3_K*VDR_Q3_K_Q8_1_MMQ];
  458. #pragma unroll
  459. for (int l = 0; l < QR3_K*VDR_Q3_K_Q8_1_MMQ; ++l) {
  460. const int kqsx = i * (WARP_SIZE + 1) + kbx*QI3_K + (QI3_K/2) * (ky/(2*QI3_K)) + ky % (QI3_K/2);
  461. const int shift = 2 * ((ky % 32) / 8);
  462. const int vll = (x_ql[kqsx + l] >> shift) & 0x03030303;
  463. const int vh = x_qh[i * (WARP_SIZE/2) + i/2 + kbx * (QI3_K/2) + (ky+l)%8] >> ((ky+l) / 8);
  464. const int vlh = (vh << 2) & 0x04040404;
  465. v[l] = __vsubss4(vll, vlh);
  466. }
  467. const int index_y = j * WARP_SIZE + (k*QR3_K) % WARP_SIZE;
  468. return vec_dot_q3_K_q8_1_impl_mmq(v, &y_qs[index_y], scales, x_dmf[i * (WARP_SIZE/QI3_K) + i/QI3_K + kbx], y_df[index_y/QI8_1]);
  469. }
  470. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q4_K(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  471. GGML_UNUSED(x_qh);
  472. __shared__ int tile_x_ql[mmq_y * (WARP_SIZE) + mmq_y];
  473. __shared__ half2 tile_x_dm[mmq_y * (WARP_SIZE/QI4_K) + mmq_y/QI4_K];
  474. __shared__ int tile_x_sc[mmq_y * (WARP_SIZE/8) + mmq_y/8];
  475. *x_ql = tile_x_ql;
  476. *x_dm = tile_x_dm;
  477. *x_sc = tile_x_sc;
  478. }
  479. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q4_K(
  480. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  481. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  482. GGML_UNUSED(x_qh);
  483. GGML_CUDA_ASSUME(i_offset >= 0);
  484. GGML_CUDA_ASSUME(i_offset < nwarps);
  485. GGML_CUDA_ASSUME(k >= 0);
  486. GGML_CUDA_ASSUME(k < WARP_SIZE);
  487. const int kbx = k / QI4_K; // == 0 if QK_K == 256
  488. const int kqsx = k % QI4_K; // == k if QK_K == 256
  489. const block_q4_K * bx0 = (const block_q4_K *) vx;
  490. #pragma unroll
  491. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  492. int i = i0 + i_offset;
  493. if (need_check) {
  494. i = min(i, i_max);
  495. }
  496. const block_q4_K * bxi = bx0 + i*blocks_per_row + kbx;
  497. x_ql[i * (WARP_SIZE + 1) + k] = get_int_from_uint8_aligned(bxi->qs, kqsx);
  498. }
  499. const int blocks_per_tile_x_row = WARP_SIZE / QI4_K; // == 1 if QK_K == 256
  500. const int kbxd = k % blocks_per_tile_x_row; // == 0 if QK_K == 256
  501. #pragma unroll
  502. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI4_K) {
  503. int i = (i0 + i_offset * QI4_K + k / blocks_per_tile_x_row) % mmq_y;
  504. if (need_check) {
  505. i = min(i, i_max);
  506. }
  507. const block_q4_K * bxi = bx0 + i*blocks_per_row + kbxd;
  508. #if QK_K == 256
  509. x_dm[i * (WARP_SIZE/QI4_K) + i / QI4_K + kbxd] = bxi->dm;
  510. #else
  511. x_dm[i * (WARP_SIZE/QI4_K) + i / QI4_K + kbxd] = {bxi->dm[0], bxi->dm[1]};
  512. #endif
  513. }
  514. #pragma unroll
  515. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * 8) {
  516. int i = (i0 + i_offset * 8 + k / (WARP_SIZE/8)) % mmq_y;
  517. if (need_check) {
  518. i = min(i, i_max);
  519. }
  520. const block_q4_K * bxi = bx0 + i*blocks_per_row + (k % (WARP_SIZE/8)) / (QI4_K/8);
  521. const int * scales = (const int *) bxi->scales;
  522. const int ksc = k % (WARP_SIZE/8);
  523. // scale arrangement after the following two lines: sc0,...,sc3, sc4,...,sc7, m0,...,m3, m4,...,m8
  524. int scales8 = (scales[(ksc%2) + (ksc!=0)] >> (4 * (ksc & (ksc/2)))) & 0x0F0F0F0F; // lower 4 bits
  525. scales8 |= (scales[ksc/2] >> (2 * (ksc % 2))) & 0x30303030; // upper 2 bits
  526. x_sc[i * (WARP_SIZE/8) + i / 8 + ksc] = scales8;
  527. }
  528. }
  529. static __device__ __forceinline__ float vec_dot_q4_K_q8_1_mul_mat(
  530. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  531. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  532. GGML_UNUSED(x_qh);
  533. const uint8_t * sc = ((const uint8_t *) &x_sc[i * (WARP_SIZE/8) + i/8 + k/16]) + 2*((k % 16) / 8);
  534. const int index_y = j * WARP_SIZE + (QR4_K*k) % WARP_SIZE;
  535. return vec_dot_q4_K_q8_1_impl_mmq(&x_ql[i * (WARP_SIZE + 1) + k], &y_qs[index_y], sc, sc+8,
  536. x_dm[i * (WARP_SIZE/QI4_K) + i/QI4_K], &y_ds[index_y/QI8_1]);
  537. }
  538. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q5_K(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  539. GGML_UNUSED(x_qh);
  540. __shared__ int tile_x_ql[mmq_y * (2*WARP_SIZE) + mmq_y];
  541. __shared__ half2 tile_x_dm[mmq_y * (WARP_SIZE/QI5_K) + mmq_y/QI5_K];
  542. __shared__ int tile_x_sc[mmq_y * (WARP_SIZE/8) + mmq_y/8];
  543. *x_ql = tile_x_ql;
  544. *x_dm = tile_x_dm;
  545. *x_sc = tile_x_sc;
  546. }
  547. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q5_K(
  548. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  549. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  550. GGML_UNUSED(x_qh);
  551. GGML_CUDA_ASSUME(i_offset >= 0);
  552. GGML_CUDA_ASSUME(i_offset < nwarps);
  553. GGML_CUDA_ASSUME(k >= 0);
  554. GGML_CUDA_ASSUME(k < WARP_SIZE);
  555. const int kbx = k / QI5_K; // == 0 if QK_K == 256
  556. const int kqsx = k % QI5_K; // == k if QK_K == 256
  557. const block_q5_K * bx0 = (const block_q5_K *) vx;
  558. #pragma unroll
  559. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  560. int i = i0 + i_offset;
  561. if (need_check) {
  562. i = min(i, i_max);
  563. }
  564. const block_q5_K * bxi = bx0 + i*blocks_per_row + kbx;
  565. const int ky = QR5_K*kqsx;
  566. const int ql = get_int_from_uint8_aligned(bxi->qs, kqsx);
  567. const int ql0 = (ql >> 0) & 0x0F0F0F0F;
  568. const int ql1 = (ql >> 4) & 0x0F0F0F0F;
  569. const int qh = get_int_from_uint8_aligned(bxi->qh, kqsx % (QI5_K/4));
  570. const int qh0 = ((qh >> (2 * (kqsx / (QI5_K/4)) + 0)) << 4) & 0x10101010;
  571. const int qh1 = ((qh >> (2 * (kqsx / (QI5_K/4)) + 1)) << 4) & 0x10101010;
  572. const int kq0 = ky - ky % (QI5_K/2) + k % (QI5_K/4) + 0;
  573. const int kq1 = ky - ky % (QI5_K/2) + k % (QI5_K/4) + (QI5_K/4);
  574. x_ql[i * (2*WARP_SIZE + 1) + kq0] = ql0 | qh0;
  575. x_ql[i * (2*WARP_SIZE + 1) + kq1] = ql1 | qh1;
  576. }
  577. const int blocks_per_tile_x_row = WARP_SIZE / QI5_K; // == 1 if QK_K == 256
  578. const int kbxd = k % blocks_per_tile_x_row; // == 0 if QK_K == 256
  579. #pragma unroll
  580. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI5_K) {
  581. int i = (i0 + i_offset * QI5_K + k / blocks_per_tile_x_row) % mmq_y;
  582. if (need_check) {
  583. i = min(i, i_max);
  584. }
  585. const block_q5_K * bxi = bx0 + i*blocks_per_row + kbxd;
  586. #if QK_K == 256
  587. x_dm[i * (WARP_SIZE/QI5_K) + i / QI5_K + kbxd] = bxi->dm;
  588. #endif
  589. }
  590. #pragma unroll
  591. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * 8) {
  592. int i = (i0 + i_offset * 8 + k / (WARP_SIZE/8)) % mmq_y;
  593. if (need_check) {
  594. i = min(i, i_max);
  595. }
  596. const block_q5_K * bxi = bx0 + i*blocks_per_row + (k % (WARP_SIZE/8)) / (QI5_K/8);
  597. const int * scales = (const int *) bxi->scales;
  598. const int ksc = k % (WARP_SIZE/8);
  599. // scale arrangement after the following two lines: sc0,...,sc3, sc4,...,sc7, m0,...,m3, m4,...,m8
  600. int scales8 = (scales[(ksc%2) + (ksc!=0)] >> (4 * (ksc & (ksc/2)))) & 0x0F0F0F0F; // lower 4 bits
  601. scales8 |= (scales[ksc/2] >> (2 * (ksc % 2))) & 0x30303030; // upper 2 bits
  602. x_sc[i * (WARP_SIZE/8) + i / 8 + ksc] = scales8;
  603. }
  604. }
  605. static __device__ __forceinline__ float vec_dot_q5_K_q8_1_mul_mat(
  606. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  607. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  608. GGML_UNUSED(x_qh);
  609. const uint8_t * sc = ((const uint8_t *) &x_sc[i * (WARP_SIZE/8) + i/8 + k/16]) + 2 * ((k % 16) / 8);
  610. const int index_x = i * (QR5_K*WARP_SIZE + 1) + QR5_K*k;
  611. const int index_y = j * WARP_SIZE + (QR5_K*k) % WARP_SIZE;
  612. return vec_dot_q5_K_q8_1_impl_mmq(&x_ql[index_x], &y_qs[index_y], sc, sc+8,
  613. x_dm[i * (WARP_SIZE/QI5_K) + i/QI5_K], &y_ds[index_y/QI8_1]);
  614. }
  615. template <int mmq_y> static __device__ __forceinline__ void allocate_tiles_q6_K(int ** x_ql, half2 ** x_dm, int ** x_qh, int ** x_sc) {
  616. GGML_UNUSED(x_qh);
  617. __shared__ int tile_x_ql[mmq_y * (2*WARP_SIZE) + mmq_y];
  618. __shared__ half2 tile_x_dm[mmq_y * (WARP_SIZE/QI6_K) + mmq_y/QI6_K];
  619. __shared__ int tile_x_sc[mmq_y * (WARP_SIZE/8) + mmq_y/8];
  620. *x_ql = tile_x_ql;
  621. *x_dm = tile_x_dm;
  622. *x_sc = tile_x_sc;
  623. }
  624. template <int mmq_y, int nwarps, bool need_check> static __device__ __forceinline__ void load_tiles_q6_K(
  625. const void * __restrict__ vx, int * __restrict__ x_ql, half2 * __restrict__ x_dm, int * __restrict__ x_qh,
  626. int * __restrict__ x_sc, const int & i_offset, const int & i_max, const int & k, const int & blocks_per_row) {
  627. GGML_UNUSED(x_qh);
  628. GGML_CUDA_ASSUME(i_offset >= 0);
  629. GGML_CUDA_ASSUME(i_offset < nwarps);
  630. GGML_CUDA_ASSUME(k >= 0);
  631. GGML_CUDA_ASSUME(k < WARP_SIZE);
  632. const int kbx = k / QI6_K; // == 0 if QK_K == 256
  633. const int kqsx = k % QI6_K; // == k if QK_K == 256
  634. const block_q6_K * bx0 = (const block_q6_K *) vx;
  635. #pragma unroll
  636. for (int i0 = 0; i0 < mmq_y; i0 += nwarps) {
  637. int i = i0 + i_offset;
  638. if (need_check) {
  639. i = min(i, i_max);
  640. }
  641. const block_q6_K * bxi = bx0 + i*blocks_per_row + kbx;
  642. const int ky = QR6_K*kqsx;
  643. const int ql = get_int_from_uint8(bxi->ql, kqsx);
  644. const int ql0 = (ql >> 0) & 0x0F0F0F0F;
  645. const int ql1 = (ql >> 4) & 0x0F0F0F0F;
  646. const int qh = get_int_from_uint8(bxi->qh, (QI6_K/4) * (kqsx / (QI6_K/2)) + kqsx % (QI6_K/4));
  647. const int qh0 = ((qh >> (2 * ((kqsx % (QI6_K/2)) / (QI6_K/4)))) << 4) & 0x30303030;
  648. const int qh1 = (qh >> (2 * ((kqsx % (QI6_K/2)) / (QI6_K/4)))) & 0x30303030;
  649. const int kq0 = ky - ky % QI6_K + k % (QI6_K/2) + 0;
  650. const int kq1 = ky - ky % QI6_K + k % (QI6_K/2) + (QI6_K/2);
  651. x_ql[i * (2*WARP_SIZE + 1) + kq0] = __vsubss4(ql0 | qh0, 0x20202020);
  652. x_ql[i * (2*WARP_SIZE + 1) + kq1] = __vsubss4(ql1 | qh1, 0x20202020);
  653. }
  654. const int blocks_per_tile_x_row = WARP_SIZE / QI6_K; // == 1 if QK_K == 256
  655. const int kbxd = k % blocks_per_tile_x_row; // == 0 if QK_K == 256
  656. float * x_dmf = (float *) x_dm;
  657. #pragma unroll
  658. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * QI6_K) {
  659. int i = (i0 + i_offset * QI6_K + k / blocks_per_tile_x_row) % mmq_y;
  660. if (need_check) {
  661. i = min(i, i_max);
  662. }
  663. const block_q6_K * bxi = bx0 + i*blocks_per_row + kbxd;
  664. x_dmf[i * (WARP_SIZE/QI6_K) + i / QI6_K + kbxd] = bxi->d;
  665. }
  666. #pragma unroll
  667. for (int i0 = 0; i0 < mmq_y; i0 += nwarps * 8) {
  668. int i = (i0 + i_offset * 8 + k / (WARP_SIZE/8)) % mmq_y;
  669. if (need_check) {
  670. i = min(i, i_max);
  671. }
  672. const block_q6_K * bxi = bx0 + i*blocks_per_row + (k % (WARP_SIZE/8)) / 4;
  673. x_sc[i * (WARP_SIZE/8) + i / 8 + k % (WARP_SIZE/8)] = get_int_from_int8(bxi->scales, k % (QI6_K/8));
  674. }
  675. }
  676. static __device__ __forceinline__ float vec_dot_q6_K_q8_1_mul_mat(
  677. const int * __restrict__ x_ql, const half2 * __restrict__ x_dm, const int * __restrict__ x_qh, const int * __restrict__ x_sc,
  678. const int * __restrict__ y_qs, const half2 * __restrict__ y_ds, const int & i, const int & j, const int & k) {
  679. GGML_UNUSED(x_qh);
  680. const float * x_dmf = (const float *) x_dm;
  681. const float * y_df = (const float *) y_ds;
  682. const int8_t * sc = ((const int8_t *) &x_sc[i * (WARP_SIZE/8) + i/8 + k/8]);
  683. const int index_x = i * (QR6_K*WARP_SIZE + 1) + QR6_K*k;
  684. const int index_y = j * WARP_SIZE + (QR6_K*k) % WARP_SIZE;
  685. return vec_dot_q6_K_q8_1_impl_mmq(&x_ql[index_x], &y_qs[index_y], sc, x_dmf[i * (WARP_SIZE/QI6_K) + i/QI6_K], &y_df[index_y/QI8_1]);
  686. }
  687. #define MMQ_X_Q4_0_RDNA2 64
  688. #define MMQ_Y_Q4_0_RDNA2 128
  689. #define NWARPS_Q4_0_RDNA2 8
  690. #define MMQ_X_Q4_0_RDNA1 64
  691. #define MMQ_Y_Q4_0_RDNA1 64
  692. #define NWARPS_Q4_0_RDNA1 8
  693. #if defined(CUDA_USE_TENSOR_CORES)
  694. #define MMQ_X_Q4_0_AMPERE 4
  695. #define MMQ_Y_Q4_0_AMPERE 32
  696. #define NWARPS_Q4_0_AMPERE 4
  697. #else
  698. #define MMQ_X_Q4_0_AMPERE 64
  699. #define MMQ_Y_Q4_0_AMPERE 128
  700. #define NWARPS_Q4_0_AMPERE 4
  701. #endif
  702. #define MMQ_X_Q4_0_PASCAL 64
  703. #define MMQ_Y_Q4_0_PASCAL 64
  704. #define NWARPS_Q4_0_PASCAL 8
  705. template <int qk, int qr, int qi, bool need_sum, typename block_q_t, int mmq_x, int mmq_y, int nwarps,
  706. allocate_tiles_cuda_t allocate_tiles, load_tiles_cuda_t load_tiles, int vdr, vec_dot_q_mul_mat_cuda_t vec_dot>
  707. static __device__ __forceinline__ void mul_mat_q(
  708. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  709. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  710. const block_q_t * x = (const block_q_t *) vx;
  711. const block_q8_1 * y = (const block_q8_1 *) vy;
  712. const int blocks_per_row_x = ncols_x / qk;
  713. const int blocks_per_col_y = nrows_y / QK8_1;
  714. const int blocks_per_warp = WARP_SIZE / qi;
  715. const int & ncols_dst = ncols_y;
  716. const int row_dst_0 = blockIdx.x*mmq_y;
  717. const int & row_x_0 = row_dst_0;
  718. const int col_dst_0 = blockIdx.y*mmq_x;
  719. const int & col_y_0 = col_dst_0;
  720. int * tile_x_ql = nullptr;
  721. half2 * tile_x_dm = nullptr;
  722. int * tile_x_qh = nullptr;
  723. int * tile_x_sc = nullptr;
  724. allocate_tiles(&tile_x_ql, &tile_x_dm, &tile_x_qh, &tile_x_sc);
  725. __shared__ int tile_y_qs[mmq_x * WARP_SIZE];
  726. __shared__ half2 tile_y_ds[mmq_x * WARP_SIZE/QI8_1];
  727. float sum[mmq_y/WARP_SIZE][mmq_x/nwarps] = {{0.0f}};
  728. for (int ib0 = 0; ib0 < blocks_per_row_x; ib0 += blocks_per_warp) {
  729. load_tiles(x + row_x_0*blocks_per_row_x + ib0, tile_x_ql, tile_x_dm, tile_x_qh, tile_x_sc,
  730. threadIdx.y, nrows_x-row_x_0-1, threadIdx.x, blocks_per_row_x);
  731. #pragma unroll
  732. for (int ir = 0; ir < qr; ++ir) {
  733. const int kqs = ir*WARP_SIZE + threadIdx.x;
  734. const int kbxd = kqs / QI8_1;
  735. #pragma unroll
  736. for (int i = 0; i < mmq_x; i += nwarps) {
  737. const int col_y_eff = min(col_y_0 + threadIdx.y + i, ncols_y-1); // to prevent out-of-bounds memory accesses
  738. const block_q8_1 * by0 = &y[col_y_eff*blocks_per_col_y + ib0 * (qk/QK8_1) + kbxd];
  739. const int index_y = (threadIdx.y + i) * WARP_SIZE + kqs % WARP_SIZE;
  740. tile_y_qs[index_y] = get_int_from_int8_aligned(by0->qs, threadIdx.x % QI8_1);
  741. }
  742. #pragma unroll
  743. for (int ids0 = 0; ids0 < mmq_x; ids0 += nwarps * QI8_1) {
  744. const int ids = (ids0 + threadIdx.y * QI8_1 + threadIdx.x / (WARP_SIZE/QI8_1)) % mmq_x;
  745. const int kby = threadIdx.x % (WARP_SIZE/QI8_1);
  746. const int col_y_eff = min(col_y_0 + ids, ncols_y-1);
  747. // if the sum is not needed it's faster to transform the scale to f32 ahead of time
  748. const half2 * dsi_src = &y[col_y_eff*blocks_per_col_y + ib0 * (qk/QK8_1) + ir*(WARP_SIZE/QI8_1) + kby].ds;
  749. half2 * dsi_dst = &tile_y_ds[ids * (WARP_SIZE/QI8_1) + kby];
  750. if (need_sum) {
  751. *dsi_dst = *dsi_src;
  752. } else {
  753. float * dfi_dst = (float *) dsi_dst;
  754. *dfi_dst = __low2float(*dsi_src);
  755. }
  756. }
  757. __syncthreads();
  758. // #pragma unroll // unrolling this loop causes too much register pressure
  759. for (int k = ir*WARP_SIZE/qr; k < (ir+1)*WARP_SIZE/qr; k += vdr) {
  760. #pragma unroll
  761. for (int j = 0; j < mmq_x; j += nwarps) {
  762. #pragma unroll
  763. for (int i = 0; i < mmq_y; i += WARP_SIZE) {
  764. sum[i/WARP_SIZE][j/nwarps] += vec_dot(
  765. tile_x_ql, tile_x_dm, tile_x_qh, tile_x_sc, tile_y_qs, tile_y_ds,
  766. threadIdx.x + i, threadIdx.y + j, k);
  767. }
  768. }
  769. }
  770. __syncthreads();
  771. }
  772. }
  773. #pragma unroll
  774. for (int j = 0; j < mmq_x; j += nwarps) {
  775. const int col_dst = col_dst_0 + j + threadIdx.y;
  776. if (col_dst >= ncols_dst) {
  777. return;
  778. }
  779. #pragma unroll
  780. for (int i = 0; i < mmq_y; i += WARP_SIZE) {
  781. const int row_dst = row_dst_0 + threadIdx.x + i;
  782. if (row_dst >= nrows_dst) {
  783. continue;
  784. }
  785. dst[col_dst*nrows_dst + row_dst] = sum[i/WARP_SIZE][j/nwarps];
  786. }
  787. }
  788. }
  789. template <bool need_check> static __global__ void
  790. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  791. #if defined(RDNA3) || defined(RDNA2)
  792. __launch_bounds__(WARP_SIZE*NWARPS_Q4_0_RDNA2, 2)
  793. #endif // defined(RDNA3) || defined(RDNA2)
  794. #endif // defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  795. mul_mat_q4_0(
  796. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  797. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  798. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  799. #if defined(RDNA3) || defined(RDNA2)
  800. const int mmq_x = MMQ_X_Q4_0_RDNA2;
  801. const int mmq_y = MMQ_Y_Q4_0_RDNA2;
  802. const int nwarps = NWARPS_Q4_0_RDNA2;
  803. #else
  804. const int mmq_x = MMQ_X_Q4_0_RDNA1;
  805. const int mmq_y = MMQ_Y_Q4_0_RDNA1;
  806. const int nwarps = NWARPS_Q4_0_RDNA1;
  807. #endif // defined(RDNA3) || defined(RDNA2)
  808. mul_mat_q<QK4_0, QR4_0, QI4_0, true, block_q4_0, mmq_x, mmq_y, nwarps, allocate_tiles_q4_0<mmq_y>,
  809. load_tiles_q4_0<mmq_y, nwarps, need_check>, VDR_Q4_0_Q8_1_MMQ, vec_dot_q4_0_q8_1_mul_mat>
  810. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  811. #elif __CUDA_ARCH__ >= CC_VOLTA
  812. const int mmq_x = MMQ_X_Q4_0_AMPERE;
  813. const int mmq_y = MMQ_Y_Q4_0_AMPERE;
  814. const int nwarps = NWARPS_Q4_0_AMPERE;
  815. mul_mat_q<QK4_0, QR4_0, QI4_0, true, block_q4_0, mmq_x, mmq_y, nwarps, allocate_tiles_q4_0<mmq_y>,
  816. load_tiles_q4_0<mmq_y, nwarps, need_check>, VDR_Q4_0_Q8_1_MMQ, vec_dot_q4_0_q8_1_mul_mat>
  817. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  818. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  819. const int mmq_x = MMQ_X_Q4_0_PASCAL;
  820. const int mmq_y = MMQ_Y_Q4_0_PASCAL;
  821. const int nwarps = NWARPS_Q4_0_PASCAL;
  822. mul_mat_q<QK4_0, QR4_0, QI4_0, true, block_q4_0, mmq_x, mmq_y, nwarps, allocate_tiles_q4_0<mmq_y>,
  823. load_tiles_q4_0<mmq_y, nwarps, need_check>, VDR_Q4_0_Q8_1_MMQ, vec_dot_q4_0_q8_1_mul_mat>
  824. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  825. #else
  826. GGML_UNUSED(vec_dot_q4_0_q8_1_mul_mat);
  827. NO_DEVICE_CODE;
  828. #endif // __CUDA_ARCH__ >= CC_VOLTA
  829. }
  830. #define MMQ_X_Q4_1_RDNA2 64
  831. #define MMQ_Y_Q4_1_RDNA2 128
  832. #define NWARPS_Q4_1_RDNA2 8
  833. #define MMQ_X_Q4_1_RDNA1 64
  834. #define MMQ_Y_Q4_1_RDNA1 64
  835. #define NWARPS_Q4_1_RDNA1 8
  836. #if defined(CUDA_USE_TENSOR_CORES)
  837. #define MMQ_X_Q4_1_AMPERE 4
  838. #define MMQ_Y_Q4_1_AMPERE 32
  839. #define NWARPS_Q4_1_AMPERE 4
  840. #else
  841. #define MMQ_X_Q4_1_AMPERE 64
  842. #define MMQ_Y_Q4_1_AMPERE 128
  843. #define NWARPS_Q4_1_AMPERE 4
  844. #endif
  845. #define MMQ_X_Q4_1_PASCAL 64
  846. #define MMQ_Y_Q4_1_PASCAL 64
  847. #define NWARPS_Q4_1_PASCAL 8
  848. template <bool need_check> static __global__ void
  849. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  850. #if defined(RDNA3) || defined(RDNA2)
  851. __launch_bounds__(WARP_SIZE*NWARPS_Q4_1_RDNA2, 2)
  852. #endif // defined(RDNA3) || defined(RDNA2)
  853. #elif __CUDA_ARCH__ < CC_VOLTA
  854. __launch_bounds__(WARP_SIZE*NWARPS_Q4_1_PASCAL, 2)
  855. #endif // __CUDA_ARCH__ < CC_VOLTA
  856. mul_mat_q4_1(
  857. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  858. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  859. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  860. #if defined(RDNA3) || defined(RDNA2)
  861. const int mmq_x = MMQ_X_Q4_1_RDNA2;
  862. const int mmq_y = MMQ_Y_Q4_1_RDNA2;
  863. const int nwarps = NWARPS_Q4_1_RDNA2;
  864. #else
  865. const int mmq_x = MMQ_X_Q4_1_RDNA1;
  866. const int mmq_y = MMQ_Y_Q4_1_RDNA1;
  867. const int nwarps = NWARPS_Q4_1_RDNA1;
  868. #endif // defined(RDNA3) || defined(RDNA2)
  869. mul_mat_q<QK4_1, QR4_1, QI4_1, true, block_q4_1, mmq_x, mmq_y, nwarps, allocate_tiles_q4_1<mmq_y>,
  870. load_tiles_q4_1<mmq_y, nwarps, need_check>, VDR_Q4_1_Q8_1_MMQ, vec_dot_q4_1_q8_1_mul_mat>
  871. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  872. #elif __CUDA_ARCH__ >= CC_VOLTA
  873. const int mmq_x = MMQ_X_Q4_1_AMPERE;
  874. const int mmq_y = MMQ_Y_Q4_1_AMPERE;
  875. const int nwarps = NWARPS_Q4_1_AMPERE;
  876. mul_mat_q<QK4_1, QR4_1, QI4_1, true, block_q4_1, mmq_x, mmq_y, nwarps, allocate_tiles_q4_1<mmq_y>,
  877. load_tiles_q4_1<mmq_y, nwarps, need_check>, VDR_Q4_1_Q8_1_MMQ, vec_dot_q4_1_q8_1_mul_mat>
  878. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  879. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  880. const int mmq_x = MMQ_X_Q4_1_PASCAL;
  881. const int mmq_y = MMQ_Y_Q4_1_PASCAL;
  882. const int nwarps = NWARPS_Q4_1_PASCAL;
  883. mul_mat_q<QK4_1, QR4_1, QI4_1, true, block_q4_1, mmq_x, mmq_y, nwarps, allocate_tiles_q4_1<mmq_y>,
  884. load_tiles_q4_1<mmq_y, nwarps, need_check>, VDR_Q4_1_Q8_1_MMQ, vec_dot_q4_1_q8_1_mul_mat>
  885. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  886. #else
  887. GGML_UNUSED(vec_dot_q4_1_q8_1_mul_mat);
  888. NO_DEVICE_CODE;
  889. #endif // __CUDA_ARCH__ >= CC_VOLTA
  890. }
  891. #define MMQ_X_Q5_0_RDNA2 64
  892. #define MMQ_Y_Q5_0_RDNA2 128
  893. #define NWARPS_Q5_0_RDNA2 8
  894. #define MMQ_X_Q5_0_RDNA1 64
  895. #define MMQ_Y_Q5_0_RDNA1 64
  896. #define NWARPS_Q5_0_RDNA1 8
  897. #if defined(CUDA_USE_TENSOR_CORES)
  898. #define MMQ_X_Q5_0_AMPERE 4
  899. #define MMQ_Y_Q5_0_AMPERE 32
  900. #define NWARPS_Q5_0_AMPERE 4
  901. #else
  902. #define MMQ_X_Q5_0_AMPERE 128
  903. #define MMQ_Y_Q5_0_AMPERE 64
  904. #define NWARPS_Q5_0_AMPERE 4
  905. #endif
  906. #define MMQ_X_Q5_0_PASCAL 64
  907. #define MMQ_Y_Q5_0_PASCAL 64
  908. #define NWARPS_Q5_0_PASCAL 8
  909. template <bool need_check> static __global__ void
  910. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  911. #if defined(RDNA3) || defined(RDNA2)
  912. __launch_bounds__(WARP_SIZE*NWARPS_Q5_0_RDNA2, 2)
  913. #endif // defined(RDNA3) || defined(RDNA2)
  914. #endif // defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  915. mul_mat_q5_0(
  916. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  917. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  918. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  919. #if defined(RDNA3) || defined(RDNA2)
  920. const int mmq_x = MMQ_X_Q5_0_RDNA2;
  921. const int mmq_y = MMQ_Y_Q5_0_RDNA2;
  922. const int nwarps = NWARPS_Q5_0_RDNA2;
  923. #else
  924. const int mmq_x = MMQ_X_Q5_0_RDNA1;
  925. const int mmq_y = MMQ_Y_Q5_0_RDNA1;
  926. const int nwarps = NWARPS_Q5_0_RDNA1;
  927. #endif // defined(RDNA3) || defined(RDNA2)
  928. mul_mat_q<QK5_0, QR5_0, QI5_0, false, block_q5_0, mmq_x, mmq_y, nwarps, allocate_tiles_q5_0<mmq_y>,
  929. load_tiles_q5_0<mmq_y, nwarps, need_check>, VDR_Q5_0_Q8_1_MMQ, vec_dot_q5_0_q8_1_mul_mat>
  930. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  931. #elif __CUDA_ARCH__ >= CC_VOLTA
  932. const int mmq_x = MMQ_X_Q5_0_AMPERE;
  933. const int mmq_y = MMQ_Y_Q5_0_AMPERE;
  934. const int nwarps = NWARPS_Q5_0_AMPERE;
  935. mul_mat_q<QK5_0, QR5_0, QI5_0, false, block_q5_0, mmq_x, mmq_y, nwarps, allocate_tiles_q5_0<mmq_y>,
  936. load_tiles_q5_0<mmq_y, nwarps, need_check>, VDR_Q5_0_Q8_1_MMQ, vec_dot_q5_0_q8_1_mul_mat>
  937. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  938. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  939. const int mmq_x = MMQ_X_Q5_0_PASCAL;
  940. const int mmq_y = MMQ_Y_Q5_0_PASCAL;
  941. const int nwarps = NWARPS_Q5_0_PASCAL;
  942. mul_mat_q<QK5_0, QR5_0, QI5_0, false, block_q5_0, mmq_x, mmq_y, nwarps, allocate_tiles_q5_0<mmq_y>,
  943. load_tiles_q5_0<mmq_y, nwarps, need_check>, VDR_Q5_0_Q8_1_MMQ, vec_dot_q5_0_q8_1_mul_mat>
  944. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  945. #else
  946. GGML_UNUSED(vec_dot_q5_0_q8_1_mul_mat);
  947. NO_DEVICE_CODE;
  948. #endif // __CUDA_ARCH__ >= CC_VOLTA
  949. }
  950. #define MMQ_X_Q5_1_RDNA2 64
  951. #define MMQ_Y_Q5_1_RDNA2 128
  952. #define NWARPS_Q5_1_RDNA2 8
  953. #define MMQ_X_Q5_1_RDNA1 64
  954. #define MMQ_Y_Q5_1_RDNA1 64
  955. #define NWARPS_Q5_1_RDNA1 8
  956. #if defined(CUDA_USE_TENSOR_CORES)
  957. #define MMQ_X_Q5_1_AMPERE 4
  958. #define MMQ_Y_Q5_1_AMPERE 32
  959. #define NWARPS_Q5_1_AMPERE 4
  960. #else
  961. #define MMQ_X_Q5_1_AMPERE 128
  962. #define MMQ_Y_Q5_1_AMPERE 64
  963. #define NWARPS_Q5_1_AMPERE 4
  964. #endif
  965. #define MMQ_X_Q5_1_PASCAL 64
  966. #define MMQ_Y_Q5_1_PASCAL 64
  967. #define NWARPS_Q5_1_PASCAL 8
  968. template <bool need_check> static __global__ void
  969. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  970. #if defined(RDNA3) || defined(RDNA2)
  971. __launch_bounds__(WARP_SIZE*NWARPS_Q5_1_RDNA2, 2)
  972. #endif // defined(RDNA3) || defined(RDNA2)
  973. #endif // defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  974. mul_mat_q5_1(
  975. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  976. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  977. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  978. #if defined(RDNA3) || defined(RDNA2)
  979. const int mmq_x = MMQ_X_Q5_1_RDNA2;
  980. const int mmq_y = MMQ_Y_Q5_1_RDNA2;
  981. const int nwarps = NWARPS_Q5_1_RDNA2;
  982. #else
  983. const int mmq_x = MMQ_X_Q5_1_RDNA1;
  984. const int mmq_y = MMQ_Y_Q5_1_RDNA1;
  985. const int nwarps = NWARPS_Q5_1_RDNA1;
  986. #endif // defined(RDNA3) || defined(RDNA2)
  987. mul_mat_q<QK5_1, QR5_1, QI5_1, true, block_q5_1, mmq_x, mmq_y, nwarps, allocate_tiles_q5_1<mmq_y>,
  988. load_tiles_q5_1<mmq_y, nwarps, need_check>, VDR_Q5_1_Q8_1_MMQ, vec_dot_q5_1_q8_1_mul_mat>
  989. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  990. #elif __CUDA_ARCH__ >= CC_VOLTA
  991. const int mmq_x = MMQ_X_Q5_1_AMPERE;
  992. const int mmq_y = MMQ_Y_Q5_1_AMPERE;
  993. const int nwarps = NWARPS_Q5_1_AMPERE;
  994. mul_mat_q<QK5_1, QR5_1, QI5_1, true, block_q5_1, mmq_x, mmq_y, nwarps, allocate_tiles_q5_1<mmq_y>,
  995. load_tiles_q5_1<mmq_y, nwarps, need_check>, VDR_Q5_1_Q8_1_MMQ, vec_dot_q5_1_q8_1_mul_mat>
  996. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  997. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  998. const int mmq_x = MMQ_X_Q5_1_PASCAL;
  999. const int mmq_y = MMQ_Y_Q5_1_PASCAL;
  1000. const int nwarps = NWARPS_Q5_1_PASCAL;
  1001. mul_mat_q<QK5_1, QR5_1, QI5_1, true, block_q5_1, mmq_x, mmq_y, nwarps, allocate_tiles_q5_1<mmq_y>,
  1002. load_tiles_q5_1<mmq_y, nwarps, need_check>, VDR_Q5_1_Q8_1_MMQ, vec_dot_q5_1_q8_1_mul_mat>
  1003. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1004. #else
  1005. GGML_UNUSED(vec_dot_q5_1_q8_1_mul_mat);
  1006. NO_DEVICE_CODE;
  1007. #endif // __CUDA_ARCH__ >= CC_VOLTA
  1008. }
  1009. #define MMQ_X_Q8_0_RDNA2 64
  1010. #define MMQ_Y_Q8_0_RDNA2 128
  1011. #define NWARPS_Q8_0_RDNA2 8
  1012. #define MMQ_X_Q8_0_RDNA1 64
  1013. #define MMQ_Y_Q8_0_RDNA1 64
  1014. #define NWARPS_Q8_0_RDNA1 8
  1015. #if defined(CUDA_USE_TENSOR_CORES)
  1016. #define MMQ_X_Q8_0_AMPERE 4
  1017. #define MMQ_Y_Q8_0_AMPERE 32
  1018. #define NWARPS_Q8_0_AMPERE 4
  1019. #else
  1020. #define MMQ_X_Q8_0_AMPERE 128
  1021. #define MMQ_Y_Q8_0_AMPERE 64
  1022. #define NWARPS_Q8_0_AMPERE 4
  1023. #endif
  1024. #define MMQ_X_Q8_0_PASCAL 64
  1025. #define MMQ_Y_Q8_0_PASCAL 64
  1026. #define NWARPS_Q8_0_PASCAL 8
  1027. template <bool need_check> static __global__ void
  1028. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1029. #if defined(RDNA3) || defined(RDNA2)
  1030. __launch_bounds__(WARP_SIZE*NWARPS_Q8_0_RDNA2, 2)
  1031. #endif // defined(RDNA3) || defined(RDNA2)
  1032. #endif // defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1033. mul_mat_q8_0(
  1034. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  1035. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  1036. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1037. #if defined(RDNA3) || defined(RDNA2)
  1038. const int mmq_x = MMQ_X_Q8_0_RDNA2;
  1039. const int mmq_y = MMQ_Y_Q8_0_RDNA2;
  1040. const int nwarps = NWARPS_Q8_0_RDNA2;
  1041. #else
  1042. const int mmq_x = MMQ_X_Q8_0_RDNA1;
  1043. const int mmq_y = MMQ_Y_Q8_0_RDNA1;
  1044. const int nwarps = NWARPS_Q8_0_RDNA1;
  1045. #endif // defined(RDNA3) || defined(RDNA2)
  1046. mul_mat_q<QK8_0, QR8_0, QI8_0, false, block_q8_0, mmq_x, mmq_y, nwarps, allocate_tiles_q8_0<mmq_y>,
  1047. load_tiles_q8_0<mmq_y, nwarps, need_check>, VDR_Q8_0_Q8_1_MMQ, vec_dot_q8_0_q8_1_mul_mat>
  1048. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1049. #elif __CUDA_ARCH__ >= CC_VOLTA
  1050. const int mmq_x = MMQ_X_Q8_0_AMPERE;
  1051. const int mmq_y = MMQ_Y_Q8_0_AMPERE;
  1052. const int nwarps = NWARPS_Q8_0_AMPERE;
  1053. mul_mat_q<QK8_0, QR8_0, QI8_0, false, block_q8_0, mmq_x, mmq_y, nwarps, allocate_tiles_q8_0<mmq_y>,
  1054. load_tiles_q8_0<mmq_y, nwarps, need_check>, VDR_Q8_0_Q8_1_MMQ, vec_dot_q8_0_q8_1_mul_mat>
  1055. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1056. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  1057. const int mmq_x = MMQ_X_Q8_0_PASCAL;
  1058. const int mmq_y = MMQ_Y_Q8_0_PASCAL;
  1059. const int nwarps = NWARPS_Q8_0_PASCAL;
  1060. mul_mat_q<QK8_0, QR8_0, QI8_0, false, block_q8_0, mmq_x, mmq_y, nwarps, allocate_tiles_q8_0<mmq_y>,
  1061. load_tiles_q8_0<mmq_y, nwarps, need_check>, VDR_Q8_0_Q8_1_MMQ, vec_dot_q8_0_q8_1_mul_mat>
  1062. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1063. #else
  1064. GGML_UNUSED(vec_dot_q8_0_q8_1_mul_mat);
  1065. NO_DEVICE_CODE;
  1066. #endif // __CUDA_ARCH__ >= CC_VOLTA
  1067. }
  1068. #define MMQ_X_Q2_K_RDNA2 64
  1069. #define MMQ_Y_Q2_K_RDNA2 128
  1070. #define NWARPS_Q2_K_RDNA2 8
  1071. #define MMQ_X_Q2_K_RDNA1 128
  1072. #define MMQ_Y_Q2_K_RDNA1 32
  1073. #define NWARPS_Q2_K_RDNA1 8
  1074. #if defined(CUDA_USE_TENSOR_CORES)
  1075. #define MMQ_X_Q2_K_AMPERE 4
  1076. #define MMQ_Y_Q2_K_AMPERE 32
  1077. #define NWARPS_Q2_K_AMPERE 4
  1078. #else
  1079. #define MMQ_X_Q2_K_AMPERE 64
  1080. #define MMQ_Y_Q2_K_AMPERE 128
  1081. #define NWARPS_Q2_K_AMPERE 4
  1082. #endif
  1083. #define MMQ_X_Q2_K_PASCAL 64
  1084. #define MMQ_Y_Q2_K_PASCAL 64
  1085. #define NWARPS_Q2_K_PASCAL 8
  1086. template <bool need_check> static __global__ void
  1087. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1088. #if defined(RDNA3) || defined(RDNA2)
  1089. __launch_bounds__(WARP_SIZE*NWARPS_Q2_K_RDNA2, 2)
  1090. #endif // defined(RDNA3) || defined(RDNA2)
  1091. #endif // defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1092. mul_mat_q2_K(
  1093. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  1094. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  1095. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1096. #if defined(RDNA3) || defined(RDNA2)
  1097. const int mmq_x = MMQ_X_Q2_K_RDNA2;
  1098. const int mmq_y = MMQ_Y_Q2_K_RDNA2;
  1099. const int nwarps = NWARPS_Q2_K_RDNA2;
  1100. #else
  1101. const int mmq_x = MMQ_X_Q2_K_RDNA1;
  1102. const int mmq_y = MMQ_Y_Q2_K_RDNA1;
  1103. const int nwarps = NWARPS_Q2_K_RDNA1;
  1104. #endif // defined(RDNA3) || defined(RDNA2)
  1105. mul_mat_q<QK_K, QR2_K, QI2_K, false, block_q2_K, mmq_x, mmq_y, nwarps, allocate_tiles_q2_K<mmq_y>,
  1106. load_tiles_q2_K<mmq_y, nwarps, need_check>, VDR_Q2_K_Q8_1_MMQ, vec_dot_q2_K_q8_1_mul_mat>
  1107. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1108. #elif __CUDA_ARCH__ >= CC_VOLTA
  1109. const int mmq_x = MMQ_X_Q2_K_AMPERE;
  1110. const int mmq_y = MMQ_Y_Q2_K_AMPERE;
  1111. const int nwarps = NWARPS_Q2_K_AMPERE;
  1112. mul_mat_q<QK_K, QR2_K, QI2_K, false, block_q2_K, mmq_x, mmq_y, nwarps, allocate_tiles_q2_K<mmq_y>,
  1113. load_tiles_q2_K<mmq_y, nwarps, need_check>, VDR_Q2_K_Q8_1_MMQ, vec_dot_q2_K_q8_1_mul_mat>
  1114. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1115. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  1116. const int mmq_x = MMQ_X_Q2_K_PASCAL;
  1117. const int mmq_y = MMQ_Y_Q2_K_PASCAL;
  1118. const int nwarps = NWARPS_Q2_K_PASCAL;
  1119. mul_mat_q<QK_K, QR2_K, QI2_K, false, block_q2_K, mmq_x, mmq_y, nwarps, allocate_tiles_q2_K<mmq_y>,
  1120. load_tiles_q2_K<mmq_y, nwarps, need_check>, VDR_Q2_K_Q8_1_MMQ, vec_dot_q2_K_q8_1_mul_mat>
  1121. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1122. #else
  1123. GGML_UNUSED(vec_dot_q2_K_q8_1_mul_mat);
  1124. NO_DEVICE_CODE;
  1125. #endif // __CUDA_ARCH__ >= CC_VOLTA
  1126. }
  1127. #define MMQ_X_Q3_K_RDNA2 128
  1128. #define MMQ_Y_Q3_K_RDNA2 64
  1129. #define NWARPS_Q3_K_RDNA2 8
  1130. #define MMQ_X_Q3_K_RDNA1 32
  1131. #define MMQ_Y_Q3_K_RDNA1 128
  1132. #define NWARPS_Q3_K_RDNA1 8
  1133. #if defined(CUDA_USE_TENSOR_CORES)
  1134. #define MMQ_X_Q3_K_AMPERE 4
  1135. #define MMQ_Y_Q3_K_AMPERE 32
  1136. #define NWARPS_Q3_K_AMPERE 4
  1137. #else
  1138. #define MMQ_X_Q3_K_AMPERE 128
  1139. #define MMQ_Y_Q3_K_AMPERE 128
  1140. #define NWARPS_Q3_K_AMPERE 4
  1141. #endif
  1142. #define MMQ_X_Q3_K_PASCAL 64
  1143. #define MMQ_Y_Q3_K_PASCAL 64
  1144. #define NWARPS_Q3_K_PASCAL 8
  1145. template <bool need_check> static __global__ void
  1146. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1147. #if defined(RDNA3) || defined(RDNA2)
  1148. __launch_bounds__(WARP_SIZE*NWARPS_Q3_K_RDNA2, 2)
  1149. #endif // defined(RDNA3) || defined(RDNA2)
  1150. #elif __CUDA_ARCH__ < CC_VOLTA
  1151. __launch_bounds__(WARP_SIZE*NWARPS_Q3_K_PASCAL, 2)
  1152. #endif // __CUDA_ARCH__ < CC_VOLTA
  1153. mul_mat_q3_K(
  1154. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  1155. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  1156. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1157. #if defined(RDNA3) || defined(RDNA2)
  1158. const int mmq_x = MMQ_X_Q3_K_RDNA2;
  1159. const int mmq_y = MMQ_Y_Q3_K_RDNA2;
  1160. const int nwarps = NWARPS_Q3_K_RDNA2;
  1161. #else
  1162. const int mmq_x = MMQ_X_Q3_K_RDNA1;
  1163. const int mmq_y = MMQ_Y_Q3_K_RDNA1;
  1164. const int nwarps = NWARPS_Q3_K_RDNA1;
  1165. #endif // defined(RDNA3) || defined(RDNA2)
  1166. mul_mat_q<QK_K, QR3_K, QI3_K, false, block_q3_K, mmq_x, mmq_y, nwarps, allocate_tiles_q3_K<mmq_y>,
  1167. load_tiles_q3_K<mmq_y, nwarps, need_check>, VDR_Q3_K_Q8_1_MMQ, vec_dot_q3_K_q8_1_mul_mat>
  1168. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1169. #elif __CUDA_ARCH__ >= CC_VOLTA
  1170. const int mmq_x = MMQ_X_Q3_K_AMPERE;
  1171. const int mmq_y = MMQ_Y_Q3_K_AMPERE;
  1172. const int nwarps = NWARPS_Q3_K_AMPERE;
  1173. mul_mat_q<QK_K, QR3_K, QI3_K, false, block_q3_K, mmq_x, mmq_y, nwarps, allocate_tiles_q3_K<mmq_y>,
  1174. load_tiles_q3_K<mmq_y, nwarps, need_check>, VDR_Q3_K_Q8_1_MMQ, vec_dot_q3_K_q8_1_mul_mat>
  1175. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1176. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  1177. const int mmq_x = MMQ_X_Q3_K_PASCAL;
  1178. const int mmq_y = MMQ_Y_Q3_K_PASCAL;
  1179. const int nwarps = NWARPS_Q3_K_PASCAL;
  1180. mul_mat_q<QK_K, QR3_K, QI3_K, false, block_q3_K, mmq_x, mmq_y, nwarps, allocate_tiles_q3_K<mmq_y>,
  1181. load_tiles_q3_K<mmq_y, nwarps, need_check>, VDR_Q3_K_Q8_1_MMQ, vec_dot_q3_K_q8_1_mul_mat>
  1182. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1183. #else
  1184. GGML_UNUSED(vec_dot_q3_K_q8_1_mul_mat);
  1185. NO_DEVICE_CODE;
  1186. #endif // __CUDA_ARCH__ >= CC_VOLTA
  1187. }
  1188. #define MMQ_X_Q4_K_RDNA2 64
  1189. #define MMQ_Y_Q4_K_RDNA2 128
  1190. #define NWARPS_Q4_K_RDNA2 8
  1191. #define MMQ_X_Q4_K_RDNA1 32
  1192. #define MMQ_Y_Q4_K_RDNA1 64
  1193. #define NWARPS_Q4_K_RDNA1 8
  1194. #if defined(CUDA_USE_TENSOR_CORES)
  1195. #define MMQ_X_Q4_K_AMPERE 4
  1196. #define MMQ_Y_Q4_K_AMPERE 32
  1197. #define NWARPS_Q4_K_AMPERE 4
  1198. #else
  1199. #define MMQ_X_Q4_K_AMPERE 64
  1200. #define MMQ_Y_Q4_K_AMPERE 128
  1201. #define NWARPS_Q4_K_AMPERE 4
  1202. #endif
  1203. #define MMQ_X_Q4_K_PASCAL 64
  1204. #define MMQ_Y_Q4_K_PASCAL 64
  1205. #define NWARPS_Q4_K_PASCAL 8
  1206. template <bool need_check> static __global__ void
  1207. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1208. #if defined(RDNA3) || defined(RDNA2)
  1209. __launch_bounds__(WARP_SIZE*NWARPS_Q4_K_RDNA2, 2)
  1210. #endif // defined(RDNA3) || defined(RDNA2)
  1211. #elif __CUDA_ARCH__ < CC_VOLTA
  1212. __launch_bounds__(WARP_SIZE*NWARPS_Q4_K_PASCAL, 2)
  1213. #endif // __CUDA_ARCH__ < CC_VOLTA
  1214. mul_mat_q4_K(
  1215. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  1216. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  1217. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1218. #if defined(RDNA3) || defined(RDNA2)
  1219. const int mmq_x = MMQ_X_Q4_K_RDNA2;
  1220. const int mmq_y = MMQ_Y_Q4_K_RDNA2;
  1221. const int nwarps = NWARPS_Q4_K_RDNA2;
  1222. #else
  1223. const int mmq_x = MMQ_X_Q4_K_RDNA1;
  1224. const int mmq_y = MMQ_Y_Q4_K_RDNA1;
  1225. const int nwarps = NWARPS_Q4_K_RDNA1;
  1226. #endif // defined(RDNA3) || defined(RDNA2)
  1227. mul_mat_q<QK_K, QR4_K, QI4_K, true, block_q4_K, mmq_x, mmq_y, nwarps, allocate_tiles_q4_K<mmq_y>,
  1228. load_tiles_q4_K<mmq_y, nwarps, need_check>, VDR_Q4_K_Q8_1_MMQ, vec_dot_q4_K_q8_1_mul_mat>
  1229. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1230. #elif __CUDA_ARCH__ >= CC_VOLTA
  1231. const int mmq_x = MMQ_X_Q4_K_AMPERE;
  1232. const int mmq_y = MMQ_Y_Q4_K_AMPERE;
  1233. const int nwarps = NWARPS_Q4_K_AMPERE;
  1234. mul_mat_q<QK_K, QR4_K, QI4_K, true, block_q4_K, mmq_x, mmq_y, nwarps, allocate_tiles_q4_K<mmq_y>,
  1235. load_tiles_q4_K<mmq_y, nwarps, need_check>, VDR_Q4_K_Q8_1_MMQ, vec_dot_q4_K_q8_1_mul_mat>
  1236. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1237. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  1238. const int mmq_x = MMQ_X_Q4_K_PASCAL;
  1239. const int mmq_y = MMQ_Y_Q4_K_PASCAL;
  1240. const int nwarps = NWARPS_Q4_K_PASCAL;
  1241. mul_mat_q<QK_K, QR4_K, QI4_K, true, block_q4_K, mmq_x, mmq_y, nwarps, allocate_tiles_q4_K<mmq_y>,
  1242. load_tiles_q4_K<mmq_y, nwarps, need_check>, VDR_Q4_K_Q8_1_MMQ, vec_dot_q4_K_q8_1_mul_mat>
  1243. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1244. #else
  1245. GGML_UNUSED(vec_dot_q4_K_q8_1_mul_mat);
  1246. NO_DEVICE_CODE;
  1247. #endif // __CUDA_ARCH__ >= CC_VOLTA
  1248. }
  1249. #define MMQ_X_Q5_K_RDNA2 64
  1250. #define MMQ_Y_Q5_K_RDNA2 128
  1251. #define NWARPS_Q5_K_RDNA2 8
  1252. #define MMQ_X_Q5_K_RDNA1 32
  1253. #define MMQ_Y_Q5_K_RDNA1 64
  1254. #define NWARPS_Q5_K_RDNA1 8
  1255. #if defined(CUDA_USE_TENSOR_CORES)
  1256. #define MMQ_X_Q5_K_AMPERE 4
  1257. #define MMQ_Y_Q5_K_AMPERE 32
  1258. #define NWARPS_Q5_K_AMPERE 4
  1259. #else
  1260. #define MMQ_X_Q5_K_AMPERE 64
  1261. #define MMQ_Y_Q5_K_AMPERE 128
  1262. #define NWARPS_Q5_K_AMPERE 4
  1263. #endif
  1264. #define MMQ_X_Q5_K_PASCAL 64
  1265. #define MMQ_Y_Q5_K_PASCAL 64
  1266. #define NWARPS_Q5_K_PASCAL 8
  1267. template <bool need_check> static __global__ void
  1268. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1269. #if defined(RDNA3) || defined(RDNA2)
  1270. __launch_bounds__(WARP_SIZE*NWARPS_Q5_K_RDNA2, 2)
  1271. #endif // defined(RDNA3) || defined(RDNA2)
  1272. #endif // defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1273. mul_mat_q5_K(
  1274. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  1275. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  1276. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1277. #if defined(RDNA3) || defined(RDNA2)
  1278. const int mmq_x = MMQ_X_Q5_K_RDNA2;
  1279. const int mmq_y = MMQ_Y_Q5_K_RDNA2;
  1280. const int nwarps = NWARPS_Q5_K_RDNA2;
  1281. #else
  1282. const int mmq_x = MMQ_X_Q5_K_RDNA1;
  1283. const int mmq_y = MMQ_Y_Q5_K_RDNA1;
  1284. const int nwarps = NWARPS_Q5_K_RDNA1;
  1285. #endif // defined(RDNA3) || defined(RDNA2)
  1286. mul_mat_q<QK_K, QR5_K, QI5_K, true, block_q5_K, mmq_x, mmq_y, nwarps, allocate_tiles_q5_K<mmq_y>,
  1287. load_tiles_q5_K<mmq_y, nwarps, need_check>, VDR_Q5_K_Q8_1_MMQ, vec_dot_q5_K_q8_1_mul_mat>
  1288. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1289. #elif __CUDA_ARCH__ >= CC_VOLTA
  1290. const int mmq_x = MMQ_X_Q5_K_AMPERE;
  1291. const int mmq_y = MMQ_Y_Q5_K_AMPERE;
  1292. const int nwarps = NWARPS_Q5_K_AMPERE;
  1293. mul_mat_q<QK_K, QR5_K, QI5_K, true, block_q5_K, mmq_x, mmq_y, nwarps, allocate_tiles_q5_K<mmq_y>,
  1294. load_tiles_q5_K<mmq_y, nwarps, need_check>, VDR_Q5_K_Q8_1_MMQ, vec_dot_q5_K_q8_1_mul_mat>
  1295. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1296. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  1297. const int mmq_x = MMQ_X_Q5_K_PASCAL;
  1298. const int mmq_y = MMQ_Y_Q5_K_PASCAL;
  1299. const int nwarps = NWARPS_Q5_K_PASCAL;
  1300. mul_mat_q<QK_K, QR5_K, QI5_K, true, block_q5_K, mmq_x, mmq_y, nwarps, allocate_tiles_q5_K<mmq_y>,
  1301. load_tiles_q5_K<mmq_y, nwarps, need_check>, VDR_Q5_K_Q8_1_MMQ, vec_dot_q5_K_q8_1_mul_mat>
  1302. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1303. #else
  1304. GGML_UNUSED(vec_dot_q5_K_q8_1_mul_mat);
  1305. NO_DEVICE_CODE;
  1306. #endif // __CUDA_ARCH__ >= CC_VOLTA
  1307. }
  1308. #define MMQ_X_Q6_K_RDNA2 64
  1309. #define MMQ_Y_Q6_K_RDNA2 128
  1310. #define NWARPS_Q6_K_RDNA2 8
  1311. #define MMQ_X_Q6_K_RDNA1 32
  1312. #define MMQ_Y_Q6_K_RDNA1 64
  1313. #define NWARPS_Q6_K_RDNA1 8
  1314. #if defined(CUDA_USE_TENSOR_CORES)
  1315. #define MMQ_X_Q6_K_AMPERE 4
  1316. #define MMQ_Y_Q6_K_AMPERE 32
  1317. #define NWARPS_Q6_K_AMPERE 4
  1318. #else
  1319. #define MMQ_X_Q6_K_AMPERE 64
  1320. #define MMQ_Y_Q6_K_AMPERE 64
  1321. #define NWARPS_Q6_K_AMPERE 4
  1322. #endif
  1323. #define MMQ_X_Q6_K_PASCAL 64
  1324. #define MMQ_Y_Q6_K_PASCAL 64
  1325. #define NWARPS_Q6_K_PASCAL 8
  1326. template <bool need_check> static __global__ void
  1327. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1328. #if defined(RDNA3) || defined(RDNA2)
  1329. __launch_bounds__(WARP_SIZE*NWARPS_Q6_K_RDNA2, 2)
  1330. #endif // defined(RDNA3) || defined(RDNA2)
  1331. #elif __CUDA_ARCH__ < CC_VOLTA
  1332. __launch_bounds__(WARP_SIZE*NWARPS_Q6_K_PASCAL, 2)
  1333. #endif // __CUDA_ARCH__ < CC_VOLTA
  1334. mul_mat_q6_K(
  1335. const void * __restrict__ vx, const void * __restrict__ vy, float * __restrict__ dst,
  1336. const int ncols_x, const int nrows_x, const int ncols_y, const int nrows_y, const int nrows_dst) {
  1337. #if defined(GGML_USE_HIPBLAS) && defined(__HIP_PLATFORM_AMD__)
  1338. #if defined(RDNA3) || defined(RDNA2)
  1339. const int mmq_x = MMQ_X_Q6_K_RDNA2;
  1340. const int mmq_y = MMQ_Y_Q6_K_RDNA2;
  1341. const int nwarps = NWARPS_Q6_K_RDNA2;
  1342. #else
  1343. const int mmq_x = MMQ_X_Q6_K_RDNA1;
  1344. const int mmq_y = MMQ_Y_Q6_K_RDNA1;
  1345. const int nwarps = NWARPS_Q6_K_RDNA1;
  1346. #endif // defined(RDNA3) || defined(RDNA2)
  1347. mul_mat_q<QK_K, QR6_K, QI6_K, false, block_q6_K, mmq_x, mmq_y, nwarps, allocate_tiles_q6_K<mmq_y>,
  1348. load_tiles_q6_K<mmq_y, nwarps, need_check>, VDR_Q6_K_Q8_1_MMQ, vec_dot_q6_K_q8_1_mul_mat>
  1349. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1350. #elif __CUDA_ARCH__ >= CC_VOLTA
  1351. const int mmq_x = MMQ_X_Q6_K_AMPERE;
  1352. const int mmq_y = MMQ_Y_Q6_K_AMPERE;
  1353. const int nwarps = NWARPS_Q6_K_AMPERE;
  1354. mul_mat_q<QK_K, QR6_K, QI6_K, false, block_q6_K, mmq_x, mmq_y, nwarps, allocate_tiles_q6_K<mmq_y>,
  1355. load_tiles_q6_K<mmq_y, nwarps, need_check>, VDR_Q6_K_Q8_1_MMQ, vec_dot_q6_K_q8_1_mul_mat>
  1356. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1357. #elif __CUDA_ARCH__ >= MIN_CC_DP4A
  1358. const int mmq_x = MMQ_X_Q6_K_PASCAL;
  1359. const int mmq_y = MMQ_Y_Q6_K_PASCAL;
  1360. const int nwarps = NWARPS_Q6_K_PASCAL;
  1361. mul_mat_q<QK_K, QR6_K, QI6_K, false, block_q6_K, mmq_x, mmq_y, nwarps, allocate_tiles_q6_K<mmq_y>,
  1362. load_tiles_q6_K<mmq_y, nwarps, need_check>, VDR_Q6_K_Q8_1_MMQ, vec_dot_q6_K_q8_1_mul_mat>
  1363. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1364. #else
  1365. GGML_UNUSED(vec_dot_q6_K_q8_1_mul_mat);
  1366. NO_DEVICE_CODE;
  1367. #endif // __CUDA_ARCH__ >= CC_VOLTA
  1368. }
  1369. static void ggml_mul_mat_q4_0_q8_1_cuda(
  1370. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1371. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1372. int id = ggml_cuda_get_device();
  1373. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1374. int mmq_x, mmq_y, nwarps;
  1375. if (compute_capability >= CC_RDNA2) {
  1376. mmq_x = MMQ_X_Q4_0_RDNA2;
  1377. mmq_y = MMQ_Y_Q4_0_RDNA2;
  1378. nwarps = NWARPS_Q4_0_RDNA2;
  1379. } else if (compute_capability >= CC_OFFSET_AMD) {
  1380. mmq_x = MMQ_X_Q4_0_RDNA1;
  1381. mmq_y = MMQ_Y_Q4_0_RDNA1;
  1382. nwarps = NWARPS_Q4_0_RDNA1;
  1383. } else if (compute_capability >= CC_VOLTA) {
  1384. mmq_x = MMQ_X_Q4_0_AMPERE;
  1385. mmq_y = MMQ_Y_Q4_0_AMPERE;
  1386. nwarps = NWARPS_Q4_0_AMPERE;
  1387. } else if (compute_capability >= MIN_CC_DP4A) {
  1388. mmq_x = MMQ_X_Q4_0_PASCAL;
  1389. mmq_y = MMQ_Y_Q4_0_PASCAL;
  1390. nwarps = NWARPS_Q4_0_PASCAL;
  1391. } else {
  1392. GGML_ASSERT(false);
  1393. }
  1394. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1395. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1396. const dim3 block_nums(block_num_x, block_num_y, 1);
  1397. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1398. if (nrows_x % mmq_y == 0) {
  1399. const bool need_check = false;
  1400. mul_mat_q4_0<need_check><<<block_nums, block_dims, 0, stream>>>
  1401. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1402. } else {
  1403. const bool need_check = true;
  1404. mul_mat_q4_0<need_check><<<block_nums, block_dims, 0, stream>>>
  1405. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1406. }
  1407. }
  1408. static void ggml_mul_mat_q4_1_q8_1_cuda(
  1409. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1410. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1411. int id = ggml_cuda_get_device();
  1412. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1413. int mmq_x, mmq_y, nwarps;
  1414. if (compute_capability >= CC_RDNA2) {
  1415. mmq_x = MMQ_X_Q4_1_RDNA2;
  1416. mmq_y = MMQ_Y_Q4_1_RDNA2;
  1417. nwarps = NWARPS_Q4_1_RDNA2;
  1418. } else if (compute_capability >= CC_OFFSET_AMD) {
  1419. mmq_x = MMQ_X_Q4_1_RDNA1;
  1420. mmq_y = MMQ_Y_Q4_1_RDNA1;
  1421. nwarps = NWARPS_Q4_1_RDNA1;
  1422. } else if (compute_capability >= CC_VOLTA) {
  1423. mmq_x = MMQ_X_Q4_1_AMPERE;
  1424. mmq_y = MMQ_Y_Q4_1_AMPERE;
  1425. nwarps = NWARPS_Q4_1_AMPERE;
  1426. } else if (compute_capability >= MIN_CC_DP4A) {
  1427. mmq_x = MMQ_X_Q4_1_PASCAL;
  1428. mmq_y = MMQ_Y_Q4_1_PASCAL;
  1429. nwarps = NWARPS_Q4_1_PASCAL;
  1430. } else {
  1431. GGML_ASSERT(false);
  1432. }
  1433. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1434. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1435. const dim3 block_nums(block_num_x, block_num_y, 1);
  1436. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1437. if (nrows_x % mmq_y == 0) {
  1438. const bool need_check = false;
  1439. mul_mat_q4_1<need_check><<<block_nums, block_dims, 0, stream>>>
  1440. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1441. } else {
  1442. const bool need_check = true;
  1443. mul_mat_q4_1<need_check><<<block_nums, block_dims, 0, stream>>>
  1444. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1445. }
  1446. }
  1447. static void ggml_mul_mat_q5_0_q8_1_cuda(
  1448. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1449. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1450. int id = ggml_cuda_get_device();
  1451. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1452. int mmq_x, mmq_y, nwarps;
  1453. if (compute_capability >= CC_RDNA2) {
  1454. mmq_x = MMQ_X_Q5_0_RDNA2;
  1455. mmq_y = MMQ_Y_Q5_0_RDNA2;
  1456. nwarps = NWARPS_Q5_0_RDNA2;
  1457. } else if (compute_capability >= CC_OFFSET_AMD) {
  1458. mmq_x = MMQ_X_Q5_0_RDNA1;
  1459. mmq_y = MMQ_Y_Q5_0_RDNA1;
  1460. nwarps = NWARPS_Q5_0_RDNA1;
  1461. } else if (compute_capability >= CC_VOLTA) {
  1462. mmq_x = MMQ_X_Q5_0_AMPERE;
  1463. mmq_y = MMQ_Y_Q5_0_AMPERE;
  1464. nwarps = NWARPS_Q5_0_AMPERE;
  1465. } else if (compute_capability >= MIN_CC_DP4A) {
  1466. mmq_x = MMQ_X_Q5_0_PASCAL;
  1467. mmq_y = MMQ_Y_Q5_0_PASCAL;
  1468. nwarps = NWARPS_Q5_0_PASCAL;
  1469. } else {
  1470. GGML_ASSERT(false);
  1471. }
  1472. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1473. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1474. const dim3 block_nums(block_num_x, block_num_y, 1);
  1475. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1476. if (nrows_x % mmq_y == 0) {
  1477. const bool need_check = false;
  1478. mul_mat_q5_0<need_check><<<block_nums, block_dims, 0, stream>>>
  1479. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1480. } else {
  1481. const bool need_check = true;
  1482. mul_mat_q5_0<need_check><<<block_nums, block_dims, 0, stream>>>
  1483. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1484. }
  1485. }
  1486. static void ggml_mul_mat_q5_1_q8_1_cuda(
  1487. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1488. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1489. int id = ggml_cuda_get_device();
  1490. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1491. int mmq_x, mmq_y, nwarps;
  1492. if (compute_capability >= CC_RDNA2) {
  1493. mmq_x = MMQ_X_Q5_1_RDNA2;
  1494. mmq_y = MMQ_Y_Q5_1_RDNA2;
  1495. nwarps = NWARPS_Q5_1_RDNA2;
  1496. } else if (compute_capability >= CC_OFFSET_AMD) {
  1497. mmq_x = MMQ_X_Q5_1_RDNA1;
  1498. mmq_y = MMQ_Y_Q5_1_RDNA1;
  1499. nwarps = NWARPS_Q5_1_RDNA1;
  1500. } else if (compute_capability >= CC_VOLTA) {
  1501. mmq_x = MMQ_X_Q5_1_AMPERE;
  1502. mmq_y = MMQ_Y_Q5_1_AMPERE;
  1503. nwarps = NWARPS_Q5_1_AMPERE;
  1504. } else if (compute_capability >= MIN_CC_DP4A) {
  1505. mmq_x = MMQ_X_Q5_1_PASCAL;
  1506. mmq_y = MMQ_Y_Q5_1_PASCAL;
  1507. nwarps = NWARPS_Q5_1_PASCAL;
  1508. } else {
  1509. GGML_ASSERT(false);
  1510. }
  1511. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1512. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1513. const dim3 block_nums(block_num_x, block_num_y, 1);
  1514. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1515. if (nrows_x % mmq_y == 0) {
  1516. const bool need_check = false;
  1517. mul_mat_q5_1<need_check><<<block_nums, block_dims, 0, stream>>>
  1518. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1519. } else {
  1520. const bool need_check = true;
  1521. mul_mat_q5_1<need_check><<<block_nums, block_dims, 0, stream>>>
  1522. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1523. }
  1524. }
  1525. static void ggml_mul_mat_q8_0_q8_1_cuda(
  1526. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1527. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1528. int id = ggml_cuda_get_device();
  1529. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1530. int mmq_x, mmq_y, nwarps;
  1531. if (compute_capability >= CC_RDNA2) {
  1532. mmq_x = MMQ_X_Q8_0_RDNA2;
  1533. mmq_y = MMQ_Y_Q8_0_RDNA2;
  1534. nwarps = NWARPS_Q8_0_RDNA2;
  1535. } else if (compute_capability >= CC_OFFSET_AMD) {
  1536. mmq_x = MMQ_X_Q8_0_RDNA1;
  1537. mmq_y = MMQ_Y_Q8_0_RDNA1;
  1538. nwarps = NWARPS_Q8_0_RDNA1;
  1539. } else if (compute_capability >= CC_VOLTA) {
  1540. mmq_x = MMQ_X_Q8_0_AMPERE;
  1541. mmq_y = MMQ_Y_Q8_0_AMPERE;
  1542. nwarps = NWARPS_Q8_0_AMPERE;
  1543. } else if (compute_capability >= MIN_CC_DP4A) {
  1544. mmq_x = MMQ_X_Q8_0_PASCAL;
  1545. mmq_y = MMQ_Y_Q8_0_PASCAL;
  1546. nwarps = NWARPS_Q8_0_PASCAL;
  1547. } else {
  1548. GGML_ASSERT(false);
  1549. }
  1550. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1551. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1552. const dim3 block_nums(block_num_x, block_num_y, 1);
  1553. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1554. if (nrows_x % mmq_y == 0) {
  1555. const bool need_check = false;
  1556. mul_mat_q8_0<need_check><<<block_nums, block_dims, 0, stream>>>
  1557. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1558. } else {
  1559. const bool need_check = true;
  1560. mul_mat_q8_0<need_check><<<block_nums, block_dims, 0, stream>>>
  1561. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1562. }
  1563. }
  1564. static void ggml_mul_mat_q2_K_q8_1_cuda(
  1565. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1566. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1567. int id = ggml_cuda_get_device();
  1568. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1569. int mmq_x, mmq_y, nwarps;
  1570. if (compute_capability >= CC_RDNA2) {
  1571. mmq_x = MMQ_X_Q2_K_RDNA2;
  1572. mmq_y = MMQ_Y_Q2_K_RDNA2;
  1573. nwarps = NWARPS_Q2_K_RDNA2;
  1574. } else if (compute_capability >= CC_OFFSET_AMD) {
  1575. mmq_x = MMQ_X_Q2_K_RDNA1;
  1576. mmq_y = MMQ_Y_Q2_K_RDNA1;
  1577. nwarps = NWARPS_Q2_K_RDNA1;
  1578. } else if (compute_capability >= CC_VOLTA) {
  1579. mmq_x = MMQ_X_Q2_K_AMPERE;
  1580. mmq_y = MMQ_Y_Q2_K_AMPERE;
  1581. nwarps = NWARPS_Q2_K_AMPERE;
  1582. } else if (compute_capability >= MIN_CC_DP4A) {
  1583. mmq_x = MMQ_X_Q2_K_PASCAL;
  1584. mmq_y = MMQ_Y_Q2_K_PASCAL;
  1585. nwarps = NWARPS_Q2_K_PASCAL;
  1586. } else {
  1587. GGML_ASSERT(false);
  1588. }
  1589. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1590. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1591. const dim3 block_nums(block_num_x, block_num_y, 1);
  1592. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1593. if (nrows_x % mmq_y == 0) {
  1594. const bool need_check = false;
  1595. mul_mat_q2_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1596. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1597. } else {
  1598. const bool need_check = true;
  1599. mul_mat_q2_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1600. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1601. }
  1602. }
  1603. static void ggml_mul_mat_q3_K_q8_1_cuda(
  1604. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1605. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1606. #if QK_K == 256
  1607. int id = ggml_cuda_get_device();
  1608. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1609. int mmq_x, mmq_y, nwarps;
  1610. if (compute_capability >= CC_RDNA2) {
  1611. mmq_x = MMQ_X_Q3_K_RDNA2;
  1612. mmq_y = MMQ_Y_Q3_K_RDNA2;
  1613. nwarps = NWARPS_Q3_K_RDNA2;
  1614. } else if (compute_capability >= CC_OFFSET_AMD) {
  1615. mmq_x = MMQ_X_Q3_K_RDNA1;
  1616. mmq_y = MMQ_Y_Q3_K_RDNA1;
  1617. nwarps = NWARPS_Q3_K_RDNA1;
  1618. } else if (compute_capability >= CC_VOLTA) {
  1619. mmq_x = MMQ_X_Q3_K_AMPERE;
  1620. mmq_y = MMQ_Y_Q3_K_AMPERE;
  1621. nwarps = NWARPS_Q3_K_AMPERE;
  1622. } else if (compute_capability >= MIN_CC_DP4A) {
  1623. mmq_x = MMQ_X_Q3_K_PASCAL;
  1624. mmq_y = MMQ_Y_Q3_K_PASCAL;
  1625. nwarps = NWARPS_Q3_K_PASCAL;
  1626. } else {
  1627. GGML_ASSERT(false);
  1628. }
  1629. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1630. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1631. const dim3 block_nums(block_num_x, block_num_y, 1);
  1632. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1633. if (nrows_x % mmq_y == 0) {
  1634. const bool need_check = false;
  1635. mul_mat_q3_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1636. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1637. } else {
  1638. const bool need_check = true;
  1639. mul_mat_q3_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1640. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1641. }
  1642. #endif
  1643. }
  1644. static void ggml_mul_mat_q4_K_q8_1_cuda(
  1645. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1646. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1647. int id = ggml_cuda_get_device();
  1648. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1649. int mmq_x, mmq_y, nwarps;
  1650. if (compute_capability >= CC_RDNA2) {
  1651. mmq_x = MMQ_X_Q4_K_RDNA2;
  1652. mmq_y = MMQ_Y_Q4_K_RDNA2;
  1653. nwarps = NWARPS_Q4_K_RDNA2;
  1654. } else if (compute_capability >= CC_OFFSET_AMD) {
  1655. mmq_x = MMQ_X_Q4_K_RDNA1;
  1656. mmq_y = MMQ_Y_Q4_K_RDNA1;
  1657. nwarps = NWARPS_Q4_K_RDNA1;
  1658. } else if (compute_capability >= CC_VOLTA) {
  1659. mmq_x = MMQ_X_Q4_K_AMPERE;
  1660. mmq_y = MMQ_Y_Q4_K_AMPERE;
  1661. nwarps = NWARPS_Q4_K_AMPERE;
  1662. } else if (compute_capability >= MIN_CC_DP4A) {
  1663. mmq_x = MMQ_X_Q4_K_PASCAL;
  1664. mmq_y = MMQ_Y_Q4_K_PASCAL;
  1665. nwarps = NWARPS_Q4_K_PASCAL;
  1666. } else {
  1667. GGML_ASSERT(false);
  1668. }
  1669. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1670. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1671. const dim3 block_nums(block_num_x, block_num_y, 1);
  1672. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1673. if (nrows_x % mmq_y == 0) {
  1674. const bool need_check = false;
  1675. mul_mat_q4_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1676. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1677. } else {
  1678. const bool need_check = true;
  1679. mul_mat_q4_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1680. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1681. }
  1682. }
  1683. static void ggml_mul_mat_q5_K_q8_1_cuda(
  1684. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1685. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1686. int id = ggml_cuda_get_device();
  1687. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1688. int mmq_x, mmq_y, nwarps;
  1689. if (compute_capability >= CC_RDNA2) {
  1690. mmq_x = MMQ_X_Q5_K_RDNA2;
  1691. mmq_y = MMQ_Y_Q5_K_RDNA2;
  1692. nwarps = NWARPS_Q5_K_RDNA2;
  1693. } else if (compute_capability >= CC_OFFSET_AMD) {
  1694. mmq_x = MMQ_X_Q5_K_RDNA1;
  1695. mmq_y = MMQ_Y_Q5_K_RDNA1;
  1696. nwarps = NWARPS_Q5_K_RDNA1;
  1697. } else if (compute_capability >= CC_VOLTA) {
  1698. mmq_x = MMQ_X_Q5_K_AMPERE;
  1699. mmq_y = MMQ_Y_Q5_K_AMPERE;
  1700. nwarps = NWARPS_Q5_K_AMPERE;
  1701. } else if (compute_capability >= MIN_CC_DP4A) {
  1702. mmq_x = MMQ_X_Q5_K_PASCAL;
  1703. mmq_y = MMQ_Y_Q5_K_PASCAL;
  1704. nwarps = NWARPS_Q5_K_PASCAL;
  1705. } else {
  1706. GGML_ASSERT(false);
  1707. }
  1708. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1709. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1710. const dim3 block_nums(block_num_x, block_num_y, 1);
  1711. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1712. if (nrows_x % mmq_y == 0) {
  1713. const bool need_check = false;
  1714. mul_mat_q5_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1715. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1716. } else {
  1717. const bool need_check = true;
  1718. mul_mat_q5_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1719. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1720. }
  1721. }
  1722. static void ggml_mul_mat_q6_K_q8_1_cuda(
  1723. const void * vx, const void * vy, float * dst, const int ncols_x, const int nrows_x,
  1724. const int ncols_y, const int nrows_y, const int nrows_dst, cudaStream_t stream) {
  1725. int id = ggml_cuda_get_device();
  1726. const int compute_capability = ggml_cuda_info().devices[id].cc;
  1727. int mmq_x, mmq_y, nwarps;
  1728. if (compute_capability >= CC_RDNA2) {
  1729. mmq_x = MMQ_X_Q6_K_RDNA2;
  1730. mmq_y = MMQ_Y_Q6_K_RDNA2;
  1731. nwarps = NWARPS_Q6_K_RDNA2;
  1732. } else if (compute_capability >= CC_OFFSET_AMD) {
  1733. mmq_x = MMQ_X_Q6_K_RDNA1;
  1734. mmq_y = MMQ_Y_Q6_K_RDNA1;
  1735. nwarps = NWARPS_Q6_K_RDNA1;
  1736. } else if (compute_capability >= CC_VOLTA) {
  1737. mmq_x = MMQ_X_Q6_K_AMPERE;
  1738. mmq_y = MMQ_Y_Q6_K_AMPERE;
  1739. nwarps = NWARPS_Q6_K_AMPERE;
  1740. } else if (compute_capability >= MIN_CC_DP4A) {
  1741. mmq_x = MMQ_X_Q6_K_PASCAL;
  1742. mmq_y = MMQ_Y_Q6_K_PASCAL;
  1743. nwarps = NWARPS_Q6_K_PASCAL;
  1744. } else {
  1745. GGML_ASSERT(false);
  1746. }
  1747. const int block_num_x = (nrows_x + mmq_y - 1) / mmq_y;
  1748. const int block_num_y = (ncols_y + mmq_x - 1) / mmq_x;
  1749. const dim3 block_nums(block_num_x, block_num_y, 1);
  1750. const dim3 block_dims(WARP_SIZE, nwarps, 1);
  1751. if (nrows_x % mmq_y == 0) {
  1752. const bool need_check = false;
  1753. mul_mat_q6_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1754. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1755. } else {
  1756. const bool need_check = true;
  1757. mul_mat_q6_K<need_check><<<block_nums, block_dims, 0, stream>>>
  1758. (vx, vy, dst, ncols_x, nrows_x, ncols_y, nrows_y, nrows_dst);
  1759. }
  1760. }
  1761. void ggml_cuda_op_mul_mat_q(
  1762. ggml_backend_cuda_context & ctx,
  1763. const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst, const char * src0_dd_i, const float * src1_ddf_i,
  1764. const char * src1_ddq_i, float * dst_dd_i, const int64_t row_low, const int64_t row_high, const int64_t src1_ncols,
  1765. const int64_t src1_padded_row_size, cudaStream_t stream) {
  1766. const int64_t ne00 = src0->ne[0];
  1767. const int64_t ne10 = src1->ne[0];
  1768. GGML_ASSERT(ne10 % QK8_1 == 0);
  1769. const int64_t ne0 = dst->ne[0];
  1770. const int64_t row_diff = row_high - row_low;
  1771. int id = ggml_cuda_get_device();
  1772. // the main device has a larger memory buffer to hold the results from all GPUs
  1773. // nrows_dst == nrows of the matrix that the kernel writes into
  1774. const int64_t nrows_dst = id == ctx.device ? ne0 : row_diff;
  1775. switch (src0->type) {
  1776. case GGML_TYPE_Q4_0:
  1777. ggml_mul_mat_q4_0_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1778. break;
  1779. case GGML_TYPE_Q4_1:
  1780. ggml_mul_mat_q4_1_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1781. break;
  1782. case GGML_TYPE_Q5_0:
  1783. ggml_mul_mat_q5_0_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1784. break;
  1785. case GGML_TYPE_Q5_1:
  1786. ggml_mul_mat_q5_1_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1787. break;
  1788. case GGML_TYPE_Q8_0:
  1789. ggml_mul_mat_q8_0_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1790. break;
  1791. case GGML_TYPE_Q2_K:
  1792. ggml_mul_mat_q2_K_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1793. break;
  1794. case GGML_TYPE_Q3_K:
  1795. ggml_mul_mat_q3_K_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1796. break;
  1797. case GGML_TYPE_Q4_K:
  1798. ggml_mul_mat_q4_K_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1799. break;
  1800. case GGML_TYPE_Q5_K:
  1801. ggml_mul_mat_q5_K_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1802. break;
  1803. case GGML_TYPE_Q6_K:
  1804. ggml_mul_mat_q6_K_q8_1_cuda(src0_dd_i, src1_ddq_i, dst_dd_i, ne00, row_diff, src1_ncols, src1_padded_row_size, nrows_dst, stream);
  1805. break;
  1806. default:
  1807. GGML_ASSERT(false);
  1808. break;
  1809. }
  1810. GGML_UNUSED(src1);
  1811. GGML_UNUSED(dst);
  1812. GGML_UNUSED(src1_ddf_i);
  1813. }
  1814. bool ggml_cuda_supports_mmq(enum ggml_type type) {
  1815. switch (type) {
  1816. case GGML_TYPE_Q4_0:
  1817. case GGML_TYPE_Q4_1:
  1818. case GGML_TYPE_Q5_0:
  1819. case GGML_TYPE_Q5_1:
  1820. case GGML_TYPE_Q8_0:
  1821. case GGML_TYPE_Q2_K:
  1822. case GGML_TYPE_Q3_K:
  1823. case GGML_TYPE_Q4_K:
  1824. case GGML_TYPE_Q5_K:
  1825. case GGML_TYPE_Q6_K:
  1826. return true;
  1827. default:
  1828. return false;
  1829. }
  1830. }