diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl index e4640cf05c..ed62091136 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl @@ -101,13 +101,12 @@ vec4 dequantize4(uint ib, uint iqs, uint a_offset) { #if defined(DATA_A_Q8_0) #if defined(A_TYPE_REPACKED) vec2 dequantize(uint ib, uint iqs, uint a_offset) { - return vec2(int(int8_t(data_a_quants[(a_offset + ib) * 32 + iqs])), - int(int8_t(data_a_quants[(a_offset + ib) * 32 + iqs + 1]))); + const i8vec2 v = unpack8(int32_t(data_a_quants16[(a_offset + ib) * 16 + iqs/2])).xy; + return vec2(v); } vec4 dequantize4(uint ib, uint iqs, uint a_offset) { - const i8vec2 v0 = unpack8(int32_t(data_a_quants16[(a_offset + ib) * 16 + iqs/2])).xy; - const i8vec2 v1 = unpack8(int32_t(data_a_quants16[(a_offset + ib) * 16 + iqs/2 + 1])).xy; - return vec4(v0.x, v0.y, v1.x, v1.y); + const i8vec4 v = unpack8(int32_t(data_a_quants32[(a_offset + ib) * 8 + iqs/4])); + return vec4(v); } #else vec2 dequantize(uint ib, uint iqs, uint a_offset) { diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl index bf5891f4b7..1634aa2653 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs_cm2.glsl @@ -27,7 +27,7 @@ float16_t dequantFuncQ1_0(const in decodeBufQ1_0 bl, const in uint blockCoords[2 #ifdef A_TYPE_REPACKED layout(buffer_reference, std430, buffer_reference_align = 16) buffer decodeBufQ4_0 { - uint16_t qs[8]; + uint32_t qs[4]; }; #else layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufQ4_0 { @@ -41,22 +41,25 @@ float16_t dequantFuncQ4_0(const in decodeBufQ4_0 bl, const in uint blockCoords[2 #ifdef A_TYPE_REPACKED const uint ib = pos_a + blockCoords[0] * (p.stride_a / QUANT_K) + blockCoords[1]; const float16_t d = data_a_deltas[p.deltas_offset + ib]; - uint32_t qs = uint32_t(bl.qs[(idx & 0xE) >> 1]); + uint32_t qs = bl.qs[(idx & 0xC) >> 2]; + const uint shift = (idx & 0x10) >> 2; + qs >>= ((idx & 3) * 8 + shift); #else const float16_t d = bl.block.d; uint32_t qs = uint32_t(bl.block.qs[(idx & 0xE) >> 1]); -#endif const uint shift = (idx & 0x10) >> 2; qs >>= shift; qs &= 0x0F0F; qs = unpack8(qs)[idx & 1]; +#endif + qs &= 0xF; float16_t ret = (float16_t(qs) - float16_t(8)) * d; return ret; } #ifdef A_TYPE_REPACKED layout(buffer_reference, std430, buffer_reference_align = 16) buffer decodeBufQ4_1 { - uint16_t qs[8]; + uint32_t qs[4]; }; #else layout(buffer_reference, std430, buffer_reference_align = 4) buffer decodeBufQ4_1 { @@ -67,21 +70,18 @@ layout(buffer_reference, std430, buffer_reference_align = 4) buffer decodeBufQ4_ float16_t dequantFuncQ4_1(const in decodeBufQ4_1 bl, const in uint blockCoords[2], const in uint coordInBlock[2]) { const uint idx = coordInBlock[1]; + const uint iqs = idx & 0xF; + const uint shift = (idx & 0x10) >> 2; #ifdef A_TYPE_REPACKED const uint ib = pos_a + blockCoords[0] * (p.stride_a / QUANT_K) + blockCoords[1]; const float16_t d = data_a_deltas[p.deltas_offset + ib * 2]; const float16_t m = data_a_deltas[p.deltas_offset + ib * 2 + 1]; - uint32_t qs = uint32_t(bl.qs[(idx & 0xE) >> 1]); + uint32_t qs = bl.qs[(idx & 0xC) >> 2]; + qs >>= ((iqs & 3) * 8 + shift); #else const float16_t d = bl.block.d; const float16_t m = bl.block.m; - uint32_t qs = bl.block.qs[idx & 0xF]; -#endif - const uint iqs = idx & 0xF; - const uint shift = (idx & 0x10) >> 2; -#ifdef A_TYPE_REPACKED - qs >>= ((iqs & 1) * 8 + shift); -#else + uint32_t qs = bl.block.qs[iqs]; qs >>= shift; #endif qs &= 0xF; @@ -136,7 +136,7 @@ float16_t dequantFuncQ5_1(const in decodeBufQ5_1 bl, const in uint blockCoords[2 #ifdef A_TYPE_REPACKED layout(buffer_reference, std430, buffer_reference_align = 16) buffer decodeBufQ8_0 { - int16_t qs[16]; + int32_t qs[8]; }; #else layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufQ8_0 { @@ -151,7 +151,7 @@ float16_t dequantFuncQ8_0(const in decodeBufQ8_0 bl, const in uint blockCoords[2 #ifdef A_TYPE_REPACKED const uint ib = pos_a + blockCoords[0] * (p.stride_a / QUANT_K) + blockCoords[1]; const float16_t d = data_a_deltas[p.deltas_offset + ib]; - int32_t qs = unpack8(bl.qs[(iqs & 0x1E) >> 1])[iqs & 1]; + int32_t qs = unpack8(bl.qs[(iqs & 0x1C) >> 2])[iqs & 3]; #else const float16_t d = bl.block.d; int32_t qs = unpack8(bl.block.qs[(iqs & 0x1E) >> 1])[iqs & 1]; @@ -701,7 +701,7 @@ float16_t dequantFuncIQ4_XS(const in decodeBufIQ4_XS bl, const in uint blockCoor #if defined(DATA_A_IQ4_NL) #ifdef A_TYPE_REPACKED layout(buffer_reference, std430, buffer_reference_align = 16) buffer decodeBufIQ4_NL { - uint16_t qs[8]; + uint32_t qs[4]; }; #else layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufIQ4_NL { @@ -715,9 +715,9 @@ float16_t dequantFuncIQ4_NL(const in decodeBufIQ4_NL bl, const in uint blockCoor #ifdef A_TYPE_REPACKED const uint ib = pos_a + blockCoords[0] * (p.stride_a / QUANT_K) + blockCoords[1]; const float16_t d = data_a_deltas[p.deltas_offset + ib]; - uint32_t qs = uint32_t(bl.qs[(idx & 0xE) >> 1]); + uint32_t qs = bl.qs[(idx & 0xC) >> 2]; const uint shift = (idx & 0x10) >> 2; - qs >>= ((idx & 1) * 8 + shift); + qs >>= ((idx & 3) * 8 + shift); #else const float16_t d = bl.block.d; const uint iqs = idx & 0xF; @@ -734,7 +734,7 @@ float16_t dequantFuncIQ4_NL(const in decodeBufIQ4_NL bl, const in uint blockCoor #if defined(DATA_A_MXFP4) #ifdef A_TYPE_REPACKED layout(buffer_reference, std430, buffer_reference_align = 16) buffer decodeBufMXFP4 { - uint16_t qs[8]; + uint32_t qs[4]; }; #else layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufMXFP4 { @@ -745,17 +745,15 @@ layout(buffer_reference, std430, buffer_reference_align = 2) buffer decodeBufMXF float16_t dequantFuncMXFP4(const in decodeBufMXFP4 bl, const in uint blockCoords[2], const in uint coordInBlock[2]) { const uint idx = coordInBlock[1]; + const uint iqs = idx & 0xF; + const uint shift = (idx & 0x10) >> 2; #ifdef A_TYPE_REPACKED const uint ib = pos_a + blockCoords[0] * (p.stride_a / QUANT_K) + blockCoords[1]; const float d = e8m0_to_fp32(data_a_scales[p.deltas_offset + ib]); - const uint iqs = idx & 0xF; - const uint shift = (idx & 0x10) >> 2; - uint32_t qs = uint32_t(bl.qs[(iqs & 0xE) >> 1]); - qs >>= ((iqs & 1) * 8 + shift); + uint32_t qs = bl.qs[(iqs & 0xC) >> 2]; + qs >>= ((iqs & 3) * 8 + shift); #else const float d = e8m0_to_fp32(bl.block.e); - const uint iqs = idx & 0xF; - const uint shift = (idx & 0x10) >> 2; uint32_t qs = bl.block.qs[iqs]; qs >>= shift; #endif diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl index fa512acf85..9e14d8fc29 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl @@ -136,14 +136,13 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin #if defined(A_TYPE_REPACKED) const float d = float(data_a_deltas[p.deltas_offset + ib]); - const i8vec2 v0 = unpack8(int32_t(data_a_quants16[ib * 16 + 2*iqs])).xy; - const i8vec2 v1 = unpack8(int32_t(data_a_quants16[ib * 16 + 2*iqs + 1])).xy; + const vec4 v = vec4(unpack8(int32_t(data_a_quants32[ib * 8 + iqs]))) * d; #else const float d = float(data_a_packed16[ib].d); const i8vec2 v0 = unpack8(int32_t(data_a_packed16[ib].qs[2*iqs])).xy; // vec4 used due to #12147 const i8vec2 v1 = unpack8(int32_t(data_a_packed16[ib].qs[2*iqs + 1])).xy; -#endif const vec4 v = vec4(v0.x, v0.y, v1.x, v1.y) * d; +#endif buf_a[buf_idx ] = FLOAT_TYPEV2(v.xy); buf_a[buf_idx + 1] = FLOAT_TYPEV2(v.zw); @@ -519,8 +518,9 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin #if defined(A_TYPE_REPACKED) const float d = e8m0_to_fp32(uint8_t(data_a_quants[p.deltas_offset + ib])) * 0.5; - const uint vui = uint(data_a_quants[ib * 16 + iqs]); - const uint vui2 = uint(data_a_quants[ib * 16 + iqs + 1]); + const uint vui16 = uint(data_a_quants16[ib * 8 + iqs/2]); + const uint vui = vui16 & 0xFF; + const uint vui2 = vui16 >> 8; #else const float d = e8m0_to_fp32(data_a[ib].e) * 0.5; const uint vui = uint(data_a[ib].qs[iqs]);