| { |
| "domain": "ai.onnx", |
| "name": "MaxPool", |
| "sinceVersion": 12, |
| "inputs": { "x": { "onnx": "X", "dtype": "T" } }, |
| "outputs": { |
| "y": { "onnx": "Y", "dtype": "T", "rank": "ranks.x", "shape": "poolOutputShape" }, |
| "indices": { |
| "onnx": "Indices", |
| "dtype": "I", |
| "rank": "ranks.x", |
| "optional": true, |
| "shape": "poolOutputShape", |
| "storage": "uint32" |
| } |
| }, |
| "attributes": { |
| "auto_pad": { "default": "NOTSET" }, |
| "storage_order": { "default": 0 }, |
| "ceil_mode": { "default": 0 }, |
| "kernel_shape": {}, |
| "strides": {}, |
| "pads": {}, |
| "dilations": {} |
| }, |
| "attributeConstraints": { |
| "auto_pad": { "values": ["NOTSET", "SAME_UPPER", "SAME_LOWER", "VALID"] }, |
| "ceil_mode": { "values": [0, 1] }, |
| "storage_order": { "values": [0, 1] }, |
| "kernel_shape": { "required": true } |
| }, |
| "typeConstraints": { "T": ["float32", "float16", "int8", "uint8"], "I": ["int64"] }, |
| "tunables": { |
| "WORKGROUP_SIZE": { "default": 128 }, |
| "TILED_D": { "default": 24 }, |
| "TILED_H": { "default": 8 }, |
| "TILED_W": { "default": 8 }, |
| "TILED_LANES_D": { "default": 4 }, |
| "SEPARABLE_MIN_WINDOW": { "default": 25 } |
| }, |
| "derive": { |
| "spatialRank": "ranks.x - 2", |
| "kernelShapeLengthOk": "not has(attrs, \"kernel_shape\") or (attrs.kernel_shape | length) == spatialRank", |
| "stridesLengthOk": "not has(attrs, \"strides\") or (attrs.strides | length) == spatialRank", |
| "dilationsLengthOk": "not has(attrs, \"dilations\") or (attrs.dilations | length) == spatialRank", |
| "padsLengthOk": "not has(attrs, \"pads\") or (attrs.pads | length) == 2 * spatialRank", |
| "kernelD": "attrs.kernel_shape[0] if kernelShapeLengthOk and has(attrs, \"kernel_shape\") and spatialRank == 3 else 1", |
| "kernelH": "attrs.kernel_shape[spatialRank - 2] if kernelShapeLengthOk and has(attrs, \"kernel_shape\") and spatialRank >= 2 else 1", |
| "kernelW": "attrs.kernel_shape[spatialRank - 1] if kernelShapeLengthOk and has(attrs, \"kernel_shape\") and spatialRank >= 1 else 1", |
| "strideD": "attrs.strides[0] if stridesLengthOk and has(attrs, \"strides\") and spatialRank == 3 else 1", |
| "strideH": "attrs.strides[spatialRank - 2] if stridesLengthOk and has(attrs, \"strides\") and spatialRank >= 2 else 1", |
| "strideW": "attrs.strides[spatialRank - 1] if stridesLengthOk and has(attrs, \"strides\") and spatialRank >= 1 else 1", |
| "dilationD": "attrs.dilations[0] if dilationsLengthOk and has(attrs, \"dilations\") and spatialRank == 3 else 1", |
| "dilationH": "attrs.dilations[spatialRank - 2] if dilationsLengthOk and has(attrs, \"dilations\") and spatialRank >= 2 else 1", |
| "dilationW": "attrs.dilations[spatialRank - 1] if dilationsLengthOk and has(attrs, \"dilations\") and spatialRank >= 1 else 1", |
| "padFront": "attrs.pads[0] if padsLengthOk and has(attrs, \"pads\") and spatialRank == 3 else 0", |
| "padTop": "attrs.pads[spatialRank - 2] if padsLengthOk and has(attrs, \"pads\") and spatialRank >= 2 else 0", |
| "padLeft": "attrs.pads[spatialRank - 1] if padsLengthOk and has(attrs, \"pads\") and spatialRank >= 1 else 0", |
| "padBack": "attrs.pads[spatialRank] if padsLengthOk and has(attrs, \"pads\") and spatialRank == 3 else 0", |
| "padBottom": "attrs.pads[2 * spatialRank - 2] if padsLengthOk and has(attrs, \"pads\") and spatialRank >= 2 else 0", |
| "padRight": "attrs.pads[2 * spatialRank - 1] if padsLengthOk and has(attrs, \"pads\") and spatialRank >= 1 else 0", |
| "poolShapeOk": "ranks.x >= 3 and ranks.x <= 5 and ranks.y == ranks.x and has(attrs, \"kernel_shape\") and kernelShapeLengthOk and stridesLengthOk and dilationsLengthOk and padsLengthOk and kernelD > 0 and kernelH > 0 and kernelW > 0 and strideD > 0 and strideH > 0 and strideW > 0 and dilationD > 0 and dilationH > 0 and dilationW > 0 and padFront >= 0 and padTop >= 0 and padLeft >= 0 and padBack >= 0 and padBottom >= 0 and padRight >= 0 and (attrs.auto_pad == \"NOTSET\" or not has(attrs, \"pads\")) and dim(shapes.y, 0) == dim(shapes.x, 0) and dim(shapes.y, 1) == dim(shapes.x, 1)", |
| "explicitOutW": "min(ceil((dim(shapes.x, ranks.x - 1) + padLeft + padRight - ((kernelW - 1) * dilationW + 1)) / strideW) + 1, ceil((dim(shapes.x, ranks.x - 1) + padLeft) / strideW)) if attrs.ceil_mode != 0 else floor((dim(shapes.x, ranks.x - 1) + padLeft + padRight - ((kernelW - 1) * dilationW + 1)) / strideW) + 1", |
| "explicitOutH": "min(ceil((dim(shapes.x, ranks.x - 2) + padTop + padBottom - ((kernelH - 1) * dilationH + 1)) / strideH) + 1, ceil((dim(shapes.x, ranks.x - 2) + padTop) / strideH)) if attrs.ceil_mode != 0 else floor((dim(shapes.x, ranks.x - 2) + padTop + padBottom - ((kernelH - 1) * dilationH + 1)) / strideH) + 1", |
| "explicitOutD": "(min(ceil((dim(shapes.x, 2) + padFront + padBack - ((kernelD - 1) * dilationD + 1)) / strideD) + 1, ceil((dim(shapes.x, 2) + padFront) / strideD)) if attrs.ceil_mode != 0 else floor((dim(shapes.x, 2) + padFront + padBack - ((kernelD - 1) * dilationD + 1)) / strideD) + 1) if ranks.x == 5 else 0", |
| "inferredOutW": "ceil(dim(shapes.x, ranks.x - 1) / strideW) if attrs.auto_pad == \"SAME_UPPER\" or attrs.auto_pad == \"SAME_LOWER\" else explicitOutW", |
| "inferredOutH": "ceil(dim(shapes.x, ranks.x - 2) / strideH) if attrs.auto_pad == \"SAME_UPPER\" or attrs.auto_pad == \"SAME_LOWER\" else explicitOutH", |
| "inferredOutD": "ceil(dim(shapes.x, 2) / strideD) if attrs.auto_pad == \"SAME_UPPER\" or attrs.auto_pad == \"SAME_LOWER\" else explicitOutD", |
| "poolOutputShape": "[dim(shapes.x, 0), dim(shapes.x, 1), inferredOutW] if ranks.x == 3 else ([dim(shapes.x, 0), dim(shapes.x, 1), inferredOutH, inferredOutW] if ranks.x == 4 else [dim(shapes.x, 0), dim(shapes.x, 1), inferredOutD, inferredOutH, inferredOutW])", |
| "pool2dOutputShapeOk": "(dim(shapes.y, 2) == explicitOutH and dim(shapes.y, 3) == explicitOutW) or ((attrs.auto_pad == \"SAME_UPPER\" or attrs.auto_pad == \"SAME_LOWER\") and dim(shapes.y, 2) == ceil(dim(shapes.x, 2) / strideH) and dim(shapes.y, 3) == ceil(dim(shapes.x, 3) / strideW))", |
| "autoPadTotalH": "max(0, (dim(shapes.y, ranks.x - 2) - 1) * strideH + ((kernelH - 1) * dilationH + 1) - dim(shapes.x, ranks.x - 2)) if ranks.x >= 4 else 0", |
| "autoPadTotalW": "max(0, (dim(shapes.y, ranks.x - 1) - 1) * strideW + ((kernelW - 1) * dilationW + 1) - dim(shapes.x, ranks.x - 1))", |
| "effectivePadH": "(floor(autoPadTotalH / 2) if attrs.auto_pad == \"SAME_UPPER\" else autoPadTotalH - floor(autoPadTotalH / 2)) if (attrs.auto_pad == \"SAME_UPPER\" or attrs.auto_pad == \"SAME_LOWER\") else padTop", |
| "effectivePadW": "(floor(autoPadTotalW / 2) if attrs.auto_pad == \"SAME_UPPER\" else autoPadTotalW - floor(autoPadTotalW / 2)) if (attrs.auto_pad == \"SAME_UPPER\" or attrs.auto_pad == \"SAME_LOWER\") else padLeft", |
| "floatDtypeOk": "(dtypes.T == \"f32\" or dtypes.T == \"f16\") and f16Ok(dtypes.T)", |
| "globalWindowParallelCovered": "not present.indices and floatDtypeOk and poolShapeOk and numel(shapes.y) == dim(shapes.x, 0) * dim(shapes.x, 1) and numel(shapes.y) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and attrs.auto_pad == \"NOTSET\" and strideD == 1 and strideH == 1 and strideW == 1 and dilationD == 1 and dilationH == 1 and dilationW == 1 and padFront == 0 and padTop == 0 and padLeft == 0 and padBack == 0 and padBottom == 0 and padRight == 0 and ((ranks.x == 3 and kernelW == dim(shapes.x, 2)) or (ranks.x == 4 and kernelH == dim(shapes.x, 2) and kernelW == dim(shapes.x, 3)) or (ranks.x == 5 and kernelD == dim(shapes.x, 2) and kernelH == dim(shapes.x, 3) and kernelW == dim(shapes.x, 4))) and inner(shapes.x, 1) >= 256", |
| "tiledStorageBytes": "(tunables.TILED_D + kernelD - 1) * (tunables.TILED_H + kernelH - 1) * (tunables.TILED_W + kernelW - 1) * 4", |
| "tiledDeviceOk": "tiledStorageBytes <= device.limits.maxComputeWorkgroupStorageSize and tunables.TILED_W * tunables.TILED_H * tunables.TILED_LANES_D <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.TILED_W <= device.limits.maxComputeWorkgroupSizeX and tunables.TILED_H <= device.limits.maxComputeWorkgroupSizeY and tunables.TILED_LANES_D <= device.limits.maxComputeWorkgroupSizeZ", |
| "tiledDispatchOk": "ceilDiv(dim(shapes.y, 4), tunables.TILED_W) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(dim(shapes.y, 3), tunables.TILED_H) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and dim(shapes.x, 0) * dim(shapes.x, 1) * ceilDiv(dim(shapes.y, 2), tunables.TILED_D) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", |
| "separableTmpElements": "dim(shapes.x, 0) * dim(shapes.x, 1) * dim(shapes.x, 2) * dim(shapes.y, 3)", |
| "separableFits": "separableTmpElements * 4 <= min(device.limits.maxStorageBufferBindingSize, device.limits.maxBufferSize) and ceilDiv(separableTmpElements, tunables.WORKGROUP_SIZE) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", |
| "autoPadTotalD": "max(0, (dim(shapes.y, 2) - 1) * strideD + ((kernelD - 1) * dilationD + 1) - dim(shapes.x, 2)) if ranks.x == 5 else 0", |
| "effectivePadD": "(floor(autoPadTotalD / 2) if attrs.auto_pad == \"SAME_UPPER\" else autoPadTotalD - floor(autoPadTotalD / 2)) if (attrs.auto_pad == \"SAME_UPPER\" or attrs.auto_pad == \"SAME_LOWER\") else padFront" |
| }, |
| "bindings": { |
| "x_2": { "name": "x", "buffer": "read-only-storage", "elementType": "$xVec4" }, |
| "params": { |
| "buffer": "uniform", |
| "struct": [ |
| { "name": "channels", "type": "u32", "value": "dim(shapes.x, 1)" }, |
| { "name": "inW", "type": "u32", "value": "dim(shapes.x, 2)" }, |
| { "name": "outW", "type": "u32", "value": "dim(shapes.y, 2)" }, |
| { "name": "kernelW", "type": "u32", "value": "kernelW" }, |
| { "name": "strideW", "type": "u32", "value": "strideW" }, |
| { "name": "dilationW", "type": "u32", "value": "dilationW" }, |
| { "name": "padW", "type": "i32", "value": "effectivePadW" }, |
| { "name": "count", "type": "u32", "value": "numel(shapes.y)" } |
| ] |
| }, |
| "params_7": { |
| "name": "params", |
| "buffer": "uniform", |
| "struct": [ |
| { "name": "channels", "type": "u32", "value": "dim(shapes.x, 1)" }, |
| { "name": "inH", "type": "u32", "value": "dim(shapes.x, 2)" }, |
| { "name": "inW", "type": "u32", "value": "dim(shapes.x, 3)" }, |
| { "name": "outH", "type": "u32", "value": "dim(shapes.y, 2)" }, |
| { "name": "outW", "type": "u32", "value": "dim(shapes.y, 3)" }, |
| { "name": "kernelH", "type": "u32", "value": "kernelH" }, |
| { "name": "kernelW", "type": "u32", "value": "kernelW" }, |
| { "name": "strideH", "type": "u32", "value": "strideH" }, |
| { "name": "strideW", "type": "u32", "value": "strideW" }, |
| { "name": "dilationH", "type": "u32", "value": "dilationH" }, |
| { "name": "dilationW", "type": "u32", "value": "dilationW" }, |
| { "name": "padH", "type": "i32", "value": "effectivePadH" }, |
| { "name": "padW", "type": "i32", "value": "effectivePadW" }, |
| { "name": "count", "type": "u32", "value": "numel(shapes.y)" } |
| ] |
| } |
| }, |
| "variants": [ |
| { |
| "id": "global_window_parallel", |
| "priority": 50, |
| "when": ["globalWindowParallelCovered"], |
| "derive": { |
| "scalar": "dtypes.T", |
| "xElem": "(\"vec4<f16>\" if dtypes.T == \"f16\" else \"vec4<f32>\") if inner(shapes.x, 1) % 4 == 0 else dtypes.T" |
| }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.GlobalWindowParallel", |
| "shader": "pool-global-reduction.wgsl.jinja", |
| "subgroupCollectivesWidth": "portable", |
| "derive": { |
| "op": "\"max\"", |
| "vec4": "inner(shapes.x, 1) % 4 == 0", |
| "workgroupSizeSpec": "tunables.WORKGROUP_SIZE", |
| "useSubgroups": "device.features.has(\"subgroups\")" |
| }, |
| "bindings": [ |
| { "arg": "x", "elementType": "$xElem" }, |
| { "arg": "y", "elementType": "$scalar" }, |
| { "name": "params", "struct": [{ "name": "spatial", "type": "u32", "value": "inner(shapes.x, 1)" }] } |
| ], |
| "dispatch": { "x": "numel(shapes.y)" } |
| } |
| ] |
| }, |
| { |
| "id": "nchw2d_k3s2p1_vec4", |
| "priority": 35, |
| "when": ["not present.indices", "f16Ok(dtypes.T)", "dtypes.T == \"f16\"", "poolShapeOk", "ranks.x == 4", "ranks.y == 4", "attrs.auto_pad == \"NOTSET\"", "kernelH == 3", "kernelW == 3", "strideH == 2", "strideW == 2", "dilationH == 1", "dilationW == 1", "padTop == 1", "padLeft == 1", "padBottom == 1", "padRight == 1", "attrs.ceil_mode == 0", "dim(shapes.x, 2) % 2 == 0", "dim(shapes.x, 3) % 8 == 0", "dim(shapes.y, 0) == dim(shapes.x, 0)", "dim(shapes.y, 1) == dim(shapes.x, 1)", "dim(shapes.y, 2) == dim(shapes.x, 2) / 2", "dim(shapes.y, 3) == dim(shapes.x, 3) / 2", "dim(shapes.y, 2) > 0", "dim(shapes.y, 3) >= 4"], |
| "derive": { "scalar": "dtypes.T", "xVec4": "\"vec4<f16>\"", "yVec4": "\"vec4<f16>\"" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Nchw2dK3s2p1Vec4", |
| "shader": "max-pool2d-nchw-k3s2p1-vec4.wgsl.jinja", |
| "derive": { |
| "inH": "dim(shapes.x, 2)", |
| "inWVec4": "dim(shapes.x, 3) / 4", |
| "inPlaneVec4": "dim(shapes.x, 2) * dim(shapes.x, 3) / 4", |
| "outH": "dim(shapes.y, 2)", |
| "outBlocks": "dim(shapes.y, 3) / 4" |
| }, |
| "bindings": [ |
| "x_2", |
| { "arg": "y", "elementType": "$yVec4" }, |
| { "name": "params", "struct": [{ "name": "tileCount", "type": "u32", "value": "numel(shapes.y) / 4" }] } |
| ], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y) / 4), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "nchw2d_k2s2_vec4", |
| "priority": 30, |
| "when": ["not present.indices", "f16Ok(dtypes.T)", "poolShapeOk", "ranks.x == 4", "ranks.y == 4", "kernelH == 2", "kernelW == 2", "strideH == 2", "strideW == 2", "dilationH == 1", "dilationW == 1", "padTop == 0", "padLeft == 0", "padBottom == 0", "padRight == 0", "attrs.ceil_mode == 0", "dim(shapes.x, 3) % 4 == 0", "dim(shapes.y, 0) == dim(shapes.x, 0)", "dim(shapes.y, 1) == dim(shapes.x, 1)", "dim(shapes.y, 2) == floor(dim(shapes.x, 2) / 2)", "dim(shapes.y, 3) == floor(dim(shapes.x, 3) / 2)", "dim(shapes.y, 2) > 0", "dim(shapes.y, 3) > 0"], |
| "derive": { "scalar": "dtypes.T", "xVec4": "\"vec4<f16>\" if dtypes.T == \"f16\" else \"vec4<f32>\"" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Nchw2dK2s2Vec4", |
| "shader": "pool2d-nchw-k2s2-vec4.wgsl.jinja", |
| "derive": { |
| "op": "\"max\"", |
| "outH": "dim(shapes.y, 2)", |
| "outW": "dim(shapes.y, 3)", |
| "outWHalf": "floor(dim(shapes.y, 3) / 2)", |
| "inWVec4": "floor(dim(shapes.x, 3) / 4)", |
| "inPlaneVec4": "floor(dim(shapes.x, 2) * dim(shapes.x, 3) / 4)" |
| }, |
| "bindings": [ |
| "x_2", |
| "y", |
| { "name": "params", "struct": [{ "name": "tileCount", "type": "u32", "value": "numel(shapes.y) / 2" }] } |
| ], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y) / 2), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y) / 2), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "nchw2d_unrolled", |
| "priority": 25, |
| "when": ["not present.indices", "floatDtypeOk", "poolShapeOk", "ranks.x == 4", "ranks.y == 4", "dim(shapes.y, 0) == dim(shapes.x, 0)", "dim(shapes.y, 1) == dim(shapes.x, 1)", "dim(shapes.y, 2) == explicitOutH", "dim(shapes.y, 3) == explicitOutW", "kernelH >= 1", "kernelH <= 3", "kernelW >= 1", "kernelW <= 3", "strideH >= 1", "strideW >= 1", "dilationH >= 1", "dilationW >= 1", "dim(shapes.y, 2) > 0", "dim(shapes.y, 3) > 0"], |
| "derive": { "scalar": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Nchw2dUnrolled", |
| "shader": "pool-window-unroll.wgsl.jinja", |
| "derive": { |
| "op": "\"max\"", |
| "kernelHSpec": "kernelH", |
| "kernelWSpec": "kernelW", |
| "strideHSpec": "strideH", |
| "strideWSpec": "strideW", |
| "dilationHSpec": "dilationH", |
| "dilationWSpec": "dilationW", |
| "padTopSpec": "padTop", |
| "padLeftSpec": "padLeft", |
| "inH": "dim(shapes.x, 2)", |
| "inW": "dim(shapes.x, 3)", |
| "outH": "dim(shapes.y, 2)", |
| "outW": "dim(shapes.y, 3)" |
| }, |
| "bindings": [ |
| "x", |
| "y", |
| { "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.y)" }] } |
| ], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "ncl1d_k2s2_x4", |
| "priority": 30, |
| "when": ["not present.indices", "floatDtypeOk", "poolShapeOk", "ranks.x == 3", "attrs.auto_pad == \"NOTSET\"", "kernelW == 2", "strideW == 2", "dilationW == 1", "padLeft == 0", "padRight == 0", "attrs.ceil_mode == 0", "dim(shapes.x, 2) % 4 == 0", "dim(shapes.y, 2) % 4 == 0", "dim(shapes.y, 2) * 2 == dim(shapes.x, 2)"], |
| "derive": { "scalar": "dtypes.T", "xVec4": "\"vec4<f16>\" if dtypes.T == \"f16\" else \"vec4<f32>\"" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Ncl1dK2s2X4", |
| "shader": "pool-ncl1d-x4.wgsl.jinja", |
| "derive": { "op": "\"max\"", "inW": "dim(shapes.x, 2)", "outW": "dim(shapes.y, 2)" }, |
| "bindings": [ |
| "x_2", |
| "y", |
| { |
| "name": "params", |
| "struct": [ |
| { |
| "name": "blockCount", |
| "type": "u32", |
| "value": "dim(shapes.x, 0) * dim(shapes.x, 1) * ceilDiv(dim(shapes.y, 2), 4)" |
| } |
| ] |
| } |
| ], |
| "dispatch": { |
| "x": "min(ceilDiv((dim(shapes.x, 0) * dim(shapes.x, 1) * ceilDiv(dim(shapes.y, 2), 4)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((dim(shapes.x, 0) * dim(shapes.x, 1) * ceilDiv(dim(shapes.y, 2), 4)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "ncl1d", |
| "priority": 10, |
| "when": ["not present.indices", "f16Ok(dtypes.T)", "poolShapeOk", "ranks.x == 3", "dim(shapes.y, 2) == inferredOutW"], |
| "derive": { "scalar": "dtypes.T", "hasIndices": false }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Ncl1d", |
| "shader": "pool-window-nd.wgsl.jinja", |
| "derive": { "op": "\"max\"", "spatialRankSpec": 1 }, |
| "bindings": ["x", "y", "params"], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "nchw2d_uint8", |
| "priority": 20, |
| "when": ["not present.indices", "tensorDtypes.x == \"uint8\"", "tensorDtypes.y == \"uint8\"", "poolShapeOk", "ranks.x == 4", "dim(shapes.y, 2) == explicitOutH", "dim(shapes.y, 3) == explicitOutW"], |
| "passes": [ |
| { |
| "id": "main", |
| "name": "nchw2d_uint8", |
| "shader": "max-pool2d-nchw-u32.wgsl.jinja", |
| "bindings": ["x", "y", "params_7"], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "nchw2d_separable", |
| "priority": 22, |
| "when": ["not present.indices", "f16Ok(dtypes.T) and dtypes.T != \"u32\"", "poolShapeOk", "ranks.x == 4", "ranks.y == 4", "pool2dOutputShapeOk", "kernelH > 1", "kernelW > 1", "kernelH * kernelW >= tunables.SEPARABLE_MIN_WINDOW", "strideH < kernelH", "separableFits"], |
| "derive": { "scalar": "dtypes.T", "pIs1": false, "pIs2": false }, |
| "intermediates": [{ "id": "tmp", "dtype": "float32", "shape": "[separableTmpElements]" }], |
| "passes": [ |
| { |
| "id": "reduce_w", |
| "name": "MaxPool.SeparableW", |
| "shader": "pool2d-nchw-separable.wgsl.jinja", |
| "derive": { "axisSpec": "\"w\"", "op": "\"max\"", "windowArea": "kernelH * kernelW" }, |
| "bindings": [ |
| "x", |
| { "name": "tmp", "buffer": "storage", "elementType": "f32" }, |
| { |
| "name": "params", |
| "struct": [ |
| { "name": "inH", "type": "u32", "value": "dim(shapes.x, 2)" }, |
| { "name": "inW", "type": "u32", "value": "dim(shapes.x, 3)" }, |
| { "name": "outW", "type": "u32", "value": "dim(shapes.y, 3)" }, |
| { "name": "kernelW", "type": "u32", "value": "kernelW" }, |
| { "name": "strideW", "type": "u32", "value": "strideW" }, |
| { "name": "dilationW", "type": "u32", "value": "dilationW" }, |
| { "name": "padW", "type": "i32", "value": "effectivePadW" }, |
| { |
| "name": "tmpCount", |
| "type": "u32", |
| "value": "dim(shapes.x, 0) * dim(shapes.x, 1) * dim(shapes.x, 2) * dim(shapes.y, 3)" |
| } |
| ] |
| } |
| ], |
| "dispatch": { |
| "x": "min(ceilDiv((separableTmpElements), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((separableTmpElements), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| }, |
| { |
| "id": "reduce_h", |
| "name": "MaxPool.SeparableH", |
| "shader": "pool2d-nchw-separable.wgsl.jinja", |
| "derive": { "axisSpec": "\"h\"", "op": "\"max\"", "windowArea": "kernelH * kernelW" }, |
| "bindings": [ |
| { "name": "tmp", "buffer": "read-only-storage", "elementType": "f32" }, |
| "y", |
| { |
| "name": "params", |
| "struct": [ |
| { "name": "inH", "type": "u32", "value": "dim(shapes.x, 2)" }, |
| { "name": "inW", "type": "u32", "value": "dim(shapes.x, 3)" }, |
| { "name": "outH", "type": "u32", "value": "dim(shapes.y, 2)" }, |
| { "name": "outW", "type": "u32", "value": "dim(shapes.y, 3)" }, |
| { "name": "kernelH", "type": "u32", "value": "kernelH" }, |
| { "name": "kernelW", "type": "u32", "value": "kernelW" }, |
| { "name": "strideH", "type": "u32", "value": "strideH" }, |
| { "name": "strideW", "type": "u32", "value": "strideW" }, |
| { "name": "dilationH", "type": "u32", "value": "dilationH" }, |
| { "name": "dilationW", "type": "u32", "value": "dilationW" }, |
| { "name": "padH", "type": "i32", "value": "effectivePadH" }, |
| { "name": "padW", "type": "i32", "value": "effectivePadW" }, |
| { "name": "count", "type": "u32", "value": "numel(shapes.y)" } |
| ] |
| } |
| ], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "nchw2d", |
| "when": ["not present.indices", "f16Ok(dtypes.T) and dtypes.T != \"u32\"", "poolShapeOk", "ranks.x == 4", "pool2dOutputShapeOk"], |
| "derive": { "scalar": "dtypes.T", "hasIndices": false }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "nchw2d", |
| "shader": "pool-window-nd.wgsl.jinja", |
| "derive": { "op": "\"max\"", "spatialRankSpec": 2 }, |
| "bindings": ["x", "y", "params_7"], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "ncdhw3d_k5s1_tiled", |
| "priority": 20, |
| "when": ["not present.indices", "dtypes.T == \"f32\"", "poolShapeOk", "ranks.x == 5", "attrs.auto_pad == \"NOTSET\"", "kernelD == 5 and kernelH == 5 and kernelW == 5", "strideD == 1 and strideH == 1 and strideW == 1", "dilationD == 1 and dilationH == 1 and dilationW == 1", "padFront == 2 and padTop == 2 and padLeft == 2 and padBack == 2 and padBottom == 2 and padRight == 2", "attrs.ceil_mode == 0", "dim(shapes.y, 2) == dim(shapes.x, 2)", "dim(shapes.y, 3) == dim(shapes.x, 3)", "dim(shapes.y, 4) == dim(shapes.x, 4)", "dim(shapes.y, 2) >= tunables.TILED_D", "dim(shapes.y, 3) >= tunables.TILED_H", "dim(shapes.y, 4) >= tunables.TILED_W", "tiledDeviceOk", "tiledDispatchOk"], |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Ncdhw3dK5s1Tiled", |
| "shader": "max-pool3d-ncdhw-k5-tiled.wgsl.jinja", |
| "bindings": [ |
| { "arg": "x", "elementType": "f32" }, |
| { "arg": "y", "elementType": "f32" }, |
| { |
| "name": "params", |
| "struct": [ |
| { "name": "inD", "type": "u32", "value": "dim(shapes.x, 2)" }, |
| { "name": "inH", "type": "u32", "value": "dim(shapes.x, 3)" }, |
| { "name": "inW", "type": "u32", "value": "dim(shapes.x, 4)" }, |
| { "name": "outD", "type": "u32", "value": "dim(shapes.y, 2)" }, |
| { "name": "outH", "type": "u32", "value": "dim(shapes.y, 3)" }, |
| { "name": "outW", "type": "u32", "value": "dim(shapes.y, 4)" } |
| ] |
| } |
| ], |
| "dispatch": { |
| "x": "ceilDiv(dim(shapes.y, 4), tunables.TILED_W)", |
| "y": "ceilDiv(dim(shapes.y, 3), tunables.TILED_H)", |
| "z": "dim(shapes.x, 0) * dim(shapes.x, 1) * ceilDiv(dim(shapes.y, 2), tunables.TILED_D)" |
| } |
| } |
| ] |
| }, |
| { |
| "id": "ncdhw3d", |
| "priority": 10, |
| "when": ["not present.indices", "floatDtypeOk", "poolShapeOk", "ranks.x == 5", "dim(shapes.y, 2) == inferredOutD", "dim(shapes.y, 3) == inferredOutH", "dim(shapes.y, 4) == inferredOutW"], |
| "supersededBy": ["global_window_parallel"], |
| "derive": { "scalar": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Ncdhw3d", |
| "shader": "pool-window-nd.wgsl.jinja", |
| "derive": { "op": "\"max\"", "spatialRankSpec": 3 }, |
| "bindings": [ |
| "x", |
| "y", |
| { |
| "name": "params", |
| "struct": [ |
| { "name": "channels", "type": "u32", "value": "dim(shapes.x, 1)" }, |
| { "name": "inD", "type": "u32", "value": "dim(shapes.x, 2)" }, |
| { "name": "inH", "type": "u32", "value": "dim(shapes.x, 3)" }, |
| { "name": "inW", "type": "u32", "value": "dim(shapes.x, 4)" }, |
| { "name": "outD", "type": "u32", "value": "dim(shapes.y, 2)" }, |
| { "name": "outH", "type": "u32", "value": "dim(shapes.y, 3)" }, |
| { "name": "outW", "type": "u32", "value": "dim(shapes.y, 4)" }, |
| { "name": "kernelD", "type": "u32", "value": "kernelD" }, |
| { "name": "kernelH", "type": "u32", "value": "kernelH" }, |
| { "name": "kernelW", "type": "u32", "value": "kernelW" }, |
| { "name": "strideD", "type": "u32", "value": "strideD" }, |
| { "name": "strideH", "type": "u32", "value": "strideH" }, |
| { "name": "strideW", "type": "u32", "value": "strideW" }, |
| { "name": "dilationD", "type": "u32", "value": "dilationD" }, |
| { "name": "dilationH", "type": "u32", "value": "dilationH" }, |
| { "name": "dilationW", "type": "u32", "value": "dilationW" }, |
| { "name": "padD", "type": "i32", "value": "effectivePadD" }, |
| { "name": "padH", "type": "i32", "value": "effectivePadH" }, |
| { "name": "padW", "type": "i32", "value": "effectivePadW" }, |
| { "name": "count", "type": "u32", "value": "numel(shapes.y)" } |
| ] |
| } |
| ], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "ncl1d_indices", |
| "priority": 15, |
| "when": ["present.indices", "f16Ok(dtypes.T)", "poolShapeOk", "ranks.x == 3", "dim(shapes.y, 2) == inferredOutW"], |
| "derive": { "scalar": "dtypes.T", "hasIndices": true }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Ncl1dIndices", |
| "shader": "pool-window-nd.wgsl.jinja", |
| "derive": { "op": "\"max\"", "spatialRankSpec": 1 }, |
| "bindings": ["x", "y", "indices", "params"], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "nchw2d_uint8_indices", |
| "priority": 16, |
| "when": ["present.indices", "tensorDtypes.x == \"uint8\"", "tensorDtypes.y == \"uint8\"", "poolShapeOk", "ranks.x == 4", "dim(shapes.y, 2) == explicitOutH", "dim(shapes.y, 3) == explicitOutW"], |
| "derive": { "scalar": "\"u32\"", "hasIndices": true, "storageOrder": "attrs.storage_order" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Nchw2dUint8Indices", |
| "shader": "pool-window-nd.wgsl.jinja", |
| "derive": { "op": "\"max\"", "spatialRankSpec": 2 }, |
| "bindings": ["x", "y", "indices", "params_7"], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "nchw2d_indices", |
| "priority": 15, |
| "when": ["present.indices", "f16Ok(dtypes.T) and dtypes.T != \"u32\"", "poolShapeOk", "ranks.x == 4", "pool2dOutputShapeOk"], |
| "derive": { "scalar": "dtypes.T", "hasIndices": true, "storageOrder": "attrs.storage_order" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "MaxPool.Nchw2dIndices", |
| "shader": "pool-window-nd.wgsl.jinja", |
| "derive": { "op": "\"max\"", "spatialRankSpec": 2 }, |
| "bindings": ["x", "y", "indices", "params_7"], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| } |
| ] |
| } |
|
|