mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
[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:
parent
bdefc6c4d8
commit
3c39fcc1fa
3 changed files with 30 additions and 42 deletions
|
|
@ -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';
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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 =
|
||||
|
|
|
|||
Loading…
Reference in a new issue