diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-common.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-common.h index 55ff365fa552..14e518b236d8 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-common.h +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-common.h @@ -211,6 +211,8 @@ inline ggml_type common_mul_mat_format_type(CommonMulMatWeightFormat format) { return GGML_TYPE_IQ3_S; case CommonMulMatWeightFormat::IQ4_NL: return GGML_TYPE_IQ4_NL; + case CommonMulMatWeightFormat::MXFP4: + return GGML_TYPE_MXFP4; case CommonMulMatWeightFormat::IQ4_XS: return GGML_TYPE_IQ4_XS; case CommonMulMatWeightFormat::Q8_0: diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-common.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-common.h index 467fda4eea73..52594a3f0f20 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-common.h +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-id-common.h @@ -52,7 +52,8 @@ inline bool common_mul_mat_id_same_shape(const Value & lhs, const Value & rhs) { } inline bool common_mul_mat_id_supported_dense_input_size(CommonMulMatWeightFormat format, int64_t input_size) { - return common_mul_mat_supported_dense_input_size(format, input_size); + // the MUL_MAT_ID kernels declare input_size mul(256), also for the 32-value block formats + return common_mul_mat_supported_dense_input_size(format, input_size) && input_size % 256 == 0; } inline bool common_mul_mat_id_supported_dense_output_size(int64_t output_size) { diff --git a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-weight-format.h b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-weight-format.h index 7da50b8b93e5..7b54134a72e7 100644 --- a/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-weight-format.h +++ b/ggml/src/ggml-hrx/dispatch_registration/common/dispatch-mul-mat-weight-format.h @@ -29,6 +29,7 @@ enum class CommonMulMatWeightFormat { IQ2_S, IQ3_S, IQ4_NL, + MXFP4, IQ4_XS, Q8_0, Q8_1, @@ -99,6 +100,9 @@ inline bool common_mul_mat_format_for_type(ggml_type type, CommonMulMatWeightFor case GGML_TYPE_IQ4_NL: format = CommonMulMatWeightFormat::IQ4_NL; return true; + case GGML_TYPE_MXFP4: + format = CommonMulMatWeightFormat::MXFP4; + return true; case GGML_TYPE_IQ4_XS: format = CommonMulMatWeightFormat::IQ4_XS; return true; @@ -168,6 +172,8 @@ inline int64_t common_mul_mat_format_config_value(CommonMulMatWeightFormat forma return 21; case CommonMulMatWeightFormat::IQ4_NL: return 20; + case CommonMulMatWeightFormat::MXFP4: + return 39; case CommonMulMatWeightFormat::IQ4_XS: return 23; case CommonMulMatWeightFormat::Q8_0: @@ -200,6 +206,7 @@ inline bool common_mul_mat_supported_dense_input_size(CommonMulMatWeightFormat f case CommonMulMatWeightFormat::Q5_0: case CommonMulMatWeightFormat::Q5_1: case CommonMulMatWeightFormat::IQ4_NL: + case CommonMulMatWeightFormat::MXFP4: case CommonMulMatWeightFormat::Q8_0: case CommonMulMatWeightFormat::Q8_1: return input_size >= 256 && input_size <= 32768 && input_size % 32 == 0; diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json index 51a868c5d594..e966a9d60890 100644 --- a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/manifest.json @@ -173,6 +173,9 @@ { "path": "motifs/dequant_prism.loom" }, + { + "path": "motifs/dequant_1bit.loom" + }, { "path": "motifs/mul_mat_f32_f32_wmma_core.loom" }, @@ -1297,6 +1300,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1313,6 +1317,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1368,6 +1373,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1384,6 +1390,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1433,6 +1440,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1449,6 +1457,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1498,6 +1507,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1514,6 +1524,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1567,6 +1578,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1583,6 +1595,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1640,6 +1653,7 @@ "motifs/rope_f32.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1656,6 +1670,7 @@ "motifs/rope_f32.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1713,6 +1728,7 @@ "motifs/rope_f32.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1729,6 +1745,7 @@ "motifs/rope_f32.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1779,6 +1796,7 @@ "motifs/llm_attention_qkv_matmul_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1794,6 +1812,7 @@ "motifs/llm_attention_qkv_matmul_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -1848,6 +1867,7 @@ "motifs/rope_f32.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -1862,6 +1882,7 @@ "motifs/rope_f32.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -1917,6 +1938,7 @@ "motifs/rope_f32.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -1931,6 +1953,7 @@ "motifs/rope_f32.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -1979,6 +2002,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -1992,6 +2016,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2042,6 +2067,7 @@ "motifs/mul_mat_id_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2055,6 +2081,7 @@ "motifs/mul_mat_id_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2149,6 +2176,7 @@ "library_sources": [ "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2161,6 +2189,7 @@ "compile_dependencies": [ "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2217,6 +2246,7 @@ "motifs/mul_mat_id_q6k_f16_wmma_projection.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2234,6 +2264,7 @@ "motifs/mul_mat_id_q6k_f16_wmma_projection.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2289,6 +2320,7 @@ "motifs/mul_mat_id_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2303,6 +2335,7 @@ "motifs/mul_mat_id_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2364,6 +2397,7 @@ "motifs/mul_mat_id_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2378,6 +2412,7 @@ "motifs/mul_mat_id_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2543,6 +2578,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2557,6 +2593,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2606,6 +2643,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2620,6 +2658,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2669,6 +2708,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2683,6 +2723,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2734,6 +2775,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2748,6 +2790,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2805,6 +2848,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2819,6 +2863,7 @@ "motifs/mul_mat_f32_f32_postops.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2859,6 +2904,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2872,6 +2918,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2933,6 +2980,7 @@ "library_sources": [ "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -2945,6 +2993,7 @@ "compile_dependencies": [ "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -3009,6 +3058,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -3022,6 +3072,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -3070,6 +3121,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3087,6 +3139,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3141,6 +3194,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3158,6 +3212,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3210,6 +3265,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3227,6 +3283,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3281,6 +3338,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3298,6 +3356,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3350,6 +3409,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3367,6 +3427,7 @@ "motifs/mul_mat_f32_f32_wmma_core.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/unary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3419,6 +3480,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/binary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -3433,6 +3495,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/binary_f32_apply.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", @@ -4266,6 +4329,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4274,6 +4338,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4315,6 +4380,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4323,6 +4389,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4364,6 +4431,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4372,6 +4440,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4413,6 +4482,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4421,6 +4491,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4450,6 +4521,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4458,6 +4530,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4493,6 +4566,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4501,6 +4575,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4542,6 +4617,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4550,6 +4626,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4593,6 +4670,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4601,6 +4679,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4638,6 +4717,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4646,6 +4726,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4699,6 +4780,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4707,6 +4789,7 @@ "motifs/q6_k_f16.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/mul_mat_quantized_f16_prefill.loom", "motifs/binary_f32_apply.loom" ] @@ -4869,6 +4952,7 @@ "library_sources": [ "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/publish_f32.loom", "motifs/quantize_q8_1_x4.loom", "motifs/q4_k_f16.loom", @@ -4883,6 +4967,7 @@ "compile_dependencies": [ "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/publish_f32.loom", "motifs/quantize_q8_1_x4.loom", "motifs/q4_k_f16.loom", @@ -4948,6 +5033,7 @@ "library_sources": [ "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/publish_f32.loom", "motifs/quantize_q8_1_x4.loom", "motifs/q4_k_f16.loom", @@ -4962,6 +5048,7 @@ "compile_dependencies": [ "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/publish_f32.loom", "motifs/quantize_q8_1_x4.loom", "motifs/q4_k_f16.loom", @@ -7532,6 +7619,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", @@ -7545,6 +7633,7 @@ "ops/mul_mat_f32_f32_decode.loom", "motifs/dequant.loom", "motifs/dequant_prism.loom", + "motifs/dequant_1bit.loom", "motifs/q4_k_f16.loom", "motifs/q6_k_f16.loom", "motifs/q8_0_f16.loom", diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant.loom index 75bcac3371c5..d54366b9f166 100644 --- a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant.loom +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant.loom @@ -3,6 +3,8 @@ func.decl @ggml_pq2_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: i func.decl @ggml_ptq1_0_f32_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf32>) func.decl @ggml_ptq1_0_f16_vector4(%weight: buffer, %row_byte_base: offset, %k: index) -> (vector<4xf16>) +func.decl @ggml_mxfp4_f16_vector4(%weight: buffer, %row_byte_base: offset, %mx_block: index, %packet: index) -> (vector<4xf16>) +func.decl @ggml_mxfp4_f32_vector4(%weight: buffer, %row_byte_base: offset, %mx_block: index, %packet: index) -> (vector<4xf32>) func.def inline @ggml_paired_weight_words(%half_packet: i1, %paired: i1, %load_up: i1, %weight: buffer, %peer: buffer, %offset: offset) -> (vector<4xi32>) { %zero = index.constant 0 : index %padding = vector.constant 0 : vector<2xi32> @@ -11167,7 +11169,11 @@ func.def inline @ggml_dequant_weight_tile_bytes(%weight_format: index) -> (offse %selected_q5_1 = scf.select %is_q5_1, %q5_1_tile_bytes, %selected_q5_0 : offset %selected_iq2_s = scf.select %is_iq2_s, %iq2_s_tile_bytes, %selected_q5_1 : offset %selected_iq4_nl = scf.select %is_iq4_nl, %iq4_nl_tile_bytes, %selected_iq2_s : offset - %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_tile_bytes, %selected_iq4_nl : offset + %mxfp4_format_t = index.constant 39 : index + %mxfp4_tile_bytes = index.constant 136 : offset + %is_mxfp4_t = index.cmp eq, %weight_format, %mxfp4_format_t : index + %selected_mxfp4_t = scf.select %is_mxfp4_t, %mxfp4_tile_bytes, %selected_iq4_nl : offset + %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_tile_bytes, %selected_mxfp4_t : offset %selected_iq4_xs = scf.select %is_iq4_xs, %iq4_xs_tile_bytes, %selected_iq3_s : offset %selected_q8_0 = scf.select %is_q8_0, %q8_0_tile_bytes, %selected_iq4_xs : offset %selected_q8_1 = scf.select %is_q8_1, %q8_1_tile_bytes, %selected_q8_0 : offset @@ -11276,7 +11282,12 @@ func.def inline @ggml_dequant_weight_row_bytes(%weight_format: index, %hidden_si %selected_q5_0 = scf.select %is_q5_0, %q5_0_row_bytes, %selected_q4_1 : offset %selected_q5_1 = scf.select %is_q5_1, %q5_1_row_bytes, %selected_q5_0 : offset %selected_iq4_nl = scf.select %is_iq4_nl, %iq4_nl_row_bytes, %selected_q5_1 : offset - %selected_q8_0 = scf.select %is_q8_0, %q8_0_row_bytes, %selected_iq4_nl : offset + %mxfp4_format_r = index.constant 39 : index + %mxfp4_block_bytes = index.constant 17 : offset + %is_mxfp4_r = index.cmp eq, %weight_format, %mxfp4_format_r : index + %mxfp4_row_bytes = index.scale %legacy_block_count, %mxfp4_block_bytes : index, offset -> offset + %selected_mxfp4_r = scf.select %is_mxfp4_r, %mxfp4_row_bytes, %selected_iq4_nl : offset + %selected_q8_0 = scf.select %is_q8_0, %q8_0_row_bytes, %selected_mxfp4_r : offset %selected_q8_1 = scf.select %is_q8_1, %q8_1_row_bytes, %selected_q8_0 : offset %selected_f16 = scf.select %is_f16, %f16_row_bytes, %selected_q8_1 : offset %selected_bf16 = scf.select %is_bf16, %f16_row_bytes, %selected_f16 : offset @@ -11460,6 +11471,14 @@ func.def inline @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_gri } else { scf.yield %zero : vector<4xf16> } + %mxfp4_format_v = index.constant 39 : index + %is_mxfp4_v = index.cmp eq, %weight_format, %mxfp4_format_v : index + %mxfp4_values = scf.if %is_mxfp4_v -> (vector<4xf16>) { + %values = func.call @ggml_mxfp4_f16_vector4(%weight, %row_byte_base, %q8_block, %packet) : (buffer, offset, index, index) -> (vector<4xf16>) + scf.yield %values : vector<4xf16> + } else { + scf.yield %zero : vector<4xf16> + } %iq4_nl_values = scf.if %is_iq4_nl -> (vector<4xf16>) { %values = func.call @ggml_iq4nl_f16_vector4(%iq4nl_table, %weight, %row_byte_base, %q8_block, %packet) : (vector<16xi8>, buffer, offset, index, index) -> (vector<4xf16>) scf.yield %values : vector<4xf16> @@ -11528,7 +11547,8 @@ func.def inline @ggml_dequant_f16_vector4(%iq4nl_table: vector<16xi8>, %iq3s_gri %selected_q5_1 = scf.select %is_q5_1, %q5_1_values, %selected_q5_0 : vector<4xf16> %selected_iq2_s = scf.select %is_iq2_s, %iq2_s_values, %selected_q5_1 : vector<4xf16> %selected_iq4_nl = scf.select %is_iq4_nl, %iq4_nl_values, %selected_iq2_s : vector<4xf16> - %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_values, %selected_iq4_nl : vector<4xf16> + %selected_mxfp4_v = scf.select %is_mxfp4_v, %mxfp4_values, %selected_iq4_nl : vector<4xf16> + %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_values, %selected_mxfp4_v : vector<4xf16> %selected_iq4_xs = scf.select %is_iq4_xs, %iq4_xs_values, %selected_iq3_s : vector<4xf16> %selected_q8_0 = scf.select %is_q8_0, %q8_0_values, %selected_iq4_xs : vector<4xf16> %selected_q8_1 = scf.select %is_q8_1, %q8_1_values, %selected_q8_0 : vector<4xf16> @@ -11706,6 +11726,14 @@ func.def inline @ggml_dequant_f32_vector4(%iq4nl_table: vector<16xi8>, %iq3s_gri } else { scf.yield %zero : vector<4xf32> } + %mxfp4_format_v = index.constant 39 : index + %is_mxfp4_v = index.cmp eq, %weight_format, %mxfp4_format_v : index + %mxfp4_values = scf.if %is_mxfp4_v -> (vector<4xf32>) { + %values = func.call @ggml_mxfp4_f32_vector4(%weight, %row_byte_base, %q8_block, %packet) : (buffer, offset, index, index) -> (vector<4xf32>) + scf.yield %values : vector<4xf32> + } else { + scf.yield %zero : vector<4xf32> + } %iq4_nl_values = scf.if %is_iq4_nl -> (vector<4xf32>) { %values = func.call @ggml_iq4nl_f32_vector4(%iq4nl_table, %weight, %row_byte_base, %q8_block, %packet) : (vector<16xi8>, buffer, offset, index, index) -> (vector<4xf32>) scf.yield %values : vector<4xf32> @@ -11771,7 +11799,8 @@ func.def inline @ggml_dequant_f32_vector4(%iq4nl_table: vector<16xi8>, %iq3s_gri %selected_q5_1 = scf.select %is_q5_1, %q5_1_values, %selected_q5_0 : vector<4xf32> %selected_iq2_s = scf.select %is_iq2_s, %iq2_s_values, %selected_q5_1 : vector<4xf32> %selected_iq4_nl = scf.select %is_iq4_nl, %iq4_nl_values, %selected_iq2_s : vector<4xf32> - %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_values, %selected_iq4_nl : vector<4xf32> + %selected_mxfp4_v = scf.select %is_mxfp4_v, %mxfp4_values, %selected_iq4_nl : vector<4xf32> + %selected_iq3_s = scf.select %is_iq3_s, %iq3_s_values, %selected_mxfp4_v : vector<4xf32> %selected_iq4_xs = scf.select %is_iq4_xs, %iq4_xs_values, %selected_iq3_s : vector<4xf32> %selected_q8_0 = scf.select %is_q8_0, %q8_0_values, %selected_iq4_xs : vector<4xf32> %selected_q8_1 = scf.select %is_q8_1, %q8_1_values, %selected_q8_0 : vector<4xf32> diff --git a/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_1bit.loom b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_1bit.loom new file mode 100644 index 000000000000..b34c87398142 --- /dev/null +++ b/ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/motifs/dequant_1bit.loom @@ -0,0 +1,134 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. +// +// Weight decoders for the shared dequantizer (motifs/dequant.loom calls these from its format switch) that are not +// part of the upstream corpus. +// +// MXFP4 (format 39, gpt-oss): 17 bytes per 32 values = e (E8M0 shared exponent), qs[16]; value j is the low nibble +// of qs[j] and value j + 16 the high nibble, as in IQ4_NL. A value is kvalues_mxfp4[code] * 2^(e - 127) / 2 +// (ggml GGML_E8M0_TO_FP32_HALF), with the E2M1 table {0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12}. +// %packet (0..7) picks four consecutive values: 4 (packet % 4)..+4, from the high nibbles when packet >= 4. + +// Nibble unpacking: four codes from four bytes, low nibbles or (uses_high) high nibbles. +func.def inline @ggml_1bit_nibble_codes4(%q_bytes: vector<4xi8>, %uses_high: i1) -> (vector<4xi8>) { + %c4_i32 = scalar.constant 4 : i32 + %c8_i32 = scalar.constant 8 : i32 + %c12_i32 = scalar.constant 12 : i32 + %c16_i32 = scalar.constant 16 : i32 + %mask0_i32 = scalar.constant 15 : i32 + %mask1_i32 = scalar.constant 240 : i32 + %mask2_i32 = scalar.constant 3840 : i32 + %mask3_i32 = scalar.constant 61440 : i32 + %q_word_vector = vector.bitcast %q_bytes : vector<4xi8> to vector<1xi32> + %word = vector.extract %q_word_vector[0] : vector<1xi32> -> i32 + %word_shr4 = scalar.shrui %word, %c4_i32 : i32 + %word_shr8 = scalar.shrui %word, %c8_i32 : i32 + %word_shr12 = scalar.shrui %word, %c12_i32 : i32 + %word_shr16 = scalar.shrui %word, %c16_i32 : i32 + %low0 = scalar.andi %word, %mask0_i32 : i32 + %low1 = scalar.andi %word_shr4, %mask1_i32 : i32 + %low2 = scalar.andi %word_shr8, %mask2_i32 : i32 + %low3 = scalar.andi %word_shr12, %mask3_i32 : i32 + %low01 = scalar.ori %low0, %low1 : i32 + %low23 = scalar.ori %low2, %low3 : i32 + %low = scalar.ori %low01, %low23 : i32 + %high0 = scalar.andi %word_shr4, %mask0_i32 : i32 + %high1 = scalar.andi %word_shr8, %mask1_i32 : i32 + %high2 = scalar.andi %word_shr12, %mask2_i32 : i32 + %high3 = scalar.andi %word_shr16, %mask3_i32 : i32 + %high01 = scalar.ori %high0, %high1 : i32 + %high23 = scalar.ori %high2, %high3 : i32 + %high = scalar.ori %high01, %high23 : i32 + %selected = scf.select %uses_high, %high, %low : i32 + %target_selected_shr8 = scalar.shrui %selected, %c8_i32 : i32 + %packed0 = scalar.trunci %selected : i32 to i8 + %packed1 = scalar.trunci %target_selected_shr8 : i32 to i8 + %packed = vector.from_elements %packed0, %packed1 : vector<2xi8> + %codes = vector.bitunpacku<4> %packed : vector<2xi8> -> vector<4xi8> + func.return %codes : vector<4xi8> +} + +func.def inline @ggml_mxfp4_table_i8() -> (vector<16xi8>) { + %v0 = scalar.constant 0 : i8 + %v1 = scalar.constant 1 : i8 + %v2 = scalar.constant 2 : i8 + %v3 = scalar.constant 3 : i8 + %v4 = scalar.constant 4 : i8 + %v5 = scalar.constant 6 : i8 + %v6 = scalar.constant 8 : i8 + %v7 = scalar.constant 12 : i8 + %v8 = scalar.constant 0 : i8 + %v9 = scalar.constant -1 : i8 + %v10 = scalar.constant -2 : i8 + %v11 = scalar.constant -3 : i8 + %v12 = scalar.constant -4 : i8 + %v13 = scalar.constant -6 : i8 + %v14 = scalar.constant -8 : i8 + %v15 = scalar.constant -12 : i8 + %table = vector.from_elements %v0, %v1, %v2, %v3, %v4, %v5, %v6, %v7, %v8, %v9, %v10, %v11, %v12, %v13, %v14, %v15 : vector<16xi8> + func.return %table : vector<16xi8> +} + +// 2^(e - 127) / 2, exact, built from its f32 bits as ggml's GGML_E8M0_TO_FP32_HALF: (e - 1) << 23 for e >= 2, +// the subnormals 0x00200000 << e below that. +func.def inline @ggml_mxfp4_half_scale(%e: i32) -> (f32) { + %c1_i32 = scalar.constant 1 : i32 + %c2_i32 = scalar.constant 2 : i32 + %c23_i32 = scalar.constant 23 : i32 + %sub_unit = scalar.constant 2097152 : i32 + %is_small = scalar.cmpi ult, %e, %c2_i32 : i32 + %e_m1 = scalar.subi %e, %c1_i32 : i32 + %normal_bits = scalar.shli %e_m1, %c23_i32 : i32 + %sub_bits = scalar.shli %sub_unit, %e : i32 + %bits = scf.select %is_small, %sub_bits, %normal_bits : i32 + %bits_v = vector.from_elements %bits : vector<1xi32> + %scale_v = vector.bitcast %bits_v : vector<1xi32> to vector<1xf32> + %scale = vector.extract %scale_v[0] : vector<1xf32> -> f32 + func.return %scale : f32 +} + +func.def inline @ggml_mxfp4_f32_vector4(%weight: buffer, %row_byte_base: offset, %mx_block: index, %packet: index) -> (vector<4xf32>) { + %c0 = index.constant 0 : index + %c4 = index.constant 4 : index + %block_bytes = index.constant 17 : offset + %code_offset = index.constant 1 : offset + %bounded_packet = index.assume %packet [range(%packet, 0, 7)] : index + %block_byte_add = index.scale %mx_block, %block_bytes : index, offset -> offset + %block_byte_base = index.add %row_byte_base, %block_byte_add : offset + %code_byte_base = index.add %block_byte_base, %code_offset : offset + %e_view = buffer.view %weight[%block_byte_base] : buffer -> view<1xi8> + %code_view = buffer.view %weight[%code_byte_base] : buffer -> view<16xi8> + %packet_page = index.rem %bounded_packet, %c4 : index + %code_index0 = index.mul %packet_page, %c4 : index + %code_index = index.assume %code_index0 [range(%code_index0, 0, 12), mul(%code_index0, 4)] : index + %uses_high = index.cmp uge, %bounded_packet, %c4 : index + %q_bytes = vector.load %code_view[%code_index] : view<16xi8> -> vector<4xi8> + %codes = func.call @ggml_1bit_nibble_codes4(%q_bytes, %uses_high) : (vector<4xi8>, i1) -> (vector<4xi8>) + %table = func.call @ggml_mxfp4_table_i8() : () -> (vector<16xi8>) + %lookup_i8 = vector.table.lookup %table[%codes] : vector<16xi8>, vector<4xi8> -> vector<4xi8> + %lookup = vector.sitofp %lookup_i8 : vector<4xi8> to vector<4xf32> + %e_i8 = view.load %e_view[%c0] : view<1xi8> -> i8 + %e = scalar.extui %e_i8 : i8 to i32 + %scale = func.call @ggml_mxfp4_half_scale(%e) : (i32) -> (f32) + %scale_v = vector.splat %scale : vector<4xf32> + %result = vector.mulf %scale_v, %lookup : vector<4xf32> + func.return %result : vector<4xf32> +} + +func.def inline @ggml_mxfp4_f16_vector4(%weight: buffer, %row_byte_base: offset, %mx_block: index, %packet: index) -> (vector<4xf16>) { + %values_f32 = func.call @ggml_mxfp4_f32_vector4(%weight, %row_byte_base, %mx_block, %packet) : (buffer, offset, index, index) -> (vector<4xf32>) + %values = vector.fptrunc %values_f32 : vector<4xf32> to vector<4xf16> + func.return %values : vector<4xf16> +} diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index ae4d35a87055..cce45a5f37fc 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -361,4 +361,9 @@ if (TARGET ggml-hrx) add_executable(test-hrx-hadamard test-hrx-hadamard.cpp) target_link_libraries(test-hrx-hadamard PRIVATE ggml) add_test(NAME test-hrx-hadamard COMMAND test-hrx-hadamard) + + add_executable(test-hrx-mxfp4 test-hrx-mxfp4.cpp) + target_link_libraries(test-hrx-mxfp4 PRIVATE ggml-hrx ggml hrx::hrx loomc::loomc) + target_include_directories(test-hrx-mxfp4 PRIVATE ../ggml/include ../ggml/src ../ggml/src/ggml-hrx) + add_test(NAME test-hrx-mxfp4 COMMAND test-hrx-mxfp4) endif() diff --git a/tests/test-hrx-mxfp4.cpp b/tests/test-hrx-mxfp4.cpp new file mode 100644 index 000000000000..d815fa129e53 --- /dev/null +++ b/tests/test-hrx-mxfp4.cpp @@ -0,0 +1,165 @@ +// Copyright 2026 bong-water-water-bong +// SPDX-License-Identifier: Apache-2.0 +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Known-answer check for MXFP4 weights on HRX: GET_ROWS of MXFP4 rows must equal ggml's dequantize_row_mxfp4 +// bit for bit (the E8M0 half scale is a power of two and every E2M1 value is exact). Rows use the exponents +// 120, 127, 134 and the edges 0, 1, 2, 254; every block holds all 16 codes in both nibbles. Exponents 0 and 1 +// have f32 subnormal scales (2^-128 and 2^-127), which the GPU kernels flush to zero, so those rows' values +// (at most 12 * 2^-127) may come back as zeros of the same sign; every other value must match exactly. The graph runs on the HRX device itself (no scheduler, so no CPU fallback), and the HRX +// dispatch plan for it must contain the get_rows kernel. + +#include "dispatch/dispatch-scheduler.h" +#include "ggml-alloc.h" +#include "ggml-backend.h" +#include "ggml.h" +#include "graph/graph.h" +#include "kernel-corpus/kernel-corpus.h" + +#include +#include +#include +#include +#include +#include +#include + +#define REQUIRE(condition) \ + do { \ + if (!(condition)) { \ + std::fprintf(stderr, "%s:%d: requirement failed: %s\n", __FILE__, __LINE__, #condition); \ + std::abort(); \ + } \ + } while (false) + +static constexpr int kExponents[] = { 120, 127, 134, 0, 1, 2, 254 }; +static constexpr int kRows = sizeof(kExponents) / sizeof(kExponents[0]); +static constexpr int kValues = 256; // eight 32-value blocks per row +static constexpr int kBlocks = kValues / 32; +static constexpr int kBlockBytes = 17; // e (E8M0), qs[16] +static constexpr int kRowBytes = kBlocks * kBlockBytes; + +static std::vector make_weights() { + std::vector weights(static_cast(kRowBytes) * kRows); + for (int r = 0; r < kRows; ++r) { + for (int b = 0; b < kBlocks; ++b) { + uint8_t * block = weights.data() + static_cast(r) * kRowBytes + b * kBlockBytes; + block[0] = static_cast(kExponents[r]); + for (int j = 0; j < 16; ++j) { + const int low = (j + b) & 15; + const int high = (15 - j + 3 * b) & 15; + block[1 + j] = static_cast(low | (high << 4)); + } + } + } + return weights; +} + +static std::string kernel_name_for_id(uint64_t kernel_id) { + const ggml::hrx::KernelResolveResult resolved = + ggml::hrx::resolve_kernel_definition(ggml::hrx::get_qwen_kernel_corpus(), "gfx1151", kernel_id); + REQUIRE(resolved.found()); + return ggml::hrx::kernel_definition_name(*resolved.definition); +} + +// The HRX dispatch plan for the graph: every node must be covered, and a get_rows kernel must run. +static void require_hrx_get_rows_plan(ggml_cgraph * graph) { + ggml::hrx::GraphImportResult imported = ggml::hrx::import_ggml_graph(*graph); + REQUIRE(imported.valid()); + ggml::hrx::DispatchScheduler scheduler; + ggml::hrx::DispatchScheduleDiagnostics diagnostics; + if (!scheduler.schedule_graph(imported.graph, { "gfx1151" }, &diagnostics)) { + std::fprintf(stderr, "unsupported: %s\n", diagnostics.unsupported_message.c_str()); + std::abort(); + } + REQUIRE(scheduler.plan().valid()); + bool found = false; + for (const ggml::hrx::Dispatch & dispatch : scheduler.plan().dispatches) { + const std::string name = kernel_name_for_id(dispatch.kernel.kernel_id); + std::printf("dispatch: %s\n", name.c_str()); + found = found || name.find("get_rows") != std::string::npos; + } + REQUIRE(found); +} + +int main() { + ggml_backend_dev_t device = ggml_backend_dev_by_name("HRX0"); + if (device == nullptr) { + ggml_backend_load_all(); + device = ggml_backend_dev_by_name("HRX0"); + } + if (device == nullptr) { + std::printf("test-hrx-mxfp4: no HRX0 device, skipped\n"); + return 0; + } + ggml_backend_t backend = ggml_backend_dev_init(device, nullptr); + REQUIRE(backend != nullptr); + + ggml_init_params params = { 16 * ggml_tensor_overhead() + ggml_graph_overhead(), nullptr, true }; + ggml_context * ctx = ggml_init(params); + REQUIRE(ctx != nullptr); + ggml_tensor * weights = ggml_new_tensor_2d(ctx, GGML_TYPE_MXFP4, kValues, kRows); + ggml_tensor * ids = ggml_new_tensor_1d(ctx, GGML_TYPE_I32, kRows); + ggml_tensor * rows = ggml_get_rows(ctx, weights, ids); + ggml_cgraph * graph = ggml_new_graph(ctx); + ggml_build_forward_expand(graph, rows); + REQUIRE(ggml_backend_supports_op(backend, rows)); + require_hrx_get_rows_plan(graph); + + ggml_backend_buffer_t buffer = ggml_backend_alloc_ctx_tensors(ctx, backend); + REQUIRE(buffer != nullptr); + const std::vector host_weights = make_weights(); + std::vector host_ids(kRows); + for (int r = 0; r < kRows; ++r) { + host_ids[r] = kRows - 1 - r; + } + ggml_backend_tensor_set(weights, host_weights.data(), 0, host_weights.size()); + ggml_backend_tensor_set(ids, host_ids.data(), 0, host_ids.size() * sizeof(int32_t)); + REQUIRE(ggml_backend_graph_compute(backend, graph) == GGML_STATUS_SUCCESS); + + std::vector got(static_cast(kValues) * kRows); + std::vector expected(kValues); + ggml_backend_tensor_get(rows, got.data(), 0, got.size() * sizeof(float)); + const ggml_type_traits * traits = ggml_get_type_traits(GGML_TYPE_MXFP4); + REQUIRE(traits->to_float != nullptr); + int mismatches = 0; + int flushed = 0; + for (int r = 0; r < kRows; ++r) { + const int source = host_ids[r]; + traits->to_float(host_weights.data() + static_cast(source) * kRowBytes, expected.data(), kValues); + for (int k = 0; k < kValues; ++k) { + const float value = got[static_cast(r) * kValues + k]; + if (std::memcmp(&value, &expected[k], sizeof(float)) == 0) { + continue; + } + if (kExponents[source] < 2 && value == 0.0f && std::signbit(value) == std::signbit(expected[k])) { + ++flushed; + continue; + } + if (mismatches < 8) { + std::fprintf(stderr, "e=%d value %d: got %.9g, expected %.9g\n", kExponents[source], k, value, + expected[k]); + } + ++mismatches; + } + } + + ggml_backend_buffer_free(buffer); + ggml_free(ctx); + ggml_backend_free(backend); + REQUIRE(mismatches == 0); + std::printf("test-hrx-mxfp4: %d rows x %d values bit-exact (%d values with a subnormal scale flushed to zero)\n", + kRows, kValues, flushed); + return 0; +}