Download build/webgpu/manifest.json from webgpu-kernels/ai.onnx.SimplifiedLayerNormalization: direct link, hf CLI and curl.
- Browser
- Download file 14.2 kB
-
https://huggingface.co/kernels/webgpu-kernels/ai.onnx.SimplifiedLayerNormalization/resolve/v1/build/webgpu/manifest.json
- Command line
-
hf download hf://webgpu-kernels/ai.onnx.SimplifiedLayerNormalization@v1/build/webgpu/manifest.json
-
curl -L -o manifest.json https://huggingface.co/kernels/webgpu-kernels/ai.onnx.SimplifiedLayerNormalization/resolve/v1/build/webgpu/manifest.json
14.2 kB
| { | |
| "domain": "ai.onnx", | |
| "name": "SimplifiedLayerNormalization", | |
| "conformance": "legacy-default-domain", | |
| "sinceVersion": 1, | |
| "inputs": { "x": { "onnx": "X", "dtype": "T" }, "scale": { "dtype": "V" } }, | |
| "outputs": { | |
| "y": { "onnx": "Y", "dtype": "V", "rank": "ranks.x", "shape": "shapes.x" }, | |
| "invStdVar": { | |
| "onnx": "inv_std_var", | |
| "dtype": "U", | |
| "rank": "ranks.x", | |
| "optional": true, | |
| "shape": "prefix(shapes.x, axisNorm) + fill(1, ranks.x - axisNorm)" | |
| } | |
| }, | |
| "attributes": { | |
| "axis": { "default": -1 }, | |
| "epsilon": { "default": 0.00001 }, | |
| "stash_type": { "default": 1 }, | |
| "keep_dims": { "default": 1 } | |
| }, | |
| "attributeConstraints": { "stash_type": { "values": [1] }, "keep_dims": { "values": [1] } }, | |
| "typeConstraints": { "T": ["float32", "float16"], "V": ["float32", "float16"], "U": ["float32"] }, | |
| "tunables": { | |
| "WORKGROUP_SIZE": { "default": 256 }, | |
| "SPLIT_MAX_ROWS": { "default": 256 }, | |
| "SPLIT_MIN_HIDDEN": { "default": 16384 }, | |
| "SPLIT_TARGET_ELEMENTS": { "default": 4096 }, | |
| "MAX_SPLITS": { "default": 64 } | |
| }, | |
| "derive": { | |
| "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)", | |
| "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32", | |
| "reportedNonWave32Adapter": "not wave32Adapter and (has(device.adapterInfo, \"subgroupMinSize\") or has(device.adapterInfo, \"subgroupMaxSize\"))", | |
| "normMaxWorkgroup": "min(tunables.WORKGROUP_SIZE, deviceWorkgroupCap)", | |
| "hasSubgroupId": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")", | |
| "axisNorm": "attrs.axis if attrs.axis >= 0 else attrs.axis + ranks.x", | |
| "normRows": "outer(shapes.x, axisNorm)", | |
| "normHidden": "dim(shapes.x, axisNorm) * inner(shapes.x, axisNorm)", | |
| "normRowStride": "max(1, min(normRows, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)))", | |
| "rowWg": "min(normMaxWorkgroup, pow2ceil(max(1, normHidden)))", | |
| "baseOk": "ranks.x >= 1 and sameShape(shapes.y, shapes.x) and ranks.scale >= 0 and ranks.scale <= ranks.x and broadcastable(shapes.scale, shapes.x) and attrs.axis + ranks.x >= 0 and attrs.axis < ranks.x and normHidden > 0 and attrs.stash_type == onnxDtypeCode(\"float32\") and f16Ok(dtypes.T) and f16Ok(dtypes.V)", | |
| "lastAxisOk": "baseOk and (attrs.axis == -1 or attrs.axis == ranks.x - 1)", | |
| "noStats": "not present.invStdVar", | |
| "statsOk": "present.invStdVar and ranks.invStdVar == ranks.x and sameShape(prefix(shapes.invStdVar, axisNorm), prefix(shapes.x, axisNorm)) and numel(suffix(shapes.invStdVar, axisNorm)) == 1", | |
| "sameDtype": "dtypes.T == dtypes.V", | |
| "splitCount": "min(tunables.MAX_SPLITS, pow2ceil(ceilDiv(normHidden, tunables.SPLIT_TARGET_ELEMENTS)))", | |
| "splitScratchBytes": "normRows * splitCount * 4", | |
| "splitFits": "normRows <= tunables.SPLIT_MAX_ROWS and splitCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and splitScratchBytes <= device.limits.maxStorageBufferBindingSize and splitScratchBytes <= device.limits.maxBufferSize" | |
| }, | |
| "bindings": { | |
| "x": { "elementType": "$xElement" }, | |
| "scale": { "elementType": "$ioElement" }, | |
| "y": { "elementType": "$ioElement" }, | |
| "params": { | |
| "struct": [ | |
| { "name": "rows", "type": "u32", "value": "normRows" }, | |
| { "name": "rowStride", "type": "u32", "value": "normRowStride" } | |
| ] | |
| }, | |
| "inv_std_out": { "arg": "invStdVar", "elementType": "f32" }, | |
| "partials_f32": { "name": "partials", "buffer": "read-only-storage", "elementType": "f32" } | |
| }, | |
| "variants": [ | |
| { | |
| "id": "last_axis", | |
| "priority": 1, | |
| "when": ["baseOk", "noStats"], | |
| "derive": { | |
| "scalar": "dtypes.V", | |
| "xElement": "dtypes.T", | |
| "ioElement": "dtypes.V", | |
| "hiddenSize": "normHidden", | |
| "workgroupSize": "rowWg", | |
| "epsilon": "attrs.epsilon" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "SimplifiedLayerNormalization.Row", | |
| "shader": "rms-normalization.wgsl.jinja", | |
| "derive": { | |
| "xShape": "shapes.x", | |
| "scaleShape": "shapes.scale", | |
| "xRank": "ranks.x", | |
| "scaleRank": "ranks.scale", | |
| "writeStats": "present.invStdVar" | |
| }, | |
| "bindings": ["x", "scale", "y", "params"], | |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "last_axis_stats", | |
| "priority": 2, | |
| "when": ["baseOk", "statsOk"], | |
| "derive": { | |
| "scalar": "dtypes.V", | |
| "xElement": "dtypes.T", | |
| "ioElement": "dtypes.V", | |
| "hiddenSize": "normHidden", | |
| "workgroupSize": "rowWg", | |
| "epsilon": "attrs.epsilon" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "SimplifiedLayerNormalization.Row", | |
| "shader": "rms-normalization.wgsl.jinja", | |
| "derive": { | |
| "xShape": "shapes.x", | |
| "scaleShape": "shapes.scale", | |
| "xRank": "ranks.x", | |
| "scaleRank": "ranks.scale", | |
| "writeStats": "present.invStdVar" | |
| }, | |
| "bindings": ["x", "scale", "y", "inv_std_out", "params"], | |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "suffix_axis_splitk", | |
| "priority": 15, | |
| "when": ["baseOk", "ranks.x >= 2", "noStats", "splitFits"], | |
| "demoteWhen": ["reportedNonWave32Adapter", "normHidden < tunables.SPLIT_MIN_HIDDEN"], | |
| "derive": { | |
| "scalar": "dtypes.V", | |
| "xElement": "dtypes.T", | |
| "ioElement": "dtypes.V", | |
| "hiddenSize": "normHidden", | |
| "workgroupSize": "normMaxWorkgroup", | |
| "split": "splitCount", | |
| "epsilon": "attrs.epsilon", | |
| "normalizeRows": "normRows" | |
| }, | |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[normRows * splitCount]" }], | |
| "passes": [ | |
| { | |
| "id": "partials", | |
| "name": "SimplifiedLayerNormalization.SplitKPartials", | |
| "shader": "rms-normalization-splitk-partials.wgsl.jinja", | |
| "bindings": ["x", { "name": "partials", "elementType": "f32" }, "params"], | |
| "dispatch": { | |
| "x": "min(normRows, DISPATCH_FOLD_WIDTH)", | |
| "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)", | |
| "z": "splitCount" | |
| } | |
| }, | |
| { | |
| "id": "normalize", | |
| "name": "SimplifiedLayerNormalization.SplitKNormalize", | |
| "shader": "rms-normalization-splitk-normalize.wgsl.jinja", | |
| "derive": { | |
| "xShape": "shapes.x", | |
| "scaleShape": "shapes.scale", | |
| "xRank": "ranks.x", | |
| "scaleRank": "ranks.scale", | |
| "writeStats": "present.invStdVar", | |
| "normalizeBlocks": "max(split, min(min(device.limits.maxComputeWorkgroupsPerDimension, 65535), ceilDiv(workgroupSize, max(1, normalizeRows)), ceilDiv(hiddenSize, workgroupSize * 4)))", | |
| "normalizeChunk": "ceilDiv(hiddenSize, normalizeBlocks)" | |
| }, | |
| "bindings": ["x", "scale", "partials_f32", "y", "params"], | |
| "dispatch": { | |
| "x": "min(normRows, DISPATCH_FOLD_WIDTH)", | |
| "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)", | |
| "z": "normalizeBlocks" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "suffix_axis_splitk_stats", | |
| "priority": 16, | |
| "when": ["baseOk", "ranks.x >= 2", "statsOk", "splitFits"], | |
| "demoteWhen": ["reportedNonWave32Adapter", "normHidden < tunables.SPLIT_MIN_HIDDEN"], | |
| "derive": { | |
| "scalar": "dtypes.V", | |
| "xElement": "dtypes.T", | |
| "ioElement": "dtypes.V", | |
| "hiddenSize": "normHidden", | |
| "workgroupSize": "normMaxWorkgroup", | |
| "split": "splitCount", | |
| "epsilon": "attrs.epsilon", | |
| "normalizeRows": "normRows" | |
| }, | |
| "intermediates": [{ "id": "partials", "dtype": "float32", "shape": "[normRows * splitCount]" }], | |
| "passes": [ | |
| { | |
| "id": "partials", | |
| "name": "SimplifiedLayerNormalization.SplitKPartials", | |
| "shader": "rms-normalization-splitk-partials.wgsl.jinja", | |
| "bindings": ["x", { "name": "partials", "elementType": "f32" }, "params"], | |
| "dispatch": { | |
| "x": "min(normRows, DISPATCH_FOLD_WIDTH)", | |
| "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)", | |
| "z": "splitCount" | |
| } | |
| }, | |
| { | |
| "id": "normalize", | |
| "name": "SimplifiedLayerNormalization.SplitKNormalize", | |
| "shader": "rms-normalization-splitk-normalize.wgsl.jinja", | |
| "derive": { | |
| "xShape": "shapes.x", | |
| "scaleShape": "shapes.scale", | |
| "xRank": "ranks.x", | |
| "scaleRank": "ranks.scale", | |
| "writeStats": "present.invStdVar", | |
| "normalizeBlocks": "max(split, min(min(device.limits.maxComputeWorkgroupsPerDimension, 65535), ceilDiv(workgroupSize, max(1, normalizeRows)), ceilDiv(hiddenSize, workgroupSize * 4)))", | |
| "normalizeChunk": "ceilDiv(hiddenSize, normalizeBlocks)" | |
| }, | |
| "bindings": ["x", "scale", "partials_f32", "y", "inv_std_out", "params"], | |
| "dispatch": { | |
| "x": "min(normRows, DISPATCH_FOLD_WIDTH)", | |
| "y": "ceilDiv(normRows, DISPATCH_FOLD_WIDTH)", | |
| "z": "normalizeBlocks" | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "last_axis_row_vec4", | |
| "priority": 110, | |
| "when": ["lastAxisOk", "sameDtype", "noStats", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "dim(shapes.x, -1) % 4 == 0"], | |
| "derive": { | |
| "scalar": "dtypes.T", | |
| "xElement": "\"vec4<\" ~ dtypes.T ~ \">\"", | |
| "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "SimplifiedLayerNormalization.LastAxisRow", | |
| "shader": "norm-row-stats.wgsl.jinja", | |
| "derive": { | |
| "vec4": true, | |
| "writeStats": "present.invStdVar", | |
| "hidden": "dim(shapes.x, -1)", | |
| "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1) / 4)))", | |
| "epsilon": "attrs.epsilon", | |
| "hiddenVec": "dim(shapes.x, -1) / 4", | |
| "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"", | |
| "combineSubgroups": "hasSubgroupId" | |
| }, | |
| "bindings": ["x", "scale", "y", "params"], | |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "last_axis_row", | |
| "priority": 100, | |
| "when": ["lastAxisOk", "sameDtype", "noStats", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)"], | |
| "derive": { "scalar": "dtypes.T", "xElement": "dtypes.T", "ioElement": "dtypes.T" }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "SimplifiedLayerNormalization.LastAxisRow", | |
| "shader": "norm-row-stats.wgsl.jinja", | |
| "derive": { | |
| "vec4": false, | |
| "writeStats": "present.invStdVar", | |
| "hidden": "dim(shapes.x, -1)", | |
| "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1))))", | |
| "epsilon": "attrs.epsilon", | |
| "hiddenVec": 1, | |
| "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"", | |
| "combineSubgroups": "hasSubgroupId" | |
| }, | |
| "bindings": ["x", "scale", "y", "params"], | |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "last_axis_row_vec4_stats", | |
| "priority": 112, | |
| "when": ["lastAxisOk", "sameDtype", "statsOk", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)", "dim(shapes.x, -1) % 4 == 0"], | |
| "derive": { | |
| "scalar": "dtypes.T", | |
| "xElement": "\"vec4<\" ~ dtypes.T ~ \">\"", | |
| "ioElement": "\"vec4<\" ~ dtypes.T ~ \">\"" | |
| }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "SimplifiedLayerNormalization.LastAxisRow", | |
| "shader": "norm-row-stats.wgsl.jinja", | |
| "derive": { | |
| "vec4": true, | |
| "writeStats": "present.invStdVar", | |
| "hidden": "dim(shapes.x, -1)", | |
| "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1) / 4)))", | |
| "epsilon": "attrs.epsilon", | |
| "hiddenVec": "dim(shapes.x, -1) / 4", | |
| "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"", | |
| "combineSubgroups": "hasSubgroupId" | |
| }, | |
| "bindings": ["x", "scale", "y", "inv_std_out", "params"], | |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "last_axis_row_stats", | |
| "priority": 102, | |
| "when": ["lastAxisOk", "sameDtype", "statsOk", "ranks.scale >= 1", "numel(shapes.scale) == dim(shapes.x, -1)", "dim(shapes.scale, -1) == dim(shapes.x, -1)"], | |
| "derive": { "scalar": "dtypes.T", "xElement": "dtypes.T", "ioElement": "dtypes.T" }, | |
| "passes": [ | |
| { | |
| "id": "main", | |
| "name": "SimplifiedLayerNormalization.LastAxisRow", | |
| "shader": "norm-row-stats.wgsl.jinja", | |
| "derive": { | |
| "vec4": false, | |
| "writeStats": "present.invStdVar", | |
| "hidden": "dim(shapes.x, -1)", | |
| "wg": "min(normMaxWorkgroup, pow2ceil(max(1, dim(shapes.x, -1))))", | |
| "epsilon": "attrs.epsilon", | |
| "hiddenVec": 1, | |
| "vecType": "\"vec4<\" ~ dtypes.T ~ \">\"", | |
| "combineSubgroups": "hasSubgroupId" | |
| }, | |
| "bindings": ["x", "scale", "y", "inv_std_out", "params"], | |
| "dispatch": { "x": "min(normRows, 65535)", "y": "ceilDiv(normRows, 65535)", "z": 1 } | |
| } | |
| ] | |
| } | |
| ] | |
| } | |