diff --git a/ggml/src/ggml-vulkan.cpp b/ggml/src/ggml-vulkan.cpp index 3569f7e0..0323bdc8 100644 --- a/ggml/src/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan.cpp @@ -1884,6 +1884,9 @@ static bool ggml_vk_matmul_shmem_support(const vk_device& device, const std::vec case GGML_TYPE_IQ4_XS: lut_size = 4*16; break; + case GGML_TYPE_IQ4_KS: + lut_size = 4*32; + break; default: break; } @@ -2344,6 +2347,8 @@ static void ggml_vk_load_shaders(vk_device& device) { CREATE_MM2(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ3_S], matmul_iq3_s_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_XS], matmul_iq4_xs_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_NL], matmul_iq4_nl_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_IQ4_KS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_KS], matmul_iq4_ks_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_IQ4_KT, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_KT], matmul_iq4_kt_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); } else { CREATE_MM(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_0].f32acc, matmul_q4_0_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_1].f32acc, matmul_q4_1_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); @@ -2365,6 +2370,8 @@ static void ggml_vk_load_shaders(vk_device& device) { CREATE_MM(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ3_S].f32acc, matmul_iq3_s_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_XS].f32acc, matmul_iq4_xs_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_NL].f32acc, matmul_iq4_nl_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM(GGML_TYPE_IQ4_KS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_KS].f32acc, matmul_iq4_ks_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM(GGML_TYPE_IQ4_KT, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_KT].f32acc, matmul_iq4_kt_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); } CREATE_MM(GGML_TYPE_F32, pipeline_matmul_id_f32, matmul_id_f32_f32, , wg_denoms, warptile, vk_mat_mat_push_constants, 4, _id); @@ -2397,6 +2404,8 @@ static void ggml_vk_load_shaders(vk_device& device) { CREATE_MM(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_S].f16acc, matmul_id_iq3_s_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_XS].f16acc, matmul_id_iq4_xs_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL].f16acc, matmul_id_iq4_nl_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); + CREATE_MM(GGML_TYPE_IQ4_KS, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_KS].f16acc, matmul_id_iq4_ks_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); + CREATE_MM(GGML_TYPE_IQ4_KT, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_KT].f16acc, matmul_id_iq4_kt_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); } else { CREATE_MM(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0].f16acc, matmul_id_q4_0_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1].f16acc, matmul_id_q4_1_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); @@ -2418,6 +2427,8 @@ static void ggml_vk_load_shaders(vk_device& device) { CREATE_MM(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_S].f16acc, matmul_id_iq3_s_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_XS].f16acc, matmul_id_iq4_xs_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL].f16acc, matmul_id_iq4_nl_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); + CREATE_MM(GGML_TYPE_IQ4_KS, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_KS].f16acc, matmul_id_iq4_ks_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); + CREATE_MM(GGML_TYPE_IQ4_KT, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_KT].f16acc, matmul_id_iq4_kt_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); } #undef CREATE_MM2 #undef CREATE_MM @@ -2485,6 +2496,8 @@ static void ggml_vk_load_shaders(vk_device& device) { CREATE_MM2(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ3_S], matmul_iq3_s_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_XS], matmul_iq4_xs_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_NL], matmul_iq4_nl_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_IQ4_KS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_KS], matmul_iq4_ks_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_IQ4_KT, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_KT], matmul_iq4_kt_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); #if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT) if (device->integer_dot_product) { @@ -2522,6 +2535,8 @@ static void ggml_vk_load_shaders(vk_device& device) { CREATE_MM(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_S].f16acc, matmul_id_iq3_s_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_XS].f16acc, matmul_id_iq4_xs_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL].f16acc, matmul_id_iq4_nl_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); + CREATE_MM(GGML_TYPE_IQ4_KS, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_KS].f16acc, matmul_id_iq4_ks_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); + CREATE_MM(GGML_TYPE_IQ4_KT, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_KT].f16acc, matmul_id_iq4_kt_f32, _f16acc, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); #undef CREATE_MM2 #undef CREATE_MMQ #undef CREATE_MM @@ -2576,6 +2591,8 @@ static void ggml_vk_load_shaders(vk_device& device) { CREATE_MM(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ3_S].f32acc, matmul_iq3_s_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_XS].f32acc, matmul_iq4_xs_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_NL].f32acc, matmul_iq4_nl_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM(GGML_TYPE_IQ4_KS, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_KS].f32acc, matmul_iq4_ks_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM(GGML_TYPE_IQ4_KT, pipeline_dequant_mul_mat_mat[GGML_TYPE_IQ4_KT].f32acc, matmul_iq4_kt_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); #if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT) if (device->integer_dot_product) { @@ -2613,6 +2630,8 @@ static void ggml_vk_load_shaders(vk_device& device) { CREATE_MM(GGML_TYPE_IQ3_S, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ3_S].f32acc, matmul_id_iq3_s_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_IQ4_XS, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_XS].f32acc, matmul_id_iq4_xs_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); CREATE_MM(GGML_TYPE_IQ4_NL, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_NL].f32acc, matmul_id_iq4_nl_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); + CREATE_MM(GGML_TYPE_IQ4_KS, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_KS].f32acc, matmul_id_iq4_ks_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); + CREATE_MM(GGML_TYPE_IQ4_KT, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_IQ4_KT].f32acc, matmul_id_iq4_kt_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, 4, _id); } // reusing CREATE_MM from the fp32 path if ((device->coopmat2 || device->coopmat_support) @@ -2671,6 +2690,8 @@ static void ggml_vk_load_shaders(vk_device& device) { ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[GGML_TYPE_IQ3_S][i], "mul_mat_vec_iq3_s_f32_f32_"+std::to_string(i+1), mul_mat_vec_iq3_s_f32_f32_len, mul_mat_vec_iq3_s_f32_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[GGML_TYPE_IQ4_XS][i], "mul_mat_vec_iq4_xs_f32_f32_"+std::to_string(i+1), mul_mat_vec_iq4_xs_f32_f32_len, mul_mat_vec_iq4_xs_f32_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[GGML_TYPE_IQ4_NL][i], "mul_mat_vec_iq4_nl_f32_f32_"+std::to_string(i+1), mul_mat_vec_iq4_nl_f32_f32_len, mul_mat_vec_iq4_nl_f32_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[GGML_TYPE_IQ4_KS][i], "mul_mat_vec_iq4_ks_f32_f32_"+std::to_string(i+1), mul_mat_vec_iq4_ks_f32_f32_len, mul_mat_vec_iq4_ks_f32_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[GGML_TYPE_IQ4_KT][i], "mul_mat_vec_iq4_kt_f32_f32_"+std::to_string(i+1), mul_mat_vec_iq4_kt_f32_f32_len, mul_mat_vec_iq4_kt_f32_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[GGML_TYPE_F32 ][i], "mul_mat_vec_f32_f16_f32_"+std::to_string(i+1), mul_mat_vec_f32_f16_f32_len, mul_mat_vec_f32_f16_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {2, 1, 1}, {device->subgroup_size, 2, i+1}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[GGML_TYPE_F16 ][i], "mul_mat_vec_f16_f16_f32_"+std::to_string(i+1), mul_mat_vec_f16_f16_f32_len, mul_mat_vec_f16_f16_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {2, 1, 1}, {device->subgroup_size, 2, i+1}, 1); @@ -2694,6 +2715,8 @@ static void ggml_vk_load_shaders(vk_device& device) { ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[GGML_TYPE_IQ3_S][i], "mul_mat_vec_iq3_s_f16_f32_"+std::to_string(i+1), mul_mat_vec_iq3_s_f16_f32_len, mul_mat_vec_iq3_s_f16_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[GGML_TYPE_IQ4_XS][i], "mul_mat_vec_iq4_xs_f16_f32_"+std::to_string(i+1), mul_mat_vec_iq4_xs_f16_f32_len, mul_mat_vec_iq4_xs_f16_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[GGML_TYPE_IQ4_NL][i], "mul_mat_vec_iq4_nl_f16_f32_"+std::to_string(i+1), mul_mat_vec_iq4_nl_f16_f32_len, mul_mat_vec_iq4_nl_f16_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[GGML_TYPE_IQ4_KS][i], "mul_mat_vec_iq4_ks_f16_f32_"+std::to_string(i+1), mul_mat_vec_iq4_ks_f16_f32_len, mul_mat_vec_iq4_ks_f16_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[GGML_TYPE_IQ4_KT][i], "mul_mat_vec_iq4_kt_f16_f32_"+std::to_string(i+1), mul_mat_vec_iq4_kt_f16_f32_len, mul_mat_vec_iq4_kt_f16_f32_data, "main", 3, sizeof(vk_mat_vec_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq, i+1}, 1, true); } ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[GGML_TYPE_F32 ], "mul_mat_vec_id_f32_f32", mul_mat_vec_id_f32_f32_len, mul_mat_vec_id_f32_f32_data, "main", 4, sizeof(vk_mat_vec_id_push_constants), {2, 1, 1}, {device->subgroup_size, 2}, 1); @@ -2718,6 +2741,8 @@ static void ggml_vk_load_shaders(vk_device& device) { ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[GGML_TYPE_IQ3_S], "mul_mat_vec_id_iq3_s_f32", mul_mat_vec_id_iq3_s_f32_len, mul_mat_vec_id_iq3_s_f32_data, "main", 4, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq}, 1, true); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[GGML_TYPE_IQ4_XS], "mul_mat_vec_id_iq4_xs_f32", mul_mat_vec_id_iq4_xs_f32_len, mul_mat_vec_id_iq4_xs_f32_data, "main", 4, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq}, 1, true); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[GGML_TYPE_IQ4_NL], "mul_mat_vec_id_iq4_nl_f32", mul_mat_vec_id_iq4_nl_f32_len, mul_mat_vec_id_iq4_nl_f32_data, "main", 4, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq}, 1, true); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[GGML_TYPE_IQ4_KS], "mul_mat_vec_id_iq4_ks_f32", mul_mat_vec_id_iq4_ks_f32_len, mul_mat_vec_id_iq4_ks_f32_data, "main", 4, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq}, 1, true); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[GGML_TYPE_IQ4_KT], "mul_mat_vec_id_iq4_kt_f32", mul_mat_vec_id_iq4_kt_f32_len, mul_mat_vec_id_iq4_kt_f32_data, "main", 4, sizeof(vk_mat_vec_id_push_constants), {rm_iq, 1, 1}, {subgroup_size_16, rm_iq}, 1, true); // dequant shaders ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_F32 ], "f32_to_f16", dequant_f32_len, dequant_f32_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); @@ -2740,6 +2765,8 @@ static void ggml_vk_load_shaders(vk_device& device) { ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_IQ3_S], "dequant_iq3_s", dequant_iq3_s_len, dequant_iq3_s_data, "main", 2, 5 * sizeof(uint32_t), {256 * 32, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_IQ4_XS], "dequant_iq4_xs", dequant_iq4_xs_len, dequant_iq4_xs_data, "main", 2, 5 * sizeof(uint32_t), {256 * 32, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_IQ4_NL], "dequant_iq4_nl", dequant_iq4_nl_len, dequant_iq4_nl_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_IQ4_KS], "dequant_iq4_ks", dequant_iq4_ks_len, dequant_iq4_ks_data, "main", 2, 5 * sizeof(uint32_t), {256 * 32, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_IQ4_KT], "dequant_iq4_kt", dequant_iq4_kt_len, dequant_iq4_kt_data, "main", 2, 5 * sizeof(uint32_t), {256 * 32, 1, 1}, {}, 1); // get_rows ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_F32 ], "get_rows_f32", get_rows_f32_len, get_rows_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), { 512, 1, 1}, {}, 1); @@ -2759,6 +2786,8 @@ static void ggml_vk_load_shaders(vk_device& device) { ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_IQ3_S], "get_rows_iq3_s", get_rows_iq3_s_len, get_rows_iq3_s_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_IQ4_XS], "get_rows_iq4_xs", get_rows_iq4_xs_len, get_rows_iq4_xs_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_IQ4_NL], "get_rows_iq4_nl", get_rows_iq4_nl_len, get_rows_iq4_nl_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_IQ4_KS], "get_rows_iq4_ks", get_rows_iq4_ks_len, get_rows_iq4_ks_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_IQ4_KT], "get_rows_iq4_kt", get_rows_iq4_kt_len, get_rows_iq4_kt_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_F32 ], "get_rows_f32_f32", get_rows_f32_f32_len, get_rows_f32_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), { 512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_F16 ], "get_rows_f16_f32", get_rows_f16_f32_len, get_rows_f16_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), { 512, 1, 1}, {}, 1); @@ -2777,6 +2806,8 @@ static void ggml_vk_load_shaders(vk_device& device) { ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_IQ3_S], "get_rows_iq3_s_f32", get_rows_iq3_s_f32_len, get_rows_iq3_s_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_IQ4_XS], "get_rows_iq4_xs_f32", get_rows_iq4_xs_f32_len, get_rows_iq4_xs_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_IQ4_NL], "get_rows_iq4_nl_f32", get_rows_iq4_nl_f32_len, get_rows_iq4_nl_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_IQ4_KS], "get_rows_iq4_ks_f32", get_rows_iq4_ks_f32_len, get_rows_iq4_ks_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_IQ4_KT], "get_rows_iq4_kt_f32", get_rows_iq4_kt_f32_len, get_rows_iq4_kt_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_matmul_split_k_reduce, "split_k_reduce", split_k_reduce_len, split_k_reduce_data, "main", 2, 2 * sizeof(uint32_t), {256 * 4, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_flash_attn_split_k_reduce, "fa_split_k_reduce", fa_split_k_reduce_len, fa_split_k_reduce_data, "main", 2, 4 * sizeof(uint32_t), {1, device->subgroup_size, 1}, {device->subgroup_size}, 1, true); @@ -4032,6 +4063,8 @@ static vk_pipeline ggml_vk_get_to_fp16(ggml_backend_vk_context * ctx, ggml_type case GGML_TYPE_IQ3_S: case GGML_TYPE_IQ4_XS: case GGML_TYPE_IQ4_NL: + case GGML_TYPE_IQ4_KS: + case GGML_TYPE_IQ4_KT: break; default: return nullptr; @@ -4040,6 +4073,11 @@ static vk_pipeline ggml_vk_get_to_fp16(ggml_backend_vk_context * ctx, ggml_type return ctx->device->pipeline_dequant[type]; } +// types that keep a per-row scale ahead of the row's blocks +static bool ggml_vk_is_row_meta_quant(ggml_type type) { + return ggml_internal_get_type_traits(type).row_meta_size != 0; +} + static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_context * ctx, ggml_type src0_type, ggml_type src1_type, ggml_prec prec) { VK_LOG_DEBUG("ggml_vk_get_mul_mat_mat_pipeline(" << ggml_type_name(src0_type) << ", " << ggml_type_name(src1_type) << ", " << prec << ")"); if (src0_type == GGML_TYPE_F32 && src1_type == GGML_TYPE_F32) { @@ -4082,6 +4120,12 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte return nullptr; } + // the coopmat2 tensor layout cannot describe a per-row scale ahead of the blocks, so + // these types have no tiled matmul there and the caller falls back to dequantizing to F16 + if (ctx->device->coopmat2 && ggml_vk_is_row_meta_quant(src0_type)) { + return nullptr; + } + switch (src0_type) { case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: @@ -4102,6 +4146,8 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte case GGML_TYPE_IQ3_S: case GGML_TYPE_IQ4_XS: case GGML_TYPE_IQ4_NL: + case GGML_TYPE_IQ4_KS: + case GGML_TYPE_IQ4_KT: break; default: return nullptr; @@ -4145,6 +4191,8 @@ static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec(ggml_backend_vk_context * case GGML_TYPE_IQ3_S: case GGML_TYPE_IQ4_XS: case GGML_TYPE_IQ4_NL: + case GGML_TYPE_IQ4_KS: + case GGML_TYPE_IQ4_KT: break; default: return nullptr; @@ -4184,6 +4232,12 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_id_pipeline(ggml_backend_vk_co GGML_ASSERT(src1_type == GGML_TYPE_F32 || (ctx->device->coopmat2 && src1_type == GGML_TYPE_F16)); + // the coopmat2 tensor layout cannot describe a per-row scale ahead of the blocks, so + // these types have no tiled matmul there and the caller falls back to dequantizing to F16 + if (ctx->device->coopmat2 && ggml_vk_is_row_meta_quant(src0_type)) { + return nullptr; + } + switch (src0_type) { case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: @@ -4204,6 +4258,8 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_id_pipeline(ggml_backend_vk_co case GGML_TYPE_IQ3_S: case GGML_TYPE_IQ4_XS: case GGML_TYPE_IQ4_NL: + case GGML_TYPE_IQ4_KS: + case GGML_TYPE_IQ4_KT: break; default: return nullptr; @@ -4239,6 +4295,8 @@ static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec_id(ggml_backend_vk_context case GGML_TYPE_IQ3_S: case GGML_TYPE_IQ4_XS: case GGML_TYPE_IQ4_NL: + case GGML_TYPE_IQ4_KS: + case GGML_TYPE_IQ4_KT: break; default: return nullptr; @@ -4965,9 +5023,14 @@ static void ggml_vk_matmul_id( } static bool ggml_vk_dim01_contiguous(const ggml_tensor * tensor) { + // a row that is not a whole number of blocks is not contiguous, and ggml_row_size() + // asserts on it + if (tensor->ne[0] % ggml_blck_size(tensor->type) != 0) { + return false; + } return tensor->nb[0] == ggml_type_size(tensor->type) && - tensor->nb[1] == (tensor->nb[0]*tensor->ne[0])/ggml_blck_size(tensor->type) && + tensor->nb[1] == ggml_row_size(tensor->type, tensor->ne[0]) && tensor->nb[3] == tensor->nb[2]*tensor->ne[2]; } @@ -5202,7 +5265,7 @@ static void ggml_vk_mul_mat_q_f16(ggml_backend_vk_context * ctx, vk_context& sub const uint32_t split_k = ggml_vk_guess_split_k(ctx, ne01, ne11, ne10, pipeline); - const uint64_t qx_sz = ggml_type_size(src0->type) * x_ne / ggml_blck_size(src0->type); + const uint64_t qx_sz = ggml_row_size(src0->type, ne00) * ne01; const uint64_t qy_sz = ggml_type_size(src1->type) * y_ne / ggml_blck_size(src1->type); const uint64_t x_sz = !qx_needs_dequant ? qx_sz : sizeof(ggml_fp16_t) * x_ne; const uint64_t y_sz = quantize_y ? (y_ne * ggml_type_size(GGML_TYPE_Q8_1) / ggml_blck_size(GGML_TYPE_Q8_1)) : (y_f32_kernel ? sizeof(float) * y_ne : sizeof(ggml_fp16_t) * y_ne); @@ -5405,7 +5468,7 @@ static void ggml_vk_mul_mat_vec_q_f16(ggml_backend_vk_context * ctx, vk_context& const uint64_t y_ne = ne11 * ne10; const uint64_t d_ne = ne11 * ne01; - const uint64_t qx_sz = ggml_vk_align_size(ggml_type_size(src0->type) * x_ne / ggml_blck_size(src0->type), ctx->device->properties.limits.minStorageBufferOffsetAlignment); + const uint64_t qx_sz = ggml_vk_align_size(ggml_row_size(src0->type, ne00) * ne01, ctx->device->properties.limits.minStorageBufferOffsetAlignment); const uint64_t qy_sz = ggml_type_size(src1->type) * y_ne / ggml_blck_size(src1->type); const uint64_t x_sz = x_non_contig ? ggml_vk_align_size(ggml_type_size(src0->type) * x_ne, ctx->device->properties.limits.minStorageBufferOffsetAlignment) : qx_sz; const uint64_t y_sz = f16_f32_kernel ? sizeof(float) * y_ne : sizeof(ggml_fp16_t) * y_ne; @@ -5818,7 +5881,7 @@ static void ggml_vk_mul_mat_id_q_f16(ggml_backend_vk_context * ctx, vk_context& const uint64_t y_ne = padded_n * ne10; const uint64_t d_ne = ne21 * ne20; - const uint64_t qx_sz = ggml_type_size(src0->type) * x_ne / ggml_blck_size(src0->type); + const uint64_t qx_sz = ggml_row_size(src0->type, ne00) * ne01; const uint64_t qy_sz = ggml_type_size(src1->type) * y_ne / ggml_blck_size(src1->type); const uint64_t x_sz = !qx_needs_dequant ? qx_sz : sizeof(ggml_fp16_t) * x_ne; const uint64_t y_sz = y_f32_kernel ? sizeof(float) * y_ne : sizeof(ggml_fp16_t) * y_ne; @@ -6012,7 +6075,7 @@ static void ggml_vk_mul_mat_vec_id_q_f16(ggml_backend_vk_context * ctx, vk_conte const uint64_t y_ne = ne11 * ne10; const uint64_t d_ne = ne21 * ne20; - const uint64_t qx_sz = ggml_vk_align_size(ggml_type_size(src0->type) * x_ne / ggml_blck_size(src0->type), ctx->device->properties.limits.minStorageBufferOffsetAlignment); + const uint64_t qx_sz = ggml_vk_align_size(ggml_row_size(src0->type, ne00) * ne01, ctx->device->properties.limits.minStorageBufferOffsetAlignment); const uint64_t qy_sz = ggml_type_size(src1->type) * y_ne / ggml_blck_size(src1->type); const uint64_t x_sz = x_non_contig ? ggml_vk_align_size(ggml_type_size(src0->type) * x_ne, ctx->device->properties.limits.minStorageBufferOffsetAlignment) : qx_sz; const uint64_t y_sz = f16_f32_kernel ? sizeof(float) * y_ne : sizeof(ggml_fp16_t) * y_ne; @@ -10431,6 +10494,8 @@ static bool ggml_backend_vk_supports_op(ggml_backend_t backend, const ggml_tenso case GGML_TYPE_IQ3_S: case GGML_TYPE_IQ4_XS: case GGML_TYPE_IQ4_NL: + case GGML_TYPE_IQ4_KS: + case GGML_TYPE_IQ4_KT: break; default: return false; @@ -10538,6 +10603,11 @@ static bool ggml_backend_vk_supports_op(ggml_backend_t backend, const ggml_tenso case GGML_TYPE_IQ2_S: case GGML_TYPE_IQ3_XXS: case GGML_TYPE_IQ3_S: + case GGML_TYPE_IQ4_KS: + case GGML_TYPE_IQ4_KT: + // these two are addressed by row rather than through nb01/02/03, + // which cannot express a per-row scale, so only contiguous src0 + return ggml_is_contiguous(op->src[0]); case GGML_TYPE_IQ4_XS: case GGML_TYPE_IQ4_NL: return true; diff --git a/ggml/src/vulkan-shaders/dequant_iq4_ks.comp b/ggml/src/vulkan-shaders/dequant_iq4_ks.comp new file mode 100644 index 00000000..9636f63a --- /dev/null +++ b/ggml/src/vulkan-shaders/dequant_iq4_ks.comp @@ -0,0 +1,39 @@ +#version 450 + +#include "dequant_head.comp" + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 1) writeonly buffer D {D_TYPE data_b[];}; + +void main() { + // Each thread handles 1 subblock (1 scale and 32 quantized values) + const uint ib = gl_WorkGroupID.x * 32 + gl_LocalInvocationID.x / 8; + + init_iq_shmem(gl_WorkGroupSize); + + if (ib >= p.nel / 256) { + return; + } + + const uint ib32 = gl_LocalInvocationID.x % 8; + + const uint blocks_per_row = p.K / QUANT_K; + const uint row_word = (ib / blocks_per_row) * a_row_words(p.K); + const uint blk = a_block_word(row_word, ib % blocks_per_row); + + uint vo; + const float dl = iq4_ks_scale(row_word, blk, ib32, vo); + + const uint b_idx = 256 * ib + 32 * ib32; + // qs[] starts at byte 8, each subblock owns 16 bytes of it + const uint q_word = blk + 2 + 4 * ib32; + [[unroll]] for (uint l = 0; l < 4; ++l) { + const uint qs = data_a_u32[q_word + l]; + [[unroll]] for (uint k = 0; k < 4; ++k) { + const uint q = (qs >> (8 * k)) & 0xFF; + data_b[b_idx + 4 * l + k ] = D_TYPE(dl * kvalues_iq4k[vo + (q & 0xF)]); + data_b[b_idx + 4 * l + k + 16] = D_TYPE(dl * kvalues_iq4k[vo + (q >> 4)]); + } + } +} diff --git a/ggml/src/vulkan-shaders/dequant_iq4_kt.comp b/ggml/src/vulkan-shaders/dequant_iq4_kt.comp new file mode 100644 index 00000000..6ec4d44d --- /dev/null +++ b/ggml/src/vulkan-shaders/dequant_iq4_kt.comp @@ -0,0 +1,36 @@ +#version 450 + +#include "dequant_head.comp" + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 1) writeonly buffer D {D_TYPE data_b[];}; + +void main() { + // Each thread handles 1 subblock (1 scale and 32 quantized values) + const uint ib = gl_WorkGroupID.x * 32 + gl_LocalInvocationID.x / 8; + + if (ib >= p.nel / 256) { + return; + } + + const uint ib32 = gl_LocalInvocationID.x % 8; + + const uint blocks_per_row = p.K / QUANT_K; + const uint row_word = (ib / blocks_per_row) * a_row_words(p.K); + const uint blk = a_block_word(row_word, ib % blocks_per_row); + + // shb[] occupies the first 8 words of the block, ql[] the next 16, qh[] the last 8 + const uint shb = data_a_u32[blk + ib32]; + uint offset; + const float sl = iq4_kt_scale(row_word, shb, offset); + + const uint b_idx = 256 * ib + 32 * ib32; + [[unroll]] for (uint ig = 0; ig < 8; ++ig) { + uint x = iq4_kt_index(blk, shb, ib32, ig, offset); + + [[unroll]] for (uint k = 0; k < 4; ++k) { + data_b[b_idx + 4 * ig + k] = D_TYPE(sl * iq4_kt_next(x)); + } + } +} diff --git a/ggml/src/vulkan-shaders/get_rows_iq4_ks.comp b/ggml/src/vulkan-shaders/get_rows_iq4_ks.comp new file mode 100644 index 00000000..cbbc6b9d --- /dev/null +++ b/ggml/src/vulkan-shaders/get_rows_iq4_ks.comp @@ -0,0 +1,48 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : enable + +#include "types.comp" +#include "generic_binary_head.comp" + +layout(local_size_x = 512, local_size_y = 1, local_size_z = 1) in; + +// src0 is assumed contiguous. p.nb01/02/03 are byte strides divided by the type size, +// which is not a whole number for a type carrying a per-row scale, so the source row is +// derived from the shape instead. Token embedding tensors are contiguous in practice. +float dequant_elem(const uint row_word, const uint i00) { + const uint blk = a_block_word(row_word, i00 / QUANT_K); + const uint pos = i00 % QUANT_K; + const uint ib32 = pos / 32; + const uint j = pos % 32; + + uint vo; + const float dl = iq4_ks_scale(row_word, blk, ib32, vo); + + // qs[] starts at byte 8; the low nibbles hold the first 16 weights of the subblock, + // the high nibbles the rest + const uint q = a_byte(blk + 2, 16 * ib32 + (j & 15)); + + return dl * float(kvalues_iq4k[vo + (j < 16 ? (q & 0xF) : (q >> 4))]); +} + +void main() { + const uint i00 = (gl_GlobalInvocationID.x)*2; + const uint i10 = gl_GlobalInvocationID.y; + const uint i11 = (gl_GlobalInvocationID.z)/p.ne12; + const uint i12 = (gl_GlobalInvocationID.z)%p.ne12; + + init_iq_shmem(gl_WorkGroupSize); + + if (i00 >= p.ne00) { + return; + } + + const uint i01 = data_b[i10*p.nb10 + i11*p.nb11 + i12*p.nb12]; + + const uint row_word = (i01 + (i11 + i12*p.ne02)*p.ne01) * a_row_words(p.ne00); + const uint d_offset = i10*p.nb21 + i11*p.nb22 + i12*p.nb23; + + data_d[d_offset + i00 ] = D_TYPE(dequant_elem(row_word, i00 )); + data_d[d_offset + i00 + 1] = D_TYPE(dequant_elem(row_word, i00 + 1)); +} diff --git a/ggml/src/vulkan-shaders/get_rows_iq4_kt.comp b/ggml/src/vulkan-shaders/get_rows_iq4_kt.comp new file mode 100644 index 00000000..e4006091 --- /dev/null +++ b/ggml/src/vulkan-shaders/get_rows_iq4_kt.comp @@ -0,0 +1,56 @@ +#version 450 + +#extension GL_EXT_control_flow_attributes : enable + +#include "types.comp" +#include "generic_binary_head.comp" + +layout(local_size_x = 512, local_size_y = 1, local_size_z = 1) in; + +// src0 is assumed contiguous. p.nb01/02/03 are byte strides divided by the type size, +// which is not a whole number for a type carrying a per-row scale, so the source row is +// derived from the shape instead. Token embedding tensors are contiguous in practice. +// The two weights handled by one invocation always fall in the same group of 4, so the +// trellis is only iterated once. +void dequant_pair(const uint row_word, const uint i00, out float v0, out float v1) { + const uint blk = a_block_word(row_word, i00 / QUANT_K); + const uint pos = i00 % QUANT_K; + const uint ib32 = pos / 32; + + const uint shb = data_a_u32[blk + ib32]; + uint offset; + const float sl = iq4_kt_scale(row_word, shb, offset); + + uint x = iq4_kt_index(blk, shb, ib32, (pos % 32) / 4, offset); + + const uint k = pos % 4; + v0 = 0.0f; + v1 = 0.0f; + [[unroll]] for (uint l = 0; l < 4; ++l) { + const float w = sl * iq4_kt_next(x); + if (l == k ) v0 = w; + if (l == k + 1) v1 = w; + } +} + +void main() { + const uint i00 = (gl_GlobalInvocationID.x)*2; + const uint i10 = gl_GlobalInvocationID.y; + const uint i11 = (gl_GlobalInvocationID.z)/p.ne12; + const uint i12 = (gl_GlobalInvocationID.z)%p.ne12; + + if (i00 >= p.ne00) { + return; + } + + const uint i01 = data_b[i10*p.nb10 + i11*p.nb11 + i12*p.nb12]; + + const uint row_word = (i01 + (i11 + i12*p.ne02)*p.ne01) * a_row_words(p.ne00); + const uint d_offset = i10*p.nb21 + i11*p.nb22 + i12*p.nb23; + + float v0, v1; + dequant_pair(row_word, i00, v0, v1); + + data_d[d_offset + i00 ] = D_TYPE(v0); + data_d[d_offset + i00 + 1] = D_TYPE(v1); +} diff --git a/ggml/src/vulkan-shaders/mul_mat_vec_iq4_ks.comp b/ggml/src/vulkan-shaders/mul_mat_vec_iq4_ks.comp new file mode 100644 index 00000000..bbaf9a8d --- /dev/null +++ b/ggml/src/vulkan-shaders/mul_mat_vec_iq4_ks.comp @@ -0,0 +1,96 @@ +#version 450 + +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require + +#include "mul_mat_vec_base.comp" + +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +FLOAT_TYPE temp[NUM_COLS][NUM_ROWS]; + +void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32, const uint i, const uint num_blocks_per_row, const uint first_row, const uint num_rows) { + const uint y_idx = i * QUANT_K + 32 * ib32; + const uint words_per_row = a_row_words(p.ncols); + + // a_offset is an element offset, so it divides exactly by the row length + uint row_word = (a_offset / p.ncols + first_row) * words_per_row; + [[unroll]] for (uint n = 0; n < num_rows; ++n) { + const uint blk = a_block_word(row_word, i); + + uint vo; + const float dl = iq4_ks_scale(row_word, blk, ib32, vo); + + FLOAT_TYPE sum[NUM_COLS]; + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + sum[j] = FLOAT_TYPE(0); + } + [[unroll]] for (uint l = 0; l < 4; ++l) { + const uint qs = data_a_u32[blk + 2 + 4 * ib32 + l]; + const vec4 vlo = vec4(kvalues_iq4k[vo + ((qs ) & 0xF)], kvalues_iq4k[vo + ((qs >> 8) & 0xF)], + kvalues_iq4k[vo + ((qs >> 16) & 0xF)], kvalues_iq4k[vo + ((qs >> 24) & 0xF)]); + const vec4 vhi = vec4(kvalues_iq4k[vo + ((qs >> 4) & 0xF)], kvalues_iq4k[vo + ((qs >> 12) & 0xF)], + kvalues_iq4k[vo + ((qs >> 20) & 0xF)], kvalues_iq4k[vo + ((qs >> 28) & 0xF)]); + + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + const uint bbase = (j*p.batch_stride_b + b_offset + y_idx) / 4; + const vec4 b0 = vec4(data_b_v4[bbase + l ]); + const vec4 b1 = vec4(data_b_v4[bbase + l + 4]); + + sum[j] = + fma(FLOAT_TYPE(b0.x), FLOAT_TYPE(vlo.x), + fma(FLOAT_TYPE(b0.y), FLOAT_TYPE(vlo.y), + fma(FLOAT_TYPE(b0.z), FLOAT_TYPE(vlo.z), + fma(FLOAT_TYPE(b0.w), FLOAT_TYPE(vlo.w), + fma(FLOAT_TYPE(b1.x), FLOAT_TYPE(vhi.x), + fma(FLOAT_TYPE(b1.y), FLOAT_TYPE(vhi.y), + fma(FLOAT_TYPE(b1.z), FLOAT_TYPE(vhi.z), + fma(FLOAT_TYPE(b1.w), FLOAT_TYPE(vhi.w), + sum[j])))))))); + } + } + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + temp[j][n] = fma(FLOAT_TYPE(dl), sum[j], temp[j][n]); + } + row_word += words_per_row; + } +} + +void compute_outputs(const uint32_t first_row, const uint32_t num_rows) { + uint a_offset, b_offset, d_offset; + get_offsets(a_offset, b_offset, d_offset); + + const uint num_blocks_per_row = p.ncols / QUANT_K; + + // 8 threads are used to process each block + const uint blocks_per_wg = gl_WorkGroupSize.x/8; + const uint tid = gl_LocalInvocationID.x; + const uint itid = tid % 8; // 0...7 + const uint ix = tid / 8; + + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + [[unroll]] for (uint i = 0; i < NUM_ROWS; ++i) { + temp[j][i] = FLOAT_TYPE(0); + } + } + + [[unroll]] for (uint i = ix; i < num_blocks_per_row; i += blocks_per_wg) + calc_superblock(a_offset, b_offset, itid, i, num_blocks_per_row, first_row, num_rows); + + reduce_result(temp, d_offset, first_row, num_rows, tid); +} + +void main() { + const uint first_row = NUM_ROWS * (gl_WorkGroupID.x + gl_NumWorkGroups.x * gl_WorkGroupID.z); + + init_iq_shmem(gl_WorkGroupSize); + + // do NUM_ROWS at a time, unless there aren't enough remaining rows + if (first_row + NUM_ROWS <= p.stride_d) { + compute_outputs(first_row, NUM_ROWS); + } else { + if (first_row >= p.stride_d) { + return; + } + compute_outputs(first_row, p.stride_d - first_row); + } +} diff --git a/ggml/src/vulkan-shaders/mul_mat_vec_iq4_kt.comp b/ggml/src/vulkan-shaders/mul_mat_vec_iq4_kt.comp new file mode 100644 index 00000000..d28bd2f2 --- /dev/null +++ b/ggml/src/vulkan-shaders/mul_mat_vec_iq4_kt.comp @@ -0,0 +1,90 @@ +#version 450 + +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require + +#include "mul_mat_vec_base.comp" + +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +FLOAT_TYPE temp[NUM_COLS][NUM_ROWS]; + +void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32, const uint i, const uint num_blocks_per_row, const uint first_row, const uint num_rows) { + const uint y_idx = i * QUANT_K + 32 * ib32; + const uint words_per_row = a_row_words(p.ncols); + + // a_offset is an element offset, so it divides exactly by the row length + uint row_word = (a_offset / p.ncols + first_row) * words_per_row; + [[unroll]] for (uint n = 0; n < num_rows; ++n) { + const uint blk = a_block_word(row_word, i); + + const uint shb = data_a_u32[blk + ib32]; + uint offset; + const float sl = iq4_kt_scale(row_word, shb, offset); + + FLOAT_TYPE sum[NUM_COLS]; + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + sum[j] = FLOAT_TYPE(0); + } + [[unroll]] for (uint ig = 0; ig < 8; ++ig) { + uint x = iq4_kt_index(blk, shb, ib32, ig, offset); + + vec4 w; + [[unroll]] for (uint k = 0; k < 4; ++k) { + w[k] = iq4_kt_next(x); + } + + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + const vec4 b0 = vec4(data_b_v4[(j*p.batch_stride_b + b_offset + y_idx) / 4 + ig]); + + sum[j] = + fma(FLOAT_TYPE(b0.x), FLOAT_TYPE(w.x), + fma(FLOAT_TYPE(b0.y), FLOAT_TYPE(w.y), + fma(FLOAT_TYPE(b0.z), FLOAT_TYPE(w.z), + fma(FLOAT_TYPE(b0.w), FLOAT_TYPE(w.w), + sum[j])))); + } + } + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + temp[j][n] = fma(FLOAT_TYPE(sl), sum[j], temp[j][n]); + } + row_word += words_per_row; + } +} + +void compute_outputs(const uint32_t first_row, const uint32_t num_rows) { + uint a_offset, b_offset, d_offset; + get_offsets(a_offset, b_offset, d_offset); + + const uint num_blocks_per_row = p.ncols / QUANT_K; + + // 8 threads are used to process each block + const uint blocks_per_wg = gl_WorkGroupSize.x/8; + const uint tid = gl_LocalInvocationID.x; + const uint itid = tid % 8; // 0...7 + const uint ix = tid / 8; + + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + [[unroll]] for (uint i = 0; i < NUM_ROWS; ++i) { + temp[j][i] = FLOAT_TYPE(0); + } + } + + [[unroll]] for (uint i = ix; i < num_blocks_per_row; i += blocks_per_wg) + calc_superblock(a_offset, b_offset, itid, i, num_blocks_per_row, first_row, num_rows); + + reduce_result(temp, d_offset, first_row, num_rows, tid); +} + +void main() { + const uint first_row = NUM_ROWS * (gl_WorkGroupID.x + gl_NumWorkGroups.x * gl_WorkGroupID.z); + + // do NUM_ROWS at a time, unless there aren't enough remaining rows + if (first_row + NUM_ROWS <= p.stride_d) { + compute_outputs(first_row, NUM_ROWS); + } else { + if (first_row >= p.stride_d) { + return; + } + compute_outputs(first_row, p.stride_d - first_row); + } +} diff --git a/ggml/src/vulkan-shaders/mul_mm.comp b/ggml/src/vulkan-shaders/mul_mm.comp index f4815499..5be38fb7 100644 --- a/ggml/src/vulkan-shaders/mul_mm.comp +++ b/ggml/src/vulkan-shaders/mul_mm.comp @@ -747,6 +747,50 @@ void main() { buf_a[buf_idx + 1 ] = FLOAT_TYPE(kvalues_iq4nl[bitfieldExtract(vui, 8, 4)]) * d; buf_a[buf_idx + 16] = FLOAT_TYPE(kvalues_iq4nl[bitfieldExtract(vui, 4, 4)]) * d; buf_a[buf_idx + 17] = FLOAT_TYPE(kvalues_iq4nl[vui >> 12]) * d; +#elif defined(DATA_A_IQ4_KS) + const uint idx = pos_a + (loadc_a + l) * p.stride_a / LOAD_VEC_A + loadr_a; + const uint buf_idx = (loadc_a + l) * SHMEM_STRIDE + loadr_a * LOAD_VEC_A; + + // recover the position from the load index; exact because batch_stride_a is ne00*ne01 + const uint elem = idx * LOAD_VEC_A; + const uint row_word = (elem / p.stride_a) * a_row_words(p.stride_a); + const uint acol = elem % p.stride_a; + + const uint blk = a_block_word(row_word, acol / QUANT_K); + const uint ib32 = (acol % QUANT_K) / 32; + const uint j = acol % 32; + + uint vo; + const float dl = iq4_ks_scale(row_word, blk, ib32, vo); + + const uint qb = 16 * ib32 + (j & 15); + const uint qw = data_a_u32[blk + 2 + qb / 4] >> (8 * (qb % 4)); + const uint qshift = j < 16 ? 0 : 4; + + buf_a[buf_idx ] = FLOAT_TYPE(dl * kvalues_iq4k[vo + ((qw >> qshift) & 0xF)]); + buf_a[buf_idx + 1] = FLOAT_TYPE(dl * kvalues_iq4k[vo + ((qw >> (8 + qshift)) & 0xF)]); +#elif defined(DATA_A_IQ4_KT) + const uint idx = pos_a + (loadc_a + l) * p.stride_a / LOAD_VEC_A + loadr_a; + const uint buf_idx = (loadc_a + l) * SHMEM_STRIDE + loadr_a * LOAD_VEC_A; + + // LOAD_VEC_A is 4 here, so one load index is one trellis group of 4 + const uint elem = idx * LOAD_VEC_A; + const uint row_word = (elem / p.stride_a) * a_row_words(p.stride_a); + const uint acol = elem % p.stride_a; + + const uint blk = a_block_word(row_word, acol / QUANT_K); + const uint ib32 = (acol % QUANT_K) / 32; + + const uint shb = data_a_u32[blk + ib32]; + uint offset; + const float sl = iq4_kt_scale(row_word, shb, offset); + + uint x = iq4_kt_index(blk, shb, ib32, (acol % 32) / 4, offset); + + buf_a[buf_idx ] = FLOAT_TYPE(sl * iq4_kt_next(x)); + buf_a[buf_idx + 1] = FLOAT_TYPE(sl * iq4_kt_next(x)); + buf_a[buf_idx + 2] = FLOAT_TYPE(sl * iq4_kt_next(x)); + buf_a[buf_idx + 3] = FLOAT_TYPE(sl * iq4_kt_next(x)); #endif } [[unroll]] for (uint l = 0; l < BN; l += loadstride_b) { diff --git a/ggml/src/vulkan-shaders/types.comp b/ggml/src/vulkan-shaders/types.comp index 3bde7178..fdb84dbe 100644 --- a/ggml/src/vulkan-shaders/types.comp +++ b/ggml/src/vulkan-shaders/types.comp @@ -1356,6 +1356,121 @@ void init_iq_shmem(uvec3 wgsize) } #endif +// IQ4_KS and IQ4_KT rows are one f32 scale followed by the row's blocks, so rows are not a +// whole number of blocks apart. These two are addressed by word through a uint32_t alias of +// binding 0; both block sizes and the prefix are multiples of 4 bytes, so strides are exact. +// The prefix is one word for these two types; a type with a different row_meta_size needs +// its own stride. +#define A_ROW_META_WORDS 1 + +#define QUANT_K_IQ4_KS 256 +#define QUANT_R_IQ4_KS 1 + +struct block_iq4_ks +{ + uint8_t scales[QUANT_K_IQ4_KS/32]; + uint8_t qs[QUANT_K_IQ4_KS/2]; +}; + +#if defined(DATA_A_IQ4_KS) +#define QUANT_K QUANT_K_IQ4_KS +#define QUANT_R QUANT_R_IQ4_KS +#define A_TYPE block_iq4_ks +#define A_TYPE_ROW_META +#define BLOCK_WORDS 34 +#endif + +#define QUANT_K_IQ4_KT 256 +#define QUANT_R_IQ4_KT 1 + +struct block_iq4_kt +{ + uint32_t qs[QUANT_K_IQ4_KT/8]; +}; + +#if defined(DATA_A_IQ4_KT) +#define QUANT_K QUANT_K_IQ4_KT +#define QUANT_R QUANT_R_IQ4_KT +#define A_TYPE block_iq4_kt +#define A_TYPE_ROW_META +#define BLOCK_WORDS 32 +#endif + +#if defined(A_TYPE_ROW_META) +layout (binding = 0) readonly buffer A_U32 {uint data_a_u32[];}; + +uint a_row_words(uint ncols) { + return A_ROW_META_WORDS + (ncols / QUANT_K) * BLOCK_WORDS; +} + +uint a_block_word(uint row_word, uint ib) { + return row_word + A_ROW_META_WORDS + ib * BLOCK_WORDS; +} + +uint a_byte(uint base_word, uint idx) { + return (data_a_u32[base_word + idx / 4] >> (8 * (idx % 4))) & 0xFF; +} + +float a_row_scale(uint row_word) { + return uintBitsToFloat(data_a_u32[row_word]); +} +#endif + +#if defined(DATA_A_IQ4_KS) +const int8_t kvalues_iq4k_const[32] = { + int8_t(-127), int8_t(-104), int8_t(-83), int8_t(-65), int8_t(-49), int8_t(-35), int8_t(-22), int8_t(-10), + int8_t(1), int8_t(13), int8_t(25), int8_t(38), int8_t(53), int8_t(69), int8_t(89), int8_t(113), + int8_t(-123), int8_t(-100), int8_t(-79), int8_t(-61), int8_t(-45), int8_t(-31), int8_t(-18), int8_t(-6), + int8_t(5), int8_t(17), int8_t(29), int8_t(42), int8_t(57), int8_t(73), int8_t(93), int8_t(117) +}; + +shared FLOAT_TYPE kvalues_iq4k[32]; + +#define NEEDS_INIT_IQ_SHMEM +void init_iq_shmem(uvec3 wgsize) +{ + // copy the table into shared memory and sync + for (uint i = gl_LocalInvocationIndex.x; i < kvalues_iq4k.length(); i += wgsize.x) { + kvalues_iq4k[i] = FLOAT_TYPE(kvalues_iq4k_const[i]); + } + barrier(); +} + +// scales[ib32]: bits 1..7 signed magnitude, bit 0 selects the codebook half (-> vo) +float iq4_ks_scale(uint row_word, uint blk, uint ib32, out uint vo) +{ + const uint sc = a_byte(blk, ib32); + vo = (sc & 1) << 4; + return a_row_scale(row_word) * float(int(sc & 254) - 127); +} +#endif + +#if defined(DATA_A_IQ4_KT) +// one trellis step; 4 weights per index. See QuantizerIQKT::set_values +float iq4_kt_next(inout uint x) +{ + x *= 0xCBAC1FED; + const uint s = x & 0x3F3F3F3F; + return float(int((s & 0xFF) + ((s >> 8) & 0xFF) + ((s >> 16) & 0xFF) + (s >> 24)) - 126); +} + +// shb low byte: bits 1..7 signed scale, bit 0 selects the index-range half (-> offset) +float iq4_kt_scale(uint row_word, uint shb, out uint offset) +{ + offset = (shb & 1) != 0 ? 4096 + 32768 : 4096; + return a_row_scale(row_word) * float(int((shb & 0xFF) >> 1) - 64); +} + +// 15-bit index: 8b ql[jj], one qh[] nibble (low for jj < 32, high above), 3b from shb +uint iq4_kt_index(uint blk, uint shb, uint ib32, uint ig, uint offset) +{ + const uint jj = 8 * ib32 + ig; + const uint ql = a_byte(blk + 8, jj); + const uint qh = a_byte(blk + 24, jj % 32); + return offset + (ql | ((qh << (8 - 4 * (jj / 32))) & 0xF00) | (((shb >> (8 + 3 * ig)) & 7) << 12)); +} +#endif + // returns the bfloat value in the low 16b. // See ggml_compute_fp32_to_bf16 uint32_t fp32_to_bf16(float f) diff --git a/ggml/src/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/vulkan-shaders/vulkan-shaders-gen.cpp index 293aa644..eb10c26e 100644 --- a/ggml/src/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/vulkan-shaders/vulkan-shaders-gen.cpp @@ -64,9 +64,18 @@ const std::vector type_names = { "iq3_s", "iq4_xs", "iq4_nl", + "iq4_ks", + "iq4_kt", "bf16", }; +// IQ4_KS and IQ4_KT keep a per-row f32 scale ahead of the row's blocks, so they are read +// through a uint32_t alias of the A buffer and need dedicated shaders. The coopmat2 path +// cannot express their addressing and skips them. +static bool is_row_meta_quant(const std::string & tname) { + return tname == "iq4_ks" || tname == "iq4_kt"; +} + namespace { void execute_command(const std::string& command, std::string& stdout_str, std::string& stderr_str) { #ifdef _WIN32 @@ -362,13 +371,19 @@ void matmul_shaders(bool fp16, bool matmul_id, bool coopmat, bool coopmat2, bool std::string load_vec_quant = "2"; if ((tname == "q4_0") || (tname == "q4_1") || (tname == "iq1_s") || (tname == "iq1_m") || (tname == "iq2_xxs") || (tname == "iq2_xs") || (tname == "iq2_s")) load_vec_quant = "8"; - else if ((tname == "q5_0") || (tname == "q5_1") || (tname == "q8_0") || (tname == "iq3_xxs") || (tname == "iq3_s") || (tname == "iq4_nl")) + // iq4_kt: 4 so that one load index is exactly one trellis group of 4 weights + else if ((tname == "q5_0") || (tname == "q5_1") || (tname == "q8_0") || (tname == "iq3_xxs") || (tname == "iq3_s") || (tname == "iq4_nl") || (tname == "iq4_kt")) load_vec_quant = "4"; if (tname == "bf16") { continue; } + // the coopmat2 tensor layout cannot describe a per-row scale ahead of the blocks + if (coopmat2 && is_row_meta_quant(tname)) { + continue; + } + std::string data_a_key = "DATA_A_" + to_uppercase(tname); // For unaligned, load one at a time for f32/f16, or two at a time for quants std::string load_vec_a_unaligned = (coopmat2 || tname == "f32" || tname == "f16" || tname == "bf16") ? "1" : load_vec_quant; @@ -431,6 +446,7 @@ void process_shaders() { continue; } if (tname == "bf16") continue; + if (is_row_meta_quant(tname)) continue; #if defined(GGML_VULKAN_COOPMAT2_GLSLC_SUPPORT) if (tname == "f16") { @@ -466,7 +482,7 @@ void process_shaders() { for (const auto& tname : type_names) { // mul mat vec std::string data_a_key = "DATA_A_" + to_uppercase(tname); - std::string shader = (string_ends_with(tname, "_k") || string_starts_with(tname, "iq1_") || string_starts_with(tname, "iq2_") || string_starts_with(tname, "iq3_")) ? "mul_mat_vec_" + tname + ".comp" : "mul_mat_vec.comp"; + std::string shader = (string_ends_with(tname, "_k") || string_starts_with(tname, "iq1_") || string_starts_with(tname, "iq2_") || string_starts_with(tname, "iq3_") || is_row_meta_quant(tname)) ? "mul_mat_vec_" + tname + ".comp" : "mul_mat_vec.comp"; string_to_spv("mul_mat_vec_" + tname + "_f32_f32", shader, merge_maps(base_dict, {{data_a_key, "1"}, {"B_TYPE", "float"}, {"B_TYPE_VEC2", "vec2"}, {"B_TYPE_VEC4", "vec4"}, {"D_TYPE", "float"}})); string_to_spv("mul_mat_vec_" + tname + "_f16_f32", shader, merge_maps(base_dict, {{data_a_key, "1"}, {"B_TYPE", "float16_t"}, {"B_TYPE_VEC2", "f16vec2"}, {"B_TYPE_VEC4", "f16vec4"}, {"D_TYPE", "float"}})); @@ -479,7 +495,8 @@ void process_shaders() { } if (!string_ends_with(tname, "_k")) { - shader = (tname == "f32" || tname == "f16" || tname == "bf16") ? "get_rows.comp" : "get_rows_quant.comp"; + shader = (tname == "f32" || tname == "f16" || tname == "bf16") ? "get_rows.comp" : + is_row_meta_quant(tname) ? "get_rows_" + tname + ".comp" : "get_rows_quant.comp"; if (tname == "f16") { string_to_spv("get_rows_" + tname, shader, merge_maps(base_dict, {{data_a_key, "1"}, {"B_TYPE", "int"}, {"D_TYPE", "float16_t"}, {"OPTIMIZATION_ERROR_WORKAROUND", "1"}})); diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 878d4a34..4da85c41 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -202,6 +202,9 @@ llama_build_and_test( ) llama_build_and_test(test-regex-partial.cpp) +llama_build_and_test(test-iq4-ks-kt-decode.cpp) +target_include_directories(test-iq4-ks-kt-decode PRIVATE ${PROJECT_SOURCE_DIR}/ggml/src) + # llama_target_and_test(test-opt.cpp) # SLOW llama_target_and_test(test-model-load-cancel.cpp LABEL "model") diff --git a/tests/test-iq4-ks-kt-decode.cpp b/tests/test-iq4-ks-kt-decode.cpp new file mode 100644 index 00000000..00d64f17 --- /dev/null +++ b/tests/test-iq4-ks-kt-decode.cpp @@ -0,0 +1,443 @@ +// Checks the IQ4_KS and IQ4_KT decode arithmetic against ggml's to_float. +// +// This is a CPU test and contains no Vulkan: it re-implements, in C++, the indexing that +// each of the four Vulkan shader families uses for these two types, and diffs the result +// against the CPU reference. It therefore validates the format transcription, not the +// compiled shaders -- it cannot catch a driver or codegen fault, and it has to be kept in +// step with ggml/src/vulkan-shaders/types.comp by hand. The shaders themselves are covered +// on device by test-backend-ops and by GGML_VULKAN_CHECK_RESULTS. +// +// dequant_iq4_k*.comp one thread per 32-weight subblock, global block index +// get_rows_iq4_k*.comp two weights per invocation, row index taken from the shape +// mul_mm.comp A-load LOAD_VEC_A weights per load index +// mul_mat_vec_iq4_k*.comp one thread per subblock, rows walked by the row stride +// +// Buffers are synthesised directly, so every bit pattern is reachable and no quantizer is +// needed, and they are laid out with ggml_row_size() so that a wrong stride in the mirrored +// a_row_words() shows up as a mismatch rather than shifting the reference too. + +#include "ggml.h" + +#define GGML_COMMON_DECL_C +#define GGML_COMMON_IMPL_C +#include "ggml-common.h" + +#include +#include +#include +#include +#include + +static const uint32_t QK_K_ = 256; +static const uint32_t KS_BLOCK_W = 34; // 8 scale bytes + 128 quant bytes +static const uint32_t KT_BLOCK_W = 32; // 8 shb words + 64 ql bytes + 32 qh bytes +static const uint32_t ROW_META_W = 1; // the row scale + +static int g_checks = 0; +static int g_failures = 0; +static int g_reported = 0; + +static float bits_to_float(uint32_t u) { + float f; + memcpy(&f, &u, sizeof(f)); + return f; +} + +static uint32_t a_row_words(uint32_t ncols, uint32_t block_words) { + return ROW_META_W + (ncols / QK_K_) * block_words; +} + +static uint32_t rd_byte(const uint32_t * w, uint32_t base_word, uint32_t byte_idx) { + return (w[base_word + byte_idx / 4] >> (8 * (byte_idx % 4))) & 0xFF; +} + +// --------------------------------------------------------------------------------------- +// IQ4_KS, once per shader family + +static float ks_scale(const uint32_t * w, uint32_t row_word, uint32_t blk, uint32_t ib32, uint32_t & vo) { + const uint32_t sc = rd_byte(w, blk, ib32); + vo = (sc & 1) << 4; + return bits_to_float(w[row_word]) * (float)((int)(sc & 254) - 127); +} + +// the 32 weights of one subblock; shared by the dequant and mul_mat_vec mirrors, which +// differ in addressing, not in how a subblock is decoded +static void ks_emit_subblock(const uint32_t * w, uint32_t blk, uint32_t ib32, float dl, uint32_t vo, float * out) { + for (uint32_t l = 0; l < 4; ++l) { + const uint32_t qs = w[blk + 2 + 4 * ib32 + l]; + for (uint32_t k = 0; k < 4; ++k) { + const uint32_t q = (qs >> (8 * k)) & 0xFF; + out[4 * l + k ] = dl * (float)iq4k_values[vo + (q & 0xF)]; + out[4 * l + k + 16] = dl * (float)iq4k_values[vo + (q >> 4)]; + } + } +} + +static void ks_dequant_shader(const uint32_t * w, uint32_t ncols, uint32_t nrows, float * out) { + const uint32_t blocks_per_row = ncols / QK_K_; + for (uint32_t ib = 0; ib < blocks_per_row * nrows; ++ib) { + const uint32_t row_word = (ib / blocks_per_row) * a_row_words(ncols, KS_BLOCK_W); + const uint32_t blk = row_word + ROW_META_W + (ib % blocks_per_row) * KS_BLOCK_W; + for (uint32_t ib32 = 0; ib32 < 8; ++ib32) { + uint32_t vo; + const float dl = ks_scale(w, row_word, blk, ib32, vo); + ks_emit_subblock(w, blk, ib32, dl, vo, out + 256 * ib + 32 * ib32); + } + } +} + +static float ks_get_rows_elem(const uint32_t * w, uint32_t row_word, uint32_t i00) { + const uint32_t blk = row_word + ROW_META_W + (i00 / QK_K_) * KS_BLOCK_W; + const uint32_t pos = i00 % QK_K_; + const uint32_t ib32 = pos / 32; + const uint32_t j = pos % 32; + + uint32_t vo; + const float dl = ks_scale(w, row_word, blk, ib32, vo); + + const uint32_t qb = 16 * ib32 + (j & 15); + const uint32_t q = rd_byte(w, blk + 2, qb); + + return dl * (float)iq4k_values[vo + (j < 16 ? (q & 0xF) : (q >> 4))]; +} + +static void ks_get_rows_shader(const uint32_t * w, uint32_t ncols, uint32_t nrows, float * out) { + for (uint32_t r = 0; r < nrows; ++r) { + const uint32_t row_word = r * a_row_words(ncols, KS_BLOCK_W); + for (uint32_t i00 = 0; i00 < ncols; i00 += 2) { + out[r * ncols + i00 ] = ks_get_rows_elem(w, row_word, i00); + out[r * ncols + i00 + 1] = ks_get_rows_elem(w, row_word, i00 + 1); + } + } +} + +static void ks_mul_mm_shader(const uint32_t * w, uint32_t ncols, uint32_t nrows, float * out) { + const uint32_t load_vec_a = 2; + for (uint32_t idx = 0; idx < nrows * ncols / load_vec_a; ++idx) { + const uint32_t elem = idx * load_vec_a; + const uint32_t row_word = (elem / ncols) * a_row_words(ncols, KS_BLOCK_W); + const uint32_t acol = elem % ncols; + + const uint32_t blk = row_word + ROW_META_W + (acol / QK_K_) * KS_BLOCK_W; + const uint32_t ib32 = (acol % QK_K_) / 32; + const uint32_t j = acol % 32; + + uint32_t vo; + const float dl = ks_scale(w, row_word, blk, ib32, vo); + + const uint32_t qb = 16 * ib32 + (j & 15); + const uint32_t qw = w[blk + 2 + qb / 4] >> (8 * (qb % 4)); + const uint32_t qshift = j < 16 ? 0 : 4; + + out[elem ] = dl * (float)iq4k_values[vo + ((qw >> qshift) & 0xF)]; + out[elem + 1] = dl * (float)iq4k_values[vo + ((qw >> (8 + qshift)) & 0xF)]; + } +} + +static void ks_mul_mat_vec_shader(const uint32_t * w, uint32_t ncols, uint32_t nrows, float * out) { + const uint32_t words_per_row = a_row_words(ncols, KS_BLOCK_W); + for (uint32_t i = 0; i < ncols / QK_K_; ++i) { + for (uint32_t ib32 = 0; ib32 < 8; ++ib32) { + uint32_t row_word = 0; + for (uint32_t n = 0; n < nrows; ++n) { + const uint32_t blk = row_word + ROW_META_W + i * KS_BLOCK_W; + uint32_t vo; + const float dl = ks_scale(w, row_word, blk, ib32, vo); + ks_emit_subblock(w, blk, ib32, dl, vo, out + n * ncols + i * QK_K_ + 32 * ib32); + row_word += words_per_row; + } + } + } +} + +// --------------------------------------------------------------------------------------- +// IQ4_KT, once per shader family + +// the single trellis step, as written in types.comp +static float kt_next(uint32_t & x) { + x *= 0xCBAC1FEDu; + const uint32_t s = x & 0x3F3F3F3Fu; + return (float)((int32_t)((s & 0xFF) + ((s >> 8) & 0xFF) + ((s >> 16) & 0xFF) + (s >> 24)) - 126); +} + +// index assembly for one group of four weights, shared by the four mirrors below +static uint32_t kt_group_x(const uint32_t * w, uint32_t blk, uint32_t ib32, uint32_t ig, float d, float & sl) { + const uint32_t shb = w[blk + ib32]; + sl = d * (float)((int)((shb & 0xFF) >> 1) - 64); + const uint32_t offset = (shb & 1) != 0 ? 4096 + 32768 : 4096; + + const uint32_t jj = 8 * ib32 + ig; + const uint32_t ql = rd_byte(w, blk + 8, jj); + const uint32_t jh = jj % 32; + const uint32_t qh = rd_byte(w, blk + 24, jh); + + return offset + (ql | ((qh << (8 - 4 * (jj / 32))) & 0xF00) | (((shb >> (8 + 3 * ig)) & 7) << 12)); +} + +// the 32 weights of one subblock; shared by the dequant and mul_mat_vec mirrors, which +// differ in addressing, not in how a subblock is decoded +static void kt_emit_subblock(const uint32_t * w, uint32_t blk, uint32_t ib32, float d, float * out) { + for (uint32_t ig = 0; ig < 8; ++ig) { + float sl; + uint32_t x = kt_group_x(w, blk, ib32, ig, d, sl); + for (uint32_t k = 0; k < 4; ++k) { + out[4 * ig + k] = sl * kt_next(x); + } + } +} + +static void kt_dequant_shader(const uint32_t * w, uint32_t ncols, uint32_t nrows, float * out) { + const uint32_t blocks_per_row = ncols / QK_K_; + for (uint32_t ib = 0; ib < blocks_per_row * nrows; ++ib) { + const uint32_t row_word = (ib / blocks_per_row) * a_row_words(ncols, KT_BLOCK_W); + const uint32_t blk = row_word + ROW_META_W + (ib % blocks_per_row) * KT_BLOCK_W; + const float d = bits_to_float(w[row_word]); + for (uint32_t ib32 = 0; ib32 < 8; ++ib32) { + kt_emit_subblock(w, blk, ib32, d, out + 256 * ib + 32 * ib32); + } + } +} + +static void kt_get_rows_shader(const uint32_t * w, uint32_t ncols, uint32_t nrows, float * out) { + for (uint32_t r = 0; r < nrows; ++r) { + const uint32_t row_word = r * a_row_words(ncols, KT_BLOCK_W); + const float d = bits_to_float(w[row_word]); + for (uint32_t i00 = 0; i00 < ncols; i00 += 2) { + const uint32_t blk = row_word + ROW_META_W + (i00 / QK_K_) * KT_BLOCK_W; + const uint32_t pos = i00 % QK_K_; + float sl; + uint32_t x = kt_group_x(w, blk, pos / 32, (pos % 32) / 4, d, sl); + + const uint32_t k0 = pos % 4; + float v0 = 0.0f; + float v1 = 0.0f; + for (uint32_t l = 0; l < 4; ++l) { + const float v = sl * kt_next(x); + if (l == k0 ) v0 = v; + if (l == k0 + 1) v1 = v; + } + out[r * ncols + i00 ] = v0; + out[r * ncols + i00 + 1] = v1; + } + } +} + +static void kt_mul_mm_shader(const uint32_t * w, uint32_t ncols, uint32_t nrows, float * out) { + const uint32_t load_vec_a = 4; + for (uint32_t idx = 0; idx < nrows * ncols / load_vec_a; ++idx) { + const uint32_t elem = idx * load_vec_a; + const uint32_t row_word = (elem / ncols) * a_row_words(ncols, KT_BLOCK_W); + const uint32_t acol = elem % ncols; + + const uint32_t blk = row_word + ROW_META_W + (acol / QK_K_) * KT_BLOCK_W; + float sl; + uint32_t x = kt_group_x(w, blk, (acol % QK_K_) / 32, (acol % 32) / 4, bits_to_float(w[row_word]), sl); + + for (uint32_t k = 0; k < 4; ++k) { + out[elem + k] = sl * kt_next(x); + } + } +} + +static void kt_mul_mat_vec_shader(const uint32_t * w, uint32_t ncols, uint32_t nrows, float * out) { + const uint32_t words_per_row = a_row_words(ncols, KT_BLOCK_W); + for (uint32_t i = 0; i < ncols / QK_K_; ++i) { + for (uint32_t ib32 = 0; ib32 < 8; ++ib32) { + uint32_t row_word = 0; + for (uint32_t n = 0; n < nrows; ++n) { + const uint32_t blk = row_word + ROW_META_W + i * KT_BLOCK_W; + kt_emit_subblock(w, blk, ib32, bits_to_float(w[row_word]), + out + n * ncols + i * QK_K_ + 32 * ib32); + row_word += words_per_row; + } + } + } +} + +// --------------------------------------------------------------------------------------- + +// Copied from QuantizerIQKT<32, 4, 15, false, true>::set_values in ggml/src/iqk/iqk_quantize.cpp. +// It reads the accumulator through int8_t, so it carries the same little-endian assumption +// the original has. Used only to cross-check kt_next() over the whole index range. +static void ref_set_values(uint32_t i, float * result, float scale, int offset) { + uint32_t x = i + offset; + const uint32_t ka = 0xCBAC1FED; + uint32_t s; + const int8_t * i8 = (const int8_t *)&s; + for (int k = 0; k < 4; ++k) { + x = ka*x; + s = x & 0x3f3f3f3f; + result[k] = scale*(i8[0] + i8[1] + i8[2] + i8[3] - 126.f); + } +} + +static void check(const char * what, const std::vector & got, const std::vector & ref) { + g_checks++; + size_t bad = 0; + size_t first = 0; + for (size_t i = 0; i < ref.size(); ++i) { + if (got[i] != ref[i]) { + if (bad == 0) { + first = i; + } + bad++; + } + } + if (bad == 0) { + return; + } + g_failures++; + if (g_reported < 12) { + g_reported++; + printf(" FAIL %-42s %zu/%zu differ, first at %zu: got %.9g want %.9g\n", + what, bad, ref.size(), first, got[first], ref[first]); + } +} + +// Fills a quantized buffer. payload < 0 gives a reproducible pseudo-random pattern, +// otherwise every payload byte is set to that value. Every byte pattern is a valid encoding +// for both types, so this reaches states a quantizer would never emit. The first eight +// bytes of every block are then overwritten with the given scale bytes. +static void fill_pattern(std::vector & w, uint32_t ncols, uint32_t nrows, uint32_t block_words, + uint32_t words_per_row, uint32_t seed, int payload, const uint8_t * scales) { + const float row_scales[4] = { 0.0125f, -0.5f, 1.0f, 3.75e-3f }; + + uint32_t rng = seed * 2654435761u + 1u; + for (uint32_t r = 0; r < nrows; ++r) { + uint32_t bits; + const float d = row_scales[r % 4]; + memcpy(&bits, &d, sizeof(bits)); + w[r * words_per_row] = bits; + + for (uint32_t k = ROW_META_W; k < words_per_row; ++k) { + if (payload >= 0) { + w[r * words_per_row + k] = 0x01010101u * (uint32_t)payload; + } else { + rng ^= rng << 13; + rng ^= rng >> 17; + rng ^= rng << 5; + w[r * words_per_row + k] = rng; + } + } + + for (uint32_t b = 0; b < ncols / QK_K_; ++b) { + const uint32_t blk = r * words_per_row + ROW_META_W + b * block_words; + for (uint32_t k = 0; k < 8; ++k) { + const uint32_t word = blk + k / 4; + const uint32_t shift = 8 * (k % 4); + w[word] = (w[word] & ~(0xFFu << shift)) | ((uint32_t)scales[k] << shift); + } + } + } +} + +static void run_shape(ggml_type type, uint32_t ncols, uint32_t nrows, uint32_t seed, + int payload, const uint8_t * scales, const char * tag) { + const uint32_t block_words = type == GGML_TYPE_IQ4_KS ? KS_BLOCK_W : KT_BLOCK_W; + // laid out with ggml's own row size, so that a wrong stride in a_row_words() shows up as + // a mismatch rather than shifting the reference along with the mirrors + const uint32_t words_per_row = (uint32_t)ggml_row_size(type, ncols) / sizeof(uint32_t); + + std::vector w(words_per_row * nrows); + fill_pattern(w, ncols, nrows, block_words, words_per_row, seed, payload, scales); + + // ggml's own dequantizer, one row at a time + std::vector ref((size_t)ncols * nrows); + ggml_type_traits_t tt = ggml_internal_get_type_traits(type); + for (uint32_t r = 0; r < nrows; ++r) { + tt.to_float(&w[r * words_per_row], &ref[(size_t)r * ncols], ncols); + } + + struct mirror { + const char * family; + void (*fn)(const uint32_t *, uint32_t, uint32_t, float *); + }; + + const mirror ks[4] = { + { "dequant", ks_dequant_shader }, + { "get_rows", ks_get_rows_shader }, + { "mul_mm", ks_mul_mm_shader }, + { "mul_mat_vec", ks_mul_mat_vec_shader }, + }; + const mirror kt[4] = { + { "dequant", kt_dequant_shader }, + { "get_rows", kt_get_rows_shader }, + { "mul_mm", kt_mul_mm_shader }, + { "mul_mat_vec", kt_mul_mat_vec_shader }, + }; + const mirror * mirrors = type == GGML_TYPE_IQ4_KS ? ks : kt; + + std::vector got((size_t)ncols * nrows); + for (int m = 0; m < 4; ++m) { + std::fill(got.begin(), got.end(), 0.0f); + mirrors[m].fn(w.data(), ncols, nrows, got.data()); + + char name[128]; + snprintf(name, sizeof(name), "%s %-11s %ux%u %s", + ggml_type_name(type), mirrors[m].family, nrows, ncols, tag); + check(name, got, ref); + } +} + +static void run_trellis_sweep(void) { + const int offsets[2] = { 4096, 4096 + 32768 }; + const float scales[3] = { 1.0f, -0.25f, 7.5f }; + + size_t bad = 0; + for (uint32_t idx = 0; idx < 32768; ++idx) { + for (int o = 0; o < 2; ++o) { + for (int s = 0; s < 3; ++s) { + float ref[4]; + ref_set_values(idx, ref, scales[s], offsets[o]); + + uint32_t x = (uint32_t)offsets[o] + idx; + for (int k = 0; k < 4; ++k) { + if (scales[s] * kt_next(x) != ref[k]) { + bad++; + } + } + } + } + } + + g_checks++; + if (bad != 0) { + g_failures++; + printf(" FAIL iq4_kt trellis sweep: %zu of %d values differ\n", bad, 32768 * 2 * 3 * 4); + } +} + +int main(void) { + // scale bytes chosen to hit both codebook halves and both ends of the scale field: + // IQ4_KS uses (b & 254) - 127, which reaches -127 and +127; IQ4_KT uses (b >> 1) - 64, + // which reaches -64 and +63. Bit 0 selects the codebook half / the index offset. + const uint8_t extremes[8] = { 0x00, 0x01, 0xFE, 0xFF, 0x80, 0x81, 0x7E, 0x7F }; + const uint8_t mid[8] = { 0x40, 0x41, 0x7A, 0x0B, 0xC0, 0xC1, 0x2A, 0x95 }; + + // several row and block counts, so the row stride is exercised and not just the + // addressing inside a block + const uint32_t shapes[5][2] = { {256, 1}, {256, 3}, {1024, 5}, {1024, 4}, {4096, 2} }; + + printf("iq4_kt trellis: all 32768 indices x 2 offsets x 3 scales\n"); + run_trellis_sweep(); + + printf("multi-row buffers against ggml to_float\n"); + for (int s = 0; s < 5; ++s) { + run_shape(GGML_TYPE_IQ4_KS, shapes[s][0], shapes[s][1], 1u + s, -1, extremes, "rand/extreme-scales"); + run_shape(GGML_TYPE_IQ4_KT, shapes[s][0], shapes[s][1], 1u + s, -1, extremes, "rand/extreme-scales"); + run_shape(GGML_TYPE_IQ4_KS, shapes[s][0], shapes[s][1], 9u + s, -1, mid, "rand/mid-scales"); + run_shape(GGML_TYPE_IQ4_KT, shapes[s][0], shapes[s][1], 9u + s, -1, mid, "rand/mid-scales"); + } + + printf("saturated payloads\n"); + for (int p = 0; p <= 0xFF; p += 0xFF) { + run_shape(GGML_TYPE_IQ4_KS, 512, 2, 0, p, extremes, p == 0 ? "all-zero" : "all-ones"); + run_shape(GGML_TYPE_IQ4_KT, 512, 2, 0, p, extremes, p == 0 ? "all-zero" : "all-ones"); + run_shape(GGML_TYPE_IQ4_KS, 512, 2, 0, p, mid, p == 0 ? "all-zero" : "all-ones"); + run_shape(GGML_TYPE_IQ4_KT, 512, 2, 0, p, mid, p == 0 ? "all-zero" : "all-ones"); + } + + printf("%d checks, %d failed\n", g_checks, g_failures); + + return g_failures == 0 ? 0 : 1; +}