diff --git a/js/web/docs/webgpu-operators.md b/js/web/docs/webgpu-operators.md index 5c8748d75c2bc..90ea46a300e33 100644 --- a/js/web/docs/webgpu-operators.md +++ b/js/web/docs/webgpu-operators.md @@ -60,6 +60,7 @@ Do not modify directly.* | GridSample | ai.onnx(16-19); com.ms.internal.nhwc(16-19) | | | GroupQueryAttention | com.microsoft(1+) | | | HardSigmoid | ai.onnx(6+) | | +| HardSwish | ai.onnx(14+) | | | If | ai.onnx(1-10,11-12,13-18,19-20,21+) | | | InstanceNormalization | ai.onnx(6+); com.ms.internal.nhwc(6+) | | | LayerNormalization | ai.onnx(1-16,17+) | | diff --git a/js/web/lib/wasm/jsep/webgpu/op-resolve-rules.ts b/js/web/lib/wasm/jsep/webgpu/op-resolve-rules.ts index 6c7afbc7365bb..7cdb3015600c9 100644 --- a/js/web/lib/wasm/jsep/webgpu/op-resolve-rules.ts +++ b/js/web/lib/wasm/jsep/webgpu/op-resolve-rules.ts @@ -111,6 +111,7 @@ export const WEBGPU_OP_RESOLVE_RULES: Map = new ['GridSample', [gridSample, parseGridSampleAttributes]], ['GroupQueryAttention', [groupQueryAttention]], ['HardSigmoid', [unaryOps.hardSigmoid, unaryOps.parseHardSigmoidAttributes]], + ['HardSwish', [unaryOps.hardSwish]], ['InstanceNormalization', [instanceNorm]], ['LayerNormalization', [layerNorm]], ['LeakyRelu', [unaryOps.leakyRelu, unaryOps.parseAlphaAttributes]], diff --git a/js/web/lib/wasm/jsep/webgpu/ops/unary-op.ts b/js/web/lib/wasm/jsep/webgpu/ops/unary-op.ts index 168d644fe064c..9907628a08dfd 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/unary-op.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/unary-op.ts @@ -357,6 +357,18 @@ export const hardSigmoid = (context: ComputeContext, attributes: HardSigmoidAttr ); }; +export const hardSwish = (context: ComputeContext): void => { + const dataType = tensorTypeToWsglValueType(context.inputs[0].dataType); + context.compute( + createElementwiseProgramInfo( + context.inputs[0], + 'HardSwish', + (a) => + `${a} * max(vec4<${dataType}>(0.0), min(vec4<${dataType}>(1.0), vec4<${dataType}>(${dataType}(1.0 / 6.0)) * ${a} + vec4<${dataType}>(0.5)))`, + ), + ); +}; + export const sin = (context: ComputeContext): void => { context.compute(createElementwiseProgramInfo(context.inputs[0], 'Sin', 'sin')); }; diff --git a/onnxruntime/core/providers/js/js_execution_provider.cc b/onnxruntime/core/providers/js/js_execution_provider.cc index a5ced8f40201f..25ca711800c27 100644 --- a/onnxruntime/core/providers/js/js_execution_provider.cc +++ b/onnxruntime/core/providers/js/js_execution_provider.cc @@ -100,6 +100,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 13, Erf class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 6, 12, Sigmoid); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 13, Sigmoid); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 6, HardSigmoid); +class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 14, HardSwish); class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 6, 12, Log); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 13, Log); @@ -440,6 +441,7 @@ std::unique_ptr RegisterKernels() { KERNEL_CREATE_INFO_VERSIONED(6, 12, Sigmoid), KERNEL_CREATE_INFO(13, Sigmoid), KERNEL_CREATE_INFO(6, HardSigmoid), + KERNEL_CREATE_INFO(14, HardSwish), KERNEL_CREATE_INFO_VERSIONED(6, 12, Log), KERNEL_CREATE_INFO(13, Log), diff --git a/onnxruntime/core/providers/js/operators/unary.cc b/onnxruntime/core/providers/js/operators/unary.cc index 26144e6ba3995..922651139941d 100644 --- a/onnxruntime/core/providers/js/operators/unary.cc +++ b/onnxruntime/core/providers/js/operators/unary.cc @@ -80,6 +80,9 @@ JSEP_ELEMENTWISE_KERNEL(Sigmoid, 13, Sigmoid) JSEP_CLASS_IMPL_ATTRIBUTE_FLOAT_2_DEFAULT(HardSigmoid, HardSigmoid, alpha, 0.2, beta, 0.5) JSEP_ELEMENTWISE_KERNEL(HardSigmoid, 6, HardSigmoid) +JSEP_KERNEL_IMPL(HardSwish, HardSwish) +JSEP_ELEMENTWISE_KERNEL(HardSwish, 14, HardSwish) + JSEP_KERNEL_IMPL(Log, Log) JSEP_ELEMENTWISE_VERSIONED_KERNEL(Log, 6, 12, Log) JSEP_ELEMENTWISE_KERNEL(Log, 13, Log) diff --git a/onnxruntime/core/providers/webgpu/math/unary_elementwise_ops.cc b/onnxruntime/core/providers/webgpu/math/unary_elementwise_ops.cc index b8a3f923e81cc..d45fe46b3cc62 100644 --- a/onnxruntime/core/providers/webgpu/math/unary_elementwise_ops.cc +++ b/onnxruntime/core/providers/webgpu/math/unary_elementwise_ops.cc @@ -132,6 +132,9 @@ class HardSigmoid final : public UnaryElementwise { WEBGPU_ELEMENTWISE_KERNEL(HardSigmoid, 6, WebGpuSupportedFloatTypes()) +WEBGPU_ELEMENTWISE_IMPL(HardSwish, "hard_swish_v(a)", HardSwishImpl, ShaderUsage::UseElementTypeAlias) +WEBGPU_ELEMENTWISE_KERNEL(HardSwish, 14, WebGpuSupportedFloatTypes()) + WEBGPU_ELEMENTWISE_IMPL(Sin, "sin(a)") WEBGPU_ELEMENTWISE_KERNEL(Sin, 7, WebGpuSupportedFloatTypes()) diff --git a/onnxruntime/core/providers/webgpu/math/unary_elementwise_ops.h b/onnxruntime/core/providers/webgpu/math/unary_elementwise_ops.h index 39cce7941a15b..e5e0b4ae33a02 100644 --- a/onnxruntime/core/providers/webgpu/math/unary_elementwise_ops.h +++ b/onnxruntime/core/providers/webgpu/math/unary_elementwise_ops.h @@ -114,6 +114,15 @@ fn hard_sigmoid_v(v: vec4) -> vec4 { } )"; +constexpr const char HardSwishImpl[] = R"( +fn hard_swish_v(v: vec4) -> vec4 { + let alpha = x_element_t(1.0 / 6.0); + let beta_v = vec4(x_element_t(0.5)); + return v * max(vec4(0.0), + min(vec4(1.0), alpha * v + beta_v)); +} +)"; + // built-in function tanh() does not work with large input (f32 88.7 or f16 11.09) // https://github.com/gpuweb/gpuweb/issues/4458 constexpr const char TanhImpl[] = R"( diff --git a/onnxruntime/core/providers/webgpu/webgpu_execution_provider.cc b/onnxruntime/core/providers/webgpu/webgpu_execution_provider.cc index 1f6e4de6fbf32..7a5b56c8aa5f4 100644 --- a/onnxruntime/core/providers/webgpu/webgpu_execution_provider.cc +++ b/onnxruntime/core/providers/webgpu/webgpu_execution_provider.cc @@ -121,6 +121,7 @@ static const BuildKernelCreateInfoFn build_kernel_create_info_function_table[] = KERNEL_CREATE_INFO_VERSIONED(6, 12, Sigmoid), KERNEL_CREATE_INFO(13, Sigmoid), KERNEL_CREATE_INFO(6, HardSigmoid), + KERNEL_CREATE_INFO(14, HardSwish), KERNEL_CREATE_INFO_VERSIONED(6, 12, Log), KERNEL_CREATE_INFO(13, Log), diff --git a/onnxruntime/test/providers/webgpu/hardswish_test.cc b/onnxruntime/test/providers/webgpu/hardswish_test.cc new file mode 100644 index 0000000000000..dd64a62a1d654 --- /dev/null +++ b/onnxruntime/test/providers/webgpu/hardswish_test.cc @@ -0,0 +1,54 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include +#include +#include +#include + +#include "gtest/gtest.h" + +#include "default_providers.h" +#include "test/providers/provider_test_utils.h" + +namespace onnxruntime { +namespace test { + +template +void RunHardSwishTest() { + auto webgpu_ep = DefaultWebGpuExecutionProvider(); + if (!webgpu_ep) { + GTEST_SKIP() << "WebGPU execution provider is not available."; + } + + const std::vector kDims{2, 5}; + const std::vector input_values{-6.0f, -3.0f, -1.0f, 0.0f, 1.0f, 3.0f, 6.0f, 8.0f, -8.0f, 0.5f}; + std::vector expected_values; + expected_values.reserve(input_values.size()); + std::transform(input_values.cbegin(), input_values.cend(), std::back_inserter(expected_values), + [](float x) { return x * std::max(0.0f, std::min(1.0f, x / 6.0f + 0.5f)); }); + + OpTester test("HardSwish", 14); + if constexpr (std::is_same_v) { + test.AddInput("X", kDims, input_values); + test.AddOutput("Y", kDims, expected_values); + } else { + test.AddInput("X", kDims, FloatsToMLFloat16s(input_values)); + test.AddOutput("Y", kDims, FloatsToMLFloat16s(expected_values)); + test.SetOutputAbsErr("Y", 0.01f); + test.SetOutputRelErr("Y", 0.01f); + } + + test.ConfigEp(std::move(webgpu_ep)).RunWithConfig(); +} + +TEST(HardSwish_WebGPU, Float32) { + RunHardSwishTest(); +} + +TEST(HardSwish_WebGPU, Float16) { + RunHardSwishTest(); +} + +} // namespace test +} // namespace onnxruntime