Skip to content
Merged
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 @@ -1107,7 +1107,7 @@ static bool match_q6_k_token1_final_projection_q8_dispatch(const DispatchMatchCo

static bool match_mul_mat_postops_dispatch(const DispatchMatchContext & context, DispatchMatch & dispatch_match) {
const MulMatPostOpsMatch match = match_mul_mat_postops(context);
if (!match.matched()) {
if (!match.matched() || match.input_size % 256 != 0) {
return false;
}

Expand Down
457 changes: 457 additions & 0 deletions ggml/src/ggml-hrx/dispatch_registration/common/dispatch-small-rows.cpp

Large diffs are not rendered by default.

Original file line number Diff line number Diff line change
Expand Up @@ -18,9 +18,9 @@

namespace ggml::hrx {

// SOFT_MAX (no mask, scale 1), SUM_ROWS, ARGSORT and narrow GET_ROWS on short F32 rows, such
// SOFT_MAX (no mask, scale 1), SUM_ROWS, NORM, ARGSORT and narrow GET_ROWS on short F32 rows, such
// as a MoE router's, CONT of strided F32 views and broadcast REPEAT: ggml_softmax_rows_f32, ggml_sum_rows_f32,
// ggml_argsort_rows_f32, ggml_get_rows_small_f32 and ggml_copy_strided_f32.
// ggml_argsort_rows_f32, ggml_norm_rows_f32, ggml_get_rows_small_f32 and ggml_copy_strided_f32.
void register_small_rows_dispatches(DispatchRegistryBuilder & registry);

} // namespace ggml::hrx
1 change: 1 addition & 0 deletions ggml/src/ggml-hrx/ggml-hrx.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -620,6 +620,7 @@ static bool eager_capability_declared(enum ggml_op op) {
case GGML_OP_MUL:
case GGML_OP_MUL_MAT:
case GGML_OP_MUL_MAT_ID:
case GGML_OP_NORM:
case GGML_OP_PERMUTE:
case GGML_OP_REPEAT: // broadcast-only, through the strided copy (dispatch-small-rows.cpp)
case GGML_OP_RESHAPE:
Expand Down
2 changes: 2 additions & 0 deletions ggml/src/ggml-hrx/graph/op-params.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -258,6 +258,7 @@ OpParams import_op_params(const ggml_tensor & tensor) {
switch (tensor.op) {
case GGML_OP_RMS_NORM:
case GGML_OP_L2_NORM:
case GGML_OP_NORM:
return RmsNormParams{ ggml_get_op_params_f32(&tensor, 0) };
case GGML_OP_SOFT_MAX:
return SoftMaxParams{
Expand Down Expand Up @@ -312,6 +313,7 @@ bool op_params_equivalent(ggml_op op, const OpParams & lhs, const OpParams & rhs
switch (op) {
case GGML_OP_RMS_NORM:
case GGML_OP_L2_NORM:
case GGML_OP_NORM:
return rms_norm_params_equivalent(lhs, rhs);
case GGML_OP_SOFT_MAX:
return soft_max_params_equivalent(lhs, rhs);
Expand Down
Loading
Loading