| { |
| "domain": "ai.onnx", |
| "name": "LessOrEqual", |
| "sinceVersion": 16, |
| "inputs": { "a": { "onnx": "A", "dtype": "T" }, "b": { "onnx": "B", "dtype": "T" } }, |
| "outputs": { |
| "c": { "onnx": "C", "dtype": "B", "rank": "max(ranks.a, ranks.b)", "shape": "broadcastShape(shapes.a, shapes.b)" } |
| }, |
| "typeConstraints": { "T": ["float32", "float16", "int32", "int16", "uint32", "int8", "uint8"], "B": ["bool"] }, |
| "tunables": { "WORKGROUP_SIZE": { "default": 256 } }, |
| "variants": [ |
| { |
| "id": "same_shape_vec4", |
| "priority": 20, |
| "when": ["sameShape(shapes.a, shapes.c)", "sameShape(shapes.b, shapes.c)", "numel(shapes.c) > 0", "numel(shapes.c) % 4 == 0", "f16Ok(dtypes.T)"], |
| "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "LessOrEqual.vec4", |
| "shader": "compare-vec4.wgsl.jinja", |
| "derive": { |
| "op": "\"lessOrEqual\"", |
| "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1" |
| }, |
| "bindings": ["a", "b", "c", "params"], |
| "dispatch": { |
| "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "scalar_b_vec4", |
| "priority": 18, |
| "when": ["sameShape(shapes.a, shapes.c)", "numel(shapes.b) == 1", "numel(shapes.c) > 0", "numel(shapes.c) % 4 == 0", "f16Ok(dtypes.T)"], |
| "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"", "scalar": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "LessOrEqual.scalarBVec4", |
| "shader": "compare-vec4.wgsl.jinja", |
| "derive": { |
| "op": "\"lessOrEqual\"", |
| "scalarOperand": "\"b\"", |
| "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1" |
| }, |
| "bindings": ["a", "b_2", "c", "params"], |
| "dispatch": { |
| "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "scalar_a_vec4", |
| "priority": 18, |
| "when": ["sameShape(shapes.b, shapes.c)", "numel(shapes.a) == 1", "numel(shapes.c) > 0", "numel(shapes.c) % 4 == 0", "f16Ok(dtypes.T)"], |
| "derive": { "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"", "scalar": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "LessOrEqual.scalarAVec4", |
| "shader": "compare-vec4.wgsl.jinja", |
| "derive": { |
| "op": "\"lessOrEqual\"", |
| "scalarOperand": "\"a\"", |
| "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1" |
| }, |
| "bindings": ["a_2", "b", "c", "params"], |
| "dispatch": { |
| "x": "min(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c) / 4, 4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "broadcast_vec4", |
| "priority": 10, |
| "when": ["ranks.a <= ranks.c", "ranks.b <= ranks.c", "ranks.c >= 1", "dim(shapes.c, ranks.c - 1) % 4 == 0", "numel(shapes.c) % 4 == 0", "numel(shapes.c) >= 4", "f16Ok(dtypes.T)"], |
| "derive": { "scalar": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "LessOrEqual", |
| "shader": "compare-broadcast-vec4.wgsl.jinja", |
| "derive": { |
| "aShape": "shapes.a", |
| "bShape": "shapes.b", |
| "cShape": "shapes.c", |
| "aRank": "ranks.a", |
| "bRank": "ranks.b", |
| "cRank": "ranks.c", |
| "op": "\"lessOrEqual\"" |
| }, |
| "bindings": ["a_2", "b_2", "c", "params"], |
| "dispatch": { |
| "x": "min(ceilDiv((numel(shapes.c) / 4), (tunables.WORKGROUP_SIZE)), min(device.limits.maxComputeWorkgroupsPerDimension, 65535))", |
| "y": 1, |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "broadcast", |
| "when": ["ranks.a <= ranks.c", "ranks.b <= ranks.c", "f16Ok(dtypes.T)"], |
| "derive": { "scalar": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "LessOrEqual", |
| "shader": "compare-broadcast.wgsl.jinja", |
| "derive": { |
| "aShape": "shapes.a", |
| "bShape": "shapes.b", |
| "cShape": "shapes.c", |
| "aRank": "ranks.a", |
| "bRank": "ranks.b", |
| "cRank": "ranks.c", |
| "op": "\"lessOrEqual\"", |
| "itemsPerInvocation": 4 |
| }, |
| "bindings": ["a_2", "b_2", "c_2", "params_2"], |
| "dispatch": { |
| "x": "min(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| }, |
| { |
| "id": "same_shape_scalar_x4", |
| "priority": 15, |
| "when": ["sameShape(shapes.a, shapes.c)", "sameShape(shapes.b, shapes.c)", "numel(shapes.c) > 0", "numel(shapes.c) % 4 != 0", "ranks.a <= ranks.c", "ranks.b <= ranks.c", "f16Ok(dtypes.T)"], |
| "derive": { "scalar": "dtypes.T" }, |
| "passes": [ |
| { |
| "id": "main", |
| "name": "LessOrEqual", |
| "shader": "compare-broadcast.wgsl.jinja", |
| "derive": { |
| "aShape": "shapes.a", |
| "bShape": "shapes.b", |
| "cShape": "shapes.c", |
| "aRank": "ranks.a", |
| "bRank": "ranks.b", |
| "cRank": "ranks.c", |
| "op": "\"lessOrEqual\"", |
| "itemsPerInvocation": 4 |
| }, |
| "bindings": ["a_2", "b_2", "c_2", "params_2"], |
| "dispatch": { |
| "x": "min(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "y": "ceilDiv(ceilDiv((ceilDiv(numel(shapes.c), 4)), (tunables.WORKGROUP_SIZE)), 65535)", |
| "z": 1 |
| } |
| } |
| ] |
| } |
| ], |
| "bindings": { |
| "a": { "buffer": "read-only-storage", "elementType": "$vectorScalar" }, |
| "b": { "buffer": "read-only-storage", "elementType": "$vectorScalar" }, |
| "c": { "buffer": "storage", "elementType": "vec4<u32>" }, |
| "params": { "buffer": "uniform", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c) / 4" }] }, |
| "b_2": { "buffer": "read-only-storage", "name": "b", "elementType": "$scalar" }, |
| "a_2": { "buffer": "read-only-storage", "name": "a", "elementType": "$scalar" }, |
| "c_2": { "buffer": "storage", "name": "c", "elementType": "u32" }, |
| "params_2": { |
| "buffer": "uniform", |
| "name": "params", |
| "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c)" }] |
| } |
| } |
| } |
|
|