Skip to content
Merged
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
59 changes: 48 additions & 11 deletions ggml/src/ggml-openvino/openvino/pass/fuse_moe_compressed.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -321,22 +321,39 @@ FuseMoeCompressedFusedGateUp::FuseMoeCompressedFusedGateUp() {
auto gate_up_w_m = any_input();
auto ids_gate_up_m = any_input();
auto bgm_fused_m = wrap_type<ov::op::internal::GatherMatmul>({ a_m, gate_up_w_m, ids_gate_up_m, any_input() });
auto gu_u_m = optional<ov::op::v0::Convert>({ wrap_type<ov::op::v0::Unsqueeze>(
{ wrap_type<ov::op::v1::Transpose>({ bgm_fused_m, any_input() }), any_input() }) });
// Two graph shapes reach here. The rank-4 (stateless) graph restores the batch dim with an
// Unsqueeze after the Transpose and may Convert afterwards:
// GatherMatmul -> Transpose -> Unsqueeze -> [Convert] -> Slice
// The rank-3 (stateful) graph never drops to a batch dim at all, so there is no Unsqueeze
// here - it appears later, between the GEGLU and the down Reshape - and the Convert sits on
// the other side of the Transpose:
// GatherMatmul -> Convert -> Transpose -> Slice
// optional<T>({a, b}) matches T(a, b) or bare a, so one pattern covers both.
auto gu_t_m =
wrap_type<ov::op::v1::Transpose>({ optional<ov::op::v0::Convert>({ bgm_fused_m }), any_input() });
auto gu_u_m =
optional<ov::op::v0::Convert>({ optional<ov::op::v0::Unsqueeze>({ gu_t_m, any_input() }) });

auto gate_slice_m = wrap_type<ov::op::v8::Slice>({ gu_u_m, any_input(), any_input(), any_input(), any_input() });
auto up_slice_m = wrap_type<ov::op::v8::Slice>({ gu_u_m, any_input(), any_input(), any_input(), any_input() });
auto gelu_m = wrap_type<ov::op::v7::Gelu>({ gate_slice_m });
auto geglu_m = wrap_type<ov::op::v1::Multiply>({ gelu_m, up_slice_m });

// The rank-3 graph inserts the Unsqueeze the gate/up branch lacked right here, before the
// Reshape that feeds the down projection; the rank-4 graph goes straight from the GEGLU
// into the Reshape.
auto d_t_m = wrap_type<ov::op::v1::Transpose>(
{ optional<ov::op::v0::Convert>({ wrap_type<ov::op::v1::Reshape>({ geglu_m, any_input() }) }),
{ optional<ov::op::v0::Convert>({ wrap_type<ov::op::v1::Reshape>(
{ optional<ov::op::v0::Unsqueeze>({ geglu_m, any_input() }), any_input() }) }),
any_input() });
auto down_w_m = any_input();
auto ids_down_m = any_input();
auto bgm_down_m = wrap_type<ov::op::internal::GatherMatmul>({ d_t_m, down_w_m, ids_down_m, any_input() });
auto down_u_m = optional<ov::op::v0::Convert>({ wrap_type<ov::op::v0::Unsqueeze>(
{ wrap_type<ov::op::v1::Transpose>({ bgm_down_m, any_input() }), any_input() }) });
// Same two shapes as the gate/up branch above.
auto down_t_m =
wrap_type<ov::op::v1::Transpose>({ optional<ov::op::v0::Convert>({ bgm_down_m }), any_input() });
auto down_u_m =
optional<ov::op::v0::Convert>({ optional<ov::op::v0::Unsqueeze>({ down_t_m, any_input() }) });

// gemma-4 applies an extra per-expert output scale to the down projection before the
// router-weight multiply (llama-graph.cpp's ffn_down_exps.scale); FuseMoeCompressed's
Expand Down Expand Up @@ -421,21 +438,32 @@ FuseMoeCompressedFusedGateUp::FuseMoeCompressedFusedGateUp() {
}
const size_t top_k = ids_pshape[ids_pshape.rank().get_length() - 1].get_length();

// routing weights arrive as [1, n_tokens, top_k, 1]; the op wants [..., top_k]
auto routing = pm.at(routing_m);
const auto routing_pshape = routing.get_partial_shape();
if (routing_pshape.rank().is_dynamic() || routing_pshape.rank().get_length() != 4 ||
routing_pshape[3] != 1) {
// Routing weights arrive as [1, n_tokens, top_k, 1] on the rank-4 (stateless) graph and
// as [n_tokens, top_k, 1] on the rank-3 (stateful) one, which is the same thing without
// the leading batch dim. Normalise the rank-3 form up to the rank-4 one so everything
// below - and the op's own config - stays in the shape that is already validated on the
// stateless path; the batch dim is taken back off the result at the end.
const auto routing_in = pm.at(routing_m);
const auto routing_pshape = routing_in.get_partial_shape();
const auto routing_rank = routing_pshape.rank();
if (routing_rank.is_dynamic() || (routing_rank.get_length() != 4 && routing_rank.get_length() != 3) ||
routing_pshape[routing_rank.get_length() - 1] != 1) {
return false;
}
const bool batchless = routing_rank.get_length() == 3;

// Fold gemma-4's per-expert output scale into the routing weights: the reduction is
// sum_e(routing[e] * scale[e] * down_out[e]), and MOECompressed only takes one
// per-expert weight, so pre-multiply it into routing here (same [.., top_k, 1] shape).
auto down_scale = pm.at(down_scale_m);
if (down_scale.get_partial_shape() != routing_pshape) {
return false;
}
routing = std::make_shared<ov::op::v1::Multiply>(routing, down_scale);
ov::Output<ov::Node> routing = std::make_shared<ov::op::v1::Multiply>(routing_in, down_scale);
if (batchless) {
routing = std::make_shared<ov::op::v0::Unsqueeze>(
routing, ov::op::v0::Constant::create(ov::element::i64, ov::Shape{ 1 }, { 0 }));
}
routing = std::make_shared<ov::op::v0::Squeeze>(
routing, ov::op::v0::Constant::create(ov::element::i64, ov::Shape{ 1 }, { 3 }));
if (ids_pshape.rank().get_length() == 2) {
Expand Down Expand Up @@ -494,10 +522,19 @@ FuseMoeCompressedFusedGateUp::FuseMoeCompressedFusedGateUp() {
auto moe = std::make_shared<ov::op::internal::MOECompressed>(args, config);

ov::Output<ov::Node> result = moe->output(0);
// The op was fed the batched form, so it produces [1, n_tokens, hidden]. On the rank-3
// graph the ReduceSum being replaced is [n_tokens, hidden], so drop the batch dim again.
if (batchless) {
result = std::make_shared<ov::op::v0::Squeeze>(
result, ov::op::v0::Constant::create(ov::element::i64, ov::Shape{ 1 }, { 0 }));
}
const auto root_type = m.get_match_root()->get_output_element_type(0);
if (result.get_element_type() != root_type) {
result = std::make_shared<ov::op::v0::Convert>(result, root_type);
}
if (result.get_partial_shape() != m.get_match_root()->get_output_partial_shape(0)) {
return false;
}

result.get_node_shared_ptr()->set_friendly_name(m.get_match_root()->get_friendly_name());
ov::copy_runtime_info(m.get_matched_nodes(), result.get_node_shared_ptr());
Expand Down