diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index b6f9965add7b..1b710218ef1f 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -4082,7 +4082,7 @@ struct GpuPipelineConfig { // Pipeline configuration for RDNA1 GPUs. static const std::unordered_map rdna1_pipelines = { {"soft_max", 64}, {"im2col", 64}, - {"argmax", 64}, {"mul_mat_vec", 64}, + {"argmax", 64}, {"mul_mat_vec", 32}, {"mul_mat_vec_f16", 32}, {"mul_mat_vec_f32_f16", 32} }; @@ -4763,6 +4763,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q1_0], matmul_q1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_PTQ1_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_PTQ1_0], matmul_ptq1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); + CREATE_MM2(GGML_TYPE_PQ2_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_PQ2_0], matmul_pq2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q2_0], matmul_q2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_0], matmul_q4_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_1], matmul_q4_1_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, ); @@ -4810,6 +4811,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_PTQ1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PTQ1_0], matmul_id_subgroup_ptq1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); + CREATE_MM2(GGML_TYPE_PQ2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PQ2_0], matmul_id_subgroup_pq2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id); @@ -4902,6 +4904,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q1_0], matmul_q1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); CREATE_MM2(GGML_TYPE_PTQ1_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_PTQ1_0], matmul_ptq1_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); + CREATE_MM2(GGML_TYPE_PQ2_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_PQ2_0], matmul_pq2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q2_0], matmul_q2_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_0], matmul_q4_0_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q4_1], matmul_q4_1_f32, mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); @@ -4952,6 +4955,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM_NODOT2(GGML_TYPE_BF16, pipeline_matmul_id_bf16, matmul_id_subgroup_bf16, , wg_denoms, warptile_id, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size_16); CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_subgroup_q1_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); CREATE_MM2(GGML_TYPE_PTQ1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PTQ1_0], matmul_id_subgroup_ptq1_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); + CREATE_MM2(GGML_TYPE_PQ2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PQ2_0], matmul_id_subgroup_pq2_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_subgroup_q2_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_subgroup_q4_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_subgroup_q4_1_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); @@ -5001,6 +5005,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM_NODOT2(GGML_TYPE_BF16, pipeline_matmul_id_bf16, matmul_id_bf16, , wg_denoms, warptile, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM2(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0], matmul_id_q1_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM2(GGML_TYPE_PTQ1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PTQ1_0], matmul_id_ptq1_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); + CREATE_MM2(GGML_TYPE_PQ2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PQ2_0], matmul_id_pq2_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM2(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0], matmul_id_q2_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM2(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0], matmul_id_q4_0_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM2(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1], matmul_id_q4_1_f32, mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); @@ -5081,6 +5086,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q1_0].f32acc, matmul_q1_0_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); CREATE_MM(GGML_TYPE_PTQ1_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_PTQ1_0].f32acc, matmul_ptq1_0_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); + CREATE_MM(GGML_TYPE_PQ2_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_PQ2_0].f32acc, matmul_pq2_0_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); CREATE_MM(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat[GGML_TYPE_Q2_0].f32acc, matmul_q2_0_f32, , mmq_wg_denoms, warptile_mmq, vk_mat_mat_push_constants, 3, , 0); 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, , 0); 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, , 0); @@ -5131,6 +5137,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0].f32acc, matmul_id_subgroup_q1_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); CREATE_MM(GGML_TYPE_PTQ1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PTQ1_0].f32acc, matmul_id_subgroup_ptq1_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); + CREATE_MM(GGML_TYPE_PQ2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PQ2_0].f32acc, matmul_id_subgroup_pq2_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); CREATE_MM(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0].f32acc, matmul_id_subgroup_q2_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); CREATE_MM(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0].f32acc, matmul_id_subgroup_q4_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); CREATE_MM(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1].f32acc, matmul_id_subgroup_q4_1_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, mul_mat_subgroup_size); @@ -5162,6 +5169,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { CREATE_MM(GGML_TYPE_Q1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q1_0].f32acc, matmul_id_q1_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM(GGML_TYPE_PTQ1_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PTQ1_0].f32acc, matmul_id_ptq1_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); + CREATE_MM(GGML_TYPE_PQ2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_PQ2_0].f32acc, matmul_id_pq2_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM(GGML_TYPE_Q2_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q2_0].f32acc, matmul_id_q2_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM(GGML_TYPE_Q4_0, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_0].f32acc, matmul_id_q4_0_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); CREATE_MM(GGML_TYPE_Q4_1, pipeline_dequant_mul_mat_mat_id[GGML_TYPE_Q4_1].f32acc, matmul_id_q4_1_f32, , mmq_wg_denoms, warptile_mmqid, vk_mat_mat_id_push_constants, mul_mat_id_param_count, _id, 0); @@ -5249,6 +5257,12 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { #define OCP_DMMV_DATA(NAME, REDUC) NAME ## _data[REDUC] #endif + // BC-250 (gfx1013): RDNA1, no coopmat/int-dot - PQ2_0 mat-vec tuned to 4 rows/WG on a 32-wide workgroup + const bool bc250_pq2_mmv = device->vendor_id == VK_VENDOR_ID_AMD && device->architecture == AMD_RDNA1 && + device->properties.deviceID == 0x13fe && !device->coopmat_support; + const uint32_t mmv_pq2_rows = bc250_pq2_mmv ? 4 : 2 * rm_stdq; + const uint32_t mmv_pq2_sg = bc250_pq2_mmv ? 32 : force_subgroup_size; + for (uint32_t w = 0; w < DMMV_WG_SIZE_COUNT; ++w) { const uint32_t wg_size_subgroup = (w == DMMV_WG_SIZE_SUBGROUP) ? subgroup_size : (subgroup_size * 4); const uint32_t wg_size_subgroup16 = (w == DMMV_WG_SIZE_SUBGROUP) ? subgroup_size16 : (subgroup_size16 * 4); @@ -5267,6 +5281,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_BF16][i], "mul_mat_vec_bf16_f32_f32", arr_dmmv_bf16_f32_f32_len[reduc], arr_dmmv_bf16_f32_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2, 1, 1}, {wg_size_subgroup, 2, i+1}, 1, false, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_Q1_0][i], "mul_mat_vec_q1_0_f32_f32", arr_dmmv_q1_0_f32_f32_len[reduc], arr_dmmv_q1_0_f32_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_PTQ1_0][i], "mul_mat_vec_ptq1_0_f32_f32", arr_dmmv_ptq1_0_f32_f32_len[reduc], arr_dmmv_ptq1_0_f32_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_PQ2_0][i], "mul_mat_vec_pq2_0_f32_f32", arr_dmmv_pq2_0_f32_f32_len[reduc], arr_dmmv_pq2_0_f32_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {mmv_pq2_rows, 1, 1}, {bc250_pq2_mmv ? 32u : wg_size_subgroup, mmv_pq2_rows, i+1}, 1, true, use_subgroups, mmv_pq2_sg); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_Q2_0][i], "mul_mat_vec_q2_0_f32_f32", arr_dmmv_q2_0_f32_f32_len[reduc], arr_dmmv_q2_0_f32_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_Q4_0][i], "mul_mat_vec_q4_0_f32_f32", arr_dmmv_q4_0_f32_f32_len[reduc], arr_dmmv_q4_0_f32_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f32_f32[w][GGML_TYPE_Q4_1][i], "mul_mat_vec_q4_1_f32_f32", arr_dmmv_q4_1_f32_f32_len[reduc], arr_dmmv_q4_1_f32_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); @@ -5296,6 +5311,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_BF16][i], "mul_mat_vec_bf16_f16_f32", arr_dmmv_bf16_f16_f32_len[reduc], arr_dmmv_bf16_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2, 1, 1}, {wg_size_subgroup, 2, i+1}, 1, false, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_Q1_0][i], "mul_mat_vec_q1_0_f16_f32", arr_dmmv_q1_0_f16_f32_len[reduc], arr_dmmv_q1_0_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_PTQ1_0][i], "mul_mat_vec_ptq1_0_f16_f32", arr_dmmv_ptq1_0_f16_f32_len[reduc], arr_dmmv_ptq1_0_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_PQ2_0][i], "mul_mat_vec_pq2_0_f16_f32", arr_dmmv_pq2_0_f16_f32_len[reduc], arr_dmmv_pq2_0_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_Q2_0][i], "mul_mat_vec_q2_0_f16_f32", arr_dmmv_q2_0_f16_f32_len[reduc], arr_dmmv_q2_0_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_Q4_0][i], "mul_mat_vec_q4_0_f16_f32", arr_dmmv_q4_0_f16_f32_len[reduc], arr_dmmv_q4_0_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_f16_f32[w][GGML_TYPE_Q4_1][i], "mul_mat_vec_q4_1_f16_f32", arr_dmmv_q4_1_f16_f32_len[reduc], arr_dmmv_q4_1_f16_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq, i+1}, 1, true, use_subgroups, force_subgroup_size); @@ -5326,6 +5342,9 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t wg_size_subgroup_int = (w == DMMV_WG_SIZE_SUBGROUP) ? subgroup_size_int : (subgroup_size_int * 4); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q2_0][i], "mul_mat_vec_q2_0_q8_1_f32", arr_dmmv_q2_0_q8_1_f32_len[reduc], arr_dmmv_q2_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); + // PTQ1_0: 4 rows per workgroup; the 4th constant is the column group size of the multi-column path + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_PTQ1_0][i], "mul_mat_vec_ptq1_0_q8_1_f32", arr_dmmv_ptq1_0_q8_1_f32_len[reduc], arr_dmmv_ptq1_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {4*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 4*rm_kq_int, i+1, 3}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_PQ2_0][i], "mul_mat_vec_pq2_0_q8_1_f32", arr_dmmv_pq2_0_q8_1_f32_len[reduc], arr_dmmv_pq2_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q4_0][i], "mul_mat_vec_q4_0_q8_1_f32", arr_dmmv_q4_0_q8_1_f32_len[reduc], arr_dmmv_q4_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q4_1][i], "mul_mat_vec_q4_1_q8_1_f32", arr_dmmv_q4_1_q8_1_f32_len[reduc], arr_dmmv_q4_1_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_q8_1_f32[w][GGML_TYPE_Q5_0][i], "mul_mat_vec_q5_0_q8_1_f32", arr_dmmv_q5_0_q8_1_f32_len[reduc], arr_dmmv_q5_0_q8_1_f32_data[reduc], "main", mul_mat_vec_num_bindings, sizeof(vk_mat_vec_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int, i+1}, 1, true, use_subgroups, subgroup_size_int); @@ -5352,6 +5371,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_BF16], "mul_mat_vec_id_bf16_f32", arr_dmmv_id_bf16_f32_f32_len[reduc], arr_dmmv_id_bf16_f32_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2, 1, 1}, {wg_size_subgroup, 2}, 1, false, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_Q1_0], "mul_mat_vec_id_q1_0_f32", arr_dmmv_id_q1_0_f32_f32_len[reduc], arr_dmmv_id_q1_0_f32_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_PTQ1_0], "mul_mat_vec_id_ptq1_0_f32", arr_dmmv_id_ptq1_0_f32_f32_len[reduc], arr_dmmv_id_ptq1_0_f32_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq}, 1, true, use_subgroups, force_subgroup_size); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_PQ2_0], "mul_mat_vec_id_pq2_0_f32", arr_dmmv_id_pq2_0_f32_f32_len[reduc], arr_dmmv_id_pq2_0_f32_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_Q2_0], "mul_mat_vec_id_q2_0_f32", arr_dmmv_id_q2_0_f32_f32_len[reduc], arr_dmmv_id_q2_0_f32_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_Q4_0], "mul_mat_vec_id_q4_0_f32", arr_dmmv_id_q4_0_f32_f32_len[reduc], arr_dmmv_id_q4_0_f32_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq}, 1, true, use_subgroups, force_subgroup_size); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_f32[w][GGML_TYPE_Q4_1], "mul_mat_vec_id_q4_1_f32", arr_dmmv_id_q4_1_f32_f32_len[reduc], arr_dmmv_id_q4_1_f32_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_stdq, 1, 1}, {wg_size_subgroup, 2*rm_stdq}, 1, true, use_subgroups, force_subgroup_size); @@ -5382,6 +5402,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { const uint32_t wg_size_subgroup_int = (w == DMMV_WG_SIZE_SUBGROUP) ? subgroup_size_int : (subgroup_size_int * 4); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q2_0], "mul_mat_vec_id_q2_0_q8_1_f32", arr_dmmv_id_q2_0_q8_1_f32_len[reduc], arr_dmmv_id_q2_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_PTQ1_0], "mul_mat_vec_id_ptq1_0_q8_1_f32", arr_dmmv_id_ptq1_0_q8_1_f32_len[reduc], arr_dmmv_id_ptq1_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); + ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_PQ2_0], "mul_mat_vec_id_pq2_0_q8_1_f32", arr_dmmv_id_pq2_0_q8_1_f32_len[reduc], arr_dmmv_id_pq2_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {2*rm_kq_int, 1, 1}, {wg_size_subgroup_int, 2*rm_kq_int}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q4_0], "mul_mat_vec_id_q4_0_q8_1_f32", arr_dmmv_id_q4_0_q8_1_f32_len[reduc], arr_dmmv_id_q4_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q4_1], "mul_mat_vec_id_q4_1_q8_1_f32", arr_dmmv_id_q4_1_q8_1_f32_len[reduc], arr_dmmv_id_q4_1_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); ggml_vk_create_pipeline(device, device->pipeline_dequant_mul_mat_vec_id_q8_1_f32[w][GGML_TYPE_Q5_0], "mul_mat_vec_id_q5_0_q8_1_f32", arr_dmmv_id_q5_0_q8_1_f32_len[reduc], arr_dmmv_id_q5_0_q8_1_f32_data[reduc], "main", mul_mat_vec_id_num_bindings, sizeof(vk_mat_vec_id_push_constants), {1*rm_stdq_int, 1, 1}, {wg_size_subgroup_int, 1*rm_stdq_int}, 1, true, use_subgroups, subgroup_size_int); @@ -5415,6 +5437,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { 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); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q1_0], "dequant_q1_0", dequant_q1_0_len, dequant_q1_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 8, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_PTQ1_0], "dequant_ptq1_0", dequant_ptq1_0_len, dequant_ptq1_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 8, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_PQ2_0], "dequant_pq2_0", dequant_pq2_0_len, dequant_pq2_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 4, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q2_0], "dequant_q2_0", dequant_q2_0_len, dequant_q2_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 4, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q4_0], "dequant_q4_0", dequant_q4_0_len, dequant_q4_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q4_1], "dequant_q4_1", dequant_q4_1_len, dequant_q4_1_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1); @@ -5446,6 +5469,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_BF16], "get_rows_bf16", get_rows_bf16_len, get_rows_bf16_data, "main", 3, sizeof(vk_op_binary_push_constants), { 512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_Q1_0], "get_rows_q1_0", get_rows_q1_0_len, get_rows_q1_0_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_PTQ1_0], "get_rows_ptq1_0", get_rows_ptq1_0_len, get_rows_ptq1_0_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_PQ2_0], "get_rows_pq2_0", get_rows_pq2_0_len, get_rows_pq2_0_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_Q2_0], "get_rows_q2_0", get_rows_q2_0_len, get_rows_q2_0_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_Q4_0], "get_rows_q4_0", get_rows_q4_0_len, get_rows_q4_0_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_get_rows[GGML_TYPE_Q4_1], "get_rows_q4_1", get_rows_q4_1_len, get_rows_q4_1_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); @@ -5476,6 +5500,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_get_rows_f32[GGML_TYPE_BF16], "get_rows_bf16_f32", get_rows_bf16_f32_len, get_rows_bf16_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_Q1_0], "get_rows_q1_0_f32", get_rows_q1_0_f32_len, get_rows_q1_0_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_PTQ1_0], "get_rows_ptq1_0_f32", get_rows_ptq1_0_f32_len, get_rows_ptq1_0_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_PQ2_0], "get_rows_pq2_0_f32", get_rows_pq2_0_f32_len, get_rows_pq2_0_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_Q2_0], "get_rows_q2_0_f32", get_rows_q2_0_f32_len, get_rows_q2_0_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_Q4_0], "get_rows_q4_0_f32", get_rows_q4_0_f32_len, get_rows_q4_0_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_Q4_1], "get_rows_q4_1_f32", get_rows_q4_1_f32_len, get_rows_q4_1_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {1024, 1, 1}, {}, 1); @@ -7694,6 +7719,7 @@ static vk_pipeline ggml_vk_get_to_fp16(ggml_backend_vk_context * ctx, ggml_type case GGML_TYPE_F32: case GGML_TYPE_Q1_0: case GGML_TYPE_PTQ1_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: @@ -7770,6 +7796,7 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_pipeline(ggml_backend_vk_conte switch (src0_type) { case GGML_TYPE_Q1_0: case GGML_TYPE_PTQ1_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: @@ -7825,7 +7852,9 @@ static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec(ggml_backend_vk_context * if (b_type == GGML_TYPE_Q8_1) { switch (a_type) { + case GGML_TYPE_PTQ1_0: case GGML_TYPE_Q2_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: @@ -7851,6 +7880,7 @@ static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec(ggml_backend_vk_context * case GGML_TYPE_BF16: case GGML_TYPE_Q1_0: case GGML_TYPE_PTQ1_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: @@ -7946,6 +7976,7 @@ static vk_matmul_pipeline ggml_vk_get_mul_mat_mat_id_pipeline(ggml_backend_vk_co switch (src0_type) { case GGML_TYPE_Q1_0: case GGML_TYPE_PTQ1_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: @@ -8000,6 +8031,8 @@ static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec_id(ggml_backend_vk_context if (b_type == GGML_TYPE_Q8_1) { switch (a_type) { case GGML_TYPE_Q2_0: + case GGML_TYPE_PTQ1_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: case GGML_TYPE_Q5_0: @@ -8025,6 +8058,7 @@ static vk_pipeline ggml_vk_get_dequantize_mul_mat_vec_id(ggml_backend_vk_context case GGML_TYPE_BF16: case GGML_TYPE_Q1_0: case GGML_TYPE_PTQ1_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: @@ -18212,6 +18246,7 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm case GGML_TYPE_BF16: case GGML_TYPE_Q1_0: case GGML_TYPE_PTQ1_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: @@ -18319,6 +18354,7 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm case GGML_TYPE_BF16: case GGML_TYPE_Q1_0: case GGML_TYPE_PTQ1_0: + case GGML_TYPE_PQ2_0: case GGML_TYPE_Q2_0: case GGML_TYPE_Q4_0: case GGML_TYPE_Q4_1: diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl index 88be1360403a..0f33e2ded7e2 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl @@ -168,6 +168,17 @@ vec4 dequantize4(uint ib, uint iqs, uint a_offset) { } #endif +#if defined(DATA_A_PQ2_0) +vec2 dequantize(uint ib, uint iqs, uint a_offset) { + const uint bits = uint(data_a[a_offset + ib].qs[iqs / 4u]) >> (2u * (iqs % 4u)); + return vec2(bits & 3u, (bits >> 2u) & 3u) - 1.0f; +} +vec4 dequantize4(uint ib, uint iqs, uint a_offset) { + const uint bits = uint(data_a[a_offset + ib].qs[iqs / 4u]); + return vec4(bits & 3u, (bits >> 2u) & 3u, (bits >> 4u) & 3u, bits >> 6u) - 1.0f; +} +#endif + #if defined(DATA_A_IQ1_S) vec2 dequantize(uint ib, uint iqs, uint a_offset) { const uint ib32 = iqs / 32; @@ -572,7 +583,7 @@ vec2 get_dm(uint ib, uint a_offset) { } #endif -#if defined(DATA_A_Q2_0) || defined(DATA_A_Q4_0) || defined(DATA_A_Q5_0) || defined(DATA_A_Q8_0) || defined(DATA_A_IQ1_S) || defined(DATA_A_IQ2_XXS) || defined(DATA_A_IQ2_XS) || defined(DATA_A_IQ2_S) || defined(DATA_A_IQ3_XXS) || defined(DATA_A_IQ3_S) || defined(DATA_A_IQ4_XS) || defined(DATA_A_IQ4_NL) +#if defined(DATA_A_Q2_0) || defined(DATA_A_PQ2_0) || defined(DATA_A_Q4_0) || defined(DATA_A_Q5_0) || defined(DATA_A_Q8_0) || defined(DATA_A_IQ1_S) || defined(DATA_A_IQ2_XXS) || defined(DATA_A_IQ2_XS) || defined(DATA_A_IQ2_S) || defined(DATA_A_IQ3_XXS) || defined(DATA_A_IQ3_S) || defined(DATA_A_IQ4_XS) || defined(DATA_A_IQ4_NL) vec2 get_dm(uint ib, uint a_offset) { return vec2(float(data_a[a_offset + ib].d), 0); } diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_pq2_0.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_pq2_0.comp new file mode 100644 index 000000000000..f18001e04e26 --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_pq2_0.comp @@ -0,0 +1,30 @@ +#version 450 + +#include "dequant_head.glsl" + +layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in; + +layout (binding = 0) readonly buffer A {block_pq2_0 data_a[];}; +layout (binding = 1) writeonly buffer D {D_TYPE data_b[];}; + +// 64 invocations cover 256 elements = two 128-weight blocks, 4 elements (one byte) each. +void main() { + const uint i = gl_WorkGroupID.x * 4 + gl_LocalInvocationID.x / 64; + + const uint tid = gl_LocalInvocationID.x % 64; + const uint il = tid % 32; + const uint ir = tid / 32; + const uint ib = 2*i + ir; + if (ib >= p.nel / QUANT_K_PQ2_0) { + return; + } + + const uint b_idx = 256*i + QUANT_K_PQ2_0*ir + 4*il; + const uint bits = uint(data_a[ib].qs[il]); + const float d = float(data_a[ib].d); + + data_b[b_idx ] = D_TYPE(d * (float(bits & 3u) - 1.0f)); + data_b[b_idx + 1] = D_TYPE(d * (float((bits >> 2u) & 3u) - 1.0f)); + data_b[b_idx + 2] = D_TYPE(d * (float((bits >> 4u) & 3u) - 1.0f)); + data_b[b_idx + 3] = D_TYPE(d * (float(bits >> 6u) - 1.0f)); +} diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp index 18d441ead40e..39774f27b91a 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq.comp @@ -11,7 +11,7 @@ layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; -#if defined(DATA_A_Q2_0) || defined(DATA_A_QUANT_K) +#if defined(DATA_A_Q2_0) || defined(DATA_A_PQ2_0) || defined(DATA_A_QUANT_K) #define K_PER_ITER 16 #elif defined(DATA_A_QUANT_LEGACY) || defined(DATA_A_MXFP4) #define K_PER_ITER 8 diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl index a5403ac82121..8b1322416009 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_funcs.glsl @@ -8,6 +8,10 @@ FLOAT_TYPE get_dm(uint ib) { return FLOAT_TYPE(data_a[ib / 2].d); } +#elif defined(DATA_A_PQ2_0) +FLOAT_TYPE get_dm(uint ib) { + return FLOAT_TYPE(data_a[ib / 4].d); +} #elif defined(DATA_A_Q4_0) || defined(DATA_A_Q5_0) || defined(DATA_A_Q8_0) || defined(DATA_A_IQ1_S) || defined(DATA_A_IQ2_XXS) || defined(DATA_A_IQ2_XS) || defined(DATA_A_IQ2_S) || defined(DATA_A_IQ3_XXS) || defined(DATA_A_IQ3_S) || defined(DATA_A_IQ4_XS) || defined(DATA_A_IQ4_NL) FLOAT_TYPE get_dm(uint ib) { return FLOAT_TYPE(data_a[ib].d); @@ -50,6 +54,27 @@ i32vec4 repack4(uint ib, uint iqs) { unpack_q2_0(bits >> 16u), unpack_q2_0(bits >> 24u)); } +FLOAT_TYPE mul_q8_1(const int32_t q_sum, const float da, const vec2 dsb, const int32_t sum_divisor) { + return FLOAT_TYPE(da * (float(q_sum) * dsb.x - dsb.y / float(sum_divisor))); +} +#endif +// PQ2_0: group-128 variant of the Q2_0 path (4 x 32-element chunks per block) +#if defined(DATA_A_PQ2_0) +uint unpack_pq2_0(uint bits) { + // Move bit pairs [1:0], [3:2], [5:4], [7:6] to [1:0], [9:8], [17:16], [25:24]. + bits &= 0xffu; + bits = (bits | (bits << 12u)) & 0x000f000fu; + return (bits | (bits << 6u)) & 0x03030303u; +} + +i32vec4 repack4(uint ib, uint iqs) { + const uint qs_idx = (ib & 3u) * 4u + iqs * 2u; + const uint bits = pack32(u16vec2(data_a_packed16[ib / 4].qs[qs_idx], + data_a_packed16[ib / 4].qs[qs_idx + 1])); + return i32vec4(unpack_pq2_0(bits), unpack_pq2_0(bits >> 8u), + unpack_pq2_0(bits >> 16u), unpack_pq2_0(bits >> 24u)); +} + FLOAT_TYPE mul_q8_1(const int32_t q_sum, const float da, const vec2 dsb, const int32_t sum_divisor) { return FLOAT_TYPE(da * (float(q_sum) * dsb.x - dsb.y / float(sum_divisor))); } @@ -157,7 +182,7 @@ FLOAT_TYPE mul_q8_1(const int32_t q_sum, const float da, const vec2 dsb, const i } #endif -#if defined(DATA_A_Q2_0) +#if defined(DATA_A_Q2_0) || defined(DATA_A_PQ2_0) FLOAT_TYPE mmvq_dot_product(const uint ib_a, const uint iqs) { int32_t q_sum = 0; const i32vec4 qs_a = repack4(ib_a, iqs); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_ptq1_0.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_ptq1_0.comp new file mode 100644 index 000000000000..fd21537bf58a --- /dev/null +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vecq_ptq1_0.comp @@ -0,0 +1,393 @@ +#version 450 + +// PTQ1_0 x q8_1 mat-vec with integer dot products. +// +// Block: 28 bytes = qs[24] (5 base-3 trits per byte), qh[2] (4 trits per byte), fp16 d. Element order is not positional: +// qs[j], j < 16: element t*16 + j (t = 0..4); qs[16 + j], j < 8: element 80 + t*8 + j; qh[h]: element 120 + 2n + h +// Trit level t of a byte: trit = (v*3) >> 8, v = (v*3) & 0xFF. This shader runs that on two bytes at a time in the 16-bit halves of a dword. +// Weight = (trit - 1) * d, so with q8_1 activations: sum_e (q_e - 1) * y_e = sum_k (d_b[k] * sum_{e in k} q_e * q8_e - s_b[k]), s_b = d_b * sum(q8). +// +// Two lanes per 128-element block: sub=0 takes qs dwords 0..2, sub=1 takes qs dwords 3..5 and qh. Same instruction stream, no divergence. +// One column (single sequence, bandwidth bound): load the activations of a block once, then decode and dot each row with the residue-tolerant trick (see TRIT_LEVEL). +// Several columns (batched decode, ALU bound): passes of CG columns; per block decode all rows once per level, pack the trits and dot them with every column of the pass. + +#extension GL_EXT_shader_explicit_arithmetic_types_int32 : require +#extension GL_EXT_integer_dot_product : require +#extension GL_EXT_shader_explicit_arithmetic_types_int16 : require + +#define B_TYPE block_q8_1_x4 + +#include "mul_mat_vec_base.glsl" + +layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in; + +// Raw dword views of the same bindings. +// A: 7 dwords per block: w[0..5] = qs[0..23], w[6] = qh[0] | qh[1] << 8 | d << 16 +// B: 36 dwords per block_q8_1_x4: ds[4] (f16 d, f16 s) then 32 dwords of int8x4 +layout (binding = 0) readonly buffer A_U32 {uint32_t data_a_u32[];}; +layout (binding = 1) readonly buffer B_U32 {uint32_t data_b_u32[];}; +layout (binding = 1) readonly buffer B_U64 {uvec2 data_b_u64[];}; +layout (binding = 1) readonly buffer B_U128 {uvec4 data_b_u128[];}; + +#define PTQ1_0_DWORDS 7u +#define Q8_1_X4_DWORDS 36u +#define Q8_1_X4_UVEC4 9u + +// columns per pass of the multi-column path (see above); 3 measured best on RDNA2 +layout (constant_id = 3) const uint CG = 3; + +uint a_offset, b_offset, d_offset; + +#if defined(PTQ1_0_MUL3_SHIFT) +#define MUL3(v) (((v) << 1u) + (v)) +#elif defined(PTQ1_0_MUL3_PK16) +#define MUL3(v) pack32(unpack16(v) * u16vec2(3us, 3us)) +#else +#define MUL3(v) ((v) * 3u) +#endif + +#ifdef PTQ1_0_DOT_NOSAT +#define DOT4(a, b, acc) ((acc) + dotPacked4x8EXT((a), (b))) +#else +#define DOT4(a, b, acc) dotPacked4x8AccSatEXT((a), (b), (acc)) +#endif + +// One trit level from the 16-bit lanes vl (bytes 0,2) and vh (bytes 1,3): after w = v*3 the trit is in byte 1/3 and the residue in byte 0/2. +// The activation words are pre-masked (bytes 0,2 zero; yh = y1,y3 in bytes 1,3; yl = y0,y2 moved to bytes 1,3), so the dot product ignores the residue. +#define TRIT_LEVEL(vl, vh, yl, yh, acc) \ + { \ + const uint wl = MUL3(vl); \ + const uint wh = MUL3(vh); \ + vl = wl & 0x00FF00FFu; \ + vh = wh & 0x00FF00FFu; \ + acc = DOT4(int32_t(wl), yl, acc); \ + acc = DOT4(int32_t(wh), yh, acc); \ + } +// last level: no residue needed afterwards +#define TRIT_LEVEL_LAST(vl, vh, yl, yh, acc) \ + { \ + acc = DOT4(int32_t(MUL3(vl)), yl, acc); \ + acc = DOT4(int32_t(MUL3(vh)), yh, acc); \ + } + +// Per-invocation state shared by the column-group passes. +uint num_blocks_per_row, sub, blk_first, blk_stride; +uint first_row_g, num_rows_g; +uint arow[NUM_ROWS]; // dword offset of the first block of each row +uint bcol[NUM_COLS]; // block_q8_1_x4 index of the first block of each column +FLOAT_TYPE temp[NUM_COLS][NUM_ROWS]; + +// One pass over the blocks for columns j0 .. j0 + CG - 1 (fewer in the last group). j0 is a literal at each call site, so column indices stay static. +void process_group(const uint j0) { + const uint nc = min(CG, NUM_COLS - j0); + + if (nc == 1u) { + // Single column: the row loop has a runtime bound, which keeps one row of decode state live (fewer registers than the level-major loop below). + for (uint blk = blk_first; blk < num_blocks_per_row; blk += blk_stride) { + const uint bx4 = bcol[j0] + blk; + const uint bbase = bx4 * Q8_1_X4_DWORDS + 4u; + + const uvec4 dsraw = data_b_u128[bx4 * Q8_1_X4_UVEC4]; + const vec2 ds0 = unpackHalf2x16(dsraw.x); + const vec2 ds1 = unpackHalf2x16(dsraw.y); + const vec2 ds2 = unpackHalf2x16(dsraw.z); + const vec2 ds3 = unpackHalf2x16(dsraw.w); + + // accA[t/2] holds sub-block t/2. accB[t/2] holds sub-block t/2 on lane 0 and sub-block (80 + 8t)/32 on lane 1 (2,2,3,3,3). Each lane subtracts the sums of its own two sub-blocks. + const float sA0 = ds0.x; + const float sA1 = ds1.x; + const float sA2 = ds2.x; + const float sB0 = (sub != 0u) ? ds2.x : ds0.x; + const float sB1 = (sub != 0u) ? ds3.x : ds1.x; + const float sB2 = (sub != 0u) ? ds3.x : ds2.x; + const float bias = (sub != 0u) ? (ds2.y + ds3.y) : (ds0.y + ds1.y); + + // Activations per trit level: both lanes issue the same two vector loads (16 bytes at dword 4t, 8 bytes at dword 20+2t) and pick their words. + int32_t yh0[5], yh1[5], yh2[5]; + int32_t yl0[5], yl1[5], yl2[5]; + [[unroll]] for (uint t = 0; t < 5; ++t) { + const uvec4 q4 = data_b_u128[(bbase >> 2u) + t]; + const uvec2 q2 = data_b_u64[(bbase >> 1u) + 10u + t]; + const uint y0 = (sub != 0u) ? q4.w : q4.x; + const uint y1 = (sub != 0u) ? q2.x : q4.y; + const uint y2 = (sub != 0u) ? q2.y : q4.z; + yh0[t] = int32_t(y0 & 0xFF00FF00u); + yh1[t] = int32_t(y1 & 0xFF00FF00u); + yh2[t] = int32_t(y2 & 0xFF00FF00u); + yl0[t] = int32_t((y0 << 8u) & 0xFF00FF00u); + yl1[t] = int32_t((y1 << 8u) & 0xFF00FF00u); + yl2[t] = int32_t((y2 << 8u) & 0xFF00FF00u); + } + // qh activations: elements 120..127 + const uvec2 yq = data_b_u64[(bbase >> 1u) + 15u]; + const int32_t yq0 = int32_t(yq.x); + const int32_t yq1 = int32_t(yq.y); + + [[unroll]] for (uint n = 0; n < num_rows_g; ++n) { + const uint abase = (a_offset + (first_row_g + n) * num_blocks_per_row + blk) * PTQ1_0_DWORDS; + const uint w0 = data_a_u32[abase + 3u*sub]; + const uint w1 = data_a_u32[abase + 3u*sub + 1u]; + const uint w2 = data_a_u32[abase + 3u*sub + 2u]; + const uint w6 = data_a_u32[abase + 6u]; + + const float da = unpackHalf2x16(w6).y; + + uint v0l = w0 & 0x00FF00FFu; + uint v0h = (w0 >> 8u) & 0x00FF00FFu; + uint v1l = w1 & 0x00FF00FFu; + uint v1h = (w1 >> 8u) & 0x00FF00FFu; + uint v2l = w2 & 0x00FF00FFu; + uint v2h = (w2 >> 8u) & 0x00FF00FFu; + + int32_t accA0 = 0, accA1 = 0, accA2 = 0; + int32_t accB0 = 0, accB1 = 0, accB2 = 0; + + TRIT_LEVEL(v0l, v0h, yl0[0], yh0[0], accA0); + TRIT_LEVEL(v1l, v1h, yl1[0], yh1[0], accB0); + TRIT_LEVEL(v2l, v2h, yl2[0], yh2[0], accB0); + + TRIT_LEVEL(v0l, v0h, yl0[1], yh0[1], accA0); + TRIT_LEVEL(v1l, v1h, yl1[1], yh1[1], accB0); + TRIT_LEVEL(v2l, v2h, yl2[1], yh2[1], accB0); + + TRIT_LEVEL(v0l, v0h, yl0[2], yh0[2], accA1); + TRIT_LEVEL(v1l, v1h, yl1[2], yh1[2], accB1); + TRIT_LEVEL(v2l, v2h, yl2[2], yh2[2], accB1); + + TRIT_LEVEL(v0l, v0h, yl0[3], yh0[3], accA1); + TRIT_LEVEL(v1l, v1h, yl1[3], yh1[3], accB1); + TRIT_LEVEL(v2l, v2h, yl2[3], yh2[3], accB1); + + TRIT_LEVEL_LAST(v0l, v0h, yl0[4], yh0[4], accA2); + TRIT_LEVEL_LAST(v1l, v1h, yl1[4], yh1[4], accB2); + TRIT_LEVEL_LAST(v2l, v2h, yl2[4], yh2[4], accB2); + + // qh: four trit levels of two bytes -> elements 120 + 2n + h, all in sub-block 3 (accB2 on lane 1). Lane 0 decodes zeros. + { + uint vq = (sub != 0u) ? ((w6 & 0xFFu) | ((w6 & 0xFF00u) << 8u)) : 0u; + uint wq; + wq = MUL3(vq); vq = wq & 0x00FF00FFu; + int32_t q01 = int32_t(((wq >> 8u) & 0x3u) | ((wq >> 16u) & 0x300u)); + wq = MUL3(vq); vq = wq & 0x00FF00FFu; + q01 |= int32_t(((wq << 8u) & 0x30000u) | (wq & 0x03000000u)); + wq = MUL3(vq); vq = wq & 0x00FF00FFu; + int32_t q23 = int32_t(((wq >> 8u) & 0x3u) | ((wq >> 16u) & 0x300u)); + wq = MUL3(vq); + q23 |= int32_t(((wq << 8u) & 0x30000u) | (wq & 0x03000000u)); + accB2 = DOT4(q01, yq0, accB2); + accB2 = DOT4(q23, yq1, accB2); + } + + const float sum = fma(sA0, float(accA0), + fma(sA1, float(accA1), + fma(sA2, float(accA2), + fma(sB0, float(accB0), + fma(sB1, float(accB1), + fma(sB2, float(accB2), -bias)))))); + temp[j0][n] = fma(da, sum, temp[j0][n]); + } + } + return; + } + + { + for (uint blk = blk_first; blk < num_blocks_per_row; blk += blk_stride) { + // qs dwords of this lane for every row, split into 16-bit lanes (l = bytes 0,2; h = bytes 1,3), plus the qh/d dword and the row scale. + uint v0l[NUM_ROWS], v0h[NUM_ROWS]; + uint v1l[NUM_ROWS], v1h[NUM_ROWS]; + uint v2l[NUM_ROWS], v2h[NUM_ROWS]; + uint w6[NUM_ROWS]; + float da[NUM_ROWS]; + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + const uint abase = arow[n] + blk * PTQ1_0_DWORDS; + const uint w0 = data_a_u32[abase + 3u*sub]; + const uint w1 = data_a_u32[abase + 3u*sub + 1u]; + const uint w2 = data_a_u32[abase + 3u*sub + 2u]; + w6[n] = data_a_u32[abase + 6u]; + da[n] = unpackHalf2x16(w6[n]).y; + v0l[n] = w0 & 0x00FF00FFu; + v0h[n] = (w0 >> 8u) & 0x00FF00FFu; + v1l[n] = w1 & 0x00FF00FFu; + v1h[n] = (w1 >> 8u) & 0x00FF00FFu; + v2l[n] = w2 & 0x00FF00FFu; + v2h[n] = (w2 >> 8u) & 0x00FF00FFu; + } + + // q8_1 block index and sub-block scales/sums of every column of the group + uint bx4[CG]; + uvec4 dsraw[CG]; + [[unroll]] for (uint jj = 0; jj < nc; ++jj) { + bx4[jj] = bcol[j0 + jj] + blk; + dsraw[jj] = data_b_u128[bx4[jj] * Q8_1_X4_UVEC4]; + } + + // accA: dword 0, always sub-block t/2. accB: dwords 1,2, sub-block t/2 on lane 0 and 2,2,3,3,3 on lane 1. Both change scale after levels 1 and 3. + int32_t accA[CG][NUM_ROWS]; + int32_t accB[CG][NUM_ROWS]; + [[unroll]] for (uint jj = 0; jj < CG; ++jj) { + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + accA[jj][n] = 0; + accB[jj][n] = 0; + } + } + + [[unroll]] for (uint t = 0; t < 5; ++t) { + // The three activation words of this lane per column: lane 0 dwords 4t..4t+2, lane 1 dwords 4t+3, 20+2t, 21+2t. Only these are loaded: the compiler hoists the loads of all levels, so registers matter. + int32_t yp0[CG], yp1[CG], yp2[CG]; + [[unroll]] for (uint jj = 0; jj < nc; ++jj) { + const uint bbase = bx4[jj] * Q8_1_X4_DWORDS + 4u; + const uint i0 = bbase + 4u*t + 3u*sub; + const uint i1 = (sub != 0u) ? (bbase + 20u + 2u*t) : (bbase + 4u*t + 1u); + yp0[jj] = int32_t(data_b_u32[i0]); + yp1[jj] = int32_t(data_b_u32[i1]); + yp2[jj] = int32_t(data_b_u32[i1 + 1u]); + } + + // Decode this level once per row, then dot with every column. + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + const uint w0l = MUL3(v0l[n]); + const uint w0h = MUL3(v0h[n]); + const uint w1l = MUL3(v1l[n]); + const uint w1h = MUL3(v1h[n]); + const uint w2l = MUL3(v2l[n]); + const uint w2h = MUL3(v2h[n]); + if (t < 4u) { + v0l[n] = w0l & 0x00FF00FFu; + v0h[n] = w0h & 0x00FF00FFu; + v1l[n] = w1l & 0x00FF00FFu; + v1h[n] = w1h & 0x00FF00FFu; + v2l[n] = w2l & 0x00FF00FFu; + v2h[n] = w2h & 0x00FF00FFu; + } + // Pack the four trits of each dword (bytes 0,2 from wl, bytes 1,3 from wh) and dot with the plain activation words: three dots per column. + const int32_t q0 = int32_t(((w0l >> 8u) & 0x00FF00FFu) | (w0h & 0xFF00FF00u)); + const int32_t q1 = int32_t(((w1l >> 8u) & 0x00FF00FFu) | (w1h & 0xFF00FF00u)); + const int32_t q2 = int32_t(((w2l >> 8u) & 0x00FF00FFu) | (w2h & 0xFF00FF00u)); + [[unroll]] for (uint jj = 0; jj < nc; ++jj) { + accA[jj][n] = DOT4(q0, yp0[jj], accA[jj][n]); + accB[jj][n] = DOT4(q1, yp1[jj], accB[jj][n]); + accB[jj][n] = DOT4(q2, yp2[jj], accB[jj][n]); + } + } + + if (t == 4u) { + // qh: four trit levels of two bytes -> elements 120 + 2n + h, all in sub-block 3 (the scale of accB on lane 1 at this level). Lane 0 decodes zeros. + int32_t yq0[CG], yq1[CG]; + [[unroll]] for (uint jj = 0; jj < nc; ++jj) { + const uint bbase = bx4[jj] * Q8_1_X4_DWORDS + 4u; + const uvec2 yq = data_b_u64[(bbase >> 1u) + 15u]; + yq0[jj] = int32_t(yq.x); + yq1[jj] = int32_t(yq.y); + } + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + uint vq = (sub != 0u) ? ((w6[n] & 0xFFu) | ((w6[n] & 0xFF00u) << 8u)) : 0u; + uint wq; + wq = MUL3(vq); vq = wq & 0x00FF00FFu; + int32_t q01 = int32_t(((wq >> 8u) & 0x3u) | ((wq >> 16u) & 0x300u)); + wq = MUL3(vq); vq = wq & 0x00FF00FFu; + q01 |= int32_t(((wq << 8u) & 0x30000u) | (wq & 0x03000000u)); + wq = MUL3(vq); vq = wq & 0x00FF00FFu; + int32_t q23 = int32_t(((wq >> 8u) & 0x3u) | ((wq >> 16u) & 0x300u)); + wq = MUL3(vq); + q23 |= int32_t(((wq << 8u) & 0x30000u) | (wq & 0x03000000u)); + [[unroll]] for (uint jj = 0; jj < nc; ++jj) { + accB[jj][n] = DOT4(q01, yq0[jj], accB[jj][n]); + accB[jj][n] = DOT4(q23, yq1[jj], accB[jj][n]); + } + } + } + + // Fold into fp32 when the sub-block scale changes (after levels 1 and 3) and at the end; the ternary offset (-sum(q8) of the owned sub-blocks) is applied once at the end. + if (t == 1u || t == 3u || t == 4u) { + [[unroll]] for (uint jj = 0; jj < nc; ++jj) { + { + float sA, sB, bias; + if (t == 1u) { + const vec2 ds0 = unpackHalf2x16(dsraw[jj].x); + const vec2 ds2 = unpackHalf2x16(dsraw[jj].z); + sA = ds0.x; + sB = (sub != 0u) ? ds2.x : ds0.x; + bias = 0.0f; + } else if (t == 3u) { + const vec2 ds1 = unpackHalf2x16(dsraw[jj].y); + const vec2 ds3 = unpackHalf2x16(dsraw[jj].w); + sA = ds1.x; + sB = (sub != 0u) ? ds3.x : ds1.x; + bias = 0.0f; + } else { + const vec2 ds0 = unpackHalf2x16(dsraw[jj].x); + const vec2 ds1 = unpackHalf2x16(dsraw[jj].y); + const vec2 ds2 = unpackHalf2x16(dsraw[jj].z); + const vec2 ds3 = unpackHalf2x16(dsraw[jj].w); + sA = ds2.x; + sB = (sub != 0u) ? ds3.x : ds2.x; + bias = (sub != 0u) ? (ds2.y + ds3.y) : (ds0.y + ds1.y); + } + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + const float sum = fma(sA, float(accA[jj][n]), fma(sB, float(accB[jj][n]), -bias)); + temp[j0 + jj][n] = fma(da[n], sum, temp[j0 + jj][n]); + accA[jj][n] = 0; + accB[jj][n] = 0; + } + } + } + } + } + } + } + +} + +void compute_outputs(const uint32_t first_row, const uint32_t num_rows) { + const uint tid = gl_LocalInvocationID.x; + + get_offsets(a_offset, b_offset, d_offset); + + num_blocks_per_row = p.ncols / QUANT_K; + sub = tid & 1u; + blk_first = tid >> 1u; + blk_stride = BLOCK_SIZE / 2u; + first_row_g = first_row; + num_rows_g = num_rows; + + // The multi-column path loops over the constant NUM_ROWS. Rows past the matrix (last workgroup) re-read the last valid row; reduce_result does not store them. + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + arow[n] = (a_offset + (first_row + min(n, num_rows - 1u)) * num_blocks_per_row) * PTQ1_0_DWORDS; + } + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + // one block_q8_1_x4 covers exactly the 128 activations of one block + bcol[j] = (j*p.batch_stride_b + b_offset) / QUANT_K; + } + + [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) { + [[unroll]] for (uint n = 0; n < NUM_ROWS; ++n) { + temp[j][n] = FLOAT_TYPE(0.0f); + } + } + + // One pass per column group, as guarded calls: the driver does not unroll a loop that contains the block loop, and one loop over all columns makes it hoist every load (spills at NUM_COLS = 8). NUM_COLS <= mul_mat_vec_max_cols (8). + process_group(0u); + if (NUM_COLS > 1u*CG) process_group(1u*CG); + if (NUM_COLS > 2u*CG) process_group(2u*CG); + if (NUM_COLS > 3u*CG) process_group(3u*CG); + if (NUM_COLS > 4u*CG) process_group(4u*CG); + if (NUM_COLS > 5u*CG) process_group(5u*CG); + if (NUM_COLS > 6u*CG) process_group(6u*CG); + if (NUM_COLS > 7u*CG) process_group(7u*CG); + + 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/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl index 3125a3fcfa93..1f60af348184 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl @@ -190,6 +190,18 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin const FLOAT_TYPE d = FLOAT_TYPE(data_a[ib].d); const uint bits = uint(data_a[ib].qs[iqs]); + const uint k_pair = row * LOAD_VEC_A / 2; + store_a(col, k_pair, d * (FLOAT_TYPEV2(bits & 3u, (bits >> 2u) & 3u) - FLOAT_TYPEV2(1.0f))); + store_a(col, k_pair + 1, d * (FLOAT_TYPEV2((bits >> 4u) & 3u, bits >> 6u) - FLOAT_TYPEV2(1.0f))); +#elif defined(DATA_A_PQ2_0) + const uint idx = pos_a + col * p.stride_a / LOAD_VEC_A + row; + + const uint ib = idx / 32; + const uint iqs = idx & 0x1fu; + + const FLOAT_TYPE d = FLOAT_TYPE(data_a[ib].d); + const uint bits = uint(data_a[ib].qs[iqs]); + const uint k_pair = row * LOAD_VEC_A / 2; store_a(col, k_pair, d * (FLOAT_TYPEV2(bits & 3u, (bits >> 2u) & 3u) - FLOAT_TYPEV2(1.0f))); store_a(col, k_pair + 1, d * (FLOAT_TYPEV2((bits >> 4u) & 3u, bits >> 6u) - FLOAT_TYPEV2(1.0f))); diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl index 469aad51ce89..f50742c8e900 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl +++ b/ggml/src/ggml-vulkan/vulkan-shaders/types.glsl @@ -232,6 +232,30 @@ struct block_ptq1_0 #define A_TYPE block_ptq1_0 #endif +// PQ2_0: Prism Q2_0 at group 128 (same 2-bit codec, one fp16 scale per 128). +#define QUANT_K_PQ2_0 128 +#define QUANT_R_PQ2_0 1 + +struct block_pq2_0 +{ + float16_t d; + uint8_t qs[QUANT_K_PQ2_0 / 4]; +}; + +struct block_pq2_0_packed16 +{ + float16_t d; + uint16_t qs[QUANT_K_PQ2_0 / 8]; +}; + +#if defined(DATA_A_PQ2_0) +#define QUANT_K QUANT_K_PQ2_0 +#define QUANT_R QUANT_R_PQ2_0 +#define QUANT_AUXF 1 +#define A_TYPE block_pq2_0 +#define A_TYPE_PACKED16 block_pq2_0_packed16 +#endif + #define QUANT_K_Q2_0 64 #define QUANT_R_Q2_0 1 diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index 83af8af068fe..f2239186d922 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -51,6 +51,7 @@ const std::vector type_names = { "f16", "q1_0", "ptq1_0", + "pq2_0", "q2_0", "q4_0", "q4_1", @@ -586,17 +587,14 @@ void matmul_shaders(bool fp16, MatMulIdType matmul_id_type, bool coopmat, bool c std::string load_vec_quant = "2"; if ((tname == "q1_0") || (tname == "ptq1_0") || (tname == "q4_0") || (tname == "q4_1") || (tname == "q5_1") || (tname == "iq1_s") || (tname == "iq1_m") || (tname == "iq2_xxs") || (tname == "iq2_xs") || (tname == "iq2_s")) load_vec_quant = "8"; - else if ((tname == "q2_0") || (tname == "q5_0") || (tname == "q8_0") || (tname == "q2_k") || (tname == "q4_k") || (tname == "q5_k") || (tname == "iq3_xxs") || (tname == "iq3_s") || (tname == "iq4_xs") || (tname == "iq4_nl") || (tname == "mxfp4") || (tname == "nvfp4")) + else if ((tname == "q2_0") || (tname == "pq2_0") || (tname == "q5_0") || (tname == "q8_0") || (tname == "q2_k") || (tname == "q4_k") || (tname == "q5_k") || (tname == "iq3_xxs") || (tname == "iq3_s") || (tname == "iq4_xs") || (tname == "iq4_nl") || (tname == "mxfp4") || (tname == "nvfp4")) load_vec_quant = "4"; if (tname == "bf16") { continue; } - // PTQ1_0 has no coopmat2 decoder: dequant_funcs_cm2.glsl carries no PTQ1_0 entry, - // so emitting mul_mm_cm2 for it fails shader compilation and takes the whole - // Vulkan build down, not just this type. Skip it; it falls back to the scalar and - // coopmat1 matmul paths, which are the ones implemented and tested. - if (coopmat2 && tname == "ptq1_0") { + // PTQ1_0 and PQ2_0 have no coopmat2 decoder in dequant_funcs_cm2.glsl, so no mul_mm_cm2 for them; they use the scalar and coopmat1 matmul paths + if (coopmat2 && (tname == "ptq1_0" || tname == "pq2_0")) { continue; } @@ -773,14 +771,16 @@ void process_shaders() { // mul mat vec with integer dot product #if defined(GGML_VULKAN_INTEGER_DOT_GLSLC_SUPPORT) - if (is_legacy_quant(tname) || tname == "mxfp4" || is_k_quant(tname) || tname == "iq1_s" || tname == "iq1_m") { - string_to_spv("mul_mat_vec_" + tname + "_q8_1_f32", "mul_mat_vecq.comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}})); - string_to_spv("mul_mat_vec_" + tname + "_q8_1_f32_subgroup", "mul_mat_vecq.comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}})); - string_to_spv("mul_mat_vec_" + tname + "_q8_1_f32_subgroup_no_shmem", "mul_mat_vecq.comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}, {"USE_SUBGROUP_ADD_NO_SHMEM", "1"}})); - - string_to_spv("mul_mat_vec_id_" + tname + "_q8_1_f32", "mul_mat_vecq.comp", merge_maps(base_dict, {{"MUL_MAT_ID", "1"}, {data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}})); - string_to_spv("mul_mat_vec_id_" + tname + "_q8_1_f32_subgroup", "mul_mat_vecq.comp", merge_maps(base_dict, {{"MUL_MAT_ID", "1"}, {data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}})); - string_to_spv("mul_mat_vec_id_" + tname + "_q8_1_f32_subgroup_no_shmem", "mul_mat_vecq.comp", merge_maps(base_dict, {{"MUL_MAT_ID", "1"}, {data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}, {"USE_SUBGROUP_ADD_NO_SHMEM", "1"}})); + if (is_legacy_quant(tname) || tname == "mxfp4" || is_k_quant(tname) || tname == "iq1_s" || tname == "iq1_m" || tname == "ptq1_0" || tname == "pq2_0") { + // PTQ1_0 has its own integer-dot shader (two lanes per 128-element block) + const std::string mmvq_shader = (tname == "ptq1_0") ? "mul_mat_vecq_ptq1_0.comp" : "mul_mat_vecq.comp"; + string_to_spv("mul_mat_vec_" + tname + "_q8_1_f32", mmvq_shader, merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}})); + string_to_spv("mul_mat_vec_" + tname + "_q8_1_f32_subgroup", mmvq_shader, merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}})); + string_to_spv("mul_mat_vec_" + tname + "_q8_1_f32_subgroup_no_shmem", mmvq_shader, merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}, {"USE_SUBGROUP_ADD_NO_SHMEM", "1"}})); + + string_to_spv("mul_mat_vec_id_" + tname + "_q8_1_f32", mmvq_shader, merge_maps(base_dict, {{"MUL_MAT_ID", "1"}, {data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}})); + string_to_spv("mul_mat_vec_id_" + tname + "_q8_1_f32_subgroup", mmvq_shader, merge_maps(base_dict, {{"MUL_MAT_ID", "1"}, {data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}})); + string_to_spv("mul_mat_vec_id_" + tname + "_q8_1_f32_subgroup_no_shmem", mmvq_shader, merge_maps(base_dict, {{"MUL_MAT_ID", "1"}, {data_a_key, "1"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV2", "vec2"}, {"ACC_TYPE", "float"}, {"USE_SUBGROUP_ADD_NO_SHMEM", "1"}})); } #endif @@ -1262,7 +1262,7 @@ void write_output_files() { for (const std::string& btype : btypes) { for (const auto& tname : type_names) { - if (btype == "q8_1" && !is_legacy_quant(tname) && tname != "mxfp4" && !is_k_quant(tname) && tname != "iq1_s" && tname != "iq1_m") { + if (btype == "q8_1" && !is_legacy_quant(tname) && tname != "mxfp4" && !is_k_quant(tname) && tname != "iq1_s" && tname != "iq1_m" && tname != "ptq1_0" && tname != "pq2_0") { continue; } hdr << "extern const void * arr_dmmv_" << tname << "_" << btype << "_f32_data[3];\n"; diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index b28bc29ef94a..14f654f966ca 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9297,6 +9297,19 @@ static std::vector> make_test_cases_eval() { } } + // PTQ1_0 / PQ2_0 integer-dot mat-vec: Bonsai-2 shapes, odd row counts (row tail), batches and multi-column B + for (int64_t n : {1, 2, 3, 4, 5, 6, 7, 8}) { + for (int64_t k : {1024, 5120, 6144, 17408}) { + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PTQ1_0, GGML_TYPE_F32, 67, n, k, {1, 1}, {1, 1})); + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PQ2_0, GGML_TYPE_F32, 67, n, k, {1, 1}, {1, 1})); + } + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PTQ1_0, GGML_TYPE_F32, 64, n, 2048, {2, 3}, {1, 1})); + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PQ2_0, GGML_TYPE_F32, 64, n, 2048, {2, 3}, {1, 1})); + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_PTQ1_0, GGML_TYPE_F32, 64, n, 2048, {1, 2}, {2, 1})); + test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_PTQ1_0, GGML_TYPE_F32, 4, 2, false, 70, n, 2048)); + test_cases.emplace_back(new test_mul_mat_id(GGML_TYPE_PQ2_0, GGML_TYPE_F32, 4, 2, false, 70, n, 2048)); + } + test_cases.emplace_back(new test_mul_mat(GGML_TYPE_Q4_0, GGML_TYPE_F32, 2880, 32, 2880, {1, 1}, {1, 1})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_Q8_0, GGML_TYPE_F32, 2880, 32, 2880, {1, 1}, {1, 1})); test_cases.emplace_back(new test_mul_mat(GGML_TYPE_MXFP4, GGML_TYPE_F32, 2880, 32, 2880, {1, 1}, {1, 1})); @@ -10276,6 +10289,18 @@ static std::vector> make_test_cases_eval() { // Test cases for performance evaluation: should be representative of real-world use cases static std::vector> make_test_cases_perf() { std::vector> test_cases; + // bandwidth comparison at Bonsai-2 shapes + for (ggml_type t : {GGML_TYPE_PTQ1_0, GGML_TYPE_PQ2_0, GGML_TYPE_Q4_0, GGML_TYPE_Q8_0, GGML_TYPE_Q2_K, GGML_TYPE_TQ2_0}) { + test_cases.emplace_back(new test_mul_mat(t, GGML_TYPE_F32, 17408, 1, 5120, {1, 1}, {1, 1})); + test_cases.emplace_back(new test_mul_mat(t, GGML_TYPE_F32, 5120, 1, 17408, {1, 1}, {1, 1})); + test_cases.emplace_back(new test_mul_mat(t, GGML_TYPE_F32, 10240, 1, 5120, {1, 1}, {1, 1})); + } + // batched decode (several sequences per step) through the mat-vec path + for (ggml_type t : {GGML_TYPE_PTQ1_0, GGML_TYPE_PQ2_0, GGML_TYPE_Q4_0}) { + for (int n : {2, 4, 8}) { + test_cases.emplace_back(new test_mul_mat(t, GGML_TYPE_F32, 17408, n, 5120, {1, 1}, {1, 1})); + } + } // SWIGLU at a 27B-class FFN width, fused [gate|up] vs split operands // note: same bytes either way, so a backend that indexes them differently shows it here