Xenova's picture
Xenova HF Staff
sync 6fdf6301e2bb
08b6b14 verified
Raw History Blame
65.3 kB
{
"domain": "com.microsoft",
"name": "LinearAttention",
"sinceVersion": 1,
"inputs": {
"queryT": { "onnx": "query", "dtype": "T", "rank": 3 },
"keyT": { "onnx": "key", "dtype": "T", "rank": 3 },
"valueT": { "onnx": "value", "dtype": "T", "rank": 3 },
"pastStateT": {
"onnx": "past_state",
"dtype": "S",
"rank": "5 if attrs.state_window > 0 else 4",
"optional": true,
"shape": "[attrs.state_window, dim(shapes.queryT, 0), attrs.kv_num_heads, dim(shapes.queryT, 2) / attrs.q_num_heads, dim(shapes.valueT, 2) / max(1, attrs.kv_num_heads)] if attrs.state_window > 0 else [dim(shapes.queryT, 0), attrs.kv_num_heads, dim(shapes.queryT, 2) / attrs.q_num_heads, dim(shapes.valueT, 2) / max(1, attrs.kv_num_heads)]"
},
"decayT": { "onnx": "decay", "dtype": "T", "rank": 3, "optional": true },
"betaT": { "onnx": "beta", "dtype": "T", "rank": 3, "optional": true }
},
"outputs": {
"outputT": {
"onnx": "output",
"dtype": "T",
"rank": 3,
"shape": "[dim(shapes.queryT, 0), dim(shapes.queryT, 1), max(attrs.q_num_heads, attrs.kv_num_heads) * (dim(shapes.valueT, 2) / max(1, attrs.kv_num_heads))]"
},
"presentStateT": {
"onnx": "present_state",
"dtype": "S",
"rank": "5 if attrs.state_window > 0 else 4",
"shape": "[attrs.state_window, dim(shapes.queryT, 0), attrs.kv_num_heads, dim(shapes.queryT, 2) / attrs.q_num_heads, dim(shapes.valueT, 2) / max(1, attrs.kv_num_heads)] if attrs.state_window > 0 else [dim(shapes.queryT, 0), attrs.kv_num_heads, dim(shapes.queryT, 2) / attrs.q_num_heads, dim(shapes.valueT, 2) / max(1, attrs.kv_num_heads)]"
}
},
"attributes": {
"chunk_size": { "default": 64 },
"scale": { "default": 0 },
"state_window": { "default": 0 },
"update_rule": { "default": "gated_delta" },
"kv_num_heads": {},
"q_num_heads": {}
},
"attributeConstraints": {
"kv_num_heads": { "required": true },
"q_num_heads": { "required": true },
"update_rule": { "values": ["linear", "gated", "delta", "gated_delta"] }
},
"typeConstraints": { "T": ["float32", "float16"], "S": ["float32", "float16"] },
"tunables": {
"tileV": { "default": 8 },
"gatedTileV": { "default": 4 },
"chunkSize": { "default": 16 },
"chunkScanTokens": { "default": 8 },
"chunkTileV": { "default": 32 },
"chunkGroups": { "default": 8 },
"chunkOutRows": { "default": 16 },
"dvGroups": { "default": 4 }
},
"derive": {
"deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)",
"effRule": "attrs.update_rule if attrs.update_rule else \"gated_delta\"",
"headDimK": "dim(shapes.queryT, 2) / max(1, attrs.q_num_heads)",
"headDimV": "dim(shapes.valueT, 2) / max(1, attrs.kv_num_heads)",
"batchSequenceOk": "dim(shapes.keyT, 0) == dim(shapes.queryT, 0) and dim(shapes.keyT, 1) == dim(shapes.queryT, 1) and dim(shapes.valueT, 0) == dim(shapes.queryT, 0) and dim(shapes.valueT, 1) == dim(shapes.queryT, 1)",
"decayOk": "not present.decayT or (ranks.decayT == 3 and dim(shapes.decayT, 0) == dim(shapes.queryT, 0) and dim(shapes.decayT, 1) == dim(shapes.queryT, 1) and (dim(shapes.decayT, 2) == attrs.kv_num_heads or dim(shapes.decayT, 2) == attrs.kv_num_heads * headDimK) and tensorDtypes.decayT == tensorDtypes.queryT)",
"betaOk": "not present.betaT or (ranks.betaT == 3 and dim(shapes.betaT, 0) == dim(shapes.queryT, 0) and dim(shapes.betaT, 1) == dim(shapes.queryT, 1) and (dim(shapes.betaT, 2) == 1 or dim(shapes.betaT, 2) == attrs.kv_num_heads) and tensorDtypes.betaT == tensorDtypes.queryT)",
"queryDtype": "tensorDtypes.queryT",
"keyDtype": "tensorDtypes.queryT",
"valueDtype": "tensorDtypes.queryT",
"decayDtype": "tensorDtypes.queryT",
"betaDtype": "tensorDtypes.queryT",
"stateDtype": "tensorDtypes.presentStateT",
"queryElem": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
"keyElem": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
"valueScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
"decayScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
"betaScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
"outputScalar": "\"f16\" if tensorDtypes.queryT == \"float16\" else \"f32\"",
"stateScalar": "\"f16\" if tensorDtypes.presentStateT == \"float16\" else \"f32\"",
"vec4Lanes": "min(deviceWorkgroupCap, pow2ceil(ceil(headDimK / 4)))",
"vec4SubgroupExact": "device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == device.adapterInfo.subgroupMaxSize and vec4Lanes == device.adapterInfo.subgroupMinSize",
"vec4SubgroupsAvailable": "device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and vec4Lanes <= device.adapterInfo.subgroupMinSize",
"vec4DvGroups": "max(1, min(tunables.dvGroups, deviceWorkgroupCap / vec4Lanes)) if (vec4SubgroupExact or not vec4SubgroupsAvailable) else 1",
"vec4WorkgroupSize": "vec4Lanes * vec4DvGroups",
"vec4HeadDimFits": "ceilDiv(headDimK, 4) <= vec4Lanes and vec4WorkgroupSize <= deviceWorkgroupCap",
"vec4BindingsNonEmpty": "dim(shapes.queryT, 0) > 0 and dim(shapes.queryT, 1) > 0",
"vec4UseSubgroups": "vec4SubgroupsAvailable and (vec4DvGroups == 1 or vec4SubgroupExact)",
"vec4TileVPlain": "max(1, min(tunables.tileV, headDimV))",
"nKeyHeads": "dim(shapes.keyT, 2) / max(1, headDimK)",
"outHeads": "max(attrs.q_num_heads, attrs.kv_num_heads)",
"headLayoutOk": "attrs.q_num_heads > 0 and attrs.kv_num_heads > 0 and dim(shapes.queryT, 2) > 0 and dim(shapes.valueT, 2) > 0 and dim(shapes.keyT, 2) > 0 and dim(shapes.queryT, 2) % attrs.q_num_heads == 0 and dim(shapes.valueT, 2) % attrs.kv_num_heads == 0 and dim(shapes.keyT, 2) % max(1, headDimK) == 0 and (attrs.q_num_heads % attrs.kv_num_heads == 0 or attrs.kv_num_heads % attrs.q_num_heads == 0) and nKeyHeads > 0 and attrs.kv_num_heads % nKeyHeads == 0",
"scalarHeadDimFits": "headDimK <= deviceWorkgroupCap and headDimK <= 256",
"serialHeadDimFits": "headDimK > 0 and headDimK <= 16",
"stateWindow": "attrs.state_window if attrs.state_window is defined else 0",
"stateSlotStride": "dim(shapes.queryT, 0) * attrs.kv_num_heads * headDimK * headDimV",
"windowed": "stateWindow > 0",
"stateWindowOk": "stateWindow >= 0 and stateWindow <= 8",
"presentStateOk": "(dim(shapes.presentStateT, 0) == dim(shapes.queryT, 0) and dim(shapes.presentStateT, 1) == attrs.kv_num_heads and dim(shapes.presentStateT, 2) == headDimK and dim(shapes.presentStateT, 3) == (headDimV) and ranks.presentStateT == 4) if not windowed else (ranks.presentStateT == 5 and dim(shapes.presentStateT, 0) == stateWindow and dim(shapes.presentStateT, 1) == dim(shapes.queryT, 0) and dim(shapes.presentStateT, 2) == attrs.kv_num_heads and dim(shapes.presentStateT, 3) == headDimK and dim(shapes.presentStateT, 4) == (headDimV))",
"ioContractOk": "dim(shapes.outputT, 0) == dim(shapes.queryT, 0) and dim(shapes.outputT, 1) == dim(shapes.queryT, 1) and dim(shapes.outputT, 2) == outHeads * headDimV and stateWindowOk and presentStateOk",
"pastStateOk": "not present.pastStateT or ((ranks.pastStateT == 4 and dim(shapes.pastStateT, 0) == dim(shapes.queryT, 0) and dim(shapes.pastStateT, 1) == attrs.kv_num_heads and dim(shapes.pastStateT, 2) == headDimK and dim(shapes.pastStateT, 3) == headDimV) if not windowed else (ranks.pastStateT == 5 and dim(shapes.pastStateT, 0) == stateWindow and dim(shapes.pastStateT, 1) == dim(shapes.queryT, 0) and dim(shapes.pastStateT, 2) == attrs.kv_num_heads and dim(shapes.pastStateT, 3) == headDimK and dim(shapes.pastStateT, 4) == headDimV))",
"needsDecay": "effRule == \"gated\" or effRule == \"gated_delta\"",
"needsBeta": "effRule == \"delta\" or effRule == \"gated_delta\"",
"gateInputsOk": "decayOk and betaOk and (not needsDecay or present.decayT) and (not needsBeta or present.betaT)",
"tensorDtypesOk": "(tensorDtypes.queryT == \"float16\" or tensorDtypes.queryT == \"float32\") and tensorDtypes.keyT == tensorDtypes.queryT and tensorDtypes.valueT == tensorDtypes.queryT and tensorDtypes.outputT == tensorDtypes.queryT and (tensorDtypes.presentStateT == \"float16\" or tensorDtypes.presentStateT == \"float32\") and (not present.pastStateT or tensorDtypes.pastStateT == tensorDtypes.presentStateT) and f16Ok(tensorDtypes.queryT) and f16Ok(tensorDtypes.presentStateT)",
"commonContract": "headLayoutOk and batchSequenceOk and gateInputsOk and tensorDtypesOk and ioContractOk and pastStateOk",
"gatedVec4TileVBudget": "tunables.gatedTileV if (vec4WorkgroupSize + vec4DvGroups) * (2 * tunables.gatedTileV + 4) * 4 <= device.limits.maxComputeWorkgroupStorageSize else (2 if (vec4WorkgroupSize + vec4DvGroups) * 32 <= device.limits.maxComputeWorkgroupStorageSize else 1)",
"gatedVec4TileV": "max(1, min(tunables.gatedTileV, headDimV)) if vec4UseSubgroups else max(1, min(gatedVec4TileVBudget, headDimV))",
"vec4SharedFitsPlain": "vec4UseSubgroups or (vec4WorkgroupSize + vec4DvGroups) * vec4TileVPlain * 4 <= device.limits.maxComputeWorkgroupStorageSize",
"vec4SharedFitsGatedBeta": "vec4UseSubgroups or (vec4WorkgroupSize + vec4DvGroups) * (2 * gatedVec4TileV + 4) * 4 <= device.limits.maxComputeWorkgroupStorageSize",
"chunkSize": "tunables.chunkSize",
"chunkTileK": "min(32, headDimK)",
"chunkUtTileK": "min(16, headDimK)",
"chunkGroups": "min(tunables.chunkGroups, 32 if (chunkSize % 32 == 0 and headDimK % 32 == 0) else (16 if (chunkSize % 16 == 0 and headDimK % 16 == 0) else (8 if (chunkSize % 8 == 0 and headDimK % 8 == 0) else (4 if (chunkSize % 4 == 0 and headDimK % 4 == 0) else (2 if (chunkSize % 2 == 0 and headDimK % 2 == 0) else 1)))))",
"chunkTileV": "min(tunables.chunkTileV, 64 if headDimV % 64 == 0 else (32 if headDimV % 32 == 0 else (16 if headDimV % 16 == 0 else (8 if headDimV % 8 == 0 else (4 if headDimV % 4 == 0 else (2 if headDimV % 2 == 0 else 1))))))",
"chunkOutRows": "max(1, min(tunables.chunkOutRows, chunkSize))",
"chunkScanTokens": "max(1, min(tunables.chunkScanTokens, chunkSize))",
"chunkScanWorkgroup": "chunkGroups * chunkTileV",
"chunkFlatWorkgroup": "min(deviceWorkgroupCap, 256)",
"chunkOutWorkgroup": "max(64, min(deviceWorkgroupCap, pow2ceil(headDimV)))",
"chunkNumChunks": "ceilDiv(dim(shapes.queryT, 1), chunkSize)",
"kvPerKeyHead": "attrs.kv_num_heads / max(1, nKeyHeads)",
"decayPerElement": "needsDecay and present.decayT and dim(shapes.decayT, 2) == attrs.kv_num_heads * headDimK",
"chunkUtBytes": "4 * (chunkSize * chunkSize + 2 * chunkSize + 2 * chunkSize * chunkUtTileK)",
"chunkOutBytes": "4 * (chunkOutRows * headDimK + chunkSize * chunkTileK + chunkOutRows * chunkSize)",
"chunkScanBytes": "4 * (headDimK * chunkTileV + chunkSize * chunkTileV + chunkScanTokens * headDimK)",
"chunkSharedOk": "max(chunkUtBytes, max(chunkOutBytes, chunkScanBytes)) <= device.limits.maxComputeWorkgroupStorageSize",
"chunkActivationBytes": "dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * (dim(shapes.queryT, 2) + dim(shapes.keyT, 2) + dim(shapes.valueT, 2)) * (2 if tensorDtypes.queryT == \"float16\" else 4)",
"chunkStatesBudget": "min(chunkActivationBytes, 134217728)",
"chunkStatesElems": "dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV",
"chunkWkElems": "dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimK",
"chunkUvecElems": "dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV",
"chunkGexpElems": "dim(shapes.queryT, 0) * dim(shapes.queryT, 1) * dim(shapes.decayT, 2) if present.decayT else 0",
"chunkScratchOk": "4 * max(chunkStatesElems, max(chunkWkElems, max(chunkUvecElems, chunkGexpElems))) <= min(device.limits.maxStorageBufferBindingSize, device.limits.maxBufferSize) and 4 * chunkStatesElems <= chunkStatesBudget",
"chunkGeometryOk": "not windowed and headDimK > 0 and headDimV > 0 and headDimV % chunkTileV == 0 and headDimK % chunkGroups == 0 and chunkSize % chunkGroups == 0 and chunkScanWorkgroup >= 64 and chunkScanWorkgroup <= deviceWorkgroupCap and chunkOutWorkgroup <= deviceWorkgroupCap and headDimV <= chunkOutWorkgroup and chunkSize % chunkScanTokens == 0 and chunkScanTokens % chunkGroups == 0 and chunkSize % chunkOutRows == 0",
"chunkedShapeOk": "dim(shapes.queryT, 1) >= 1024 and chunkGeometryOk and chunkSharedOk and chunkScratchOk"
},
"bindings": {
"key": { "arg": "keyT", "elementType": "$keyElem" },
"value": { "arg": "valueT", "elementType": "$valueScalar" },
"states": { "buffer": "storage", "elementType": "f32" },
"present_state": { "arg": "presentStateT", "elementType": "$stateScalar" },
"params_chunk_linear": {
"name": "params",
"struct": [
{ "name": "batchSize", "type": "u32", "value": "dim(shapes.queryT, 0)" },
{ "name": "seqLength", "type": "u32", "value": "dim(shapes.queryT, 1)" },
{ "name": "qNumHeads", "type": "u32", "value": "attrs.q_num_heads" },
{ "name": "kvNumHeads", "type": "u32", "value": "attrs.kv_num_heads" },
{ "name": "qPackedDim", "type": "u32", "value": "dim(shapes.queryT, 2)" },
{ "name": "kPackedDim", "type": "u32", "value": "dim(shapes.keyT, 2)" },
{ "name": "vPackedDim", "type": "u32", "value": "dim(shapes.valueT, 2)" },
{ "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" }
]
},
"past_state": { "arg": "pastStateT", "elementType": "$stateScalar" },
"query": { "arg": "queryT", "elementType": "$queryElem" },
"states_f32": { "name": "states", "buffer": "read-only-storage", "elementType": "f32" },
"output": { "arg": "outputT", "elementType": "$outputScalar" },
"decay": { "arg": "decayT", "elementType": "$decayScalar" },
"gexp": { "buffer": "storage", "elementType": "f32" },
"params_chunk_gated": {
"name": "params",
"struct": [
{ "name": "batchSize", "type": "u32", "value": "dim(shapes.queryT, 0)" },
{ "name": "seqLength", "type": "u32", "value": "dim(shapes.queryT, 1)" },
{ "name": "qNumHeads", "type": "u32", "value": "attrs.q_num_heads" },
{ "name": "kvNumHeads", "type": "u32", "value": "attrs.kv_num_heads" },
{ "name": "qPackedDim", "type": "u32", "value": "dim(shapes.queryT, 2)" },
{ "name": "kPackedDim", "type": "u32", "value": "dim(shapes.keyT, 2)" },
{ "name": "vPackedDim", "type": "u32", "value": "dim(shapes.valueT, 2)" },
{ "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decayT, 2) if present.decayT else 0" },
{ "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" }
]
},
"gexp_ut": { "name": "gexp", "buffer": "read-only-storage", "elementType": "f32" },
"beta": { "arg": "betaT", "elementType": "$betaScalar" },
"wk": { "buffer": "storage", "elementType": "f32" },
"uvec": { "buffer": "storage", "elementType": "f32" },
"params_chunk_delta": {
"name": "params",
"struct": [
{ "name": "batchSize", "type": "u32", "value": "dim(shapes.queryT, 0)" },
{ "name": "seqLength", "type": "u32", "value": "dim(shapes.queryT, 1)" },
{ "name": "qNumHeads", "type": "u32", "value": "attrs.q_num_heads" },
{ "name": "kvNumHeads", "type": "u32", "value": "attrs.kv_num_heads" },
{ "name": "qPackedDim", "type": "u32", "value": "dim(shapes.queryT, 2)" },
{ "name": "kPackedDim", "type": "u32", "value": "dim(shapes.keyT, 2)" },
{ "name": "vPackedDim", "type": "u32", "value": "dim(shapes.valueT, 2)" },
{ "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.betaT, 2) if present.betaT else 0" },
{ "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" }
]
},
"wk_f32": { "name": "wk", "buffer": "read-only-storage", "elementType": "f32" },
"uvec_f32": { "name": "uvec", "buffer": "read-only-storage", "elementType": "f32" },
"deltas": { "buffer": "storage", "elementType": "f32" },
"deltas_f32": { "name": "deltas", "buffer": "read-only-storage", "elementType": "f32" },
"params_chunk_gated_delta": {
"name": "params",
"struct": [
{ "name": "batchSize", "type": "u32", "value": "dim(shapes.queryT, 0)" },
{ "name": "seqLength", "type": "u32", "value": "dim(shapes.queryT, 1)" },
{ "name": "qNumHeads", "type": "u32", "value": "attrs.q_num_heads" },
{ "name": "kvNumHeads", "type": "u32", "value": "attrs.kv_num_heads" },
{ "name": "qPackedDim", "type": "u32", "value": "dim(shapes.queryT, 2)" },
{ "name": "kPackedDim", "type": "u32", "value": "dim(shapes.keyT, 2)" },
{ "name": "vPackedDim", "type": "u32", "value": "dim(shapes.valueT, 2)" },
{ "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decayT, 2) if present.decayT else 0" },
{ "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.betaT, 2) if present.betaT else 0" },
{ "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" }
]
},
"params_win_linear": {
"name": "params",
"struct": [
{ "name": "batchSize", "type": "u32", "value": "dim(shapes.queryT, 0)" },
{ "name": "seqLength", "type": "u32", "value": "dim(shapes.queryT, 1)" },
{ "name": "qNumHeads", "type": "u32", "value": "attrs.q_num_heads" },
{ "name": "kvNumHeads", "type": "u32", "value": "attrs.kv_num_heads" },
{ "name": "qPackedDim", "type": "u32", "value": "dim(shapes.queryT, 2)" },
{ "name": "kPackedDim", "type": "u32", "value": "dim(shapes.keyT, 2)" },
{ "name": "vPackedDim", "type": "u32", "value": "dim(shapes.valueT, 2)" },
{ "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" },
{ "name": "stateWindow", "type": "u32", "value": "stateWindow" },
{ "name": "stateSlotStride", "type": "u32", "value": "stateSlotStride" }
]
},
"params_win_gated_delta": {
"name": "params",
"struct": [
{ "name": "batchSize", "type": "u32", "value": "dim(shapes.queryT, 0)" },
{ "name": "seqLength", "type": "u32", "value": "dim(shapes.queryT, 1)" },
{ "name": "qNumHeads", "type": "u32", "value": "attrs.q_num_heads" },
{ "name": "kvNumHeads", "type": "u32", "value": "attrs.kv_num_heads" },
{ "name": "qPackedDim", "type": "u32", "value": "dim(shapes.queryT, 2)" },
{ "name": "kPackedDim", "type": "u32", "value": "dim(shapes.keyT, 2)" },
{ "name": "vPackedDim", "type": "u32", "value": "dim(shapes.valueT, 2)" },
{ "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decayT, 2) if present.decayT else 0" },
{ "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.betaT, 2) if present.betaT else 0" },
{ "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" },
{ "name": "stateWindow", "type": "u32", "value": "stateWindow" },
{ "name": "stateSlotStride", "type": "u32", "value": "stateSlotStride" }
]
},
"params_win_gated": {
"name": "params",
"struct": [
{ "name": "batchSize", "type": "u32", "value": "dim(shapes.queryT, 0)" },
{ "name": "seqLength", "type": "u32", "value": "dim(shapes.queryT, 1)" },
{ "name": "qNumHeads", "type": "u32", "value": "attrs.q_num_heads" },
{ "name": "kvNumHeads", "type": "u32", "value": "attrs.kv_num_heads" },
{ "name": "qPackedDim", "type": "u32", "value": "dim(shapes.queryT, 2)" },
{ "name": "kPackedDim", "type": "u32", "value": "dim(shapes.keyT, 2)" },
{ "name": "vPackedDim", "type": "u32", "value": "dim(shapes.valueT, 2)" },
{ "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decayT, 2) if present.decayT else 0" },
{ "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" },
{ "name": "stateWindow", "type": "u32", "value": "stateWindow" },
{ "name": "stateSlotStride", "type": "u32", "value": "stateSlotStride" }
]
},
"params_win_delta": {
"name": "params",
"struct": [
{ "name": "batchSize", "type": "u32", "value": "dim(shapes.queryT, 0)" },
{ "name": "seqLength", "type": "u32", "value": "dim(shapes.queryT, 1)" },
{ "name": "qNumHeads", "type": "u32", "value": "attrs.q_num_heads" },
{ "name": "kvNumHeads", "type": "u32", "value": "attrs.kv_num_heads" },
{ "name": "qPackedDim", "type": "u32", "value": "dim(shapes.queryT, 2)" },
{ "name": "kPackedDim", "type": "u32", "value": "dim(shapes.keyT, 2)" },
{ "name": "vPackedDim", "type": "u32", "value": "dim(shapes.valueT, 2)" },
{ "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.betaT, 2) if present.betaT else 0" },
{ "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" },
{ "name": "stateWindow", "type": "u32", "value": "stateWindow" },
{ "name": "stateSlotStride", "type": "u32", "value": "stateSlotStride" }
]
}
},
"variants": [
{
"id": "linear_zero_chunked",
"priority": 30,
"when": ["effRule == \"linear\"", "not present.pastStateT", "commonContract", "chunkedShapeOk"],
"derive": { "hasPastState": "present.pastStateT", "usesDecay": "needsDecay", "usesBeta": "needsBeta" },
"intermediates": [{ "id": "states", "dtype": "float32", "shape": "[chunkStatesElems]" }],
"passes": [
{
"id": "scan",
"name": "LinearAttention.ChunkScan",
"shader": "chunk-scan.wgsl.jinja",
"derive": { "workgroupSize": "chunkScanWorkgroup" },
"bindings": ["key", "value", "states", "present_state", "params_chunk_linear"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"z": 1
}
},
{
"id": "out",
"name": "LinearAttention.ChunkOutput",
"shader": "chunk-out.wgsl.jinja",
"derive": { "workgroupSize": "chunkOutWorkgroup" },
"bindings": ["query", "key", "states_f32", "value", "output", "params_chunk_linear"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"z": 1
}
}
]
},
{
"id": "linear_zero_scalar",
"priority": 0,
"when": ["effRule == \"linear\"", "not present.pastStateT", "commonContract", "scalarHeadDimFits"],
"derive": {
"updateRule": "\"linear\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16",
"workgroupSize": "min(256, pow2ceil(dim(shapes.queryT, 2) / attrs.q_num_heads))",
"tileV": "max(1, min(tunables.tileV, headDimV))",
"outputDtype": "tensorDtypes.queryT"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.scalar.wgsl.jinja",
"bindings": ["query", "key", "value", "output", "present_state", "params_win_linear"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"z": 1
}
}
]
},
{
"id": "linear_zero_vec4",
"priority": 10,
"when": ["effRule == \"linear\"", "not present.pastStateT", "commonContract", "headDimK % 4 == 0", "vec4SharedFitsPlain", "vec4HeadDimFits", "vec4BindingsNonEmpty"],
"derive": {
"updateRule": "\"linear\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "vec4UseSubgroups",
"tileV": "vec4TileVPlain",
"queryElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"keyElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"dvGroups": "vec4DvGroups"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.vec4.wgsl.jinja",
"bindings": ["query", "key", "value", "output", "present_state", "params_win_linear"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"z": 1
}
}
]
},
{
"id": "linear_state_chunked",
"priority": 30,
"when": ["effRule == \"linear\"", "present.pastStateT", "commonContract", "chunkedShapeOk"],
"derive": { "hasPastState": "present.pastStateT", "usesDecay": "needsDecay", "usesBeta": "needsBeta" },
"intermediates": [{ "id": "states", "dtype": "float32", "shape": "[chunkStatesElems]" }],
"passes": [
{
"id": "scan",
"name": "LinearAttention.ChunkScan",
"shader": "chunk-scan.wgsl.jinja",
"derive": { "workgroupSize": "chunkScanWorkgroup" },
"bindings": ["key", "value", "past_state", "states", "present_state", "params_chunk_linear"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"z": 1
}
},
{
"id": "out",
"name": "LinearAttention.ChunkOutput",
"shader": "chunk-out.wgsl.jinja",
"derive": { "workgroupSize": "chunkOutWorkgroup" },
"bindings": ["query", "key", "states_f32", "value", "output", "params_chunk_linear"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"z": 1
}
}
]
},
{
"id": "linear_state_scalar",
"priority": 0,
"when": ["effRule == \"linear\"", "present.pastStateT", "commonContract", "scalarHeadDimFits"],
"derive": {
"updateRule": "\"linear\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16",
"workgroupSize": "min(256, pow2ceil(dim(shapes.queryT, 2) / attrs.q_num_heads))",
"tileV": "max(1, min(tunables.tileV, headDimV))",
"outputDtype": "tensorDtypes.queryT"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.scalar.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "output", "present_state", "params_win_linear"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"z": 1
}
}
]
},
{
"id": "linear_state_vec4",
"priority": 10,
"when": ["effRule == \"linear\"", "present.pastStateT", "commonContract", "headDimK % 4 == 0", "vec4SharedFitsPlain", "vec4HeadDimFits", "vec4BindingsNonEmpty"],
"derive": {
"updateRule": "\"linear\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "vec4UseSubgroups",
"tileV": "vec4TileVPlain",
"queryElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"keyElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"dvGroups": "vec4DvGroups"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.vec4.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "output", "present_state", "params_win_linear"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_zero_chunked",
"priority": 30,
"when": ["effRule == \"gated\"", "not present.pastStateT", "commonContract", "chunkedShapeOk"],
"derive": { "hasPastState": "present.pastStateT", "usesDecay": "needsDecay", "usesBeta": "needsBeta" },
"intermediates": [
{ "id": "gexp", "dtype": "float32", "shape": "[chunkGexpElems]" },
{ "id": "states", "dtype": "float32", "shape": "[chunkStatesElems]" }
],
"passes": [
{
"id": "prep",
"name": "LinearAttention.ChunkPrep",
"shader": "chunk-prep.wgsl.jinja",
"derive": { "workgroupSize": "chunkFlatWorkgroup" },
"bindings": ["decay", "gexp", "params_chunk_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * chunkNumChunks, 65535)",
"z": 1
}
},
{
"id": "scan",
"name": "LinearAttention.ChunkScan",
"shader": "chunk-scan.wgsl.jinja",
"derive": { "workgroupSize": "chunkScanWorkgroup" },
"bindings": ["key", "gexp_ut", "value", "states", "present_state", "params_chunk_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"z": 1
}
},
{
"id": "out",
"name": "LinearAttention.ChunkOutput",
"shader": "chunk-out.wgsl.jinja",
"derive": { "workgroupSize": "chunkOutWorkgroup" },
"bindings": ["query", "key", "gexp_ut", "states_f32", "value", "output", "params_chunk_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"z": 1
}
}
]
},
{
"id": "gated_zero_scalar",
"priority": 0,
"when": ["effRule == \"gated\"", "not present.pastStateT", "commonContract", "scalarHeadDimFits"],
"derive": {
"updateRule": "\"gated\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16",
"workgroupSize": "min(256, pow2ceil(dim(shapes.queryT, 2) / attrs.q_num_heads))",
"tileV": "max(1, min(tunables.tileV, headDimV))",
"outputDtype": "tensorDtypes.queryT"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.scalar.wgsl.jinja",
"bindings": ["query", "key", "value", "decay", "output", "present_state", "params_win_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_zero_vec4",
"priority": 10,
"when": ["effRule == \"gated\"", "not present.pastStateT", "commonContract", "headDimK % 4 == 0", "vec4SharedFitsPlain", "vec4HeadDimFits", "vec4BindingsNonEmpty"],
"derive": {
"updateRule": "\"gated\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "vec4UseSubgroups",
"tileV": "vec4TileVPlain",
"queryElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"keyElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"dvGroups": "vec4DvGroups"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.vec4.wgsl.jinja",
"bindings": ["query", "key", "value", "decay", "output", "present_state", "params_win_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_state_chunked",
"priority": 30,
"when": ["effRule == \"gated\"", "present.pastStateT", "commonContract", "chunkedShapeOk"],
"derive": { "hasPastState": "present.pastStateT", "usesDecay": "needsDecay", "usesBeta": "needsBeta" },
"intermediates": [
{ "id": "gexp", "dtype": "float32", "shape": "[chunkGexpElems]" },
{ "id": "states", "dtype": "float32", "shape": "[chunkStatesElems]" }
],
"passes": [
{
"id": "prep",
"name": "LinearAttention.ChunkPrep",
"shader": "chunk-prep.wgsl.jinja",
"derive": { "workgroupSize": "chunkFlatWorkgroup" },
"bindings": ["decay", "gexp", "params_chunk_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * chunkNumChunks, 65535)",
"z": 1
}
},
{
"id": "scan",
"name": "LinearAttention.ChunkScan",
"shader": "chunk-scan.wgsl.jinja",
"derive": { "workgroupSize": "chunkScanWorkgroup" },
"bindings": ["key", "gexp_ut", "value", "past_state", "states", "present_state", "params_chunk_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"z": 1
}
},
{
"id": "out",
"name": "LinearAttention.ChunkOutput",
"shader": "chunk-out.wgsl.jinja",
"derive": { "workgroupSize": "chunkOutWorkgroup" },
"bindings": ["query", "key", "gexp_ut", "states_f32", "value", "output", "params_chunk_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"z": 1
}
}
]
},
{
"id": "gated_state_scalar",
"priority": 0,
"when": ["effRule == \"gated\"", "present.pastStateT", "commonContract", "scalarHeadDimFits"],
"derive": {
"updateRule": "\"gated\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16",
"workgroupSize": "min(256, pow2ceil(dim(shapes.queryT, 2) / attrs.q_num_heads))",
"tileV": "max(1, min(tunables.tileV, headDimV))",
"outputDtype": "tensorDtypes.queryT"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.scalar.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "decay", "output", "present_state", "params_win_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_state_vec4",
"priority": 10,
"when": ["effRule == \"gated\"", "present.pastStateT", "commonContract", "headDimK % 4 == 0", "vec4SharedFitsPlain", "vec4HeadDimFits", "vec4BindingsNonEmpty"],
"derive": {
"updateRule": "\"gated\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "vec4UseSubgroups",
"tileV": "vec4TileVPlain",
"queryElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"keyElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"dvGroups": "vec4DvGroups"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.vec4.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "decay", "output", "present_state", "params_win_gated"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"z": 1
}
}
]
},
{
"id": "delta_zero_chunked",
"priority": 30,
"when": ["effRule == \"delta\"", "not present.pastStateT", "commonContract", "chunkedShapeOk"],
"derive": { "hasPastState": "present.pastStateT", "usesDecay": "needsDecay", "usesBeta": "needsBeta" },
"intermediates": [
{ "id": "wk", "dtype": "float32", "shape": "[chunkWkElems]" },
{ "id": "uvec", "dtype": "float32", "shape": "[chunkUvecElems]" },
{ "id": "states", "dtype": "float32", "shape": "[chunkStatesElems]" },
{ "id": "deltas", "dtype": "float32", "shape": "[chunkUvecElems]" }
],
"passes": [
{
"id": "ut",
"name": "LinearAttention.ChunkTransform",
"shader": "chunk-ut.wgsl.jinja",
"derive": { "workgroupSize": "chunkFlatWorkgroup", "chunkTileK": "chunkUtTileK" },
"bindings": ["key", "value", "beta", "wk", "uvec", "params_chunk_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks, 65535)",
"z": 1
}
},
{
"id": "scan",
"name": "LinearAttention.ChunkScan",
"shader": "chunk-scan.wgsl.jinja",
"derive": { "workgroupSize": "chunkScanWorkgroup" },
"bindings": ["key", "wk_f32", "uvec_f32", "states", "deltas", "present_state", "params_chunk_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"z": 1
}
},
{
"id": "out",
"name": "LinearAttention.ChunkOutput",
"shader": "chunk-out.wgsl.jinja",
"derive": { "workgroupSize": "chunkOutWorkgroup" },
"bindings": ["query", "key", "states_f32", "deltas_f32", "output", "params_chunk_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"z": 1
}
}
]
},
{
"id": "delta_zero_scalar",
"priority": 0,
"when": ["effRule == \"delta\"", "not present.pastStateT", "commonContract", "scalarHeadDimFits"],
"derive": {
"updateRule": "\"delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16",
"workgroupSize": "min(256, pow2ceil(dim(shapes.queryT, 2) / attrs.q_num_heads))",
"tileV": "max(1, min(tunables.gatedTileV, headDimV))",
"outputDtype": "tensorDtypes.queryT"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.scalar.wgsl.jinja",
"bindings": ["query", "key", "value", "beta", "output", "present_state", "params_win_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"z": 1
}
}
]
},
{
"id": "delta_zero_vec4",
"priority": 10,
"when": ["effRule == \"delta\"", "not present.pastStateT", "commonContract", "headDimK % 4 == 0", "vec4SharedFitsGatedBeta", "vec4HeadDimFits", "vec4BindingsNonEmpty"],
"derive": {
"updateRule": "\"delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "vec4UseSubgroups",
"tileV": "gatedVec4TileV",
"queryElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"keyElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"dvGroups": "vec4DvGroups"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.vec4.wgsl.jinja",
"bindings": ["query", "key", "value", "beta", "output", "present_state", "params_win_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"z": 1
}
}
]
},
{
"id": "delta_state_chunked",
"priority": 30,
"when": ["effRule == \"delta\"", "present.pastStateT", "commonContract", "chunkedShapeOk"],
"derive": { "hasPastState": "present.pastStateT", "usesDecay": "needsDecay", "usesBeta": "needsBeta" },
"intermediates": [
{ "id": "wk", "dtype": "float32", "shape": "[chunkWkElems]" },
{ "id": "uvec", "dtype": "float32", "shape": "[chunkUvecElems]" },
{ "id": "states", "dtype": "float32", "shape": "[chunkStatesElems]" },
{ "id": "deltas", "dtype": "float32", "shape": "[chunkUvecElems]" }
],
"passes": [
{
"id": "ut",
"name": "LinearAttention.ChunkTransform",
"shader": "chunk-ut.wgsl.jinja",
"derive": { "workgroupSize": "chunkFlatWorkgroup", "chunkTileK": "chunkUtTileK" },
"bindings": ["key", "value", "beta", "wk", "uvec", "params_chunk_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks, 65535)",
"z": 1
}
},
{
"id": "scan",
"name": "LinearAttention.ChunkScan",
"shader": "chunk-scan.wgsl.jinja",
"derive": { "workgroupSize": "chunkScanWorkgroup" },
"bindings": ["key", "wk_f32", "uvec_f32", "past_state", "states", "deltas", "present_state", "params_chunk_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"z": 1
}
},
{
"id": "out",
"name": "LinearAttention.ChunkOutput",
"shader": "chunk-out.wgsl.jinja",
"derive": { "workgroupSize": "chunkOutWorkgroup" },
"bindings": ["query", "key", "states_f32", "deltas_f32", "output", "params_chunk_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"z": 1
}
}
]
},
{
"id": "delta_state_scalar",
"priority": 0,
"when": ["effRule == \"delta\"", "present.pastStateT", "commonContract", "scalarHeadDimFits"],
"derive": {
"updateRule": "\"delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16",
"workgroupSize": "min(256, pow2ceil(dim(shapes.queryT, 2) / attrs.q_num_heads))",
"tileV": "max(1, min(tunables.gatedTileV, headDimV))",
"outputDtype": "tensorDtypes.queryT"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.scalar.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "beta", "output", "present_state", "params_win_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"z": 1
}
}
]
},
{
"id": "delta_state_vec4",
"priority": 10,
"when": ["effRule == \"delta\"", "present.pastStateT", "commonContract", "headDimK % 4 == 0", "vec4SharedFitsGatedBeta", "vec4HeadDimFits", "vec4BindingsNonEmpty"],
"derive": {
"updateRule": "\"delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "vec4UseSubgroups",
"tileV": "gatedVec4TileV",
"queryElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"keyElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"dvGroups": "vec4DvGroups"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.vec4.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "beta", "output", "present_state", "params_win_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_delta_zero_chunked",
"priority": 30,
"when": ["effRule == \"gated_delta\"", "not present.pastStateT", "commonContract", "chunkedShapeOk"],
"derive": { "hasPastState": "present.pastStateT", "usesDecay": "needsDecay", "usesBeta": "needsBeta" },
"intermediates": [
{ "id": "gexp", "dtype": "float32", "shape": "[chunkGexpElems]" },
{ "id": "wk", "dtype": "float32", "shape": "[chunkWkElems]" },
{ "id": "uvec", "dtype": "float32", "shape": "[chunkUvecElems]" },
{ "id": "states", "dtype": "float32", "shape": "[chunkStatesElems]" },
{ "id": "deltas", "dtype": "float32", "shape": "[chunkUvecElems]" }
],
"passes": [
{
"id": "prep",
"name": "LinearAttention.ChunkPrep",
"shader": "chunk-prep.wgsl.jinja",
"derive": { "workgroupSize": "chunkFlatWorkgroup" },
"bindings": ["decay", "gexp", "params_chunk_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * chunkNumChunks, 65535)",
"z": 1
}
},
{
"id": "ut",
"name": "LinearAttention.ChunkTransform",
"shader": "chunk-ut.wgsl.jinja",
"derive": { "workgroupSize": "chunkFlatWorkgroup", "chunkTileK": "chunkUtTileK" },
"bindings": ["key", "value", "beta", "gexp_ut", "wk", "uvec", "params_chunk_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks, 65535)",
"z": 1
}
},
{
"id": "scan",
"name": "LinearAttention.ChunkScan",
"shader": "chunk-scan.wgsl.jinja",
"derive": { "workgroupSize": "chunkScanWorkgroup" },
"bindings": ["key", "gexp_ut", "wk_f32", "uvec_f32", "states", "deltas", "present_state", "params_chunk_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"z": 1
}
},
{
"id": "out",
"name": "LinearAttention.ChunkOutput",
"shader": "chunk-out.wgsl.jinja",
"derive": { "workgroupSize": "chunkOutWorkgroup" },
"bindings": ["query", "key", "gexp_ut", "states_f32", "deltas_f32", "output", "params_chunk_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"z": 1
}
}
]
},
{
"id": "gated_delta_zero_scalar",
"priority": 0,
"when": ["effRule == \"gated_delta\"", "not present.pastStateT", "commonContract", "scalarHeadDimFits"],
"derive": {
"updateRule": "\"gated_delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16",
"workgroupSize": "min(256, pow2ceil(dim(shapes.queryT, 2) / attrs.q_num_heads))",
"tileV": "max(1, min(tunables.gatedTileV, headDimV))",
"outputDtype": "tensorDtypes.queryT"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.scalar.wgsl.jinja",
"bindings": ["query", "key", "value", "decay", "beta", "output", "present_state", "params_win_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_delta_zero_vec4",
"priority": 10,
"when": ["effRule == \"gated_delta\"", "not present.pastStateT", "commonContract", "headDimK % 4 == 0", "vec4SharedFitsGatedBeta", "vec4HeadDimFits", "vec4BindingsNonEmpty"],
"derive": {
"updateRule": "\"gated_delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "vec4UseSubgroups",
"tileV": "gatedVec4TileV",
"queryElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"keyElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"dvGroups": "vec4DvGroups"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.vec4.wgsl.jinja",
"bindings": ["query", "key", "value", "decay", "beta", "output", "present_state", "params_win_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_delta_state_chunked",
"priority": 30,
"when": ["effRule == \"gated_delta\"", "present.pastStateT", "commonContract", "chunkedShapeOk"],
"derive": { "hasPastState": "present.pastStateT", "usesDecay": "needsDecay", "usesBeta": "needsBeta" },
"intermediates": [
{ "id": "gexp", "dtype": "float32", "shape": "[chunkGexpElems]" },
{ "id": "wk", "dtype": "float32", "shape": "[chunkWkElems]" },
{ "id": "uvec", "dtype": "float32", "shape": "[chunkUvecElems]" },
{ "id": "states", "dtype": "float32", "shape": "[chunkStatesElems]" },
{ "id": "deltas", "dtype": "float32", "shape": "[chunkUvecElems]" }
],
"passes": [
{
"id": "prep",
"name": "LinearAttention.ChunkPrep",
"shader": "chunk-prep.wgsl.jinja",
"derive": { "workgroupSize": "chunkFlatWorkgroup" },
"bindings": ["decay", "gexp", "params_chunk_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * chunkNumChunks, 65535)",
"z": 1
}
},
{
"id": "ut",
"name": "LinearAttention.ChunkTransform",
"shader": "chunk-ut.wgsl.jinja",
"derive": { "workgroupSize": "chunkFlatWorkgroup", "chunkTileK": "chunkUtTileK" },
"bindings": ["key", "value", "beta", "gexp_ut", "wk", "uvec", "params_chunk_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * chunkNumChunks, 65535)",
"z": 1
}
},
{
"id": "scan",
"name": "LinearAttention.ChunkScan",
"shader": "chunk-scan.wgsl.jinja",
"derive": { "workgroupSize": "chunkScanWorkgroup" },
"bindings": ["key", "gexp_ut", "wk_f32", "uvec_f32", "past_state", "states", "deltas", "present_state", "params_chunk_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV / chunkTileV), 65535)",
"z": 1
}
},
{
"id": "out",
"name": "LinearAttention.ChunkOutput",
"shader": "chunk-out.wgsl.jinja",
"derive": { "workgroupSize": "chunkOutWorkgroup" },
"bindings": ["query", "key", "gexp_ut", "states_f32", "deltas_f32", "output", "params_chunk_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * outHeads * chunkNumChunks, 65535)",
"z": 1
}
}
]
},
{
"id": "gated_delta_state_scalar",
"priority": 0,
"when": ["effRule == \"gated_delta\"", "present.pastStateT", "commonContract", "scalarHeadDimFits"],
"derive": {
"updateRule": "\"gated_delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16",
"workgroupSize": "min(256, pow2ceil(dim(shapes.queryT, 2) / attrs.q_num_heads))",
"tileV": "max(1, min(tunables.gatedTileV, headDimV))",
"outputDtype": "tensorDtypes.queryT"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.scalar.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "decay", "beta", "output", "present_state", "params_win_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(headDimV, tileV), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_delta_state_vec4",
"priority": 10,
"when": ["effRule == \"gated_delta\"", "present.pastStateT", "commonContract", "headDimK % 4 == 0", "vec4SharedFitsGatedBeta", "vec4HeadDimFits", "vec4BindingsNonEmpty"],
"derive": {
"updateRule": "\"gated_delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"useSubgroups": "vec4UseSubgroups",
"tileV": "gatedVec4TileV",
"queryElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"keyElem": "\"vec4<f16>\" if tensorDtypes.queryT == \"float16\" else \"vec4<f32>\"",
"dvGroups": "vec4DvGroups"
},
"passes": [
{
"id": "main",
"name": "LinearAttention",
"shader": "linear-attention.vec4.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "decay", "beta", "output", "present_state", "params_win_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(headDimV, tileV), dvGroups), 65535)",
"z": 1
}
}
]
},
{
"id": "linear_zero_serial_small_dk",
"priority": 20,
"when": ["effRule == \"linear\"", "not present.pastStateT", "commonContract", "serialHeadDimFits"],
"derive": {
"updateRule": "\"linear\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"outputDtype": "tensorDtypes.queryT",
"serialBatchSize": "(dim(shapes.queryT, 0)) if updateRule == \"linear\" else 0",
"serialSeqLength": "(dim(shapes.queryT, 1)) if updateRule == \"linear\" else 0",
"serialQNumHeads": "(attrs.q_num_heads) if updateRule == \"linear\" else 0",
"serialKvNumHeads": "(attrs.kv_num_heads) if updateRule == \"linear\" else 0",
"serialQPackedDim": "(dim(shapes.queryT, 2)) if updateRule == \"linear\" else 0",
"serialKPackedDim": "(dim(shapes.keyT, 2)) if updateRule == \"linear\" else 0",
"serialVPackedDim": "(dim(shapes.valueT, 2)) if updateRule == \"linear\" else 0",
"serialScale": "(attrs.scale if attrs.scale else 0) if updateRule == \"linear\" else 0",
"serialStateWindow": "(stateWindow) if updateRule == \"linear\" else 0",
"serialStateSlotStride": "stateSlotStride if updateRule == \"linear\" else 0"
},
"passes": [
{
"id": "main",
"name": "LinearAttention.SerialSmallDk",
"shader": "linear-attention.serial.wgsl.jinja",
"bindings": ["query", "key", "value", "output", "present_state"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV), 65535)",
"z": 1
}
}
]
},
{
"id": "linear_state_serial_small_dk",
"priority": 20,
"when": ["effRule == \"linear\"", "present.pastStateT", "commonContract", "serialHeadDimFits"],
"derive": {
"updateRule": "\"linear\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"outputDtype": "tensorDtypes.queryT",
"serialBatchSize": "(dim(shapes.queryT, 0)) if updateRule == \"linear\" else 0",
"serialSeqLength": "(dim(shapes.queryT, 1)) if updateRule == \"linear\" else 0",
"serialQNumHeads": "(attrs.q_num_heads) if updateRule == \"linear\" else 0",
"serialKvNumHeads": "(attrs.kv_num_heads) if updateRule == \"linear\" else 0",
"serialQPackedDim": "(dim(shapes.queryT, 2)) if updateRule == \"linear\" else 0",
"serialKPackedDim": "(dim(shapes.keyT, 2)) if updateRule == \"linear\" else 0",
"serialVPackedDim": "(dim(shapes.valueT, 2)) if updateRule == \"linear\" else 0",
"serialScale": "(attrs.scale if attrs.scale else 0) if updateRule == \"linear\" else 0",
"serialStateWindow": "(stateWindow) if updateRule == \"linear\" else 0",
"serialStateSlotStride": "stateSlotStride if updateRule == \"linear\" else 0"
},
"passes": [
{
"id": "main",
"name": "LinearAttention.SerialSmallDk",
"shader": "linear-attention.serial.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "output", "present_state"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_delta_zero_serial_small_dk",
"priority": 20,
"when": ["effRule == \"gated_delta\"", "not present.pastStateT", "commonContract", "serialHeadDimFits"],
"derive": {
"updateRule": "\"gated_delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"outputDtype": "tensorDtypes.queryT",
"serialBatchSize": "(dim(shapes.queryT, 0)) if updateRule == \"linear\" else 0",
"serialSeqLength": "(dim(shapes.queryT, 1)) if updateRule == \"linear\" else 0",
"serialQNumHeads": "(attrs.q_num_heads) if updateRule == \"linear\" else 0",
"serialKvNumHeads": "(attrs.kv_num_heads) if updateRule == \"linear\" else 0",
"serialQPackedDim": "(dim(shapes.queryT, 2)) if updateRule == \"linear\" else 0",
"serialKPackedDim": "(dim(shapes.keyT, 2)) if updateRule == \"linear\" else 0",
"serialVPackedDim": "(dim(shapes.valueT, 2)) if updateRule == \"linear\" else 0",
"serialScale": "(attrs.scale if attrs.scale else 0) if updateRule == \"linear\" else 0",
"serialStateWindow": "(stateWindow) if updateRule == \"linear\" else 0",
"serialStateSlotStride": "stateSlotStride if updateRule == \"linear\" else 0"
},
"passes": [
{
"id": "main",
"name": "LinearAttention.SerialSmallDk",
"shader": "linear-attention.serial.wgsl.jinja",
"bindings": ["query", "key", "value", "decay", "beta", "output", "present_state", "params_win_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV), 65535)",
"z": 1
}
}
]
},
{
"id": "gated_delta_state_serial_small_dk",
"priority": 20,
"when": ["effRule == \"gated_delta\"", "present.pastStateT", "commonContract", "serialHeadDimFits"],
"derive": {
"updateRule": "\"gated_delta\"",
"hasPastState": "present.pastStateT",
"hasStateWindow": "windowed",
"outputDtype": "tensorDtypes.queryT",
"serialBatchSize": "(dim(shapes.queryT, 0)) if updateRule == \"linear\" else 0",
"serialSeqLength": "(dim(shapes.queryT, 1)) if updateRule == \"linear\" else 0",
"serialQNumHeads": "(attrs.q_num_heads) if updateRule == \"linear\" else 0",
"serialKvNumHeads": "(attrs.kv_num_heads) if updateRule == \"linear\" else 0",
"serialQPackedDim": "(dim(shapes.queryT, 2)) if updateRule == \"linear\" else 0",
"serialKPackedDim": "(dim(shapes.keyT, 2)) if updateRule == \"linear\" else 0",
"serialVPackedDim": "(dim(shapes.valueT, 2)) if updateRule == \"linear\" else 0",
"serialScale": "(attrs.scale if attrs.scale else 0) if updateRule == \"linear\" else 0",
"serialStateWindow": "(stateWindow) if updateRule == \"linear\" else 0",
"serialStateSlotStride": "stateSlotStride if updateRule == \"linear\" else 0"
},
"passes": [
{
"id": "main",
"name": "LinearAttention.SerialSmallDk",
"shader": "linear-attention.serial.wgsl.jinja",
"bindings": ["query", "key", "value", "past_state", "decay", "beta", "output", "present_state", "params_win_gated_delta"],
"dispatch": {
"x": "min(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV), 65535)",
"y": "ceilDiv(dim(shapes.queryT, 0) * attrs.kv_num_heads * (headDimV), 65535)",
"z": 1
}
}
]
}
]
}