diff --git a/js/web/lib/wasm/jsep/webgpu/ops/transpose.ts b/js/web/lib/wasm/jsep/webgpu/ops/transpose.ts index 11f29e60b3..7a213e5d52 100644 --- a/js/web/lib/wasm/jsep/webgpu/ops/transpose.ts +++ b/js/web/lib/wasm/jsep/webgpu/ops/transpose.ts @@ -23,8 +23,9 @@ const validateInputs = (inputs: readonly TensorView[]): void => { throw new Error('Transpose requires 1 input.'); } - if (inputs[0].dataType !== DataType.float) { - throw new Error('input should be float tensor'); + if (inputs[0].dataType !== DataType.float && inputs[0].dataType !== DataType.int32 && + inputs[0].dataType !== DataType.uint32) { + throw new Error('Transpose only support float, int32, and uint32 data types'); } }; @@ -45,7 +46,9 @@ const permFunctionBody = (perm: number[], rank: number): string => { }; export const createTransposeProgramInfo = (input: TensorView, permAttr: number[]): ProgramInfo => { - const dataType = 'f32'; // TODO: support other data type + // We currently only support 4-byte element tensors, so using f32 here is safe + // TODO: support other data types for Transpose + const dataType = 'f32'; const inputShape = input.dims; const perm = getAdjustedPerm(inputShape, permAttr); const outputShape = getOutputShape(inputShape, perm); diff --git a/js/web/test/data/ops/transpose_int32_uint32.jsonc b/js/web/test/data/ops/transpose_int32_uint32.jsonc new file mode 100644 index 0000000000..fedc3d41ca --- /dev/null +++ b/js/web/test/data/ops/transpose_int32_uint32.jsonc @@ -0,0 +1,50 @@ +[ + { + "name": "Transpose int32", + "operator": "Transpose", + "attributes": [{ "name": "perm", "data": [1, 0, 2], "type": "ints" }], + "cases": [ + { + "name": "T[2,3]", + "inputs": [ + { + "data": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24], + "dims": [2, 3, 4], + "type": "int32" + } + ], + "outputs": [ + { + "data": [1, 2, 3, 4, 13, 14, 15, 16, 5, 6, 7, 8, 17, 18, 19, 20, 9, 10, 11, 12, 21, 22, 23, 24], + "dims": [3, 2, 4], + "type": "int32" + } + ] + } + ] + }, + { + "name": "Transpose uint32", + "operator": "Transpose", + "attributes": [{ "name": "perm", "data": [1, 0, 2], "type": "ints" }], + "cases": [ + { + "name": "T[2,3]", + "inputs": [ + { + "data": [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24], + "dims": [2, 3, 4], + "type": "uint32" + } + ], + "outputs": [ + { + "data": [1, 2, 3, 4, 13, 14, 15, 16, 5, 6, 7, 8, 17, 18, 19, 20, 9, 10, 11, 12, 21, 22, 23, 24], + "dims": [3, 2, 4], + "type": "uint32" + } + ] + } + ] + } +] diff --git a/js/web/test/suite-test-list.jsonc b/js/web/test/suite-test-list.jsonc index d420971a90..c894f0c58f 100644 --- a/js/web/test/suite-test-list.jsonc +++ b/js/web/test/suite-test-list.jsonc @@ -1359,7 +1359,8 @@ "sqrt.jsonc", "sub.jsonc", "tan.jsonc", - "transpose.jsonc" + "transpose.jsonc", + "transpose_int32_uint32.jsonc" //"xor.jsonc" ] }, diff --git a/onnxruntime/core/providers/js/operators/transpose.cc b/onnxruntime/core/providers/js/operators/transpose.cc index 763bcafc05..ef1e49046a 100644 --- a/onnxruntime/core/providers/js/operators/transpose.cc +++ b/onnxruntime/core/providers/js/operators/transpose.cc @@ -12,7 +12,9 @@ ONNX_OPERATOR_VERSIONED_KERNEL_EX( 1, 12, kJsExecutionProvider, (*KernelDefBuilder::Create()) - .TypeConstraint("T", DataTypeImpl::GetTensorType()), + .TypeConstraint("T", {DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType()}), Transpose); ONNX_OPERATOR_KERNEL_EX( @@ -21,7 +23,9 @@ ONNX_OPERATOR_KERNEL_EX( 13, kJsExecutionProvider, (*KernelDefBuilder::Create()) - .TypeConstraint("T", DataTypeImpl::GetTensorType()), + .TypeConstraint("T", {DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType(), + DataTypeImpl::GetTensorType()}), Transpose); } // namespace js