Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
8aaab97 verified
Raw History Blame
19.4 kB
{
"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<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\"",
"outputVec4": "\"vec4<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\"",
"weightElem": "(\"vec4<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\") 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<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\"",
"outputVec4": "\"vec4<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\"",
"weightElem": "(\"vec4<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\") 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<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\"",
"outputVec4": "\"vec4<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\"",
"weightElem": "(\"vec4<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\") 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<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\"",
"outputVec4": "\"vec4<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\"",
"weightElem": "(\"vec4<f16>\" if tensorDtypes.inputT == \"float16\" else \"vec4<f32>\") 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
}
}
]
}
]
}