Skip to content
Merged
Changes from 2 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
25 changes: 16 additions & 9 deletions onnxruntime/core/providers/webgpu/reduction/reduction_ops.cc
Original file line number Diff line number Diff line change
Expand Up @@ -372,7 +372,22 @@ Status ReduceKernel<allow_multi_axes>::ComputeInternal(ComputeContext& context)
return Status::OK();
}

bool use_naive_reduction = name_ == "ArgMin" || name_ == "ArgMax" || (reduce_size < 32 && output_size > 1024) || is_input_empty || input_tensor->Shape().NumDimensions() == 0;
// Prefer the naive ReduceMean path when the shared path would pay extra overhead (transposing non-innermost reduce axes).
Comment thread
guschmue marked this conversation as resolved.
Outdated
constexpr size_t kReduceMeanNaiveMaxReduceSize = 128;
constexpr size_t kReduceMeanNaiveMinOutputSize = 20000;
bool are_axes_innermost = true;
size_t axes_rank = input_axes.size();
for (size_t i = 0; i < input_axes.size() && are_axes_innermost; ++i) {
if (input_axes[axes_rank - 1 - i] != rank - 1 - i) {
are_axes_innermost = false;
break;
}
}
bool use_naive_reduction = name_ == "ArgMin" || name_ == "ArgMax" || (reduce_size < 32 && output_size > 1024) ||
(name_ == "ReduceMean" && !are_axes_innermost &&
reduce_size <= kReduceMeanNaiveMaxReduceSize &&
output_size > kReduceMeanNaiveMinOutputSize) ||
is_input_empty || input_tensor->Shape().NumDimensions() == 0;

if (use_naive_reduction) {
ReduceNaiveProgram program(name_, reduce_op_type, keepdims_, noop_with_empty_axes_, input_axes, is_input_empty);
Expand All @@ -395,14 +410,6 @@ Status ReduceKernel<allow_multi_axes>::ComputeInternal(ComputeContext& context)

return context.RunProgram(program);
} else {
bool are_axes_innermost = true;
size_t axes_rank = input_axes.size();
for (size_t i = 0; i < input_axes.size() && are_axes_innermost; ++i) {
if (input_axes[axes_rank - 1 - i] != rank - 1 - i) {
are_axes_innermost = false;
break;
}
}
Tensor input_transpose;
if (!are_axes_innermost) {
InlinedVector<size_t> perm;
Expand Down
Loading