{ "domain": "ai.onnx", "name": "Add", "sinceVersion": 14, "inputs": { "a": { "onnx": "A", "dtype": "T" }, "b": { "onnx": "B", "dtype": "T" } }, "outputs": { "c": { "onnx": "C", "dtype": "T", "rank": "max(ranks.a, ranks.b)", "shape": "broadcastShape(shapes.a, shapes.b)" } }, "typeConstraints": { "T": ["float32", "float16", "int32", "uint32", "int8", "uint8"] }, "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": { "scalar": "dtypes.T", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" }, "passes": [ { "id": "main", "name": "Add.vec4", "shader": "binary-vec4.wgsl.jinja", "derive": { "op": "\"add\"", "cDtype": "tensorDtypes.c", "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1" }, "bindings": ["a", "b", "c_binary", "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": { "scalar": "dtypes.T", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" }, "passes": [ { "id": "main", "name": "Add.scalarBVec4", "shader": "binary-vec4.wgsl.jinja", "derive": { "op": "\"add\"", "cDtype": "tensorDtypes.c", "scalarOperand": "\"b\"", "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1" }, "bindings": ["a", "b_2", "c_binary", "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": { "scalar": "dtypes.T", "vectorScalar": "\"vec4<\" ~ dtypes.T ~ \">\"" }, "passes": [ { "id": "main", "name": "Add.scalarAVec4", "shader": "binary-vec4.wgsl.jinja", "derive": { "op": "\"add\"", "cDtype": "tensorDtypes.c", "scalarOperand": "\"a\"", "vec4PerThread": "4 if numel(shapes.c) * dtypeBytes(tensorDtypes.c) <= 16777216 else 1" }, "bindings": ["a_2", "b", "c_binary", "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", "numel(shapes.c) % 4 == 0", "numel(shapes.c) >= 4", "f16Ok(dtypes.T)"], "derive": { "scalar": "dtypes.T", "aElement": "\"vec4<\" ~ dtypes.T ~ \">\" if sameShape(shapes.a, shapes.c) else dtypes.T", "bElement": "\"vec4<\" ~ dtypes.T ~ \">\" if sameShape(shapes.b, shapes.c) else dtypes.T", "vec4Scalar": "\"vec4<\" ~ dtypes.T ~ \">\"" }, "passes": [ { "id": "main", "name": "Add", "shader": "binary-broadcast-vec4.wgsl.jinja", "derive": { "aShape": "shapes.a", "bShape": "shapes.b", "cShape": "shapes.c", "aRank": "ranks.a", "bRank": "ranks.b", "cRank": "ranks.c", "op": "\"add\"", "cDtype": "tensorDtypes.c" }, "bindings": ["a_3", "b_3", "c_2_binary", "params"], "dispatch": { "x": "min(ceilDiv((numel(shapes.c) / 4), (tunables.WORKGROUP_SIZE)), 65535)", "y": "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", "f16Ok(dtypes.T)"], "derive": { "scalar": "dtypes.T" }, "passes": [ { "id": "main", "name": "Add", "shader": "binary-broadcast.wgsl.jinja", "derive": { "aShape": "shapes.a", "bShape": "shapes.b", "cShape": "shapes.c", "aRank": "ranks.a", "bRank": "ranks.b", "cRank": "ranks.c", "op": "\"add\"", "cDtype": "tensorDtypes.c", "itemsPerInvocation": 4 }, "bindings": ["a_2", "b_2", "c_3", "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": "broadcast", "when": ["ranks.a <= ranks.c", "ranks.b <= ranks.c", "f16Ok(dtypes.T)"], "derive": { "scalar": "dtypes.T" }, "passes": [ { "id": "main", "name": "Add", "shader": "binary-broadcast.wgsl.jinja", "derive": { "aShape": "shapes.a", "bShape": "shapes.b", "cShape": "shapes.c", "aRank": "ranks.a", "bRank": "ranks.b", "cRank": "ranks.c", "op": "\"add\"", "cDtype": "tensorDtypes.c", "itemsPerInvocation": 4 }, "bindings": ["a_2", "b_2", "c_3", "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_binary": { "buffer": "storage", "elementType": "$vectorScalar", "name": "c" }, "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" }, "a_3": { "buffer": "read-only-storage", "name": "a", "elementType": "$aElement" }, "b_3": { "buffer": "read-only-storage", "name": "b", "elementType": "$bElement" }, "c_2_binary": { "buffer": "storage", "name": "c", "elementType": "$vec4Scalar" }, "c_3": { "buffer": "storage", "name": "c", "elementType": "$scalar" }, "params_2": { "buffer": "uniform", "name": "params", "struct": [{ "name": "count", "type": "u32", "value": "numel(shapes.c)" }] } } }