mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-19 19:00:47 +00:00
[js/WebGPU] Support int32 Transpose in WebGPU (#16952)
This commit is contained in:
parent
6361b22103
commit
506ddb3d5d
4 changed files with 64 additions and 6 deletions
|
|
@ -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);
|
||||
|
|
|
|||
50
js/web/test/data/ops/transpose_int32_uint32.jsonc
Normal file
50
js/web/test/data/ops/transpose_int32_uint32.jsonc
Normal file
|
|
@ -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"
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
|
|
@ -1359,7 +1359,8 @@
|
|||
"sqrt.jsonc",
|
||||
"sub.jsonc",
|
||||
"tan.jsonc",
|
||||
"transpose.jsonc"
|
||||
"transpose.jsonc",
|
||||
"transpose_int32_uint32.jsonc"
|
||||
//"xor.jsonc"
|
||||
]
|
||||
},
|
||||
|
|
|
|||
|
|
@ -12,7 +12,9 @@ ONNX_OPERATOR_VERSIONED_KERNEL_EX(
|
|||
1, 12,
|
||||
kJsExecutionProvider,
|
||||
(*KernelDefBuilder::Create())
|
||||
.TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
.TypeConstraint("T", {DataTypeImpl::GetTensorType<float>(),
|
||||
DataTypeImpl::GetTensorType<int32_t>(),
|
||||
DataTypeImpl::GetTensorType<uint32_t>()}),
|
||||
Transpose);
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(
|
||||
|
|
@ -21,7 +23,9 @@ ONNX_OPERATOR_KERNEL_EX(
|
|||
13,
|
||||
kJsExecutionProvider,
|
||||
(*KernelDefBuilder::Create())
|
||||
.TypeConstraint("T", DataTypeImpl::GetTensorType<float>()),
|
||||
.TypeConstraint("T", {DataTypeImpl::GetTensorType<float>(),
|
||||
DataTypeImpl::GetTensorType<int32_t>(),
|
||||
DataTypeImpl::GetTensorType<uint32_t>()}),
|
||||
Transpose);
|
||||
|
||||
} // namespace js
|
||||
|
|
|
|||
Loading…
Reference in a new issue