{ "domain": "com.microsoft", "name": "CausalConvWithState", "sinceVersion": 1, "inputs": { "inputT": { "onnx": "input", "dtype": "T", "rank": 3 }, "weightT": { "onnx": "weight", "dtype": "T", "rank": 3 }, "biasT": { "onnx": "bias", "dtype": "T", "rank": 1, "optional": true }, "pastStateT": { "onnx": "past_state", "dtype": "T", "rank": "3 if attrs.state_window == 0 else 4", "optional": true } }, "outputs": { "outputT": { "onnx": "output", "dtype": "T", "rank": 3, "shape": "shapes.inputT" }, "presentStateT": { "onnx": "present_state", "dtype": "T", "rank": "3 if attrs.state_window == 0 else 4", "shape": "([stateWindow] + stateShape) if windowed else stateShape" } }, "attributes": { "activation": { "default": "none" }, "channels_last": { "default": 0 }, "dilation": { "default": 1 }, "ndim": { "default": 1 }, "state_window": { "default": 0 } }, "attributeConstraints": { "activation": { "values": ["none", "silu", "swish"] }, "channels_last": { "values": [0, 1] }, "ndim": { "values": [1] } }, "typeConstraints": { "T": ["float32", "float16"] }, "tunables": { "workgroupSize": { "default": 256 }, "tiledWorkgroupSize": { "default": 128 } }, "derive": { "channels": "dim(shapes.inputT, 2 if attrs.channels_last == 1 else 1)", "inputLength": "dim(shapes.inputT, 1 if attrs.channels_last == 1 else 2)", "causalDilation": "attrs.dilation", "channelsLast": "attrs.channels_last == 1", "stateWindow": "attrs.state_window", "windowed": "stateWindow > 0", "stateWindowOk": "stateWindow >= 0 and stateWindow <= 8", "kernelSize": "dim(shapes.weightT, ranks.weightT - 1)", "kernelSizePadded": "ceilDiv(kernelSize, 4) * 4", "weightRankOk": "ranks.weightT == 3 and dim(shapes.weightT, 1) == 1", "stateLength": "(kernelSize - 1) * causalDilation", "stateSlotStride": "dim(shapes.inputT, 0) * channels * stateLength", "windowedLengthOk": "not windowed or inputLength > 0", "stateShape": "[dim(shapes.inputT, 0), stateLength, channels] if channelsLast else [dim(shapes.inputT, 0), channels, stateLength]", "presentStateOk": "sameShape(shapes.presentStateT, ([stateWindow] + stateShape) if windowed else stateShape)", "pastStateShapeOk": "present.pastStateT and sameShape(shapes.pastStateT, ([stateWindow] + stateShape) if windowed else stateShape)", "commonContract": "ranks.inputT == 3 and weightRankOk and ranks.outputT == 3 and (tensorDtypes.inputT == \"float32\" or tensorDtypes.inputT == \"float16\") and tensorDtypes.weightT == tensorDtypes.inputT and tensorDtypes.outputT == tensorDtypes.inputT and tensorDtypes.presentStateT == tensorDtypes.inputT and f16Ok(dtypes.T) and channels == dim(shapes.weightT, 0) and sameShape(shapes.outputT, shapes.inputT) and stateWindowOk and windowedLengthOk and presentStateOk and kernelSize >= 1 and causalDilation >= 1 and floor(causalDilation) == causalDilation", "zeroStateContract": "commonContract and not present.pastStateT and not present.biasT", "biasNoStateContract": "commonContract and not present.pastStateT and present.biasT and ranks.biasT == 1 and tensorDtypes.biasT == tensorDtypes.inputT and dim(shapes.biasT, 0) == channels", "stateNoBiasContract": "commonContract and present.pastStateT and not present.biasT and tensorDtypes.pastStateT == tensorDtypes.inputT and pastStateShapeOk", "stateBiasContract": "commonContract and present.pastStateT and present.biasT and ranks.biasT == 1 and tensorDtypes.pastStateT == tensorDtypes.inputT and tensorDtypes.biasT == tensorDtypes.inputT and pastStateShapeOk and dim(shapes.biasT, 0) == channels", "useSilu": "attrs.activation == \"silu\" or attrs.activation == \"swish\"", "inputScalar": "dtypes.T", "outputScalar": "dtypes.T", "hasStateWindow": "windowed", "hasBias": "present.biasT", "hasState": "present.pastStateT" }, "bindings": { "input": { "arg": "inputT", "elementType": "$inputVec4" }, "weight": { "arg": "weightT", "elementType": "$weightElem" }, "output": { "arg": "outputT", "elementType": "$outputVec4" }, "present_state": { "arg": "presentStateT", "elementType": "$outputScalar" }, "params": { "struct": [ { "name": "batchSize", "type": "u32", "value": "dim(shapes.inputT, 0)" }, { "name": "channels", "type": "u32", "value": "channels" }, { "name": "length", "type": "u32", "value": "inputLength" }, { "name": "stateWindow", "type": "u32", "value": "stateWindow" }, { "name": "stateSlotStride", "type": "u32", "value": "stateSlotStride" } ] }, "bias": { "arg": "biasT", "elementType": "$inputScalar" }, "past_state": { "arg": "pastStateT", "elementType": "$inputScalar" }, "input_main": { "arg": "inputT", "name": "input", "elementType": "$inputScalar" }, "weight_main": { "arg": "weightT", "name": "weight", "elementType": "$inputScalar" }, "output_main": { "arg": "outputT", "name": "output", "elementType": "$outputScalar" }, "params_main": { "name": "params", "struct": [ { "name": "batchSize", "type": "u32", "value": "dim(shapes.inputT, 0)" }, { "name": "channels", "type": "u32", "value": "channels" }, { "name": "length", "type": "u32", "value": "inputLength" }, { "name": "kernelSize", "type": "u32", "value": "kernelSize" }, { "name": "stateWindow", "type": "u32", "value": "stateWindow" }, { "name": "stateSlotStride", "type": "u32", "value": "stateSlotStride" } ] } }, "variants": [ { "id": "zero_state_vec4", "priority": 20, "when": ["zeroStateContract", "kernelSize >= 2", "kernelSize <= 4", "dim(shapes.inputT, 2) >= 4", "dim(shapes.inputT, 2) % 4 == 0", "not channelsLast", "causalDilation == 1"], "derive": { "workgroupSize": 256, "inputVec4": "\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\"", "outputVec4": "\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\"", "weightElem": "(\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\") if kernelSize == 4 else dtypes.T" }, "passes": [ { "id": "main", "name": "CausalConvWithState.Vec4", "shader": "causal-conv-with-state-vec4.wgsl.jinja", "bindings": ["input", "weight", "output", "present_state", "params"], "dispatch": { "x": "min(ceilDiv((dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * (dim(shapes.outputT, 2) / 4)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * (dim(shapes.outputT, 2) / 4)), (workgroupSize)), 65535)", "z": 1 } } ] }, { "id": "zero_state_tiled_large_kernel", "priority": 10, "when": ["zeroStateContract", "kernelSize >= 32", "dim(shapes.inputT, 2) >= 256", "dim(shapes.inputT, 2) % 8 == 0", "tunables.tiledWorkgroupSize >= 1", "floor(tunables.tiledWorkgroupSize) == tunables.tiledWorkgroupSize", "tunables.tiledWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup", "tunables.tiledWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX", "(tunables.tiledWorkgroupSize * 8 + kernelSizePadded + (kernelSizePadded - 1) * causalDilation) * 4 <= device.limits.maxComputeWorkgroupStorageSize", "not channelsLast"], "derive": { "workgroupSize": "tunables.tiledWorkgroupSize", "tileSize": "tunables.tiledWorkgroupSize * 8", "inputTileSize": "tunables.tiledWorkgroupSize * 8 + (kernelSizePadded - 1) * causalDilation" }, "passes": [ { "id": "main", "name": "CausalConvWithState.TiledLargeKernel", "shader": "causal-conv-with-state-tiled.wgsl.jinja", "bindings": ["input_main", "weight_main", "output_main", "present_state", "params"], "dispatch": { "x": "min(dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * ceilDiv(dim(shapes.outputT, 2), tileSize), 65535)", "y": "ceilDiv(dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * ceilDiv(dim(shapes.outputT, 2), tileSize), 65535)", "z": 1 } } ] }, { "id": "zero_state", "priority": 0, "when": ["zeroStateContract"], "derive": { "workgroupSize": "tunables.workgroupSize" }, "passes": [ { "id": "main", "name": "CausalConvWithState", "shader": "causal-conv-with-state.wgsl.jinja", "bindings": ["input_main", "weight_main", "output_main", "present_state", "params_main"], "dispatch": { "x": "min(ceilDiv((dim(shapes.inputT, 0) * channels * max(1, inputLength)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((dim(shapes.inputT, 0) * channels * max(1, inputLength)), (workgroupSize)), 65535)", "z": 1 } } ] }, { "id": "state_bias_vec4", "priority": 20, "when": ["stateBiasContract", "kernelSize >= 2", "kernelSize <= 4", "dim(shapes.inputT, 2) >= 4", "dim(shapes.inputT, 2) % 4 == 0", "not channelsLast", "causalDilation == 1"], "derive": { "workgroupSize": 256, "inputVec4": "\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\"", "outputVec4": "\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\"", "weightElem": "(\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\") if kernelSize == 4 else dtypes.T" }, "passes": [ { "id": "main", "name": "CausalConvWithState.Vec4", "shader": "causal-conv-with-state-vec4.wgsl.jinja", "bindings": ["input", "weight", "bias", "past_state", "output", "present_state", "params"], "dispatch": { "x": "min(ceilDiv((dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * (dim(shapes.outputT, 2) / 4)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * (dim(shapes.outputT, 2) / 4)), (workgroupSize)), 65535)", "z": 1 } } ] }, { "id": "state_bias_tiled_large_kernel", "priority": 10, "when": ["stateBiasContract", "kernelSize >= 32", "dim(shapes.inputT, 2) >= 256", "dim(shapes.inputT, 2) % 8 == 0", "tunables.tiledWorkgroupSize >= 1", "floor(tunables.tiledWorkgroupSize) == tunables.tiledWorkgroupSize", "tunables.tiledWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup", "tunables.tiledWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX", "(tunables.tiledWorkgroupSize * 8 + kernelSizePadded + (kernelSizePadded - 1) * causalDilation) * 4 <= device.limits.maxComputeWorkgroupStorageSize", "not channelsLast"], "derive": { "workgroupSize": "tunables.tiledWorkgroupSize", "tileSize": "tunables.tiledWorkgroupSize * 8", "inputTileSize": "tunables.tiledWorkgroupSize * 8 + (kernelSizePadded - 1) * causalDilation" }, "passes": [ { "id": "main", "name": "CausalConvWithState.TiledLargeKernel", "shader": "causal-conv-with-state-tiled.wgsl.jinja", "bindings": ["input_main", "weight_main", "bias", "past_state", "output_main", "present_state", "params"], "dispatch": { "x": "min(dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * ceilDiv(dim(shapes.outputT, 2), tileSize), 65535)", "y": "ceilDiv(dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * ceilDiv(dim(shapes.outputT, 2), tileSize), 65535)", "z": 1 } } ] }, { "id": "state_bias", "priority": 0, "when": ["stateBiasContract"], "derive": { "workgroupSize": "tunables.workgroupSize" }, "passes": [ { "id": "main", "name": "CausalConvWithState", "shader": "causal-conv-with-state.wgsl.jinja", "bindings": ["input_main", "weight_main", "bias", "past_state", "output_main", "present_state", "params_main"], "dispatch": { "x": "min(ceilDiv((dim(shapes.inputT, 0) * channels * max(1, inputLength)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((dim(shapes.inputT, 0) * channels * max(1, inputLength)), (workgroupSize)), 65535)", "z": 1 } } ] }, { "id": "bias_no_state_vec4", "priority": 20, "when": ["biasNoStateContract", "kernelSize >= 2", "kernelSize <= 4", "dim(shapes.inputT, 2) >= 4", "dim(shapes.inputT, 2) % 4 == 0", "not channelsLast", "causalDilation == 1"], "derive": { "workgroupSize": 256, "inputVec4": "\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\"", "outputVec4": "\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\"", "weightElem": "(\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\") if kernelSize == 4 else dtypes.T" }, "passes": [ { "id": "main", "name": "CausalConvWithState.Vec4", "shader": "causal-conv-with-state-vec4.wgsl.jinja", "bindings": ["input", "weight", "bias", "output", "present_state", "params"], "dispatch": { "x": "min(ceilDiv((dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * (dim(shapes.outputT, 2) / 4)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * (dim(shapes.outputT, 2) / 4)), (workgroupSize)), 65535)", "z": 1 } } ] }, { "id": "bias_no_state_tiled_large_kernel", "priority": 10, "when": ["biasNoStateContract", "kernelSize >= 32", "dim(shapes.inputT, 2) >= 256", "dim(shapes.inputT, 2) % 8 == 0", "tunables.tiledWorkgroupSize >= 1", "floor(tunables.tiledWorkgroupSize) == tunables.tiledWorkgroupSize", "tunables.tiledWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup", "tunables.tiledWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX", "(tunables.tiledWorkgroupSize * 8 + kernelSizePadded + (kernelSizePadded - 1) * causalDilation) * 4 <= device.limits.maxComputeWorkgroupStorageSize", "not channelsLast"], "derive": { "workgroupSize": "tunables.tiledWorkgroupSize", "tileSize": "tunables.tiledWorkgroupSize * 8", "inputTileSize": "tunables.tiledWorkgroupSize * 8 + (kernelSizePadded - 1) * causalDilation" }, "passes": [ { "id": "main", "name": "CausalConvWithState.TiledLargeKernel", "shader": "causal-conv-with-state-tiled.wgsl.jinja", "bindings": ["input_main", "weight_main", "bias", "output_main", "present_state", "params"], "dispatch": { "x": "min(dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * ceilDiv(dim(shapes.outputT, 2), tileSize), 65535)", "y": "ceilDiv(dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * ceilDiv(dim(shapes.outputT, 2), tileSize), 65535)", "z": 1 } } ] }, { "id": "bias_no_state", "priority": 0, "when": ["biasNoStateContract"], "derive": { "workgroupSize": "tunables.workgroupSize" }, "passes": [ { "id": "main", "name": "CausalConvWithState", "shader": "causal-conv-with-state.wgsl.jinja", "bindings": ["input_main", "weight_main", "bias", "output_main", "present_state", "params_main"], "dispatch": { "x": "min(ceilDiv((dim(shapes.inputT, 0) * channels * max(1, inputLength)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((dim(shapes.inputT, 0) * channels * max(1, inputLength)), (workgroupSize)), 65535)", "z": 1 } } ] }, { "id": "state_no_bias_vec4", "priority": 20, "when": ["stateNoBiasContract", "kernelSize >= 2", "kernelSize <= 4", "dim(shapes.inputT, 2) >= 4", "dim(shapes.inputT, 2) % 4 == 0", "not channelsLast", "causalDilation == 1"], "derive": { "workgroupSize": 256, "inputVec4": "\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\"", "outputVec4": "\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\"", "weightElem": "(\"vec4\" if tensorDtypes.inputT == \"float16\" else \"vec4\") if kernelSize == 4 else dtypes.T" }, "passes": [ { "id": "main", "name": "CausalConvWithState.Vec4", "shader": "causal-conv-with-state-vec4.wgsl.jinja", "bindings": ["input", "weight", "past_state", "output", "present_state", "params"], "dispatch": { "x": "min(ceilDiv((dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * (dim(shapes.outputT, 2) / 4)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * (dim(shapes.outputT, 2) / 4)), (workgroupSize)), 65535)", "z": 1 } } ] }, { "id": "state_no_bias_tiled_large_kernel", "priority": 10, "when": ["stateNoBiasContract", "kernelSize >= 32", "dim(shapes.inputT, 2) >= 256", "dim(shapes.inputT, 2) % 8 == 0", "tunables.tiledWorkgroupSize >= 1", "floor(tunables.tiledWorkgroupSize) == tunables.tiledWorkgroupSize", "tunables.tiledWorkgroupSize <= device.limits.maxComputeInvocationsPerWorkgroup", "tunables.tiledWorkgroupSize <= device.limits.maxComputeWorkgroupSizeX", "(tunables.tiledWorkgroupSize * 8 + kernelSizePadded + (kernelSizePadded - 1) * causalDilation) * 4 <= device.limits.maxComputeWorkgroupStorageSize", "not channelsLast"], "derive": { "workgroupSize": "tunables.tiledWorkgroupSize", "tileSize": "tunables.tiledWorkgroupSize * 8", "inputTileSize": "tunables.tiledWorkgroupSize * 8 + (kernelSizePadded - 1) * causalDilation" }, "passes": [ { "id": "main", "name": "CausalConvWithState.TiledLargeKernel", "shader": "causal-conv-with-state-tiled.wgsl.jinja", "bindings": ["input_main", "weight_main", "past_state", "output_main", "present_state", "params"], "dispatch": { "x": "min(dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * ceilDiv(dim(shapes.outputT, 2), tileSize), 65535)", "y": "ceilDiv(dim(shapes.outputT, 0) * dim(shapes.outputT, 1) * ceilDiv(dim(shapes.outputT, 2), tileSize), 65535)", "z": 1 } } ] }, { "id": "state_no_bias", "priority": 0, "when": ["stateNoBiasContract"], "derive": { "workgroupSize": "tunables.workgroupSize" }, "passes": [ { "id": "main", "name": "CausalConvWithState", "shader": "causal-conv-with-state.wgsl.jinja", "bindings": ["input_main", "weight_main", "past_state", "output_main", "present_state", "params_main"], "dispatch": { "x": "min(ceilDiv((dim(shapes.inputT, 0) * channels * max(1, inputLength)), (workgroupSize)), 65535)", "y": "ceilDiv(ceilDiv((dim(shapes.inputT, 0) * channels * max(1, inputLength)), (workgroupSize)), 65535)", "z": 1 } } ] } ] }