{ "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\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "keyElem": "\"vec4\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "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\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "keyElem": "\"vec4\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "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\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "keyElem": "\"vec4\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "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\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "keyElem": "\"vec4\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "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\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "keyElem": "\"vec4\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "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\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "keyElem": "\"vec4\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "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\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "keyElem": "\"vec4\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "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\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "keyElem": "\"vec4\" if tensorDtypes.queryT == \"float16\" else \"vec4\"", "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 } } ] } ] }