onnxruntime/js/web/lib/wasm/jsep/webgpu/ops/argminmax.ts
Guenther Schmuelling 0df2e14038
js/webgpu: argmax,argmin,softmax support (#16882)
argmax and argmin are similar to reduce. Eventually we need to add
optimized flavors of the shader.

softmax is optimized but only works on the last axis for now which
should be the common use case.

todo: enable more ut for argmax/argmin
2023-08-02 18:16:19 -07:00

156 lines
6.6 KiB
TypeScript

// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
// TODO: this is the same naive implementation we use for reduce that has
// performance limitations when the reduced axis is long. Need to add
// a optimized codepath for this.
import {DataType} from '../../../wasm-common';
import {TensorView} from '../../tensor';
import {ShapeUtil} from '../../util';
import {AttributeWithCacheKey, createAttributeWithCacheKey} from '../attribute-with-cache-key';
import {ComputeContext, GpuDataType, ProgramInfo, ProgramInfoLoader, ProgramMetadata} from '../types';
import {createIndicesHelper, ShaderHelper} from './common';
const validateInputs = (inputs: readonly TensorView[]): void => {
if (!inputs || inputs.length === 0 || inputs.length > 2) {
throw new Error('ArgMinMaxOp op requires 1 or 2 inputs.');
}
if (inputs[0].dataType !== DataType.float) {
throw new Error('Invalid input type.');
}
};
export interface ArgMinMaxAttributes extends AttributeWithCacheKey {
keepDims: boolean;
axes: number;
selectLastIndex: number;
}
type ArgMinMaxOp = (inputs: readonly TensorView[], axes: number[]) => string[];
const createReduceProgramInfo =
(metadata: ProgramMetadata, inputs: readonly TensorView[], attributes: ArgMinMaxAttributes,
argMinMaxOp: ArgMinMaxOp): ProgramInfo => {
const outputShape: number[] = [];
const inputShape = inputs[0].dims;
const idxCopy: string[] = []; // copy output indexes to input indexes
const axes = ShapeUtil.normalizeAxes([attributes.axes], inputs[0].dims.length);
const outputDimsLength = inputs[0].dims.length - (attributes.keepDims ? 0 : axes.length);
const ops = argMinMaxOp(inputs, axes);
const inputIndicesHelper = createIndicesHelper('input', inputShape);
const initInputIdx = (ops[1] === '') ? '' : `let inputIdx = ${inputIndicesHelper.i2oExpression('inputIndices')};`;
let reduceOps = `
let inputIdx = ${inputIndicesHelper.i2oExpression('inputIndices')};
${ops[2]};`;
for (let k = 0; k < inputs[0].dims.length; k++) {
// if this axis is reduced
if (axes.indexOf(k) >= 0) {
if (attributes.keepDims) {
outputShape.push(1);
}
// loop over the d-th axis
reduceOps = `for(var j${k}: u32 = 0; j${k} < ${inputs[0].dims[k]}; j${k}++) {
let lastIndex = j${k};
inputIndices[${k}] = lastIndex;
${reduceOps}
}`;
} else {
if (outputDimsLength > 1) {
idxCopy.push(`inputIndices[${k}] = outputIndices[${outputShape.length}];`);
} else {
idxCopy.push(`inputIndices[${k}] = outputIndices;`);
}
outputShape.push(inputs[0].dims[k]);
}
}
const outputIndicesHelper = createIndicesHelper('output', outputShape);
const outputSize = ShapeUtil.size(outputShape);
const dataType = 'f32';
const getShaderSource = (shaderHelper: ShaderHelper) => `
@group(0) @binding(0) var<storage, read> _A : array<${dataType}>;
@group(0) @binding(1) var<storage, read_write> output : array<i32>;
${outputIndicesHelper.o2iImpl}
${inputIndicesHelper.i2oImpl}
${shaderHelper.mainStart()}
${shaderHelper.guardAgainstOutOfBoundsWorkgroupSizes(outputSize)}
${inputIndicesHelper.indicesVariableDeclaration('inputIndices')}
${outputIndicesHelper.indicesVariableDeclaration('outputIndices')}
${outputIndicesHelper.o2iCall('global_idx', 'outputIndices')}
${idxCopy.join('\n')}
${ops[0]} // init ops
${initInputIdx}
${ops[1]}
${reduceOps}
${ops[3]} // final values
output[global_idx*2] = bestIndex; // result it int64
}`;
return {
...metadata,
getShaderSource,
outputs: [{dims: outputShape, dataType: DataType.int64, gpuDataType: GpuDataType.default}],
dispatchGroup: () => ({x: Math.ceil(outputSize / 64)})
};
};
const createArgMinMaxAttributesFromInputs =
(inputs: readonly TensorView[], attributes: ArgMinMaxAttributes): ArgMinMaxAttributes =>
createAttributeWithCacheKey(
{axes: attributes.axes, keepDims: attributes.keepDims, selectLastIndex: attributes.selectLastIndex});
const createReduceProgramInfoLoader =
(inputs: readonly TensorView[], name: string, attributes: ArgMinMaxAttributes, reduceOp: ArgMinMaxOp):
ProgramInfoLoader => {
const updatedAttributes: ArgMinMaxAttributes =
inputs.length === 1 ? attributes : createArgMinMaxAttributesFromInputs(inputs, attributes);
const metadata:
ProgramMetadata = {name, inputTypes: [GpuDataType.default], cacheHint: updatedAttributes.cacheKey};
return {...metadata, get: () => createReduceProgramInfo(metadata, [inputs[0]], updatedAttributes, reduceOp)};
};
export const argMin = (context: ComputeContext, attributes: ArgMinMaxAttributes): void => {
validateInputs(context.inputs);
const argMinMaxOp: ArgMinMaxOp = (inputs: TensorView[], axes: number[]): string[] => {
const idxZero = [];
for (let k = 0; k < inputs[0].dims.length; k++) {
if (axes.indexOf(k) >= 0 || axes.length === 0) {
idxZero.push(`inputIndices[${k}] = 0;`); // first element
}
}
return [
`${idxZero.join('\n')}`, 'var value = _A[inputIdx];\nvar bestIndex : i32 = 0;',
'if (_A[inputIdx] < value) {value = _A[inputIdx]; bestIndex = i32(lastIndex);} ', ''
];
};
context.compute(createReduceProgramInfoLoader(context.inputs, 'ArgMin', attributes, argMinMaxOp), {inputs: [0]});
};
export const argMax = (context: ComputeContext, attributes: ArgMinMaxAttributes): void => {
validateInputs(context.inputs);
const argMinMaxOp: ArgMinMaxOp = (inputs: TensorView[], axes: number[]): string[] => {
const idxZero = [];
for (let k = 0; k < inputs[0].dims.length; k++) {
if (axes.indexOf(k) >= 0 || axes.length === 0) {
idxZero.push(`inputIndices[${k}] = 0;`); // first element
}
}
return [
`${idxZero.join('\n')}`, 'var value = _A[inputIdx];\nvar bestIndex : i32 = 0;',
'if (_A[inputIdx] > value) {value = _A[inputIdx]; bestIndex = i32(lastIndex);} ', ''
];
};
context.compute(createReduceProgramInfoLoader(context.inputs, 'argMax', attributes, argMinMaxOp), {inputs: [0]});
};
export const parseArgMinMaxAttributes = (attributes: Record<string, unknown>): ArgMinMaxAttributes =>
createAttributeWithCacheKey(attributes as Omit<ArgMinMaxAttributes, keyof AttributeWithCacheKey>);