From 28c23aed041010726eea9ef7ed58528e98701b8c Mon Sep 17 00:00:00 2001 From: Jiajia Qin Date: Sat, 11 Nov 2023 04:54:35 +0800 Subject: [PATCH] [js/webgpu] Fix conv2d with activation (#18388) ### Description Fix #18297 With PR #17766, conv2d activation in mobilenetv2-12 will not be empty. However, activation is not supported yet in [biasActivationSnippet](https://github.com/microsoft/onnxruntime/blob/main/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/activation_util.ts#L48C14-L48C36). This PR makes all places unify to use [getActivationSnippet](https://github.com/microsoft/onnxruntime/blob/main/js/web/lib/wasm/jsep/webgpu/ops/fuse-utils.ts#L13) to fix this issue. --- .../webgpu/ops/3rd-party/activation_util.ts | 15 +------ .../webgpu/ops/3rd-party/conv2d_mm_webgpu.ts | 16 ++++--- .../ops/3rd-party/conv_backprop_mm_webgpu.ts | 17 ++++---- .../ops/3rd-party/matmul_packed_webgpu.ts | 2 +- .../lib/wasm/jsep/webgpu/ops/conv-grouped.ts | 3 +- js/web/lib/wasm/jsep/webgpu/ops/fuse-utils.ts | 43 +++++++++---------- 6 files changed, 41 insertions(+), 55 deletions(-) diff --git a/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/activation_util.ts b/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/activation_util.ts index 6481a6b21d..a121bf3892 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/activation_util.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/activation_util.ts @@ -19,8 +19,6 @@ // // modified to fit the needs of the project -export declare type Activation = 'linear' | 'relu' | 'prelu' | 'elu' | 'relu6' | 'leakyrelu' | 'sigmoid' | 'gelu'; - export const typeSnippet = (component: number, dataType: string) => { switch (component) { case 1: @@ -36,17 +34,6 @@ export const typeSnippet = (component: number, dataType: string) => { } }; -export const activationFnSnippet = - (activation?: Activation, _hasPreluActivationWeights = false, _packed = false, _coordsLength = 3): string => { - if (!activation) { - return ''; - } - // TODO: add implementations - return ''; - }; - -export const biasActivationSnippet = (hasBias: boolean, activation?: Activation): string => ` +export const biasSnippet = (hasBias: boolean): string => ` ${hasBias ? 'value = value + getBiasByOutputCoords(coords);' : ''} - // TODO uncomment the following line when activation is supported above. - // ${activation ? 'value = activation(value, coords);' : ''} `; diff --git a/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/conv2d_mm_webgpu.ts b/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/conv2d_mm_webgpu.ts index fbb936a045..089e783d7e 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/conv2d_mm_webgpu.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/conv2d_mm_webgpu.ts @@ -25,15 +25,16 @@ import {ShapeUtil} from '../../../util'; import {ProgramInfo} from '../../types'; import {tensorTypeToWsglStorageType} from '../common'; import {ConvAttributes} from '../conv'; +import {getActivationSnippet} from '../fuse-utils'; -import {Activation, activationFnSnippet, biasActivationSnippet, typeSnippet} from './activation_util'; +import {biasSnippet, typeSnippet} from './activation_util'; import {utilFunctions} from './conv_util'; import {makeMatMulPackedSource, makeMatMulPackedVec4Source} from './matmul_packed_webgpu'; const conv2dCommonSnippet = (isChannelsLast: boolean, fitAOuter: boolean, fitBOuter: boolean, fitInner: boolean, addBias = false, - activation?: Activation, hasPreluActivationWeights = false, innerElementSizeX = 4, innerElementSizeW = 4, - innerElementSize = 4, dataType = 'f32'): string => { + attributes: ConvAttributes, innerElementSizeX = 4, innerElementSizeW = 4, innerElementSize = 4, + dataType = 'f32'): string => { const getXSnippet = (innerElementSize: number) => { switch (innerElementSize) { case 1: @@ -129,8 +130,9 @@ const conv2dCommonSnippet = isChannelsLast ? typeSnippet(innerElementSizeX, dataType) : typeSnippet(innerElementSizeW, dataType); const bType = isChannelsLast ? typeSnippet(innerElementSizeW, dataType) : typeSnippet(innerElementSizeX, dataType); + const {activationFunction, applyActivation} = getActivationSnippet(attributes, resType); const userCode = ` - ${activationFnSnippet(activation, hasPreluActivationWeights, innerElementSize === 4, 4)} + ${activationFunction} fn mm_readA(batch: i32, row : i32, colIn : i32) -> ${aType} { ${isChannelsLast ? sampleX : sampleW} } @@ -146,7 +148,8 @@ const conv2dCommonSnippet = var value = valueIn; let outWidth = ${isChannelsLast ? 'outShape[2]' : 'outShape[3]'}; ${coordResSnippet} - ${biasActivationSnippet(addBias, activation)} + ${biasSnippet(addBias)} + ${applyActivation} setOutputAtCoords(coords[0], coords[1], coords[2], coords[3], value); } }`; @@ -242,8 +245,7 @@ export const createConv2DMatMulProgramInfo = ${declareFunctions} ${ conv2dCommonSnippet( - isChannelsLast, fitAOuter, fitBOuter, fitInner, hasBias, - attributes.activation.toLowerCase() as Activation, false, elementsSize[0], elementsSize[1], + isChannelsLast, fitAOuter, fitBOuter, fitInner, hasBias, attributes, elementsSize[0], elementsSize[1], elementsSize[2], t)} ${ isVec4 ? diff --git a/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/conv_backprop_mm_webgpu.ts b/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/conv_backprop_mm_webgpu.ts index a95d3830f3..85cf7bf87f 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/conv_backprop_mm_webgpu.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/conv_backprop_mm_webgpu.ts @@ -24,14 +24,14 @@ import {TensorView} from '../../../tensor-view'; import {ShapeUtil} from '../../../util'; import {ProgramInfo} from '../../types'; import {ConvTransposeAttributes} from '../conv-transpose'; +import {getActivationSnippet} from '../fuse-utils'; -import {Activation, activationFnSnippet, biasActivationSnippet, typeSnippet} from './activation_util'; +import {biasSnippet, typeSnippet} from './activation_util'; import {utilFunctions} from './conv_util'; import {makeMatMulPackedSource, makeMatMulPackedVec4Source} from './matmul_packed_webgpu'; const conv2dTransposeCommonSnippet = - (isChannelsLast: boolean, addBias = false, activation?: Activation, hasPreluActivationWeights = false, - innerElementSize = 4): string => { + (isChannelsLast: boolean, addBias = false, attributes: ConvTransposeAttributes, innerElementSize = 4): string => { const type = typeSnippet(innerElementSize, 'f32'); const getWSnippet = (innerElementSize: number) => { switch (innerElementSize) { @@ -129,9 +129,9 @@ const conv2dTransposeCommonSnippet = return ${type}(0.0); `; - + const {activationFunction, applyActivation} = getActivationSnippet(attributes, type); const userCode = ` - ${activationFnSnippet(activation, hasPreluActivationWeights, innerElementSize === 4, 4)} + ${activationFunction} fn mm_readA(batch: i32, row : i32, colIn : i32) -> ${type} { ${isChannelsLast ? sampleA : sampleW} } @@ -146,7 +146,8 @@ const conv2dTransposeCommonSnippet = var value = valueInput; let outWidth = ${isChannelsLast ? 'outShape[2]' : 'outShape[3]'}; ${coordResSnippet} - ${biasActivationSnippet(addBias, activation)} + ${biasSnippet(addBias)} + ${applyActivation} result[getIndexFromCoords4D(coords, outShape)/${innerElementSize}] = value; } }`; @@ -236,9 +237,7 @@ export const createConv2DTransposeMatMulProgramInfo = const dimBOuter : i32 = ${dimBOuter}; const dimInner : i32 = ${dimInner}; ${declareFunctions} - ${ - conv2dTransposeCommonSnippet( - isChannelsLast, hasBias, attributes.activation.toLowerCase() as Activation, false, innerElementSize)} + ${conv2dTransposeCommonSnippet(isChannelsLast, hasBias, attributes, innerElementSize)} ${ isVec4 ? makeMatMulPackedVec4Source( elementsPerThread, workGroupSize, 'f32', undefined, !isChannelsLast, tileInner) : diff --git a/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/matmul_packed_webgpu.ts b/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/matmul_packed_webgpu.ts index 0a0f29db6a..335de01c59 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/matmul_packed_webgpu.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/3rd-party/matmul_packed_webgpu.ts @@ -440,7 +440,6 @@ export const createMatmulProgramInfo = const dimInner = aShape[aShape.length - 1]; const dimBOuter = bShape[bShape.length - 1]; const isVec4 = dimInner % 4 === 0 && dimBOuter % 4 === 0; - const {activationFunction, applyActivation} = getActivationSnippet(activationAttributes, isVec4); // TODO: fine tune size const elementsPerThread = dimAOuter <= 8 ? [4, 1, 1] : [4, 4, 1]; @@ -462,6 +461,7 @@ export const createMatmulProgramInfo = variables.push(output); const inputVariables = [A, B]; const hasBias = inputs.length > 2; + const {activationFunction, applyActivation} = getActivationSnippet(activationAttributes, output.type.value); const declareFunctions = matMulReadWriteFnSource(components, hasBias, applyActivation, variables, batchShapes, isChannelsLast); if (hasBias) { diff --git a/js/web/lib/wasm/jsep/webgpu/ops/conv-grouped.ts b/js/web/lib/wasm/jsep/webgpu/ops/conv-grouped.ts index 8bfa722dd0..14482272ba 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/conv-grouped.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/conv-grouped.ts @@ -22,14 +22,13 @@ export const createGroupedConvProgramInfo = const wShape = inputs[1].dims; const outputChannelsPerGroup = wShape[0] / attributes.group; - const {activationFunction, applyActivation} = getActivationSnippet(attributes); - const isChannelLast = attributes.format === 'NHWC'; const outputShape = calculateOutputShape( xShape, wShape, attributes.dilations, attributes.pads, attributes.strides, isChannelLast); const outputSize = ShapeUtil.size(outputShape); const output = outputVariable('output', inputs[0].dataType, outputShape); + const {activationFunction, applyActivation} = getActivationSnippet(attributes, output.type.value); const x = inputVariable('x', inputs[0].dataType, xShape); const w = inputVariable('w', inputs[1].dataType, wShape); const inputVars = [x, w]; diff --git a/js/web/lib/wasm/jsep/webgpu/ops/fuse-utils.ts b/js/web/lib/wasm/jsep/webgpu/ops/fuse-utils.ts index a920875be1..0b5c0db2b5 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/fuse-utils.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/fuse-utils.ts @@ -10,28 +10,27 @@ export interface InternalActivationAttributes { readonly activationCacheKey: string; } -export const getActivationSnippet = (attributes: InternalActivationAttributes, isVec4 = false): { - activationFunction: string; applyActivation: string; -} => { - switch (attributes.activation) { - case 'Relu': - return { - activationFunction: '', - applyActivation: isVec4 ? 'value = max(value, vec4(0.0));' : 'value = max(value, 0.0);' - }; - case 'Sigmoid': - return {activationFunction: '', applyActivation: 'value = (1.0 / (1.0 + exp(-value)));'}; - case 'Clip': - return { - activationFunction: `const clip_min_=f32(${attributes.clipMin!});const clip_max_=f32(${attributes.clipMax!});`, - applyActivation: isVec4 ? 'value = clamp(value, vec4(clip_min_), vec4(clip_max_));' : - 'value = clamp(value, clip_min_, clip_max_);' - }; - // TODO: adding other activations that can be fused. - default: - return {activationFunction: '', applyActivation: ''}; - } -}; +export const getActivationSnippet = (attributes: InternalActivationAttributes, valueType: string): + {activationFunction: string; applyActivation: string} => { + switch (attributes.activation) { + case 'Relu': + return {activationFunction: '', applyActivation: `value = max(value, ${valueType}(0.0));`}; + case 'Sigmoid': + return { + activationFunction: '', + applyActivation: `value = (${valueType}(1.0) / (${valueType}(1.0) + exp(-value)));` + }; + case 'Clip': + return { + activationFunction: `const clip_min_=${valueType}(${attributes.clipMin!});const clip_max_=${valueType}(${ + attributes.clipMax!});`, + applyActivation: 'value = clamp(value, clip_min_, clip_max_);' + }; + // TODO: adding other activations that can be fused. + default: + return {activationFunction: '', applyActivation: ''}; + } + }; export const parseInternalActivationAttributes = (attributes: Record|undefined): InternalActivationAttributes => {