mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
[js/web] FP16 Conv, ConvTranspose and MatMul (#17514)
### Description Another three ops for fp16 --------- Co-authored-by: Guenther Schmuelling <guschmue@microsoft.com> Co-authored-by: Yulong Wang <7679871+fs-eire@users.noreply.github.com>
This commit is contained in:
parent
9aad78721c
commit
a941dd583e
14 changed files with 148 additions and 174 deletions
|
|
@ -21,16 +21,16 @@
|
|||
|
||||
export declare type Activation = 'linear' | 'relu' | 'prelu' | 'elu' | 'relu6' | 'leakyrelu' | 'sigmoid' | 'gelu';
|
||||
|
||||
export const typeSnippet = (component: number) => {
|
||||
export const typeSnippet = (component: number, dataType: string) => {
|
||||
switch (component) {
|
||||
case 1:
|
||||
return 'f32';
|
||||
return dataType;
|
||||
case 2:
|
||||
return 'vec2<f32>';
|
||||
return `vec2<${dataType}>`;
|
||||
case 3:
|
||||
return 'vec3<f32>';
|
||||
return `vec3<${dataType}>`;
|
||||
case 4:
|
||||
return 'vec4<f32>';
|
||||
return `vec4<${dataType}>`;
|
||||
default:
|
||||
throw new Error(`${component}-component is not supported.`);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ import {LOG_DEBUG} from '../../../log';
|
|||
import {TensorView} from '../../../tensor-view';
|
||||
import {ShapeUtil} from '../../../util';
|
||||
import {GpuDataType, ProgramInfo, ProgramMetadata} from '../../types';
|
||||
import {tensorTypeToWsglStorageType} from '../common';
|
||||
import {ConvAttributes} from '../conv';
|
||||
|
||||
import {Activation, activationFnSnippet, biasActivationSnippet, typeSnippet} from './activation_util';
|
||||
|
|
@ -32,13 +33,13 @@ import {makeMatMulPackedSource, makeMatMulPackedVec4Source} from './matmul_packe
|
|||
const conv2dCommonSnippet =
|
||||
(isChannelsLast: boolean, fitAOuter: boolean, fitBOuter: boolean, fitInner: boolean, addBias = false,
|
||||
activation?: Activation, hasPreluActivationWeights = false, innerElementSizeX = 4, innerElementSizeW = 4,
|
||||
innerElementSize = 4): string => {
|
||||
innerElementSize = 4, dataType = 'f32'): string => {
|
||||
const getXSnippet = (innerElementSize: number) => {
|
||||
switch (innerElementSize) {
|
||||
case 1:
|
||||
return 'resData = x[xIndex];';
|
||||
case 3:
|
||||
return 'resData = vec3<f32>(x[xIndex], x[xIndex + 1], x[xIndex + 2]);';
|
||||
return `resData = vec3<${dataType}>(x[xIndex], x[xIndex + 1], x[xIndex + 2]);`;
|
||||
case 4:
|
||||
return 'resData = x[xIndex / 4];';
|
||||
default:
|
||||
|
|
@ -92,7 +93,7 @@ const conv2dCommonSnippet =
|
|||
let xRow = outRow * stride[0] + dilation[0] * WRow - pad[0];
|
||||
let xCol = outCol * stride[1] + dilation[1] * WCol - pad[1];
|
||||
let xCh = ${col} % inChannels;
|
||||
var resData = ${typeSnippet(innerElementSizeX)}(0.0);
|
||||
var resData = ${typeSnippet(innerElementSizeX, dataType)}(0.0);
|
||||
// The bounds checking is always needed since we use it to pad zero for
|
||||
// the 'same' padding type.
|
||||
if (xRow >= 0 && xRow < ${xHeight} && xCol >= 0 && xCol < ${xWidth}) {
|
||||
|
|
@ -110,7 +111,7 @@ const conv2dCommonSnippet =
|
|||
if (row < dimAOuter && col < dimInner) {
|
||||
${readXSnippet}
|
||||
}
|
||||
return ${typeSnippet(innerElementSizeX)}(0.0);`) :
|
||||
return ${typeSnippet(innerElementSizeX, dataType)}(0.0);`) :
|
||||
(fitInner && fitBOuter ? `
|
||||
let col = colIn * ${innerElementSizeX};
|
||||
${readXSnippet}` :
|
||||
|
|
@ -119,13 +120,15 @@ const conv2dCommonSnippet =
|
|||
if (row < dimInner && col < dimBOuter) {
|
||||
${readXSnippet}
|
||||
}
|
||||
return ${typeSnippet(innerElementSizeX)}(0.0);`);
|
||||
return ${typeSnippet(innerElementSizeX, dataType)}(0.0);`);
|
||||
|
||||
const sampleW = `${getWSnippet(innerElementSizeW)}`;
|
||||
|
||||
const resType = typeSnippet(innerElementSize);
|
||||
const aType = isChannelsLast ? typeSnippet(innerElementSizeX) : typeSnippet(innerElementSizeW);
|
||||
const bType = isChannelsLast ? typeSnippet(innerElementSizeW) : typeSnippet(innerElementSizeX);
|
||||
const resType = typeSnippet(innerElementSize, dataType);
|
||||
const aType =
|
||||
isChannelsLast ? typeSnippet(innerElementSizeX, dataType) : typeSnippet(innerElementSizeW, dataType);
|
||||
const bType =
|
||||
isChannelsLast ? typeSnippet(innerElementSizeW, dataType) : typeSnippet(innerElementSizeX, dataType);
|
||||
const userCode = `
|
||||
${activationFnSnippet(activation, hasPreluActivationWeights, innerElementSize === 4, 4)}
|
||||
fn mm_readA(batch: i32, row : i32, colIn : i32) -> ${aType} {
|
||||
|
|
@ -190,23 +193,24 @@ export const createConv2DMatMulProgramInfo =
|
|||
const fitInner = dimInner % tileInner === 0;
|
||||
|
||||
const elementsSize = isVec4 ? [innerElementSize, 4, 4] : [1, 1, 1];
|
||||
const t = tensorTypeToWsglStorageType(inputs[0].dataType);
|
||||
|
||||
const declareInputs = [
|
||||
`@group(0) @binding(0) var<storage, read> x: array<${isVec4 && innerElementSize === 4 ? 'vec4<f32>' : 'f32'}>;`,
|
||||
`@group(0) @binding(1) var<storage, read> w: array<${isVec4 ? 'vec4<f32>' : 'f32'}>;`
|
||||
`@group(0) @binding(0) var<storage, read> x: array<${isVec4 && innerElementSize === 4 ? `vec4<${t}>` : t}>;`,
|
||||
`@group(0) @binding(1) var<storage, read> w: array<${isVec4 ? `vec4<${t}>` : t}>;`
|
||||
];
|
||||
let declareFunctions = `
|
||||
fn setOutputAtIndex(flatIndex : i32, value : ${isVec4 ? 'vec4<f32>' : 'f32'}) {
|
||||
result[flatIndex] = ${isVec4 ? 'vec4<f32>' : 'f32'}(value);
|
||||
fn setOutputAtIndex(flatIndex : i32, value : ${isVec4 ? `vec4<${t}>` : t}) {
|
||||
result[flatIndex] = ${isVec4 ? `vec4<${t}>` : t}(value);
|
||||
}
|
||||
fn setOutputAtCoords(d0 : i32, d1 : i32, d2 : i32, d3 : i32, value : ${isVec4 ? 'vec4<f32>' : 'f32'}) {
|
||||
fn setOutputAtCoords(d0 : i32, d1 : i32, d2 : i32, d3 : i32, value : ${isVec4 ? `vec4<${t}>` : t}) {
|
||||
let flatIndex = getOutputIndexFromCoords(vec4<i32>(d0, d1, d2, d3));
|
||||
setOutputAtIndex(flatIndex ${isVec4 ? '/ 4' : ''}, value);
|
||||
}`;
|
||||
if (hasBias) {
|
||||
declareInputs.push(`@group(0) @binding(2) var<storage, read> bias: array<${isVec4 ? 'vec4<f32>' : 'f32'}>;`);
|
||||
declareInputs.push(`@group(0) @binding(2) var<storage, read> bias: array<${isVec4 ? `vec4<${t}>` : t}>;`);
|
||||
declareFunctions += `
|
||||
fn getBiasByOutputCoords(coords : vec4<i32>) -> ${isVec4 ? 'vec4<f32>' : 'f32'} {
|
||||
fn getBiasByOutputCoords(coords : vec4<i32>) -> ${isVec4 ? `vec4<${t}>` : t} {
|
||||
return bias[coords.${isChannelsLast ? 'w' : 'y'}${isVec4 ? '/ 4' : ''}];
|
||||
}`;
|
||||
}
|
||||
|
|
@ -222,7 +226,7 @@ export const createConv2DMatMulProgramInfo =
|
|||
// dilation : vec2<i32>, dimAOuter : i32, dimBOuter : i32, dimInner : i32 };
|
||||
${declareInputs.join('')}
|
||||
@group(0) @binding(${declareInputs.length}) var<storage, read_write> result: array<${
|
||||
isVec4 ? 'vec4<f32>' : 'f32'}>;
|
||||
isVec4 ? `vec4<${t}>` : t}>;
|
||||
//@group(0) @binding(${declareInputs.length + 1}) var<uniform> uniforms: Uniforms;
|
||||
|
||||
const xShape : vec4<i32> = vec4<i32>(${inputs[0].dims.join(',')});
|
||||
|
|
@ -240,12 +244,12 @@ export const createConv2DMatMulProgramInfo =
|
|||
${
|
||||
conv2dCommonSnippet(
|
||||
isChannelsLast, fitAOuter, fitBOuter, fitInner, hasBias, undefined, false, elementsSize[0],
|
||||
elementsSize[1], elementsSize[2])}
|
||||
elementsSize[1], elementsSize[2], t)}
|
||||
${
|
||||
isVec4 ?
|
||||
makeMatMulPackedVec4Source(elementsPerThread, workGroupSize, undefined, !isChannelsLast, tileInner) :
|
||||
makeMatMulPackedVec4Source(elementsPerThread, workGroupSize, t, undefined, !isChannelsLast, tileInner) :
|
||||
makeMatMulPackedSource(
|
||||
elementsPerThread, workGroupSize, undefined, !isChannelsLast, tileInner, false, undefined,
|
||||
elementsPerThread, workGroupSize, t, undefined, !isChannelsLast, tileInner, false, undefined,
|
||||
sequentialAccessByThreads)}`
|
||||
};
|
||||
};
|
||||
|
|
|
|||
|
|
@ -32,6 +32,7 @@ import {makeMatMulPackedSource, makeMatMulPackedVec4Source} from './matmul_packe
|
|||
const conv2dTransposeCommonSnippet =
|
||||
(isChannelsLast: boolean, addBias = false, activation?: Activation, hasPreluActivationWeights = false,
|
||||
innerElementSize = 4): string => {
|
||||
const type = typeSnippet(innerElementSize, 'f32');
|
||||
const getWSnippet = (innerElementSize: number) => {
|
||||
switch (innerElementSize) {
|
||||
case 1:
|
||||
|
|
@ -89,10 +90,10 @@ const conv2dTransposeCommonSnippet =
|
|||
let xR = f32(outRow - pads[0] + dilation[0] * WRow) / f32(strides[0]);
|
||||
let xC = f32(outCol - pads[1] + dilation[1] * WCol) / f32(strides[1]);
|
||||
if (xR < 0.0 || xR >= f32(${xHeight}) || fract(xR) > 0.0) {
|
||||
return ${typeSnippet(innerElementSize)}(0.0);
|
||||
return ${type}(0.0);
|
||||
}
|
||||
if (xC < 0.0 || xC >= f32(${xWidth}) || fract(xC) > 0.0) {
|
||||
return ${typeSnippet(innerElementSize)}(0.0);
|
||||
return ${type}(0.0);
|
||||
}
|
||||
let iXR = i32(xR);
|
||||
let iXC = i32(xC);
|
||||
|
|
@ -105,13 +106,13 @@ const conv2dTransposeCommonSnippet =
|
|||
if (row < dimAOuter && col < dimInner) {
|
||||
${readASnippet}
|
||||
}
|
||||
return ${typeSnippet(innerElementSize)}(0.0);` :
|
||||
return ${type}(0.0);` :
|
||||
`
|
||||
let col = colIn * ${innerElementSize};
|
||||
if (row < dimInner && col < dimBOuter) {
|
||||
${readASnippet}
|
||||
}
|
||||
return ${typeSnippet(innerElementSize)}(0.0);`;
|
||||
return ${type}(0.0);`;
|
||||
|
||||
const sampleW = `
|
||||
let col = colIn * ${innerElementSize};
|
||||
|
|
@ -125,21 +126,21 @@ const conv2dTransposeCommonSnippet =
|
|||
let coord = vec4<i32>(coordX, coordY, col, rowInner);
|
||||
${getWSnippet(innerElementSize)}
|
||||
}
|
||||
return ${typeSnippet(innerElementSize)}(0.0);
|
||||
return ${type}(0.0);
|
||||
`;
|
||||
|
||||
|
||||
const userCode = `
|
||||
${activationFnSnippet(activation, hasPreluActivationWeights, innerElementSize === 4, 4)}
|
||||
fn mm_readA(batch: i32, row : i32, colIn : i32) -> ${typeSnippet(innerElementSize)} {
|
||||
fn mm_readA(batch: i32, row : i32, colIn : i32) -> ${type} {
|
||||
${isChannelsLast ? sampleA : sampleW}
|
||||
}
|
||||
|
||||
fn mm_readB(batch: i32, row : i32, colIn : i32) -> ${typeSnippet(innerElementSize)} {
|
||||
fn mm_readB(batch: i32, row : i32, colIn : i32) -> ${type} {
|
||||
${isChannelsLast ? sampleW : sampleA}
|
||||
}
|
||||
|
||||
fn mm_write(batch: i32, row : i32, colIn : i32, valueInput : ${typeSnippet(innerElementSize)}) {
|
||||
fn mm_write(batch: i32, row : i32, colIn : i32, valueInput : ${type}) {
|
||||
let col = colIn * ${innerElementSize};
|
||||
if (row < dimAOuter && col < dimBOuter) {
|
||||
var value = valueInput;
|
||||
|
|
@ -234,10 +235,10 @@ export const createConv2DTransposeMatMulProgramInfo =
|
|||
${declareFunctions}
|
||||
${conv2dTransposeCommonSnippet(isChannelsLast, hasBias, undefined, false, innerElementSize)}
|
||||
${
|
||||
isVec4 ?
|
||||
makeMatMulPackedVec4Source(elementsPerThread, workGroupSize, undefined, !isChannelsLast, tileInner) :
|
||||
makeMatMulPackedSource(
|
||||
elementsPerThread, workGroupSize, undefined, !isChannelsLast, tileInner, false, undefined,
|
||||
sequentialAccessByThreads)}`
|
||||
isVec4 ? makeMatMulPackedVec4Source(
|
||||
elementsPerThread, workGroupSize, 'f32', undefined, !isChannelsLast, tileInner) :
|
||||
makeMatMulPackedSource(
|
||||
elementsPerThread, workGroupSize, 'f32', undefined, !isChannelsLast, tileInner, false,
|
||||
undefined, sequentialAccessByThreads)}`
|
||||
};
|
||||
};
|
||||
|
|
|
|||
|
|
@ -21,12 +21,13 @@ import {LOG_DEBUG} from '../../../log';
|
|||
import {TensorView} from '../../../tensor-view';
|
||||
import {ShapeUtil} from '../../../util';
|
||||
import {GpuDataType, ProgramInfo, ProgramMetadata} from '../../types';
|
||||
import {inputVariable, outputVariable, ShaderHelper} from '../common';
|
||||
import {inputVariable, outputVariable, ShaderHelper, tensorTypeToWsglStorageType} from '../common';
|
||||
import {ConvTransposeAttributes} from '../conv-transpose';
|
||||
|
||||
const createConvTranspose2DOpProgramShaderSource =
|
||||
(shaderHelper: ShaderHelper, inputs: readonly TensorView[], attributes: ConvTransposeAttributes,
|
||||
outputShape: readonly number[], hasBias: boolean, is1DimensionDispatch: boolean, isVec4 = false): string => {
|
||||
outputShape: readonly number[], hasBias: boolean, is1DimensionDispatch: boolean, isVec4 = false,
|
||||
dataType: string): string => {
|
||||
const isChannelsLast = attributes.format === 'NHWC';
|
||||
const rowDim = isChannelsLast ? 1 : 2;
|
||||
const colDim = isChannelsLast ? 2 : 3;
|
||||
|
|
@ -39,12 +40,12 @@ const createConvTranspose2DOpProgramShaderSource =
|
|||
const outputChannelsPerGroup = wShape[1];
|
||||
|
||||
let declareFunctions = `
|
||||
fn setOutputAtIndex(flatIndex : u32, value : ${isVec4 ? 'vec4<f32>' : 'f32'}) {
|
||||
result[flatIndex] = ${isVec4 ? 'vec4<f32>' : 'f32'}(value);
|
||||
fn setOutputAtIndex(flatIndex : u32, value : ${isVec4 ? `vec4<${dataType}>` : dataType}) {
|
||||
result[flatIndex] = ${isVec4 ? `vec4<${dataType}>` : dataType}(value);
|
||||
}`;
|
||||
if (hasBias) {
|
||||
declareFunctions += `
|
||||
fn getBiasByOutputCoords(coords : vec4<u32>) -> ${isVec4 ? 'vec4<f32>' : 'f32'} {
|
||||
fn getBiasByOutputCoords(coords : vec4<u32>) -> ${isVec4 ? `vec4<${dataType}>` : dataType} {
|
||||
return bias[coords.${isChannelsLast ? 'w' : 'y'}${isVec4 ? '/ 4' : ''}];
|
||||
}`;
|
||||
}
|
||||
|
|
@ -66,33 +67,33 @@ const createConvTranspose2DOpProgramShaderSource =
|
|||
|
||||
// Convolve dy(?, ?, d2) with w(:, :, d1, d2) to compute dx(xR, xC, d1).
|
||||
// ? = to be determined. : = across all values in that axis.
|
||||
var dotProd: array<vec4<f32>, ${workPerThread}>;
|
||||
var dotProd: array<vec4<${dataType}>, ${workPerThread}>;
|
||||
for (var i = 0; i < ${workPerThread}; i++) {
|
||||
dotProd[i] = vec4<f32>(0.0);
|
||||
dotProd[i] = vec4<${dataType}>(0.0);
|
||||
}
|
||||
for (var wR: u32 = 0; wR < filterDims[0]; wR = wR + 1) {
|
||||
var dyR = (f32(dyCorner.x) + f32(wR)) / f32(strides.x);
|
||||
var dyR = (${dataType}(dyCorner.x) + ${dataType}(wR)) / ${dataType}(strides.x);
|
||||
let wRPerm = filterDims[0] - 1 - wR;
|
||||
if (dyR < 0.0 || dyR >= f32(outBackprop[1]) ||
|
||||
if (dyR < 0.0 || dyR >= ${dataType}(outBackprop[1]) ||
|
||||
fract(dyR) > 0.0 || wRPerm < 0) {
|
||||
continue;
|
||||
}
|
||||
let idyR: u32 = u32(dyR);
|
||||
|
||||
for (var wC: u32 = 0; wC < filterDims[1]; wC = wC + 1) {
|
||||
let dyC = (f32(dyCorner.y) + f32(wC)) / f32(strides.y);
|
||||
let dyC2 = (f32(dyCorner.y) + 1.0 + f32(wC)) / f32(strides.y);
|
||||
let dyC = (${dataType}(dyCorner.y) + ${dataType}(wC)) / ${dataType}(strides.y);
|
||||
let dyC2 = (${dataType}(dyCorner.y) + 1.0 + ${dataType}(wC)) / ${dataType}(strides.y);
|
||||
let wCPerm = filterDims[1] - 1 - wC;
|
||||
if (wCPerm < 0) {
|
||||
continue;
|
||||
}
|
||||
var bDyCVal = true;
|
||||
var bDyCVal2 = true;
|
||||
if (dyC < 0.0 || dyC >= f32(outBackprop[2]) ||
|
||||
if (dyC < 0.0 || dyC >= ${dataType}(outBackprop[2]) ||
|
||||
fract(dyC) > 0.0) {
|
||||
bDyCVal = false;
|
||||
}
|
||||
if (dyC2 < 0.0 || dyC2 >= f32(outBackprop[2]) ||
|
||||
if (dyC2 < 0.0 || dyC2 >= ${dataType}(outBackprop[2]) ||
|
||||
fract(dyC2) > 0.0) {
|
||||
bDyCVal2 = false;
|
||||
}
|
||||
|
|
@ -108,7 +109,7 @@ const createConvTranspose2DOpProgramShaderSource =
|
|||
let wValue3 = ${w.get('u32(wRPerm)', 'u32(wCPerm)', 'd1 + 3', 'd2')};
|
||||
|
||||
var xValue = ${dy.get('batch', 'idyR', 'idyC', 'd2')};
|
||||
let tmpval = vec4<f32>(dot(xValue, wValue0),
|
||||
let tmpval = vec4<${dataType}>(dot(xValue, wValue0),
|
||||
dot(xValue, wValue1),
|
||||
dot(xValue, wValue2),
|
||||
dot(xValue, wValue3));
|
||||
|
|
@ -116,7 +117,7 @@ const createConvTranspose2DOpProgramShaderSource =
|
|||
|
||||
xValue = ${dy.get('batch', 'idyR', 'idyC2', 'd2')};
|
||||
|
||||
dotProd[1] = dotProd[1] + vec4<f32>(dot(xValue, wValue0),
|
||||
dotProd[1] = dotProd[1] + vec4<${dataType}>(dot(xValue, wValue0),
|
||||
dot(xValue, wValue1),
|
||||
dot(xValue, wValue2),
|
||||
dot(xValue, wValue3));
|
||||
|
|
@ -130,7 +131,7 @@ const createConvTranspose2DOpProgramShaderSource =
|
|||
let wValue3 = ${w.get('u32(wRPerm)', 'u32(wCPerm)', 'd1 + 3', 'd2')};
|
||||
|
||||
var xValue = ${dy.get('batch', 'idyR', 'idyC', 'd2')};
|
||||
let tmpval = vec4<f32>(dot(xValue, wValue0),
|
||||
let tmpval = vec4<${dataType}>(dot(xValue, wValue0),
|
||||
dot(xValue, wValue1),
|
||||
dot(xValue, wValue2),
|
||||
dot(xValue, wValue3));
|
||||
|
|
@ -145,7 +146,7 @@ const createConvTranspose2DOpProgramShaderSource =
|
|||
let wValue3 = ${w.get('u32(wRPerm)', 'u32(wCPerm)', 'd1 + 3', 'd2')};
|
||||
|
||||
var xValue = ${dy.get('batch', 'idyR', 'idyC2', 'd2')};
|
||||
let tmpval = vec4<f32>(dot(xValue, wValue0),
|
||||
let tmpval = vec4<${dataType}>(dot(xValue, wValue0),
|
||||
dot(xValue, wValue1),
|
||||
dot(xValue, wValue2),
|
||||
dot(xValue, wValue3));
|
||||
|
|
@ -178,9 +179,9 @@ const createConvTranspose2DOpProgramShaderSource =
|
|||
if (wR % dilations.x != 0) {
|
||||
continue;
|
||||
}
|
||||
let dyR = (f32(dyRCorner) + f32(wR)) / f32(strides[0]);
|
||||
let dyR = (${dataType}(dyRCorner) + ${dataType}(wR)) / ${dataType}(strides[0]);
|
||||
let wRPerm = filterDims.x - 1 - wR / dilations.x;
|
||||
if (dyR < 0.0 || dyR >= f32(outBackprop[${rowDim}]) || fract(dyR) > 0.0 ||
|
||||
if (dyR < 0.0 || dyR >= ${dataType}(outBackprop[${rowDim}]) || fract(dyR) > 0.0 ||
|
||||
wRPerm < 0) {
|
||||
continue;
|
||||
}
|
||||
|
|
@ -190,9 +191,9 @@ const createConvTranspose2DOpProgramShaderSource =
|
|||
if (wC % dilations.y != 0) {
|
||||
continue;
|
||||
}
|
||||
let dyC = (f32(dyCCorner) + f32(wC)) / f32(strides.y);
|
||||
let dyC = (${dataType}(dyCCorner) + ${dataType}(wC)) / ${dataType}(strides.y);
|
||||
let wCPerm = filterDims.y - 1 - wC / dilations.y;
|
||||
if (dyC < 0.0 || dyC >= f32(outBackprop[${colDim}]) ||
|
||||
if (dyC < 0.0 || dyC >= ${dataType}(outBackprop[${colDim}]) ||
|
||||
fract(dyC) > 0.0 || wCPerm < 0) {
|
||||
continue;
|
||||
}
|
||||
|
|
@ -256,6 +257,7 @@ export const createConvTranspose2DProgramInfo =
|
|||
];
|
||||
LOG_DEBUG('verbose', () => `[conv2d_backprop_webgpu] dispatch = ${dispatch}`);
|
||||
|
||||
const dataType = tensorTypeToWsglStorageType(inputs[0].dataType);
|
||||
return {
|
||||
...metadata,
|
||||
outputs: [{
|
||||
|
|
@ -265,6 +267,7 @@ export const createConvTranspose2DProgramInfo =
|
|||
}],
|
||||
dispatchGroup: () => ({x: dispatch[0], y: dispatch[1], z: dispatch[2]}),
|
||||
getShaderSource: (shaderHelper: ShaderHelper) => createConvTranspose2DOpProgramShaderSource(
|
||||
shaderHelper, inputs, attributes, outputShape, hasBias, dispatch[1] === 1 && dispatch[2] === 1),
|
||||
shaderHelper, inputs, attributes, outputShape, hasBias, dispatch[1] === 1 && dispatch[2] === 1, false,
|
||||
dataType),
|
||||
};
|
||||
};
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@
|
|||
import {TensorView} from '../../../tensor-view';
|
||||
import {ShapeUtil} from '../../../util';
|
||||
import {GpuDataType, ProgramInfo, ProgramMetadata} from '../../types';
|
||||
import {getBroadcastDims, IndicesHelper, inputVariable, outputVariable, ShaderHelper} from '../common';
|
||||
import {getBroadcastDims, IndicesHelper, inputVariable, outputVariable, ShaderHelper, tensorTypeToWsglStorageType} from '../common';
|
||||
import {getActicationSnippet, InternalActivationAttributes} from '../fuse-utils';
|
||||
|
||||
import {typeSnippet} from './activation_util';
|
||||
|
|
@ -70,8 +70,8 @@ const calculateResultSnippet = (transposeA: boolean, innerElementSize: number) =
|
|||
};
|
||||
|
||||
export const makeMatMulPackedVec4Source =
|
||||
(workPerThread: number[], workgroupSize: [number, number, number], batchDims?: IndicesHelper, transposeA = false,
|
||||
tileInner = 32, splitK = false, splitedDimInner = 32): string => {
|
||||
(workPerThread: number[], workgroupSize: [number, number, number], type = 'f32', batchDims?: IndicesHelper,
|
||||
transposeA = false, tileInner = 32, splitK = false, splitedDimInner = 32): string => {
|
||||
const tileAOuter = workgroupSize[1] * workPerThread[1];
|
||||
const tileBOuter = workgroupSize[0] * workPerThread[0];
|
||||
const tileAWidth = transposeA ? tileAOuter : tileInner;
|
||||
|
|
@ -90,8 +90,8 @@ export const makeMatMulPackedVec4Source =
|
|||
workPerThread[0]} must be 4.`);
|
||||
}
|
||||
return `
|
||||
var<workgroup> mm_Asub : array<array<vec${innerElementSize}<f32>, ${tileAWidth / innerElementSize}>, ${tileAHight}>;
|
||||
var<workgroup> mm_Bsub : array<array<vec4<f32>, ${tileBOuter / workPerThread[0]}>, ${tileInner}>;
|
||||
var<workgroup> mm_Asub : array<array<vec${innerElementSize}<${type}>, ${tileAWidth / innerElementSize}>, ${tileAHight}>;
|
||||
var<workgroup> mm_Bsub : array<array<vec4<${type}>, ${tileBOuter / workPerThread[0]}>, ${tileInner}>;
|
||||
|
||||
const rowPerThread = ${workPerThread[1]};
|
||||
const colPerThread = ${workPerThread[0]};
|
||||
|
|
@ -115,7 +115,7 @@ fn main(@builtin(local_invocation_id) localId : vec3<u32>,
|
|||
let numTiles = ${splitK ? `${Math.ceil(splitedDimInner / tileInner)}` : '(dimInner - 1) / tileInner + 1'};
|
||||
var kStart = ${splitK ? `i32(globalId.z) * ${splitedDimInner}` : '0'};
|
||||
|
||||
var acc: array<vec4<f32>, rowPerThread>;
|
||||
var acc: array<vec4<${type}>, rowPerThread>;
|
||||
|
||||
// Loop over shared dimension.
|
||||
let tileRowB = localRow * ${rowPerThreadB};
|
||||
|
|
@ -179,8 +179,9 @@ const readDataFromSubASnippet = (transposeA: boolean) =>
|
|||
// sequentialAccessByThreads means sequential data in memory is accessed by
|
||||
// threads, instead of a single thread (default behavior).
|
||||
export const makeMatMulPackedSource =
|
||||
(workPerThread: number[], workgroupSize: [number, number, number], batchDims?: IndicesHelper, transposeA = false,
|
||||
tileInner = 32, splitK = false, splitedDimInner = 32, sequentialAccessByThreads = false): string => {
|
||||
(workPerThread: number[], workgroupSize: [number, number, number], type = 'f32', batchDims?: IndicesHelper,
|
||||
transposeA = false, tileInner = 32, splitK = false, splitedDimInner = 32,
|
||||
sequentialAccessByThreads = false): string => {
|
||||
const tileAOuter = workPerThread[1] * workgroupSize[1];
|
||||
const tileBOuter = workPerThread[0] * workgroupSize[0];
|
||||
const tileAWidth = transposeA ? tileAOuter : tileInner;
|
||||
|
|
@ -222,7 +223,7 @@ export const makeMatMulPackedSource =
|
|||
workgroupBarrier();
|
||||
|
||||
// Compute acc values for a single thread.
|
||||
var BCached : array<f32, colPerThread>;
|
||||
var BCached : array<${type}, colPerThread>;
|
||||
for (var k = 0; k < tileInner; k = k + 1) {
|
||||
for (var inner = 0; inner < colPerThread; inner = inner + 1) {
|
||||
BCached[inner] = mm_Bsub[k][localCol + inner * ${workgroupSize[0]}];
|
||||
|
|
@ -283,7 +284,7 @@ for (var t = 0; t < numTiles; t = t + 1) {
|
|||
workgroupBarrier();
|
||||
|
||||
// Compute acc values for a single thread.
|
||||
var BCached : array<f32, colPerThread>;
|
||||
var BCached : array<${type}, colPerThread>;
|
||||
for (var k = 0; k < tileInner; k = k + 1) {
|
||||
for (var inner = 0; inner < colPerThread; inner = inner + 1) {
|
||||
BCached[inner] = mm_Bsub[k][tileCol + inner];
|
||||
|
|
@ -309,8 +310,8 @@ for (var innerRow = 0; innerRow < rowPerThread; innerRow = innerRow + 1) {
|
|||
`;
|
||||
|
||||
return `
|
||||
var<workgroup> mm_Asub : array<array<f32, ${tileAWidth}>, ${tileAHight}>;
|
||||
var<workgroup> mm_Bsub : array<array<f32, ${tileBOuter}>, ${tileInner}>;
|
||||
var<workgroup> mm_Asub : array<array<${type}, ${tileAWidth}>, ${tileAHight}>;
|
||||
var<workgroup> mm_Bsub : array<array<${type}, ${tileBOuter}>, ${tileInner}>;
|
||||
const rowPerThread = ${workPerThread[1]};
|
||||
const colPerThread = ${workPerThread[0]};
|
||||
const tileInner = ${tileInner};
|
||||
|
|
@ -324,7 +325,7 @@ fn main(@builtin(local_invocation_id) localId : vec3<u32>,
|
|||
let numTiles = ${splitK ? `${Math.ceil(splitedDimInner / tileInner)}` : '(dimInner - 1) / tileInner + 1'};
|
||||
var kStart = ${splitK ? `i32(globalId.z) * ${splitedDimInner}` : '0'};
|
||||
|
||||
var acc : array<array<f32, colPerThread>, rowPerThread>;
|
||||
var acc : array<array<${type}, colPerThread>, rowPerThread>;
|
||||
|
||||
// Without this initialization strange values show up in acc.
|
||||
for (var innerRow = 0; innerRow < rowPerThread; innerRow = innerRow + 1) {
|
||||
|
|
@ -347,6 +348,7 @@ const matMulReadWriteFnSource =
|
|||
const outputVariable = variables[5];
|
||||
const broadCastADims = getBroadcastDims(batchAVariable.shape, batchVariable.shape);
|
||||
const broadCastBDims = getBroadcastDims(batchBVariable.shape, batchVariable.shape);
|
||||
const dataType = tensorTypeToWsglStorageType(variables[0].type.tensor);
|
||||
const getAIndices = () => {
|
||||
const aRank = aVariable.shape.length;
|
||||
const batchRank = batchVariable.shape.length;
|
||||
|
|
@ -377,8 +379,8 @@ const matMulReadWriteFnSource =
|
|||
};
|
||||
const source = `
|
||||
fn mm_readA(batch: i32, row: i32, colIn: i32, batchIndices: ${batchVariable.type.indices}) -> ${
|
||||
typeSnippet(component)} {
|
||||
var value = ${typeSnippet(component)}(0.0);
|
||||
typeSnippet(component, dataType)} {
|
||||
var value = ${typeSnippet(component, dataType)}(0.0);
|
||||
let col = colIn * ${component};
|
||||
if(row < dimAOuter && col < dimInner)
|
||||
{
|
||||
|
|
@ -389,8 +391,8 @@ const matMulReadWriteFnSource =
|
|||
}
|
||||
|
||||
fn mm_readB(batch: i32, row: i32, colIn: i32, batchIndices: ${batchVariable.type.indices}) -> ${
|
||||
typeSnippet(component)} {
|
||||
var value = ${typeSnippet(component)}(0.0);
|
||||
typeSnippet(component, dataType)} {
|
||||
var value = ${typeSnippet(component, dataType)}(0.0);
|
||||
let col = colIn * ${component};
|
||||
if(row < dimInner && col < dimBOuter)
|
||||
{
|
||||
|
|
@ -400,7 +402,7 @@ const matMulReadWriteFnSource =
|
|||
return value;
|
||||
}
|
||||
|
||||
fn mm_write(batch: i32, row: i32, colIn: i32, valueIn: ${typeSnippet(component)}) {
|
||||
fn mm_write(batch: i32, row: i32, colIn: i32, valueIn: ${typeSnippet(component, dataType)}) {
|
||||
let col = colIn * ${component};
|
||||
if (row < dimAOuter && col < dimBOuter) {
|
||||
var value = valueIn;
|
||||
|
|
@ -444,6 +446,7 @@ export const createMatmulProgramInfo =
|
|||
Math.ceil(batchSize / workgroupSize[2] / elementsPerThread[2])
|
||||
];
|
||||
|
||||
const dataType = tensorTypeToWsglStorageType(inputs[0].dataType);
|
||||
const components = isVec4 ? 4 : 1;
|
||||
const A = inputVariable('a', inputs[0].dataType, [...outerDimsA, dimAOuter, dimInner / components], components);
|
||||
const B = inputVariable('b', inputs[1].dataType, [...outerDimsB, dimInner, dimBOuter / components], components);
|
||||
|
|
@ -466,8 +469,8 @@ export const createMatmulProgramInfo =
|
|||
${declareFunctions}
|
||||
${activationFunction}
|
||||
${
|
||||
isVec4 ? makeMatMulPackedVec4Source(elementsPerThread, workgroupSize, batchDims) :
|
||||
makeMatMulPackedSource(elementsPerThread, workgroupSize, batchDims)}
|
||||
isVec4 ? makeMatMulPackedVec4Source(elementsPerThread, workgroupSize, dataType, batchDims) :
|
||||
makeMatMulPackedSource(elementsPerThread, workgroupSize, dataType, batchDims)}
|
||||
${batchDims.impl()}`;
|
||||
return {
|
||||
...metadata,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
import {DataType} from '../../../wasm-common';
|
||||
import {TensorView} from '../../tensor-view';
|
||||
import {createAttributeWithCacheKey} from '../attribute-with-cache-key';
|
||||
import {ComputeContext, GpuDataType, ProgramInfoLoader, ProgramMetadata} from '../types';
|
||||
|
|
@ -201,15 +200,6 @@ const validateInputs = (inputs: readonly TensorView[], attributes: ConvTranspose
|
|||
if (attributes.outputShape.length !== 0 && attributes.outputShape.length !== inputs[0].dims.length - 2) {
|
||||
throw new Error('invalid output shape');
|
||||
}
|
||||
|
||||
// TODO : Need to add support for float64
|
||||
if (inputs[0].dataType !== DataType.float || inputs[1].dataType !== DataType.float) {
|
||||
throw new Error('ConvTranspose input(X,W) should be float tensor');
|
||||
}
|
||||
|
||||
if (inputs.length === 3 && inputs[2].dataType !== DataType.float) {
|
||||
throw new Error('ConvTranspose input(bias) should be float tensor');
|
||||
}
|
||||
};
|
||||
|
||||
const createConvTranspose2DProgramMetadata = (hasBias: boolean, cacheHint: string): ProgramMetadata => ({
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
import {DataType} from '../../../wasm-common';
|
||||
import {TensorView} from '../../tensor-view';
|
||||
import {PoolConvUtil} from '../../util';
|
||||
import {AttributeWithCacheKey, createAttributeWithCacheKey} from '../attribute-with-cache-key';
|
||||
|
|
@ -93,15 +92,6 @@ const validateInputs = (inputs: readonly TensorView[], attributes: ConvAttribute
|
|||
if (attributes.kernelShape.length !== 0 && attributes.kernelShape.length !== inputs[1].dims.length - 2) {
|
||||
throw new Error('invalid kernel shape');
|
||||
}
|
||||
|
||||
// TODO : Need to add support for float64
|
||||
if (inputs[0].dataType !== DataType.float || inputs[1].dataType !== DataType.float) {
|
||||
throw new Error('Conv input(X,W) should be float tensor');
|
||||
}
|
||||
|
||||
if (inputs.length === 3 && inputs[2].dataType !== DataType.float) {
|
||||
throw new Error('Conv input(bias) should be float tensor');
|
||||
}
|
||||
};
|
||||
|
||||
const getAdjustedConvAttributes = <T extends ConvAttributes>(attributes: T, inputs: readonly TensorView[]): T => {
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
import {DataType} from '../../../wasm-common';
|
||||
import {TensorView} from '../../tensor-view';
|
||||
import {BroadcastUtil} from '../../util';
|
||||
import {ComputeContext, GpuDataType, ProgramInfoLoader} from '../types';
|
||||
|
|
@ -9,7 +8,6 @@ import {ComputeContext, GpuDataType, ProgramInfoLoader} from '../types';
|
|||
import {createMatmulProgramInfo} from './3rd-party/matmul_packed_webgpu';
|
||||
import {InternalActivationAttributes} from './fuse-utils';
|
||||
|
||||
|
||||
const createMatmulProgramMetadata = (hasBias: boolean, cacheHint: string) => ({
|
||||
name: 'MatMul',
|
||||
inputTypes: hasBias ? [GpuDataType.default, GpuDataType.default, GpuDataType.default] :
|
||||
|
|
@ -35,10 +33,6 @@ const validateInputs = (inputs: readonly TensorView[]): void => {
|
|||
if (inputs[0].dims[inputs[0].dims.length - 1] !== inputs[1].dims[inputs[1].dims.length - 2]) {
|
||||
throw new Error('shared dimension does not match.');
|
||||
}
|
||||
|
||||
if (inputs[0].dataType !== DataType.float || inputs[1].dataType !== DataType.float) {
|
||||
throw new Error('inputs should be float type');
|
||||
}
|
||||
};
|
||||
|
||||
export const matMul = (context: ComputeContext): void => {
|
||||
|
|
|
|||
|
|
@ -232,18 +232,18 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 13, Uns
|
|||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 12, Transpose);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 13, Transpose);
|
||||
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, float, Conv);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, float, ConvTranspose);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, Conv);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, ConvTranspose);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, 11, MaxPool);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 12, MaxPool);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, AveragePool);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 1, GlobalAveragePool);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 1, GlobalMaxPool);
|
||||
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 10, float, Conv);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, float, Conv);
|
||||
class ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 10, float, ConvTranspose);
|
||||
class ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, float, ConvTranspose);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 10, Conv);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, Conv);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 10, ConvTranspose);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, ConvTranspose);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 7, 8, Gemm);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 9, 10, Gemm);
|
||||
class ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, 12, Gemm);
|
||||
|
|
@ -496,18 +496,18 @@ std::unique_ptr<KernelRegistry> RegisterKernels() {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 12, Transpose)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 13, Transpose)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, float, Conv)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, float, ConvTranspose)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, Conv)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, ConvTranspose)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, 11, MaxPool)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 12, MaxPool)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 11, AveragePool)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 1, GlobalAveragePool)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kMSInternalNHWCDomain, 1, GlobalMaxPool)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 10, float, Conv)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, float, Conv)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 10, float, ConvTranspose)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_TYPED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, float, ConvTranspose)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 10, Conv)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, Conv)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 1, 10, ConvTranspose)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, ConvTranspose)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 7, 8, Gemm)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 9, 10, Gemm)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_VERSIONED_KERNEL_CLASS_NAME(kJsExecutionProvider, kOnnxDomain, 11, 12, Gemm)>,
|
||||
|
|
|
|||
|
|
@ -9,33 +9,27 @@
|
|||
namespace onnxruntime {
|
||||
namespace js {
|
||||
|
||||
#define REGISTER_KERNEL_TYPED(T) \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
Conv, \
|
||||
kMSInternalNHWCDomain, \
|
||||
11, \
|
||||
T, \
|
||||
kJsExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
Conv<T, true>); \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
Conv, \
|
||||
kOnnxDomain, \
|
||||
11, \
|
||||
T, \
|
||||
kJsExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
Conv<T, false>); \
|
||||
ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_EX( \
|
||||
Conv, \
|
||||
kOnnxDomain, \
|
||||
1, 10, \
|
||||
T, \
|
||||
kJsExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
Conv<T, false>);
|
||||
|
||||
REGISTER_KERNEL_TYPED(float)
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
Conv,
|
||||
kMSInternalNHWCDomain,
|
||||
11,
|
||||
kJsExecutionProvider,
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", JsepSupportedFloatTypes()),
|
||||
Conv<true>);
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
Conv,
|
||||
kOnnxDomain,
|
||||
11,
|
||||
kJsExecutionProvider,
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", JsepSupportedFloatTypes()),
|
||||
Conv<false>);
|
||||
ONNX_OPERATOR_VERSIONED_KERNEL_EX(
|
||||
Conv,
|
||||
kOnnxDomain,
|
||||
1, 10,
|
||||
kJsExecutionProvider,
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", JsepSupportedFloatTypes()),
|
||||
Conv<false>);
|
||||
|
||||
} // namespace js
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@
|
|||
namespace onnxruntime {
|
||||
namespace js {
|
||||
|
||||
template <typename T, bool is_channels_last>
|
||||
template <bool is_channels_last>
|
||||
class Conv : public JsKernel {
|
||||
public:
|
||||
Conv(const OpKernelInfo& info) : JsKernel(info), conv_attrs_(info), w_is_const_(false) {
|
||||
|
|
|
|||
|
|
@ -7,33 +7,28 @@
|
|||
#include "conv_transpose.h"
|
||||
namespace onnxruntime {
|
||||
namespace js {
|
||||
#define REGISTER_KERNEL_TYPED(T) \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
ConvTranspose, \
|
||||
kMSInternalNHWCDomain, \
|
||||
11, \
|
||||
T, \
|
||||
kJsExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
ConvTranspose<T, true>); \
|
||||
ONNX_OPERATOR_TYPED_KERNEL_EX( \
|
||||
ConvTranspose, \
|
||||
kOnnxDomain, \
|
||||
11, \
|
||||
T, \
|
||||
kJsExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
ConvTranspose<T, false>); \
|
||||
ONNX_OPERATOR_VERSIONED_TYPED_KERNEL_EX( \
|
||||
ConvTranspose, \
|
||||
kOnnxDomain, \
|
||||
1, 10, \
|
||||
T, \
|
||||
kJsExecutionProvider, \
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::GetTensorType<T>()), \
|
||||
ConvTranspose<T, false>);
|
||||
|
||||
REGISTER_KERNEL_TYPED(float)
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
ConvTranspose,
|
||||
kMSInternalNHWCDomain,
|
||||
11,
|
||||
kJsExecutionProvider,
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", JsepSupportedFloatTypes()),
|
||||
ConvTranspose<true>);
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
ConvTranspose,
|
||||
kOnnxDomain,
|
||||
11,
|
||||
kJsExecutionProvider,
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", JsepSupportedFloatTypes()),
|
||||
ConvTranspose<false>);
|
||||
ONNX_OPERATOR_VERSIONED_KERNEL_EX(
|
||||
ConvTranspose,
|
||||
kOnnxDomain,
|
||||
1, 10,
|
||||
kJsExecutionProvider,
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", JsepSupportedFloatTypes()),
|
||||
ConvTranspose<false>);
|
||||
|
||||
} // namespace js
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@
|
|||
#include "core/providers/js/js_kernel.h"
|
||||
namespace onnxruntime {
|
||||
namespace js {
|
||||
template <typename T, bool is_channels_last>
|
||||
template <bool is_channels_last>
|
||||
class ConvTranspose : public JsKernel {
|
||||
public:
|
||||
ConvTranspose(const OpKernelInfo& info) : JsKernel(info), conv_transpose_attrs_(info), w_is_const_(false) {
|
||||
|
|
|
|||
|
|
@ -9,11 +9,11 @@ namespace js {
|
|||
JSEP_KERNEL_IMPL(MatMul, MatMul)
|
||||
|
||||
ONNX_OPERATOR_VERSIONED_KERNEL_EX(MatMul, kOnnxDomain, 1, 12, kJsExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
KernelDefBuilder().TypeConstraint("T", JsepSupportedFloatTypes()),
|
||||
MatMul);
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(MatMul, kOnnxDomain, 13, kJsExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
KernelDefBuilder().TypeConstraint("T", JsepSupportedFloatTypes()),
|
||||
MatMul);
|
||||
|
||||
} // namespace js
|
||||
|
|
|
|||
Loading…
Reference in a new issue