Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -13,11 +13,12 @@
// See the License for the specific language governing permissions and
// limitations under the License.

// NVFP4 prompt matmuls (256..2048 tokens in multiples of 256, input sizes in multiples of 256) on AMD's
// q8_1 x4 int8 WMMA kernel (ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256) at weight format 43: activations are
// quantized to q8_1 x4 once, NVFP4 blocks are staged as signed E2M1 codes with one scale per 16 values
// (motifs/nvfp4_q8_1_x4.loom). Without this, NVFP4 prompts take the generic f16 WMMA kernels
// (common.mul_mat.f32_f32_wmma, common.mul_mat_swiglu.f32_f32_wmma), which dequantize every weight to f16.
// NVFP4 and Q2_0 prompt matmuls (256..2048 tokens in multiples of 256, input sizes in multiples of 256) on AMD's
// q8_1 x4 int8 WMMA kernel (ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256) at weight format 43 / 42:
// activations are quantized to q8_1 x4 once; NVFP4 blocks are staged as signed E2M1 codes with one scale per 16
// values, Q2_0 blocks as signed -1..2 under their block scale (motifs/nvfp4_q8_1_x4.loom). Without this, they take
// the generic f16 WMMA kernels (common.mul_mat.f32_f32_wmma, common.mul_mat_swiglu.f32_f32_wmma), which dequantize
// every weight to f16.

#include "dispatch-mul-mat-nvfp4.h"

Expand All @@ -30,18 +31,20 @@ namespace {
static constexpr KernelCatalogRef kNvfp4Q8_1X4Kernel =
GGML_HRX_KERNEL_REF("loom_libs", "ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256");

static bool match_nvfp4_q8_1_x4_prefill_dispatch(const DispatchMatchContext & context,
DispatchMatch & dispatch_match) {
static bool match_own_q8_1_x4_prefill_dispatch(const DispatchMatchContext & context,
DispatchMatch & dispatch_match,
ggml_type type,
CommonMulMatWeightFormat format) {
if (context.root_node == nullptr || context.root_node->inputs.empty()) {
return false;
}
const Value * weight = common_graph_value(context.graph, context.root_node->inputs[0]);
if (weight == nullptr || weight->type != GGML_TYPE_NVFP4) {
if (weight == nullptr || weight->type != type) {
return false;
}
const CommonMulMatMatch match =
common_match_mul_mat_any_format(context.graph, context.root_node, kNvfp4Q8_1X4Kernel, false);
if (!match.matched() || match.weight_format != CommonMulMatWeightFormat::NVFP4 || match.token_count < 256 ||
if (!match.matched() || match.weight_format != format || match.token_count < 256 ||
match.token_count > 2048 || match.token_count % 256 != 0 || match.input_size % 256 != 0 ||
match.output_size % 64 != 0) {
return false;
Expand All @@ -63,8 +66,7 @@ static bool match_nvfp4_q8_1_x4_prefill_dispatch(const DispatchMatchContext & co
dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q8_1_x4.token_capacity",
common_to_config_value(match.token_count));
dispatch.kernel.compile_parameters.emplace("ggml.mul_mat_q8_1_x4.weight_format",
common_to_config_value(common_mul_mat_format_config_value(
CommonMulMatWeightFormat::NVFP4)));
common_to_config_value(common_mul_mat_format_config_value(format)));
dispatch.bindings.push_back(activation);
dispatch.bindings.push_back({ match.weight->id, 0, match.weight->byte_count });
dispatch.bindings.push_back({ match.output->id, 0, match.output->byte_count });
Expand All @@ -74,6 +76,18 @@ static bool match_nvfp4_q8_1_x4_prefill_dispatch(const DispatchMatchContext & co
return true;
}

static bool match_nvfp4_q8_1_x4_prefill_dispatch(const DispatchMatchContext & context,
DispatchMatch & dispatch_match) {
return match_own_q8_1_x4_prefill_dispatch(context, dispatch_match, GGML_TYPE_NVFP4,
CommonMulMatWeightFormat::NVFP4);
}

static bool match_q2_0_q8_1_x4_prefill_dispatch(const DispatchMatchContext & context,
DispatchMatch & dispatch_match) {
return match_own_q8_1_x4_prefill_dispatch(context, dispatch_match, GGML_TYPE_Q2_0,
CommonMulMatWeightFormat::Q2_0);
}

} // namespace

void register_nvfp4_prefill_dispatches(DispatchRegistryBuilder & registry) {
Expand All @@ -86,6 +100,14 @@ void register_nvfp4_prefill_dispatches(DispatchRegistryBuilder & registry) {
DispatchSource::Common,
match_nvfp4_q8_1_x4_prefill_dispatch,
});
registry.add({
"common.mul_mat.q2_0_q8_1_x4_prefill",
GGML_OP_MUL_MAT,
DispatchMatchKind::Fused,
300,
DispatchSource::Common,
match_q2_0_q8_1_x4_prefill_dispatch,
});
}

} // namespace ggml::hrx
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@

namespace ggml::hrx {

// NVFP4 prompt matmuls on the int8 q8_1 x4 WMMA kernel (motifs/nvfp4_q8_1_x4.loom).
// NVFP4 and Q2_0 prompt matmuls on the int8 q8_1 x4 WMMA kernel (motifs/nvfp4_q8_1_x4.loom).
void register_nvfp4_prefill_dispatches(DispatchRegistryBuilder & registry);

} // namespace ggml::hrx
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@
// See the License for the specific language governing permissions and
// limitations under the License.
//
// NVFP4 (HRX format 43) prompt matmuls on the int8 WMMA path of ops/mul_mat_q5_k_q8_plane_wmma.loom
// NVFP4 (HRX format 43) and Q2_0 (format 42, see below) prompt matmuls on the int8 WMMA path of ops/mul_mat_q5_k_q8_plane_wmma.loom
// (ggml_mul_mat_q5_k_iq4_xs_q8_1_x4_wmma_token256): activations are the q8_1 x4 groups (one d per 32 values),
// weights are staged per 64-value K chunk, which is exactly one 36-byte NVFP4 block: d[4] (UE4M3, one scale per 16
// values), then qs[32]; sub-block s reads qs[8 s .. 8 s + 7], low nibbles first.
Expand Down Expand Up @@ -174,18 +174,135 @@ func.def inline @ggml_nvfp4_x4_stage_row(%weight: buffer, %scratch: buffer, %w_o
func.return
}

// The q8_1 x4 kernel's weight staging hook: at NVFP4, staging threads (tid < rows) stage their row of the chunk
// here. Returns whether the kernel's own staging runs (tid < rows and not NVFP4).
func.def inline @ggml_nvfp4_x4_stage(%is_nvfp4: i1, %weight: buffer, %scratch: buffer, %w_off: offset, %ws_off: offset, %wc_off: offset, %input_size: index, %row_base: index, %tid: index, %rows: index, %chunk: index) -> (i1) {
// Q2_0 (format 42): 18-byte blocks of 64 values, d (f16) then qs[16]; value j is ((qs[j / 4] >> 2 (j % 4)) & 3) - 1
// times d. A K chunk is one block; it is staged as signed values -1..2 (so the kernel needs no offset term) under the
// same scale d for both K32 blocks, on the kernel's ordinary contraction.
// The 16 values of one 4-byte code word, four per byte in order.
func.def inline @ggml_q2_0_x4_values16(%codes_word: i32) -> (vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>) {
%c255 = scalar.constant 255 : i32
%c8 = scalar.constant 8 : i32
%c16 = scalar.constant 16 : i32
%c24 = scalar.constant 24 : i32
%vm1 = scalar.constant -1 : i8
%v0 = scalar.constant 0 : i8
%v1 = scalar.constant 1 : i8
%v2 = scalar.constant 2 : i8
%table = vector.from_elements %vm1, %v0, %v1, %v2, %v0, %v0, %v0, %v0, %v0, %v0, %v0, %v0, %v0, %v0, %v0, %v0 : vector<16xi8>
%b1s = scalar.shrui %codes_word, %c8 : i32
%b2s = scalar.shrui %codes_word, %c16 : i32
%b3s = scalar.shrui %codes_word, %c24 : i32
%b0 = scalar.andi %codes_word, %c255 : i32
%b1 = scalar.andi %b1s, %c255 : i32
%b2 = scalar.andi %b2s, %c255 : i32
%b3 = scalar.andi %b3s, %c255 : i32
%v0s = func.call @ggml_q2_0_x4_byte_values4(%b0, %table) : (i32, vector<16xi8>) -> (vector<4xi8>)
%v1s = func.call @ggml_q2_0_x4_byte_values4(%b1, %table) : (i32, vector<16xi8>) -> (vector<4xi8>)
%v2s = func.call @ggml_q2_0_x4_byte_values4(%b2, %table) : (i32, vector<16xi8>) -> (vector<4xi8>)
%v3s = func.call @ggml_q2_0_x4_byte_values4(%b3, %table) : (i32, vector<16xi8>) -> (vector<4xi8>)
func.return %v0s, %v1s, %v2s, %v3s : vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>
}

// The four values of one Q2_0 code byte: its 2-bit fields spread to bytes 0..3, then -1..2 by table.
func.def inline @ggml_q2_0_x4_byte_values4(%byte: i32, %table: vector<16xi8>) -> (vector<4xi8>) {
%c3 = scalar.constant 3 : i32
%c6 = scalar.constant 6 : i32
%c12 = scalar.constant 12 : i32
%c18 = scalar.constant 18 : i32
%m1 = scalar.constant 768 : i32
%m2 = scalar.constant 196608 : i32
%m3 = scalar.constant 50331648 : i32
%f0 = scalar.andi %byte, %c3 : i32
%s1 = scalar.shli %byte, %c6 : i32
%s2 = scalar.shli %byte, %c12 : i32
%s3 = scalar.shli %byte, %c18 : i32
%f1 = scalar.andi %s1, %m1 : i32
%f2 = scalar.andi %s2, %m2 : i32
%f3 = scalar.andi %s3, %m3 : i32
%f01 = scalar.ori %f0, %f1 : i32
%f23 = scalar.ori %f2, %f3 : i32
%spread = scalar.ori %f01, %f23 : i32
%spread_v = vector.from_elements %spread : vector<1xi32>
%codes = vector.bitcast %spread_v : vector<1xi32> to vector<4xi8>
%values = vector.table.lookup %table[%codes] : vector<16xi8>, vector<4xi8> -> vector<4xi8>
func.return %values : vector<4xi8>
}

// Stages K chunk %chunk (64 values = Q2_0 block %chunk) of weight row %row as staged row %lane_row: 64 signed values
// at %w_off + 80 lane_row, the block scale d at both scale slots of the row, zero corrections.
func.def inline @ggml_q2_0_x4_stage_row(%weight: buffer, %scratch: buffer, %w_off: offset, %ws_off: offset, %wc_off: offset, %input_size: index, %row: index, %lane_row: index, %chunk: index) {
%c0 = index.constant 0 : index
%c1 = index.constant 1 : index
%c2 = index.constant 2 : index
%c4 = index.constant 4 : index
%c8 = index.constant 8 : index
%c12 = index.constant 12 : index
%c16 = index.constant 16 : index
%c64 = index.constant 64 : index
%c80 = index.constant 80 : index
%zero = scalar.constant 0.0 : f32
%block_bytes = index.constant 18 : offset
%qs_offset = index.constant 2 : offset
%blocks = index.div %input_size, %c64 : index
%row_bytes = index.scale %blocks, %block_bytes : index, offset -> offset
%row_byte_base = index.scale %row, %row_bytes : index, offset -> offset
%block_add = index.scale %chunk, %block_bytes : index, offset -> offset
%block_base = index.add %row_byte_base, %block_add : offset
%qs_base = index.add %block_base, %qs_offset : offset
%d_view = buffer.view %weight[%block_base] : buffer -> view<1xf16>
%qs_view = buffer.view %weight[%qs_base] : buffer -> view<16xi8>
%d_vector = vector.load %d_view[%c0] : view<1xf16> -> vector<1xf16>
%d_f16 = vector.extract %d_vector[0] : vector<1xf16> -> f16
%d = scalar.extf %d_f16 : f16 to f32
%weights = buffer.view %scratch[%w_off] : buffer -> view<5120xi8>
%scales = buffer.view %scratch[%ws_off] : buffer -> view<128xf32>
%corrections = buffer.view %scratch[%wc_off] : buffer -> view<128xf32>
%bounded_lane_row = index.assume %lane_row [range(%lane_row, 0, 63)] : index
%dst0_0 = index.mul %bounded_lane_row, %c80 : index
%dst0 = index.assume %dst0_0 [range(%dst0_0, 0, 5040), mul(%dst0_0, 16)] : index
scf.for %word = [%c0 to %c4 step %c1] {
%byte_index = index.mul %word, %c4 : index
%code_bytes = vector.load %qs_view[%byte_index] : view<16xi8> -> vector<4xi8>
%code_word_v = vector.bitcast %code_bytes : vector<4xi8> to vector<1xi32>
%code_word = vector.extract %code_word_v[0] : vector<1xi32> -> i32
%v0, %v1, %v2, %v3 = func.call @ggml_q2_0_x4_values16(%code_word) : (i32) -> (vector<4xi8>, vector<4xi8>, vector<4xi8>, vector<4xi8>)
%word16 = index.mul %word, %c16 : index
%d0 = index.add %dst0, %word16 : index
%d4 = index.add %d0, %c4 : index
%d8 = index.add %d0, %c8 : index
%d12 = index.add %d0, %c12 : index
vector.store %v0, %weights[%d0] : vector<4xi8>, view<5120xi8>
vector.store %v1, %weights[%d4] : vector<4xi8>, view<5120xi8>
vector.store %v2, %weights[%d8] : vector<4xi8>, view<5120xi8>
vector.store %v3, %weights[%d12] : vector<4xi8>, view<5120xi8>
}
%meta0_0 = index.mul %bounded_lane_row, %c2 : index
%meta0 = index.assume %meta0_0 [range(%meta0_0, 0, 126)] : index
%meta1_0 = index.add %meta0, %c1 : index
%meta1 = index.assume %meta1_0 [range(%meta1_0, 1, 127)] : index
view.store %d, %scales[%meta0] : f32, view<128xf32>
view.store %d, %scales[%meta1] : f32, view<128xf32>
view.store %zero, %corrections[%meta0] : f32, view<128xf32>
view.store %zero, %corrections[%meta1] : f32, view<128xf32>
func.return
}

// The q8_1 x4 kernel's weight staging hook: at NVFP4 or Q2_0, staging threads (tid < rows) stage their row of the
// chunk here. Returns whether the kernel's own staging runs (tid < rows and neither format).
func.def inline @ggml_own_x4_stage(%is_nvfp4: i1, %is_q2_0: i1, %weight: buffer, %scratch: buffer, %w_off: offset, %ws_off: offset, %wc_off: offset, %input_size: index, %row_base: index, %tid: index, %rows: index, %chunk: index) -> (i1) {
%true = scalar.constant true : i1
%lane_active = index.cmp ult, %tid, %rows : index
%stage = scalar.andi %lane_active, %is_nvfp4 : i1
scf.if %stage {
%row0 = index.add %row_base, %tid : index
%row = index.assume %row0 [range(%row0, 0, 262143)] : index
%row0 = index.add %row_base, %tid : index
%row = index.assume %row0 [range(%row0, 0, 262143)] : index
%stage_nvfp4 = scalar.andi %lane_active, %is_nvfp4 : i1
scf.if %stage_nvfp4 {
func.call @ggml_nvfp4_x4_stage_row(%weight, %scratch, %w_off, %ws_off, %wc_off, %input_size, %row, %tid, %chunk) : (buffer, buffer, offset, offset, offset, index, index, index, index)
}
%other = scalar.xori %is_nvfp4, %true : i1
%stage_q2_0 = scalar.andi %lane_active, %is_q2_0 : i1
scf.if %stage_q2_0 {
func.call @ggml_q2_0_x4_stage_row(%weight, %scratch, %w_off, %ws_off, %wc_off, %input_size, %row, %tid, %chunk) : (buffer, buffer, offset, offset, offset, index, index, index, index)
}
%ours = scalar.ori %is_nvfp4, %is_q2_0 : i1
%other = scalar.xori %ours, %true : i1
%own = scalar.andi %lane_active, %other : i1
func.return %own : i1
}
Expand Down
Loading
Loading