diff --git a/ggml/src/ggml-backend-meta.cpp b/ggml/src/ggml-backend-meta.cpp --- a/ggml/src/ggml-backend-meta.cpp +++ b/ggml/src/ggml-backend-meta.cpp @@ ggml_cgraph * cgraph_main = nullptr; int offset = 0; // Node offset vs. original graph int offset_end = 0; // Exclusive node offset vs. original graph int fused_add = -1; + int fused_add_via = -1; int fused_to = -1; @@ int i_stop = i; int fused_add = -1; + int fused_add_via = -1; int fused_to = -1; @@ fused_add = i + 1; fused_allreduce_add_count++; } + + if (fused_add < 0 && i + 2 < cgraph->n_nodes && + next->op == GGML_OP_RESHAPE && + next->src[0] == node && + ggml_node_get_use_count(cgraph, i) == 1 && + ggml_node_get_use_count(cgraph, i + 1) == 1 && + next->type == GGML_TYPE_F32 && + node->type == GGML_TYPE_F32 && + ggml_nbytes(next) == ggml_nbytes(node) && + ggml_nbytes(node) <= 64*1024 && + next_split_state.axis == GGML_BACKEND_SPLIT_AXIS_MIRRORED) { + ggml_tensor * add = cgraph->nodes[i + 2]; + const bool reshaped_is_src0 = add->op == GGML_OP_ADD && add->src[0] == next; + const bool reshaped_is_src1 = add->op == GGML_OP_ADD && add->src[1] == next; + ggml_tensor * add_residual = reshaped_is_src0 ? add->src[1] : (reshaped_is_src1 ? add->src[0] : nullptr); + const ggml_backend_meta_split_state add_residual_split_state = add_residual ? + ggml_backend_meta_get_split_state(add_residual, /*assume_sync =*/ false) : + ggml_backend_meta_split_state{GGML_BACKEND_SPLIT_AXIS_UNKNOWN, {0}, 1}; + const ggml_backend_meta_split_state add_split_state = + ggml_backend_meta_get_split_state(add, /*assume_sync =*/ false); + + if ((reshaped_is_src0 || reshaped_is_src1) && + add->type == GGML_TYPE_F32 && + add_residual != nullptr && add_residual->type == GGML_TYPE_F32 && + ggml_nbytes(add) == ggml_nbytes(node) && + ggml_nbytes(add_residual) == ggml_nbytes(node) && + add_residual_split_state.axis == GGML_BACKEND_SPLIT_AXIS_MIRRORED && + add_split_state.axis == GGML_BACKEND_SPLIT_AXIS_MIRRORED) { + fused_add = i + 2; + fused_add_via = i + 1; + fused_allreduce_add_count++; + } + } } @@ - fprintf(stderr, " fused_to=%d", fused_to); + fprintf(stderr, " fused_add_via=%d fused_to=%d", fused_add_via, fused_to); @@ bcj.cgraphs[n_subgraphs].offset = i_start; bcj.cgraphs[n_subgraphs].offset_end = i_stop + 1; bcj.cgraphs[n_subgraphs].fused_add = fused_add; + bcj.cgraphs[n_subgraphs].fused_add_via = fused_add_via; bcj.cgraphs[n_subgraphs].fused_to = fused_to; @@ ggml_cgraph * cgraph_i0 = backend_ctx->backend_configs[0].cgraphs[i].cgraph_main; const int fused_add_i = backend_ctx->backend_configs[0].cgraphs[i].fused_add; + const int fused_add_via_i = backend_ctx->backend_configs[0].cgraphs[i].fused_add_via; const int fused_to_i = backend_ctx->backend_configs[0].cgraphs[i].fused_to; @@ auto & bcj = backend_ctx->backend_configs[j]; ggml_tensor * add_node = bcj.nodes[fused_add_i]; + ggml_tensor * add_input = nodes[j]; if (add_node->op != GGML_OP_ADD) { valid_fused_add = false; break; } - if (add_node->src[0] == nodes[j]) { + if (fused_add_via_i >= 0) { + ggml_tensor * via_node = bcj.nodes[fused_add_via_i]; + if (via_node->op != GGML_OP_RESHAPE || + via_node->src[0] != nodes[j] || + ggml_nbytes(via_node) != ggml_nbytes(nodes[j])) { + valid_fused_add = false; + break; + } + add_input = via_node; + } + if (add_node->src[0] == add_input) { residuals.push_back(add_node->src[1]); - } else if (add_node->src[1] == nodes[j]) { + } else if (add_node->src[1] == add_input) { residuals.push_back(add_node->src[0]); } else { valid_fused_add = false; @@ } if (fused_add_i >= 0 && !backend_allreduce_add_success) { + if (fused_add_via_i >= 0) { + const ggml_status status = compute_aux_node(fused_add_via_i); + if (status != GGML_STATUS_SUCCESS) { + return status; + } + } const ggml_status status = compute_aux_node(fused_add_i); if (status != GGML_STATUS_SUCCESS) { return status;