diff --git a/js/web/lib/wasm/jsep/webgpu/ops/instance-norm.ts b/js/web/lib/wasm/jsep/webgpu/ops/instance-norm.ts index 7b6140f3b1..859bd85086 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/instance-norm.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/instance-norm.ts @@ -4,18 +4,17 @@ import { DataType } from '../../../wasm-common'; import { TensorView } from '../../tensor-view'; import { ShapeUtil } from '../../util'; -import { ComputeContext, ProgramInfo, ProgramInputTensorInfoDependency, ProgramUniform } from '../types'; +import { ComputeContext, ProgramInputTensorInfoDependency, ProgramUniform } from '../types'; +import { createTransposeProgramInfo } from './transpose'; import { createTensorShapeVariables, - fillVector, getMaxComponents, inputVariable, outputVariable, ShaderHelper, sumVector, tensorTypeToWsglStorageType, - UniformsArrayType, } from './common'; export interface InstanceNormAttributes { @@ -23,117 +22,7 @@ export interface InstanceNormAttributes { format: 'NHWC' | 'NCHW'; } -const createInstanceNormProgramInfo = ( - inputs: readonly TensorView[], - attributes: InstanceNormAttributes, -): ProgramInfo => { - const xShape = inputs[0].dims; - const outputShape = xShape; - const axis = 2; - const normCount = ShapeUtil.sizeToDimension(xShape, axis); - const normSize = ShapeUtil.sizeFromDimension(xShape, axis); - const components = getMaxComponents(normSize); - const normPackedSize = normSize / components; - const inputShape = [xShape[0], xShape[1], normPackedSize]; - const inputDependencies: ProgramInputTensorInfoDependency[] = ['rank', 'type', 'type']; - const programUniforms: ProgramUniform[] = [ - { type: DataType.uint32, data: normSize }, - { type: DataType.uint32, data: normPackedSize }, - ]; - programUniforms.push(...createTensorShapeVariables(inputShape, inputShape)); - - const getShaderSource = (shaderHelper: ShaderHelper) => { - const x = inputVariable('x', inputs[0].dataType, inputShape.length, components); - const scale = inputVariable('scale', inputs[1].dataType, inputs[1].dims); - const bias = inputVariable('bias', inputs[2].dataType, inputs[2].dims); - const output = outputVariable('output', inputs[0].dataType, inputShape.length, components); - const variables = [x, scale, bias, output]; - const dataType = x.type.value; - const f32Type = components === 1 ? 'f32' : `vec${components}`; - const workgroupSize = 64; - - const uniforms: UniformsArrayType = [ - { name: 'normSize', type: 'u32' }, - { name: 'normPackedSize', type: 'u32' }, - ]; - return ` - var meanShared : f32; - var squaredNormShared : f32; - var workgroupShared : array<${f32Type}, ${workgroupSize}>; - const workgroupSize = ${workgroupSize}u; - ${shaderHelper.registerUniforms(uniforms).declareVariables(...variables)} - ${shaderHelper.mainStart(workgroupSize)} - let norm = global_idx / workgroupSize; - let batch = norm / uniforms.x_shape[1]; - let channel = norm % uniforms.x_shape[1]; - let localIndex = local_id.x; - - // initialize workgroup memory - var initial = ${f32Type}(0); - for (var h = localIndex; h < uniforms.normPackedSize; h += workgroupSize) { - initial = initial + ${f32Type}(${x.get('batch', 'channel', 'h')}); - } - workgroupShared[localIndex] = initial; - workgroupBarrier(); - - // Calculate the mean of current channel data. - for (var currSize = workgroupSize >> 1; currSize > 0; currSize = currSize >> 1) { - if (localIndex < currSize) { - workgroupShared[localIndex] = workgroupShared[localIndex] + workgroupShared[localIndex + currSize]; - } - workgroupBarrier(); - } - if (localIndex == 0) { - meanShared = ${sumVector('workgroupShared[0]', components)} / f32(uniforms.normSize); - } - workgroupBarrier(); - - // reinitialize workgroup memory. - initial = ${f32Type}(0); - for (var h = localIndex; h < uniforms.normPackedSize; h += workgroupSize) { - let deviation = ${f32Type}(${x.get('batch', 'channel', 'h')}) - ${f32Type}(meanShared); - initial = initial + deviation * deviation; - } - workgroupShared[localIndex] = initial; - workgroupBarrier(); - - // Calculate the sum of square of deviation of current channel data. - for (var currSize = workgroupSize >> 1; currSize > 0; currSize = currSize >> 1) { - if (localIndex < currSize) { - workgroupShared[localIndex] = workgroupShared[localIndex] + workgroupShared[localIndex + currSize]; - } - workgroupBarrier(); - } - if (localIndex == 0) { - squaredNormShared = ${sumVector('workgroupShared[0]', components)}; - } - workgroupBarrier(); - - let invStdDev = inverseSqrt(squaredNormShared / f32(uniforms.normSize) + f32(${attributes.epsilon})); - let channelScale = invStdDev * f32(${scale.getByOffset('channel')}); - let channelShift = f32(${bias.getByOffset('channel')}) - meanShared * channelScale; - for (var h = localIndex; h < uniforms.normPackedSize; h += workgroupSize) { - let value = ${x.get('batch', 'channel', 'h')} * ${dataType}(${f32Type}(channelScale)) + ${dataType}(${ - f32Type - }(channelShift)); - ${output.set('batch', 'channel', 'h', 'value')}; - } - }`; - }; - return { - ...{ name: 'InstanceNormalization' }, - // TODO: use epsilon as uniform. Currently epsilon as uniform fails test_instancenorm_epsilon. - shaderCache: { hint: `${attributes.epsilon};${components}`, inputDependencies }, - getRunData: () => ({ - outputs: [{ dims: outputShape, dataType: inputs[0].dataType }], - dispatchGroup: { x: normCount }, - programUniforms, - }), - getShaderSource, - }; -}; - -const computeMean = ( +const computeChannelScaleShift = ( context: ComputeContext, input: TensorView, scale: TensorView, @@ -143,123 +32,142 @@ const computeMean = ( c: number, epsilon: number, ) => { - const components = getMaxComponents(c); - const WG = 64; - // we will store channel scale and channel shift in [2, components] matrix - // or in vec2 when components == 1 - const outputType = components === 1 ? 'vec2f' : `mat2x${components}f`; - const sumCastType = components === 1 ? 'f32' : `vec${components}f`; - const setOutputValue = (var1: string, var2: string) => `${outputType}(${var1}, ${var2})`; - const unitsOfWork = (n * c) / components; - const wgSize = Math.ceil(h / WG); + const components = getMaxComponents(h); + const f32Type = components === 1 ? 'f32' : `vec${components}f`; + const wgType = components === 1 ? 'vec2f' : `mat2x${components}f`; + const unitsOfWork = n * c; - const meanInputDependencies: ProgramInputTensorInfoDependency[] = ['type']; - const meanProgramUniforms: ProgramUniform[] = [ - { type: DataType.uint32, data: wgSize }, - { type: DataType.uint32, data: h }, - { type: DataType.uint32, data: Math.floor(c / components) }, - { type: DataType.uint32, data: Math.floor((h * c) / components) }, - ]; + const inputShape = [n, c, h / components]; + const outputShape = [n, c, 2]; + const inputDependencies: ProgramInputTensorInfoDependency[] = ['rank', 'type', 'type']; + const programUniforms: ProgramUniform[] = []; + programUniforms.push(...createTensorShapeVariables(inputShape, outputShape)); - const getMeanShaderSource = (shaderHelper: ShaderHelper) => { - const inputHelper = inputVariable('input', input.dataType, input.dims, components); - return ` - ${shaderHelper.declareVariables(inputHelper)} - @group(0) @binding(1) var output : array<${outputType}>; - struct Uniforms {wg_size:u32, H:u32, C:u32, image_size:u32}; - @group(0) @binding(2) var uniforms: Uniforms; - - ${shaderHelper.mainStart(WG)} - let currentImageNumber = global_idx / ${WG} / uniforms.C; - let currentChannelNumber = (global_idx / ${WG}) % uniforms.C; - let wgOffset = local_id.x * uniforms.wg_size; - if (wgOffset >= uniforms.H) { - return; - } - let wgMax = min(wgOffset + uniforms.wg_size, uniforms.H); - - let offset = currentImageNumber * uniforms.image_size + currentChannelNumber; - var sum = ${fillVector('f32', components)}; - var squaredSum = ${fillVector('f32', components)}; - for (var i: u32 = wgOffset; i < wgMax; i++) { - let value = ${sumCastType}(input[offset + i * uniforms.C]); - sum += value; - squaredSum += value * value; - } - output[global_idx] = ${setOutputValue('sum', 'squaredSum')}; - }`; - }; - - const meanValues = context.compute( - { - name: 'InstanceNormComputeMean', - shaderCache: { hint: `${components}`, inputDependencies: meanInputDependencies }, - getRunData: () => ({ - outputs: [{ dims: [n, c, WG, 2], dataType: DataType.float }], - dispatchGroup: { x: (n * c) / components }, - programUniforms: meanProgramUniforms, - }), - getShaderSource: getMeanShaderSource, - }, - { inputs: [input], outputs: [-1] }, - )[0]; - - const programUniforms: ProgramUniform[] = [ - { type: DataType.uint32, data: unitsOfWork }, - { type: DataType.uint32, data: h }, - { type: DataType.uint32, data: Math.floor(c / components) }, - { type: DataType.uint32, data: Math.floor((WG * c) / components) }, - ]; - const inputDependencies: ProgramInputTensorInfoDependency[] = ['type', 'type', 'type']; const getShaderSource = (shaderHelper: ShaderHelper) => { - const scaleHelper = inputVariable('scale', scale.dataType, scale.dims, components); - const biasHelper = inputVariable('bias', bias.dataType, bias.dims, components); + const x = inputVariable('x', input.dataType, 3, components); + const s = inputVariable('scale', scale.dataType, scale.dims); + const b = inputVariable('bias', bias.dataType, bias.dims); + const output = outputVariable('output', DataType.float, 3, 2); + const variables = [x, s, b, output]; + const workgroupSize = 64; return ` - @group(0) @binding(0) var input : array<${outputType}>; - @group(0) @binding(1) var scale : array<${scaleHelper.type.storage}>; - @group(0) @binding(2) var bias : array<${biasHelper.type.storage}>; - @group(0) @binding(3) var output : array<${outputType}>; - struct Uniforms {units_of_work : u32, H: u32, C : u32, image_size : u32}; - @group(0) @binding(4) var uniforms: Uniforms; - - ${shaderHelper.mainStart()} - ${shaderHelper.guardAgainstOutOfBoundsWorkgroupSizes('uniforms.units_of_work')} - let currentImageNumber = global_idx / uniforms.C; - let currentChannelNumber = global_idx % uniforms.C; - - let offset = currentImageNumber * uniforms.image_size; - var sum = ${fillVector('f32', components)}; - var squaredSum = ${fillVector('f32', components)}; - for (var i: u32 = 0; i < min(${WG}, uniforms.H); i++) { - let value = input[offset + i + currentChannelNumber * ${WG}]; - sum += value[0]; - squaredSum += value[1]; + var workgroup_shared : array<${wgType}, ${workgroupSize}>; + const workgroup_size = ${workgroupSize}u; + ${shaderHelper.declareVariables(...variables)} + ${shaderHelper.mainStart(workgroupSize)} + let batch = workgroup_index / uniforms.x_shape[1]; + let channel = workgroup_index % uniforms.x_shape[1]; + let hight = uniforms.x_shape[2]; + // initialize workgroup memory + var sum = ${f32Type}(0); + var squared_sum = ${f32Type}(0); + for (var h = local_idx; h < hight; h += workgroup_size) { + let value = ${f32Type}(${x.get('batch', 'channel', 'h')}); + sum += value; + squared_sum += value * value; } - sum = sum / f32(uniforms.H); - squaredSum = squaredSum / f32(uniforms.H); - let invStdDev = inverseSqrt(squaredSum - sum * sum + f32(${epsilon})); - let channelScale = invStdDev * ${sumCastType}(scale[currentChannelNumber]); - let channelShift = ${sumCastType}(bias[currentChannelNumber]) - sum * channelScale; + workgroup_shared[local_idx] = ${wgType}(sum, squared_sum); + workgroupBarrier(); - output[global_idx] = ${setOutputValue('channelScale', 'channelShift')}; + for (var currSize = workgroup_size >> 1; currSize > 0; currSize = currSize >> 1) { + if (local_idx < currSize) { + workgroup_shared[local_idx] = workgroup_shared[local_idx] + workgroup_shared[local_idx + currSize]; + } + workgroupBarrier(); + } + if (local_idx == 0) { + let sum_final = ${sumVector('workgroup_shared[0][0]', components)} / f32(hight * ${components}); + let squared_sum_final = ${sumVector('workgroup_shared[0][1]', components)} / f32(hight * ${components}); + + let inv_std_dev = inverseSqrt(squared_sum_final - sum_final * sum_final + f32(${epsilon})); + let channel_scale = inv_std_dev * f32(scale[channel]); + let channel_shift = f32(bias[channel]) - sum_final * channel_scale; + output[workgroup_index] = vec2f(channel_scale, channel_shift); + } }`; }; + return context.compute( { name: 'InstanceNormComputeChannelScaleShift', // TODO: use epsilon as uniform. Currently epsilon as uniform fails test_instancenorm_epsilon. shaderCache: { hint: `${components};${epsilon}`, inputDependencies }, getRunData: () => ({ - outputs: [{ dims: [n, c, 2], dataType: DataType.float }], - dispatchGroup: { x: Math.ceil(unitsOfWork / 64 /* workgroup size */) }, + outputs: [{ dims: outputShape, dataType: DataType.float }], + dispatchGroup: { x: unitsOfWork }, programUniforms, }), getShaderSource, }, - { inputs: [meanValues, scale, bias], outputs: [-1] }, + { inputs: [input, scale, bias], outputs: [-1] }, )[0]; }; +const createInstanceNormProgramInfo = ( + context: ComputeContext, + inputs: readonly TensorView[], + attributes: InstanceNormAttributes, +) => { + const xShape = inputs[0].dims; + const outputShape = xShape; + const axis = 2; + const N = xShape[0]; + const C = xShape[1]; + const H = ShapeUtil.sizeFromDimension(xShape, axis); + const components = getMaxComponents(H); + const outputSize = ShapeUtil.size(outputShape) / components; + // compute channel scale and channel shift. + const channelScaleShift = computeChannelScaleShift( + context, + inputs[0], + inputs[1], + inputs[2], + N, + H, + C, + attributes.epsilon, + ); + + const inputShape = [N, C, H / components]; + const scaleShape = [N, C]; + const inputDependencies: ProgramInputTensorInfoDependency[] = ['type', 'none']; + + const getShaderSource = (shaderHelper: ShaderHelper) => { + const x = inputVariable('x', inputs[0].dataType, inputShape.length, components); + const scale = inputVariable('scale_shift', DataType.float, scaleShape.length, 2); + const output = outputVariable('output', inputs[0].dataType, inputShape.length, components); + const variables = [x, scale, output]; + return ` + ${shaderHelper.registerUniform('output_size', 'u32').declareVariables(...variables)} + ${shaderHelper.mainStart()} + ${shaderHelper.guardAgainstOutOfBoundsWorkgroupSizes('uniforms.output_size')} + let outputIndices = ${output.offsetToIndices('global_idx')}; + let batch = outputIndices[0]; + let channel = outputIndices[1]; + let scale_shift = ${scale.getByIndices('vec2(batch, channel)')}; + let value = ${x.getByOffset('global_idx')} * ${output.type.value}(scale_shift.x) + ${output.type.value}(scale_shift.y); + ${output.setByOffset('global_idx', 'value')}; + }`; + }; + + context.compute( + { + name: 'InstanceNormalization', + shaderCache: { hint: `${components}`, inputDependencies }, + getRunData: () => ({ + outputs: [{ dims: outputShape, dataType: inputs[0].dataType }], + dispatchGroup: { x: Math.ceil(outputSize / 64 /* workgroup size */) }, + programUniforms: [ + { type: DataType.uint32, data: outputSize }, + ...createTensorShapeVariables(inputShape, scaleShape, inputShape), + ], + }), + getShaderSource, + }, + { inputs: [inputs[0], channelScaleShift] }, + ); +}; + const createInstanceNormNHWCProgramInfo = ( context: ComputeContext, inputs: readonly TensorView[], @@ -277,30 +185,61 @@ const createInstanceNormNHWCProgramInfo = ( { type: DataType.uint32, data: Math.floor(C / components) }, ]; const inputDependencies: ProgramInputTensorInfoDependency[] = ['type', 'type']; - // first compute mean - const channelScaleShift = computeMean(context, inputs[0], inputs[1], inputs[2], N, H, C, attributes.epsilon); + + // 1. transpose x from NHWC to NCHW + const transposedXPerm = [0, xShape.length - 1]; + for (let i = 0; i < xShape.length - 2; i++) { + transposedXPerm.push(i + 1); + } + const transposedX = context.compute(createTransposeProgramInfo(context.inputs[0], transposedXPerm), { + inputs: [context.inputs[0]], + outputs: [-1], + })[0]; + // 2. compute channel scale and channel shift. + const channelScaleShift = computeChannelScaleShift( + context, + transposedX, + inputs[1], + inputs[2], + N, + H, + C, + attributes.epsilon, + ); const getShaderSource = (shaderHelper: ShaderHelper) => { const dataType = tensorTypeToWsglStorageType(inputs[0].dataType); - const scaleType = components === 1 ? 'vec2f' : `mat2x${components}f`; - const scaleCastType = components === 1 ? dataType : `vec${components}<${dataType}>`; - + const scaleType = components === 1 ? 'vec2f' : `mat${components}x2f`; + const scaleData = (num: number) => { + const index = num === 0 ? 'x' : 'y'; + const f32Type = components === 1 ? 'f32' : `vec${components}f`; + switch (components) { + case 1: + return `${dataType}(${f32Type}(scale.${index}))`; + case 2: + return `vec2<${dataType}>(${f32Type}(scale[0].${index}, scale[1].${index}))`; + case 4: + return `vec4<${dataType}>(${f32Type}(scale[0].${index}, scale[1].${index}, scale[2].${index}, scale[3].${index}))`; + default: + throw new Error(`Not supported compoents ${components}`); + } + }; const inputHelper = inputVariable('input', inputs[0].dataType, inputs[0].dims, components); const outputHelper = outputVariable('output', inputs[0].dataType, outputShape, components); return ` @group(0) @binding(0) var input : array<${inputHelper.type.storage}>; - @group(0) @binding(1) var scaleInput : array<${scaleType}>; + @group(0) @binding(1) var scale_input : array<${scaleType}>; @group(0) @binding(2) var output : array<${outputHelper.type.storage}>; struct Uniforms {H: u32, C : u32}; @group(0) @binding(3) var uniforms: Uniforms; ${shaderHelper.mainStart()} - let currentImageNumber = global_idx / (uniforms.C * uniforms.H); - let currentChannelNumber = global_idx % uniforms.C; + let current_image_number = global_idx / (uniforms.C * uniforms.H); + let current_channel_number = global_idx % uniforms.C; - let scaleOffset = currentImageNumber * uniforms.C + currentChannelNumber; - let scale = scaleInput[scaleOffset]; - output[global_idx] = fma(input[global_idx], ${scaleCastType}(scale[0]), ${scaleCastType}(scale[1])); + let scale_offset = current_image_number * uniforms.C + current_channel_number; + let scale = scale_input[scale_offset]; + output[global_idx] = fma(input[global_idx], ${scaleData(0)}, ${scaleData(1)}); }`; }; context.compute( @@ -322,6 +261,6 @@ export const instanceNorm = (context: ComputeContext, attributes: InstanceNormAt if (attributes.format === 'NHWC') { createInstanceNormNHWCProgramInfo(context, context.inputs, attributes); } else { - context.compute(createInstanceNormProgramInfo(context.inputs, attributes)); + createInstanceNormProgramInfo(context, context.inputs, attributes); } };