{ "domain": "com.microsoft", "name": "SkipLayerNormalization", "sinceVersion": 1, "inputs": { "inputT": { "onnx": "input", "dtype": "T" }, "skipT": { "onnx": "skip", "dtype": "T" }, "gammaT": { "onnx": "gamma", "dtype": "T", "rank": 1 }, "betaT": { "onnx": "beta", "dtype": "T", "rank": 1, "optional": true }, "biasT": { "onnx": "bias", "dtype": "T", "rank": 1, "optional": true } }, "outputs": { "outputT": { "onnx": "output", "dtype": "T", "rank": "ranks.inputT", "shape": "shapes.inputT" }, "meanT": { "onnx": "mean", "dtype": "U", "optional": true, "rank": "ranks.inputT", "shape": "prefix(shapes.inputT, ranks.inputT - 1) + [1]" }, "invStdT": { "onnx": "inv_std_var", "dtype": "U", "optional": true, "rank": "ranks.inputT", "shape": "prefix(shapes.inputT, ranks.inputT - 1) + [1]" }, "residualT": { "onnx": "input_skip_bias_sum", "dtype": "T", "rank": "ranks.inputT", "optional": true, "shape": "shapes.inputT" } }, "attributes": { "epsilon": { "default": 9.999999960041972e-13 } }, "typeConstraints": { "T": ["float32", "float16"], "U": ["float32"] }, "tunables": { "MAX_WORKGROUP_SIZE": { "default": 256 } }, "derive": { "rowCount": "numel(shapes.inputT) / max(1, dim(shapes.inputT, -1))", "hiddenSize": "dim(shapes.inputT, -1)", "skipWg": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)) if hiddenSize == 1 else (max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(hiddenSize))))", "skipWgVec4": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(ceilDiv(hiddenSize, 4))))", "rowDispatchFits": "rowCount <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "normResourcesFit": "skipWg * 8 <= device.limits.maxComputeWorkgroupStorageSize and skipWgVec4 * 8 <= device.limits.maxComputeWorkgroupStorageSize", "epsilonOk": "attrs.epsilon >= 0", "coreContract": "epsilonOk and (ranks.inputT == 2 or ranks.inputT == 3) and ranks.skipT == ranks.inputT and ranks.gammaT == 1 and ranks.outputT == ranks.inputT and sameShape(shapes.inputT, shapes.skipT) and sameShape(shapes.outputT, shapes.inputT) and dim(shapes.inputT, -1) > 0 and dim(shapes.gammaT, 0) == dim(shapes.inputT, -1)", "residualOutputContract": "present.residualT and sameShape(shapes.residualT, shapes.inputT)", "outputOnlyContract": "not present.residualT", "f32MainDtypes": "tensorDtypes.inputT == \"float32\" and tensorDtypes.skipT == \"float32\" and tensorDtypes.gammaT == \"float32\" and tensorDtypes.outputT == \"float32\"", "f16MainDtypes": "tensorDtypes.inputT == \"float16\" and tensorDtypes.skipT == \"float16\" and tensorDtypes.gammaT == \"float16\" and tensorDtypes.outputT == \"float16\"", "f32ResidualDtypes": "f32MainDtypes and tensorDtypes.residualT == \"float32\" if present.residualT else false", "f16ResidualDtypes": "f16MainDtypes and tensorDtypes.residualT == \"float16\" if present.residualT else false", "vec4Aligned": "dim(shapes.inputT, -1) % 4 == 0", "hasSubgroups": "device.features.has(\"subgroups\") and device.wgslLanguageFeatures.has(\"subgroup_id\")", "hasF16": "device.features.has(\"shader-f16\")", "statsRequested": "present.meanT or present.invStdT", "portableWideExecution": "not has(device.adapterInfo, \"subgroupMinSize\") or device.adapterInfo.subgroupMinSize >= 32", "broadcastRows": "dim(shapes.inputT, 0) * dim(shapes.inputT, 1)", "broadcastHiddenSize": "dim(shapes.inputT, 2)", "broadcastSkipWgVec4": "max(1, min(tunables.MAX_WORKGROUP_SIZE, device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX, pow2ceil(ceilDiv(broadcastHiddenSize, 4))))", "broadcastDispatchFits": "broadcastRows <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) * min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", "broadcastResourcesFit": "broadcastSkipWgVec4 * 8 <= device.limits.maxComputeWorkgroupStorageSize", "betaContract": "false if not present.betaT else (ranks.betaT == 1 and dim(shapes.betaT, 0) == dim(shapes.inputT, -1))", "noBetaContract": "not present.betaT", "broadcastSkipShapeOk": "(ranks.skipT == 2 and dim(shapes.skipT, 0) == dim(shapes.inputT, 1) and dim(shapes.skipT, 1) == dim(shapes.inputT, 2)) or (ranks.skipT == 3 and ((dim(shapes.skipT, 0) == 1 and dim(shapes.skipT, 1) == dim(shapes.inputT, 1) and dim(shapes.skipT, 2) == dim(shapes.inputT, 2)) or sameShape(shapes.skipT, shapes.inputT)))", "broadcastOutputOnlyContract": "false if ranks.inputT != 3 or not present.betaT else (epsilonOk and not present.biasT and not present.residualT and dim(shapes.inputT, 2) % 4 == 0 and broadcastSkipShapeOk and ranks.gammaT == 1 and ranks.betaT == 1 and ranks.outputT == 3 and tensorDtypes.inputT == \"float32\" and tensorDtypes.skipT == \"float32\" and tensorDtypes.gammaT == \"float32\" and tensorDtypes.betaT == \"float32\" and tensorDtypes.outputT == \"float32\" and dim(shapes.inputT, 2) > 0 and dim(shapes.gammaT, 0) == dim(shapes.inputT, 2) and dim(shapes.betaT, 0) == dim(shapes.inputT, 2) and sameShape(shapes.outputT, shapes.inputT))", "f32_beta_no_bias_residual_contract": "coreContract and residualOutputContract and betaContract and f32ResidualDtypes and not present.biasT and tensorDtypes.betaT == \"float32\"", "f32_beta_bias_residual_contract": "false if not present.biasT else (coreContract and residualOutputContract and betaContract and f32ResidualDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float32\" and tensorDtypes.biasT == \"float32\" and dim(shapes.biasT, 0) == hiddenSize)", "f16_beta_bias_residual_contract": "false if not present.biasT else (hasF16 and coreContract and residualOutputContract and betaContract and f16ResidualDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float16\" and tensorDtypes.biasT == \"float16\" and dim(shapes.biasT, 0) == hiddenSize)", "f32_no_beta_output_contract": "coreContract and outputOnlyContract and noBetaContract and f32MainDtypes and not present.biasT", "f32_beta_no_bias_output_only_contract": "coreContract and outputOnlyContract and betaContract and f32MainDtypes and not present.biasT and tensorDtypes.betaT == \"float32\"", "f32_beta_bias_output_only_contract": "false if not present.biasT else (coreContract and outputOnlyContract and betaContract and f32MainDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float32\" and tensorDtypes.biasT == \"float32\" and dim(shapes.biasT, 0) == hiddenSize)", "f16_no_beta_output_contract": "hasF16 and coreContract and outputOnlyContract and noBetaContract and f16MainDtypes and not present.biasT", "f16_beta_no_bias_output_only_contract": "hasF16 and coreContract and outputOnlyContract and betaContract and f16MainDtypes and not present.biasT and tensorDtypes.betaT == \"float16\"", "f16_beta_bias_output_only_contract": "false if not present.biasT else (hasF16 and coreContract and outputOnlyContract and betaContract and f16MainDtypes and ranks.biasT == 1 and tensorDtypes.betaT == \"float16\" and tensorDtypes.biasT == \"float16\" and dim(shapes.biasT, 0) == hiddenSize)", "statsContract": "statsRequested and coreContract and f16Ok(dtypes.T) and (not present.biasT or (ranks.biasT == 1 and dim(shapes.biasT, 0) == hiddenSize)) and (not present.residualT or sameShape(shapes.residualT, shapes.inputT)) and (not present.betaT or (ranks.betaT == 1 and dim(shapes.betaT, 0) == hiddenSize))" }, "bindings": { "input": { "arg": "inputT", "elementType": "$vectorScalar" }, "skip": { "arg": "skipT", "elementType": "$vectorScalar" }, "gamma": { "arg": "gammaT", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" }, "beta": { "arg": "betaT", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" }, "output": { "arg": "outputT", "elementType": "$vectorScalar" }, "bias": { "arg": "biasT", "elementType": "$vectorScalar", "length": "$HIDDEN_LEN" }, "input_skip_bias_sum": { "arg": "residualT", "elementType": "$vectorScalar" }, "params_main": { "name": "params", "struct": [ { "name": "rows", "type": "u32", "value": "rowCount" }, { "name": "rowStride", "type": "u32", "value": "max(1, min(rowCount, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)))" }, { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" } ] }, "input_main": { "arg": "inputT", "name": "input", "elementType": "$scalar" }, "skip_main": { "arg": "skipT", "name": "skip", "elementType": "$scalar" }, "bias_main": { "arg": "biasT", "name": "bias", "elementType": "$scalar", "length": "$HIDDEN_LEN" }, "gamma_main": { "arg": "gammaT", "name": "gamma", "elementType": "$scalar", "length": "$HIDDEN_LEN" }, "beta_main": { "arg": "betaT", "name": "beta", "elementType": "$scalar", "length": "$HIDDEN_LEN" }, "output_main": { "arg": "outputT", "name": "output", "elementType": "$scalar" }, "input_skip_bias_sum_main": { "arg": "residualT", "name": "input_skip_bias_sum", "elementType": "$scalar" }, "params__uniform": { "name": "params", "struct": [ { "name": "rows", "type": "u32", "value": "rowCount" }, { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" } ] }, "mean": { "arg": "meanT", "elementType": "f32" }, "inv_std_var": { "arg": "invStdT", "elementType": "f32" }, "row_stats": { "scratch": "rowStats", "elementType": "vec2" }, "row_stats_read": { "scratch": "rowStats", "name": "row_stats", "buffer": "read-only-storage", "elementType": "vec2" }, "params_stats": { "name": "params", "struct": [{ "name": "rows", "type": "u32", "value": "rowCount" }] } }, "variants": [ { "id": "hidden1_f32_no_beta", "priority": 20, "when": ["f32_no_beta_output_contract", "hiddenSize == 1", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": false, "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": false, "scalar": "dtypes.T", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Hidden1", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["gamma_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "hidden1_f32_beta", "priority": 20, "when": ["(f32_beta_no_bias_output_only_contract or f32_beta_bias_output_only_contract)", "hiddenSize == 1", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": false, "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": false, "scalar": "dtypes.T", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Hidden1", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["gamma_main", "beta_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "hidden1_f16_no_beta", "priority": 20, "when": ["f16_no_beta_output_contract", "hiddenSize == 1", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": false, "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": false, "scalar": "dtypes.T", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Hidden1", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["gamma_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "hidden1_f16_beta", "priority": 20, "when": ["(f16_beta_no_bias_output_only_contract or f16_beta_bias_output_only_contract)", "hiddenSize == 1", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": false, "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": false, "scalar": "dtypes.T", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Hidden1", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["gamma_main", "beta_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "beta_output_only_vec4_broadcast", "priority": 19, "when": ["broadcastOutputOnlyContract", "broadcastResourcesFit", "broadcastDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "HIDDEN_LEN": "broadcastHiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.BroadcastSkip", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBias": false, "hasBeta": true, "writeResidualSum": false, "broadcastSkip": true, "hidden": "broadcastHiddenSize", "hiddenVec": "broadcastHiddenSize / 4", "wg": "broadcastSkipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": [ "input", "skip", "gamma", "beta", "output", { "name": "params", "struct": [ { "name": "rows", "type": "u32", "value": "dim(shapes.inputT, 0) * dim(shapes.inputT, 1)" }, { "name": "rowStride", "type": "u32", "value": "max(1, min(broadcastRows, min(device.limits.maxComputeWorkgroupsPerDimension, 65535)))" }, { "name": "epsilon", "type": "f32", "value": "attrs.epsilon" }, { "name": "skipRows", "type": "u32", "value": "numel(shapes.skipT) / broadcastHiddenSize" } ] } ], "dispatch": { "x": "min(broadcastRows, 65535)", "y": "ceilDiv(broadcastRows, 65535)", "z": 1 } } ] }, { "id": "beta_bias_vec4", "priority": 15, "when": ["f32_beta_bias_residual_contract", "vec4Aligned", "hasSubgroups or \"bias\" == \"bias\" or portableWideExecution", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "hasBias": "\"bias\" == \"bias\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "normalize", "name": "SkipLayerNormalization.Vec4.Normalize", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBeta": true, "writeResidualSum": true, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "bias", "gamma", "beta", "output", "input_skip_bias_sum", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "beta_bias_row", "priority": 5, "when": ["f32_beta_bias_residual_contract", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "useSubgroups": "hasSubgroups", "hasBeta": true, "writeResidualSum": true, "hasBias": "\"bias\" == \"bias\"", "scalar": "\"f32\"", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "normalize", "name": "SkipLayerNormalization.Row.Normalize", "shader": "norm-skip-row.wgsl.jinja", "derive": { "writeResidualSum": true }, "bindings": ["input_main", "skip_main", "bias_main", "gamma_main", "beta_main", "output_main", "input_skip_bias_sum_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "beta_no_bias_vec4", "priority": 20, "when": ["f32_beta_no_bias_residual_contract", "vec4Aligned", "hasSubgroups or portableWideExecution", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "hasBias": "\"no_bias\" == \"bias\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Vec4", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBeta": true, "writeResidualSum": true, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "gamma", "beta", "output", "input_skip_bias_sum", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "beta_no_bias_row", "priority": 10, "when": ["f32_beta_no_bias_residual_contract", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "useSubgroups": "hasSubgroups", "hasBeta": true, "writeResidualSum": true, "hasBias": "\"no_bias\" == \"bias\"", "scalar": "\"f32\"", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Row", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "input_skip_bias_sum_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "beta_bias_vec4_f16", "priority": 21, "when": ["f16_beta_bias_residual_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Vec4", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBias": true, "hasBeta": true, "writeResidualSum": true, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "bias", "gamma", "beta", "output", "input_skip_bias_sum", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "no_beta_output_only_vec4", "priority": 20, "when": ["f32_no_beta_output_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.NoBetaOutputOnly.Vec4", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "gamma", "output", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "no_beta_output_only_row", "priority": 10, "when": ["f32_no_beta_output_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": "hasSubgroups", "scalar": "\"f32\"", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.NoBetaOutputOnly.Row", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "no_beta_output_only_vec4_f16", "priority": 20, "when": ["f16_no_beta_output_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.NoBetaOutputOnly.Vec4.F16", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "gamma", "output", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "no_beta_output_only_row_f16", "priority": 10, "when": ["f16_no_beta_output_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": "hasSubgroups", "scalar": "\"f16\"", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.NoBetaOutputOnly.Row.F16", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "beta_no_bias_output_only_vec4", "priority": 20, "when": ["f32_beta_no_bias_output_only_contract", "vec4Aligned", "hasSubgroups or portableWideExecution", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Vec4", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "gamma", "beta", "output", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "beta_no_bias_output_only_row", "priority": 10, "when": ["f32_beta_no_bias_output_only_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": "hasSubgroups", "scalar": "\"f32\"", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Row", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "beta_no_bias_output_only_vec4_f16", "priority": 20, "when": ["f16_beta_no_bias_output_only_contract", "vec4Aligned", "hasSubgroups or portableWideExecution", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Vec4.F16", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "gamma", "beta", "output", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "beta_no_bias_output_only_row_f16", "priority": 10, "when": ["f16_beta_no_bias_output_only_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": "hasSubgroups", "scalar": "\"f16\"", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Row.F16", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "beta_bias_output_only_vec4", "priority": 20, "when": ["f32_beta_bias_output_only_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Vec4", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "gamma", "beta", "bias", "output", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "beta_bias_output_only_row", "priority": 10, "when": ["f32_beta_bias_output_only_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": "hasSubgroups", "scalar": "\"f32\"", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Row", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "bias_main", "gamma_main", "beta_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "beta_bias_output_only_vec4_f16", "priority": 20, "when": ["f16_beta_bias_output_only_contract", "vec4Aligned", "normResourcesFit", "rowDispatchFits", "not present.meanT and not present.invStdT"], "derive": { "vectorScalar": "\"vec4\"", "HIDDEN_LEN": "hiddenSize / 4" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Vec4.F16", "shader": "norm-skip-row-vec4.wgsl.jinja", "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "hidden": "hiddenSize", "hiddenVec": "hiddenSize / 4", "wg": "skipWgVec4", "vecType": "\"vec4\"", "useSubgroups": "hasSubgroups" }, "bindings": ["input", "skip", "gamma", "beta", "bias", "output", "params_main"], "dispatch": { "x": "min(rowCount, 65535)", "y": "ceilDiv(rowCount, 65535)", "z": 1 } } ] }, { "id": "beta_bias_output_only_row_f16", "priority": 10, "when": ["f16_beta_bias_output_only_contract", "normResourcesFit", "rowDispatchFits", "hiddenSize != 1", "not present.meanT and not present.invStdT"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": false, "useSubgroups": "hasSubgroups", "scalar": "\"f16\"", "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize" }, "passes": [ { "id": "main", "name": "SkipLayerNormalization.Row.F16", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "bias_main", "gamma_main", "beta_main", "output_main", "params__uniform"], "dispatch": { "x": "min(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "y": "ceilDiv(ceilDiv((rowCount * (1 if hiddenSize == 1 else skipWg)), (skipWg)), 65535)", "z": 1 } } ] }, { "id": "stats_mean_plain", "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "mean", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_mean_residual", "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "mean", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_mean_beta", "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "mean", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_mean_beta_residual", "when": ["statsContract", "present.meanT", "not present.invStdT", "not present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "input_skip_bias_sum_main", "output_main", "mean", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_mean_bias", "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "mean", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_mean_bias_residual", "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "mean", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_mean_bias_beta", "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "output_main", "mean", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_mean_bias_beta_residual", "when": ["statsContract", "present.meanT", "not present.invStdT", "present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "input_skip_bias_sum_main", "output_main", "mean", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_inv_plain", "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "inv_std_var", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_inv_residual", "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_inv_beta", "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "inv_std_var", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_inv_beta_residual", "when": ["statsContract", "not present.meanT", "present.invStdT", "not present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_inv_bias", "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "inv_std_var", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_inv_bias_residual", "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_inv_bias_beta", "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "output_main", "inv_std_var", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_inv_bias_beta_residual", "when": ["statsContract", "not present.meanT", "present.invStdT", "present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": false }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "input_skip_bias_sum_main", "output_main", "inv_std_var", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } } ], "intermediates": [] }, { "id": "stats_both_plain", "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": true }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "output_main", "row_stats", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } }, { "id": "statistics", "shader": "norm-stats-copy.wgsl.jinja", "derive": { "workgroupSize": 64 }, "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"], "dispatch": { "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)", "z": 1 } } ], "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }] }, { "id": "stats_both_residual", "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": true }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "input_skip_bias_sum_main", "output_main", "row_stats", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } }, { "id": "statistics", "shader": "norm-stats-copy.wgsl.jinja", "derive": { "workgroupSize": 64 }, "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"], "dispatch": { "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)", "z": 1 } } ], "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }] }, { "id": "stats_both_beta", "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": true }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "output_main", "row_stats", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } }, { "id": "statistics", "shader": "norm-stats-copy.wgsl.jinja", "derive": { "workgroupSize": 64 }, "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"], "dispatch": { "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)", "z": 1 } } ], "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }] }, { "id": "stats_both_beta_residual", "when": ["statsContract", "present.meanT", "present.invStdT", "not present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": true }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "beta_main", "input_skip_bias_sum_main", "output_main", "row_stats", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } }, { "id": "statistics", "shader": "norm-stats-copy.wgsl.jinja", "derive": { "workgroupSize": 64 }, "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"], "dispatch": { "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)", "z": 1 } } ], "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }] }, { "id": "stats_both_bias", "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "not present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": true }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "output_main", "row_stats", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } }, { "id": "statistics", "shader": "norm-stats-copy.wgsl.jinja", "derive": { "workgroupSize": 64 }, "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"], "dispatch": { "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)", "z": 1 } } ], "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }] }, { "id": "stats_both_bias_residual", "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "not present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": true }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "input_skip_bias_sum_main", "output_main", "row_stats", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } }, { "id": "statistics", "shader": "norm-stats-copy.wgsl.jinja", "derive": { "workgroupSize": 64 }, "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"], "dispatch": { "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)", "z": 1 } } ], "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }] }, { "id": "stats_both_bias_beta", "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "present.betaT", "not present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": true }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "output_main", "row_stats", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } }, { "id": "statistics", "shader": "norm-stats-copy.wgsl.jinja", "derive": { "workgroupSize": 64 }, "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"], "dispatch": { "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)", "z": 1 } } ], "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }] }, { "id": "stats_both_bias_beta_residual", "when": ["statsContract", "present.meanT", "present.invStdT", "present.biasT", "present.betaT", "present.residualT", "normResourcesFit", "rowDispatchFits"], "derive": { "simplified": false, "hasBias": "present.biasT", "hasBeta": "present.betaT", "writeResidualSum": "present.residualT", "writeMean": "present.meanT", "writeInvStd": "present.invStdT", "scalar": "dtypes.T", "useSubgroups": false, "workgroupSize": "skipWg", "HIDDEN_LEN": "hiddenSize", "packedStatistics": true }, "passes": [ { "id": "main", "shader": "norm-skip-row.wgsl.jinja", "bindings": ["input_main", "skip_main", "gamma_main", "bias_main", "beta_main", "input_skip_bias_sum_main", "output_main", "row_stats", "params__uniform"], "dispatch": { "x": "min(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "y": "ceilDiv(ceilDiv(rowCount, skipWg) if hiddenSize == 1 else rowCount, 65535)", "z": 1 } }, { "id": "statistics", "shader": "norm-stats-copy.wgsl.jinja", "derive": { "workgroupSize": 64 }, "bindings": ["row_stats_read", "mean", "inv_std_var", "params_stats"], "dispatch": { "x": "min(ceilDiv((rowCount), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((rowCount), (workgroupSize)), 65535)", "z": 1 } } ], "intermediates": [{ "id": "rowStats", "dtype": "float32", "shape": "[rowCount, 2]" }] } ] }