[js/web] port fixes for packed concat over to ort repo (#7605)

* port fixes for packed concat over to ort repo

* fix format
This commit is contained in:
Tixxx 2021-05-07 13:04:53 -07:00 committed by GitHub
parent bdefc6c4d8
commit 3c39fcc1fa
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 30 additions and 42 deletions

View file

@ -89,7 +89,7 @@ export class CoordsGlslLib extends GlslLib {
* Generates code for packed output sampler.
*/
protected getPackedOutputSamplingSnippet(outputLayout: TextureLayout): {[name: string]: GlslLibRoutine} {
const outShape = outputLayout.shape;
const outShape = outputLayout.unpackedShape;
const outTexShape = [outputLayout.width, outputLayout.height];
const result: {[name: string]: GlslLibRoutine} = {};
const funcName = 'getOutputCoords';
@ -231,7 +231,7 @@ export class CoordsGlslLib extends GlslLib {
const packedTexShape = texShape;
// texels needed to accommodate a logical row
const texelsInLogicalRow = shape[1];
const texelsInLogicalRow = Math.ceil(shape[1] / 2);
/**
* getOutputCoords
@ -264,8 +264,8 @@ export class CoordsGlslLib extends GlslLib {
*/
protected getOutputPacked3DCoords(shape: [number, number, number], texShape: [number, number]): GlslLibRoutine {
const packedTexShape = [texShape[0], texShape[1]];
const texelsInLogicalRow = shape[2];
const texelsInBatch = texelsInLogicalRow * shape[1];
const texelsInLogicalRow = Math.ceil(shape[2] / 2);
const texelsInBatch = texelsInLogicalRow * Math.ceil(shape[1] / 2);
const source = `
ivec3 getOutputCoords() {
ivec2 resTexRC = ivec2(TexCoords.xy *
@ -291,8 +291,8 @@ export class CoordsGlslLib extends GlslLib {
protected getOutputPackedNDCoords(shape: readonly number[], texShape: [number, number]): GlslLibRoutine {
const packedTexShape = [texShape[0], texShape[1]];
const texelsInLogicalRow = shape[shape.length - 1];
const texelsInBatch = texelsInLogicalRow * shape[shape.length - 2];
const texelsInLogicalRow = Math.ceil(shape[shape.length - 1] / 2);
const texelsInBatch = texelsInLogicalRow * Math.ceil(shape[shape.length - 2] / 2);
let texelsInBatchN = texelsInBatch;
let batches = '';
let coords = 'b, r, c';

View file

@ -45,7 +45,8 @@ export class WebGLPackedConcat extends Concat implements WebGLOperator {
const unpackChannel = unpackFromChannel();
const shapes = inputs.map(i => i.dims);
const channels = ['x', 'y', 'z', 'w', 'u', 'v'].slice(0, rank);
const allGlChannels = ['x', 'y', 'z', 'w', 'u', 'v'];
const channels = allGlChannels.slice(0, rank);
const offsets: number[] = new Array(shapes.length - 1);
const samplers = inputs.map((v, i) => `X${i}`);
@ -88,6 +89,10 @@ export class WebGLPackedConcat extends Concat implements WebGLOperator {
void main() {
${dtype} coords = getOutputCoords();
int lastDim = coords.${allGlChannels[rank - 1]};
coords.${allGlChannels[rank - 1]} = coords.${allGlChannels[rank - 2]};
coords.${allGlChannels[rank - 2]} = lastDim;
vec4 result = vec4(getValue(${coords}), 0., 0., 0.);
${coords[rank - 1]} = ${coords[rank - 1]} + 1;
@ -110,8 +115,9 @@ export class WebGLPackedConcat extends Concat implements WebGLOperator {
`;
return {
inputLayouts: inputs.map(t => handler.getOrCreateTextureLayout(t)),
outputLayout: handler.createTextureLayoutFromShape(outputShape),
inputLayouts: inputs.map(t => handler.getOrCreateTextureLayout(t, 4, true, t.dims, true)),
outputLayout:
handler.createTextureLayoutFromShape(outputShape, 4, outputShape, {isPacked: true, reverseWH: true}),
samplers,
shaderSource,
hasMain: true,
@ -120,7 +126,7 @@ export class WebGLPackedConcat extends Concat implements WebGLOperator {
};
}
createRunData(handler: WebGLInferenceHandler, programInfo: ProgramInfo, inputs: Tensor[]): RunData {
const inputTDs = inputs.map((t, i) => handler.getOrCreateTextureData(t, programInfo.inputLayouts[i]));
const inputTDs = inputs.map((t, i) => handler.getOrCreateTextureData(t, programInfo.inputLayouts[i], true));
return {
inputTextureDatas: inputTDs,
outputTextureData: handler.createTextureDataFromLayout(programInfo.outputLayout, inputTDs[0].tensor.type),

View file

@ -58,24 +58,7 @@ function getTestData(): TestData[] {
outputShape: [4, 4],
inputTextureShape: [2, 1],
outputTextureShape: [2, 2],
expectedOutput: new Float32Array([
1,
2,
5,
6,
3,
4,
7,
8,
1,
2,
5,
6,
3,
4,
7,
8,
]),
expectedOutput: new Float32Array([1, 2, 5, 6, 3, 4, 7, 8, 1, 2, 5, 6, 3, 4, 7, 8]),
},
{
elementCount: 8,
@ -83,7 +66,7 @@ function getTestData(): TestData[] {
inputShape: [2, 4],
outputShape: [2, 8],
inputTextureShape: [2, 1],
outputTextureShape: [2, 4],
outputTextureShape: [4, 2],
expectedOutput: new Float32Array([
1,
2,
@ -182,8 +165,8 @@ function getTestData(): TestData[] {
outputTextureShape: [8, 4],
expectedOutput: new Float32Array([
1, 2, 5, 6, 3, 4, 7, 8, 9, 10, 13, 14, 11, 12, 15, 16, 1, 2, 5, 6, 3, 4,
7, 8, 9, 10, 13, 14, 11, 12, 15, 16, 25, 26, 29, 30, 27, 28, 31, 32, 25, 26, 29, 30,
27, 28, 31, 32, 25, 26, 29, 30, 27, 28, 31, 32, 25, 26, 29, 30, 27, 28, 31, 32
7, 8, 9, 10, 13, 14, 11, 12, 15, 16, 17, 18, 21, 22, 19, 20, 23, 24, 25, 26, 29, 30,
27, 28, 31, 32, 17, 18, 21, 22, 19, 20, 23, 24, 25, 26, 29, 30, 27, 28, 31, 32
])
},
@ -195,9 +178,9 @@ function getTestData(): TestData[] {
inputTextureShape: [2, 4],
outputTextureShape: [8, 4],
expectedOutput: new Float32Array([
1, 2, 5, 6, 3, 4, 7, 8, 1, 2, 5, 6, 3, 4, 7, 8, 17, 18, 21, 22, 19, 20,
23, 24, 17, 18, 21, 22, 19, 20, 23, 24, 25, 26, 29, 30, 27, 28, 31, 32, 25, 26, 29, 30,
27, 28, 31, 32, 25, 26, 29, 30, 27, 28, 31, 32, 25, 26, 29, 30, 27, 28, 31, 32
1, 2, 5, 6, 3, 4, 7, 8, 1, 2, 5, 6, 3, 4, 7, 8, 9, 10, 13, 14, 11, 12,
15, 16, 9, 10, 13, 14, 11, 12, 15, 16, 17, 18, 21, 22, 19, 20, 23, 24, 17, 18, 21, 22,
19, 20, 23, 24, 25, 26, 29, 30, 27, 28, 31, 32, 25, 26, 29, 30, 27, 28, 31, 32
])
},
{
@ -208,9 +191,9 @@ function getTestData(): TestData[] {
inputTextureShape: [2, 4],
outputTextureShape: [8, 4],
expectedOutput: new Float32Array([
1, 2, 5, 6, 1, 2, 5, 6, 3, 4, 7, 8, 3, 4, 7, 8, 17, 18, 21, 22, 17, 18,
21, 22, 19, 20, 23, 24, 19, 20, 23, 24, 25, 26, 29, 30, 25, 26, 29, 30, 27, 28, 31, 32,
27, 28, 31, 32, 25, 26, 29, 30, 25, 26, 29, 30, 27, 28, 31, 32, 27, 28, 31, 32
1, 2, 5, 6, 1, 2, 5, 6, 3, 4, 7, 8, 3, 4, 7, 8, 9, 10, 13, 14, 9, 10,
13, 14, 11, 12, 15, 16, 11, 12, 15, 16, 17, 18, 21, 22, 17, 18, 21, 22, 19, 20, 23, 24,
19, 20, 23, 24, 25, 26, 29, 30, 25, 26, 29, 30, 27, 28, 31, 32, 27, 28, 31, 32
])
},
];
@ -257,7 +240,6 @@ describe('#UnitTest# - packed concat - Tensor concat', () => {
const elementCount = testData.elementCount;
const inputTensorShape = testData.inputShape;
const inputTextureShape = testData.inputTextureShape;
const outputTensorShape = testData.outputShape;
// create input data and tensor. The input data will be used to verify if the output tensor contains the
// same value but possibly different order depending on our packing algorithm.
@ -284,7 +266,7 @@ describe('#UnitTest# - packed concat - Tensor concat', () => {
isPacked: true,
shape: packedShape,
strides: ShapeUtil.computeStrides(packedShape),
unpackedShape: outputTensorShape,
unpackedShape: inputTensorShape,
tensor: inputTensorA,
texture: webglTextureA!
};
@ -295,13 +277,13 @@ describe('#UnitTest# - packed concat - Tensor concat', () => {
isPacked: true,
shape: packedShape,
strides: ShapeUtil.computeStrides(packedShape),
unpackedShape: outputTensorShape,
unpackedShape: inputTensorShape,
tensor: inputTensorB,
texture: webglTextureB!
};
webglInferenceHandler.setTextureData(inputTensorA.dataId, textureDataA);
webglInferenceHandler.setTextureData(inputTensorB.dataId, textureDataB);
webglInferenceHandler.setTextureData(inputTensorA.dataId, textureDataA, true);
webglInferenceHandler.setTextureData(inputTensorB.dataId, textureDataB, true);
// compile shader code
const programInfo =