Xenova's picture
Xenova HF Staff
sync 91d990483a17
8b49b15 verified
Raw
History Blame
22.4 kB
{
"domain": "ai.onnx",
"name": "QuantizeLinear",
"sinceVersion": 25,
"inputs": {
"x": { "dtype": "TX" },
"y_scale": { "dtype": "TS" },
"y_zero_point": { "dtype": "TQ", "optional": true }
},
"outputs": { "y": { "dtype": "TQ", "rank": "ranks.x", "shape": "shapes.x" } },
"attributes": {
"axis": { "default": 1 },
"block_size": { "default": 0 },
"output_dtype": { "default": 0 },
"precision": { "default": 0 },
"saturate": { "default": 1 }
},
"attributeConstraints": { "precision": { "values": [0, 1, 10] }, "saturate": { "values": [1] } },
"typeConstraints": { "TX": ["float32", "float16"], "TS": ["float32", "float16"], "TQ": ["uint8", "int8"] },
"tunables": { "WORKGROUP_SIZE": { "default": 256 }, "VEC4_TAIL_MIN_ELEMENTS": { "default": 4096 } },
"derive": {
"deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
"workgroupOk": "tunables.WORKGROUP_SIZE <= deviceWorkgroupCap",
"outputDispatchFits": "workgroupOk and ceilDiv(ceilDiv(numel(shapes.y), tunables.WORKGROUP_SIZE), min(device.limits.maxComputeWorkgroupsPerDimension, 65535)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)",
"outputDtypeOk": "(attrs.output_dtype == 0 and (present.y_zero_point or tensorDtypes.y == \"uint8\")) or attrs.output_dtype == onnxDtypeCode(logicalDtypes.TQ)",
"quantizeDivisionF16": "attrs.precision == onnxDtypeCode(\"float16\") or (attrs.precision == 0 and tensorDtypes.y_scale == \"float16\")",
"sameShapeOk": "ranks.y == ranks.x and numel(shapes.x) == numel(shapes.y) and outputDispatchFits and outputDtypeOk",
"quantDtypesOk": "(tensorDtypes.x != \"float16\" and tensorDtypes.y_scale != \"float16\" and not quantizeDivisionF16) or device.features.has(\"shader-f16\")",
"blockedScaleOk": "attrs.block_size > 0 and ranks.x >= 2 and ranks.y_scale == ranks.x",
"elementCount": "numel(shapes.y)",
"elementCount4": "floor(numel(shapes.y) / 4)",
"scaleSize": "1 if ranks.y_scale == 0 else dim(shapes.y_scale, 0)",
"scaleInner": "1 if ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1 else inner(shapes.x, attrs.axis)",
"blockedAxisDim": "dim(shapes.x, attrs.axis) if ranks.x >= 2 else 1",
"blockedScaleAxisDim": "dim(shapes.y_scale, attrs.axis) if ranks.y_scale >= 2 else 1",
"blockedInner": "inner(shapes.x, attrs.axis) if ranks.x >= 2 else 1",
"xVec4": "\"vec4<\" ~ dtypes.TX ~ \">\"",
"scaleVec4": "\"vec4<\" ~ dtypes.TS ~ \">\"",
"yScalar": "dtypes.TQ",
"yVec4": "\"vec4<\" ~ dtypes.TQ ~ \">\"",
"yUnsigned": "tensorDtypes.y == \"uint8\"",
"qMin": "0 if tensorDtypes.y == \"uint8\" else 0 - 128",
"qMax": "255 if tensorDtypes.y == \"uint8\" else 127",
"divisionF16": "quantizeDivisionF16"
},
"when": ["sameShapeOk"],
"bindings": {
"x": { "buffer": "read-only-storage", "elementType": "$xVec4" },
"y_scale": { "buffer": "read-only-storage", "elementType": "$scaleVec4" },
"y_zero_point": { "buffer": "read-only-storage", "elementType": "$yVec4" },
"y": { "buffer": "storage", "elementType": "$yVec4" },
"params": {
"buffer": "uniform",
"struct": [
{ "name": "count4", "type": "u32", "value": "elementCount4" },
{ "name": "scaleSize", "type": "u32", "value": "scaleSize" }
]
},
"params_2": {
"name": "params",
"buffer": "uniform",
"struct": [
{ "name": "count4", "type": "u32", "value": "elementCount4" },
{ "name": "scaleSize", "type": "u32", "value": "scaleSize" },
{ "name": "inner", "type": "u32", "value": "scaleInner" }
]
},
"params_3": {
"name": "params",
"buffer": "uniform",
"struct": [{ "name": "count4", "type": "u32", "value": "elementCount4" }]
},
"params_4": {
"name": "params",
"buffer": "uniform",
"struct": [{ "name": "count", "type": "u32", "value": "elementCount" }]
},
"x_3": { "name": "x", "buffer": "read-only-storage", "elementType": "vec4<f32>" },
"y_scale_3": { "name": "y_scale", "buffer": "read-only-storage", "elementType": "f32" },
"params_5": {
"name": "params",
"buffer": "uniform",
"struct": [{ "name": "count", "type": "u32", "value": "elementCount4" }]
},
"y_scale_4": { "name": "y_scale", "buffer": "read-only-storage", "elementType": "vec4<f32>" },
"params_6": {
"name": "params",
"buffer": "uniform",
"struct": [
{ "name": "count", "type": "u32", "value": "elementCount4" },
{ "name": "axisDim", "type": "u32", "value": "blockedAxisDim" },
{ "name": "scaleAxisDim", "type": "u32", "value": "blockedScaleAxisDim" },
{ "name": "inner", "type": "u32", "value": "blockedInner" },
{ "name": "blockSize", "type": "u32", "value": "attrs.block_size" }
]
},
"params_7": {
"name": "params",
"buffer": "uniform",
"struct": [
{ "name": "count", "type": "u32", "value": "elementCount" },
{ "name": "axisDim", "type": "u32", "value": "blockedAxisDim" },
{ "name": "scaleAxisDim", "type": "u32", "value": "blockedScaleAxisDim" },
{ "name": "inner", "type": "u32", "value": "blockedInner" },
{ "name": "blockSize", "type": "u32", "value": "attrs.block_size" }
]
},
"params_8": {
"name": "params",
"buffer": "uniform",
"struct": [
{ "name": "count", "type": "u32", "value": "elementCount" },
{ "name": "scaleSize", "type": "u32", "value": "scaleSize" },
{ "name": "inner", "type": "u32", "value": "scaleInner" }
]
}
},
"variants": [
{
"id": "innermost_axis_vec4_with_zero_point",
"priority": 20,
"when": ["present.y_zero_point", "ranks.x >= 1", "attrs.axis == -1 or attrs.axis == ranks.x - 1", "ranks.y_scale == 1", "dim(shapes.y_scale, 0) == dim(shapes.x, ranks.x - 1)", "dim(shapes.y_scale, 0) % 4 == 0", "ranks.y_zero_point == 1", "dim(shapes.y_zero_point, 0) == dim(shapes.y_scale, 0)", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "quantDtypesOk"],
"derive": { "hasZero": true },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.InnermostAxisVec4",
"shader": "quant-linear-vec4.wgsl.jinja",
"derive": { "op": "\"quantize\"", "perAxis": true, "vectorParams": true },
"bindings": ["x", "y_scale", "y_zero_point", "y", "params"],
"dispatch": {
"x": "min(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "innermost_axis_vec4_no_zero_point",
"priority": 20,
"when": ["not present.y_zero_point", "ranks.x >= 1", "attrs.axis == -1 or attrs.axis == ranks.x - 1", "ranks.y_scale == 1", "dim(shapes.y_scale, 0) == dim(shapes.x, ranks.x - 1)", "dim(shapes.y_scale, 0) % 4 == 0", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "quantDtypesOk"],
"derive": { "hasZero": false },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.InnermostAxisVec4NoZero",
"shader": "quant-linear-vec4.wgsl.jinja",
"derive": { "op": "\"quantize\"", "perAxis": true, "vectorParams": true },
"bindings": ["x", "y_scale", "y", "params"],
"dispatch": {
"x": "min(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "vec4_cross_axis_with_zero_point",
"priority": 16,
"when": ["present.y_zero_point", "attrs.block_size == 0", "ranks.x >= 1", "attrs.axis >= 0", "attrs.axis < ranks.x", "ranks.y_scale == 1", "dim(shapes.y_scale, 0) == dim(shapes.x, attrs.axis)", "inner(shapes.x, attrs.axis) % 4 != 0", "ranks.y_zero_point == 1", "dim(shapes.y_zero_point, 0) == dim(shapes.y_scale, 0)", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "quantDtypesOk"],
"derive": { "hasZero": true },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.Vec4CrossAxis",
"shader": "quant-linear-vec4.wgsl.jinja",
"derive": { "op": "\"quantize\"", "perAxis": true, "crossingParams": true },
"bindings": ["x", { "arg": "y_scale" }, { "arg": "y_zero_point" }, "y", "params_2"],
"dispatch": {
"x": "min(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "vec4_with_zero_point",
"priority": 15,
"when": ["present.y_zero_point", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "(ranks.y_scale == 0 or ranks.y_scale == 1)", "(ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1 or (ranks.x >= 1 and dim(shapes.y_scale, 0) == dim(shapes.x, attrs.axis) and inner(shapes.x, attrs.axis) % 4 == 0))", "(ranks.y_zero_point == 0 or ranks.y_zero_point == 1)", "(ranks.y_zero_point == 0 or dim(shapes.y_zero_point, 0) == 1 or (ranks.x >= 1 and dim(shapes.y_zero_point, 0) == dim(shapes.x, attrs.axis)))", "quantDtypesOk"],
"derive": { "hasZero": true },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.Vec4",
"shader": "quant-linear-vec4.wgsl.jinja",
"derive": { "op": "\"quantize\"", "perAxis": "not (ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1)" },
"bindings": ["x", { "arg": "y_scale" }, { "arg": "y_zero_point" }, "y", "params_2"],
"dispatch": {
"x": "min(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "vec4_no_zero_point",
"priority": 15,
"when": ["not present.y_zero_point", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "(ranks.y_scale == 0 or ranks.y_scale == 1)", "(ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1 or (ranks.x >= 1 and dim(shapes.y_scale, 0) == dim(shapes.x, attrs.axis) and inner(shapes.x, attrs.axis) % 4 == 0))", "quantDtypesOk"],
"derive": { "hasZero": false },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.Vec4",
"shader": "quant-linear-vec4.wgsl.jinja",
"derive": { "op": "\"quantize\"", "perAxis": "not (ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1)" },
"bindings": ["x", { "arg": "y_scale" }, "y", "params_2"],
"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": "vec4_tail_with_zero_point",
"priority": 12,
"when": ["present.y_zero_point", "numel(shapes.y) >= tunables.VEC4_TAIL_MIN_ELEMENTS", "numel(shapes.y) % 4 != 0", "(ranks.y_scale == 0 or (ranks.y_scale == 1 and dim(shapes.y_scale, 0) == 1))", "(ranks.y_zero_point == 0 or (ranks.y_zero_point == 1 and dim(shapes.y_zero_point, 0) == 1))", "quantDtypesOk"],
"derive": { "hasZero": true },
"passes": [
{
"id": "bulk",
"name": "QuantizeLinear.Vec4Bulk",
"shader": "quant-linear-vec4.wgsl.jinja",
"derive": { "op": "\"quantize\"", "perAxis": false },
"bindings": ["x", { "arg": "y_scale" }, { "arg": "y_zero_point" }, "y", "params_3"],
"dispatch": {
"x": "min(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
},
{
"id": "tail",
"name": "QuantizeLinear.ScalarTail",
"shader": "quant-linear-scalar.wgsl.jinja",
"derive": { "op": "\"quantize\"", "x4": true, "perAxis": false },
"bindings": [{ "arg": "x" }, { "arg": "y_scale" }, { "arg": "y_zero_point" }, { "arg": "y" }, "params_4"],
"dispatch": {
"x": "min(ceilDiv((1), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((1), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "vec4_tail_no_zero_point",
"priority": 12,
"when": ["not present.y_zero_point", "numel(shapes.y) >= tunables.VEC4_TAIL_MIN_ELEMENTS", "numel(shapes.y) % 4 != 0", "(ranks.y_scale == 0 or (ranks.y_scale == 1 and dim(shapes.y_scale, 0) == 1))", "quantDtypesOk"],
"derive": { "hasZero": false },
"passes": [
{
"id": "bulk",
"name": "QuantizeLinear.Vec4Bulk",
"shader": "quant-linear-vec4.wgsl.jinja",
"derive": { "op": "\"quantize\"", "perAxis": false },
"bindings": ["x", { "arg": "y_scale" }, "y", "params_3"],
"dispatch": {
"x": "min(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((floor(numel(shapes.y) / 4)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
},
{
"id": "tail",
"name": "QuantizeLinear.ScalarTail",
"shader": "quant-linear-scalar.wgsl.jinja",
"derive": { "op": "\"quantize\"", "x4": true, "perAxis": false },
"bindings": [{ "arg": "x" }, { "arg": "y_scale" }, { "arg": "y" }, "params_4"],
"dispatch": {
"x": "min(ceilDiv((1), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((1), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "blocked_last_axis_vec4_with_zero_point",
"priority": 24,
"when": ["present.y_zero_point", "blockedScaleOk", "attrs.block_size % 4 == 0", "ranks.y_zero_point == ranks.x", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "inner(shapes.x, attrs.axis) == 1", "dim(shapes.x, attrs.axis) % attrs.block_size == 0", "dim(shapes.y_scale, attrs.axis) == dim(shapes.x, attrs.axis) / attrs.block_size", "numel(shapes.y_zero_point) == numel(shapes.y_scale)", "tensorDtypes.x == \"float32\"", "tensorDtypes.y_scale == \"float32\""],
"derive": { "hasZero": true, "blockVectors": "attrs.block_size / 4" },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.BlockedLastAxisVec4",
"shader": "quant-linear-blocked-axis.wgsl.jinja",
"derive": { "lastAxisVectorized": true },
"bindings": ["x_3", "y_scale_3", { "arg": "y_zero_point" }, "y", "params_5"],
"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": "blocked_last_axis_vec4_no_zero_point",
"priority": 23,
"when": ["not present.y_zero_point", "blockedScaleOk", "attrs.block_size % 4 == 0", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "inner(shapes.x, attrs.axis) == 1", "dim(shapes.x, attrs.axis) % attrs.block_size == 0", "dim(shapes.y_scale, attrs.axis) == dim(shapes.x, attrs.axis) / attrs.block_size", "tensorDtypes.x == \"float32\"", "tensorDtypes.y_scale == \"float32\""],
"derive": { "hasZero": false, "blockVectors": "attrs.block_size / 4" },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.BlockedLastAxisVec4",
"shader": "quant-linear-blocked-axis.wgsl.jinja",
"derive": { "lastAxisVectorized": true },
"bindings": ["x_3", "y_scale_3", "y", "params_5"],
"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": "blocked_vec4_with_zero_point",
"priority": 22,
"when": ["present.y_zero_point", "blockedScaleOk", "ranks.y_zero_point == ranks.x", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "numel(shapes.y_scale) % 4 == 0", "numel(shapes.y_zero_point) == numel(shapes.y_scale)", "inner(shapes.x, attrs.axis) % 4 == 0", "dim(shapes.y_scale, attrs.axis) == ceil(dim(shapes.x, attrs.axis) / attrs.block_size)", "dim(shapes.y_zero_point, attrs.axis) == dim(shapes.y_scale, attrs.axis)", "tensorDtypes.x == \"float32\"", "tensorDtypes.y_scale == \"float32\""],
"derive": { "hasZero": true },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.BlockedVec4WithZeroPoint",
"shader": "quant-linear-blocked-axis.wgsl.jinja",
"derive": { "vectorizedSpec": true },
"bindings": ["x_3", "y_scale_4", "y_zero_point", "y", "params_6"],
"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": "blocked_vec4_no_zero_point",
"priority": 21,
"when": ["not present.y_zero_point", "blockedScaleOk", "numel(shapes.y) > 0", "numel(shapes.y) % 4 == 0", "numel(shapes.y_scale) % 4 == 0", "inner(shapes.x, attrs.axis) % 4 == 0", "dim(shapes.y_scale, attrs.axis) == ceil(dim(shapes.x, attrs.axis) / attrs.block_size)", "tensorDtypes.x == \"float32\"", "tensorDtypes.y_scale == \"float32\""],
"derive": { "hasZero": false },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.BlockedVec4",
"shader": "quant-linear-blocked-axis.wgsl.jinja",
"derive": { "vectorizedSpec": true },
"bindings": ["x_3", "y_scale_4", "y", "params_6"],
"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": "blocked_with_zero_point",
"priority": 19,
"when": ["present.y_zero_point", "blockedScaleOk", "ranks.y_zero_point == ranks.x", "dim(shapes.y_scale, attrs.axis) == ceil(dim(shapes.x, attrs.axis) / attrs.block_size)", "dim(shapes.y_zero_point, attrs.axis) == dim(shapes.y_scale, attrs.axis)", "quantDtypesOk"],
"derive": { "hasZero": true },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.Blocked",
"shader": "quant-linear-blocked-axis.wgsl.jinja",
"bindings": [{ "arg": "x" }, { "arg": "y_scale" }, { "arg": "y_zero_point" }, { "arg": "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": "blocked_no_zero_point",
"priority": 19,
"when": ["not present.y_zero_point", "blockedScaleOk", "dim(shapes.y_scale, attrs.axis) == ceil(dim(shapes.x, attrs.axis) / attrs.block_size)", "quantDtypesOk"],
"derive": { "hasZero": false },
"passes": [
{
"id": "main",
"name": "QuantizeLinear.BlockedNoZero",
"shader": "quant-linear-blocked-axis.wgsl.jinja",
"bindings": [{ "arg": "x" }, { "arg": "y_scale" }, { "arg": "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": "with_zero_point",
"priority": 10,
"when": ["present.y_zero_point", "(ranks.y_scale == 0 or ranks.y_scale == 1)", "(ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1 or (ranks.x >= 1 and dim(shapes.y_scale, 0) == dim(shapes.x, attrs.axis)))", "(ranks.y_zero_point == 0 or ranks.y_zero_point == 1)", "(ranks.y_zero_point == 0 or dim(shapes.y_zero_point, 0) == 1 or (ranks.x >= 1 and dim(shapes.y_zero_point, 0) == dim(shapes.x, attrs.axis)))", "quantDtypesOk"],
"derive": { "hasZero": true },
"passes": [
{
"id": "main",
"name": "QuantizeLinear",
"shader": "quant-linear-scalar.wgsl.jinja",
"derive": {
"op": "\"quantize\"",
"x4": "false",
"perAxis": "not (ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1)"
},
"bindings": [{ "arg": "x" }, { "arg": "y_scale" }, { "arg": "y_zero_point" }, { "arg": "y" }, "params_8"],
"dispatch": {
"x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
},
{
"id": "no_zero_point",
"when": ["not present.y_zero_point", "(ranks.y_scale == 0 or ranks.y_scale == 1)", "(ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1 or (ranks.x >= 1 and dim(shapes.y_scale, 0) == dim(shapes.x, attrs.axis)))", "quantDtypesOk"],
"derive": { "hasZero": false },
"passes": [
{
"id": "main",
"name": "QuantizeLinear",
"shader": "quant-linear-scalar.wgsl.jinja",
"derive": {
"op": "\"quantize\"",
"x4": "false",
"perAxis": "not (ranks.y_scale == 0 or dim(shapes.y_scale, 0) == 1)"
},
"bindings": [{ "arg": "x" }, { "arg": "y_scale" }, { "arg": "y" }, "params_8"],
"dispatch": {
"x": "min(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)",
"y": "ceilDiv(ceilDiv((numel(shapes.y)), (tunables.WORKGROUP_SIZE)), 65535)",
"z": 1
}
}
]
}
]
}