diff --git "a/build/webgpu/manifest.json" "b/build/webgpu/manifest.json" new file mode 100644--- /dev/null +++ "b/build/webgpu/manifest.json" @@ -0,0 +1,3118 @@ +{ + "domain": "com.microsoft", + "name": "LinearAttention", + "sinceVersion": 1, + "description": "Recurrent linear attention for packed `[B, T, H*D]` decode and prefill. It supports all four update rules, standard and inverse GQA, shared-key heads, and rollback states through `state_window`. Activations and state may independently use float16 or float32; bfloat16 is not implemented. `past_state` is optional for every update rule and defaults to zeros.", + "inputs": [ + { + "role": "query", + "dtype": "T", + "rank": 3, + "description": "Query vectors with 3D packed shape `(B, T, H_q * d_k)`; heads are packed into the last dimension." + }, + { + "role": "key", + "dtype": "T", + "rank": 3, + "description": "Key vectors with 3D packed shape `(B, T, H_k * d_k)`, where positive `H_k` divides `H_kv`; `H_k < H_kv` shares each key head across multiple KV-state heads. Keys should be L2-normalized for `delta`/`gated_delta` modes." + }, + { + "role": "value", + "dtype": "T", + "rank": 3, + "description": "Value vectors with 3D packed shape `(B, T, H_kv * d_v)`." + }, + { + "role": "past_state", + "dtype": "S", + "rank": "5 if attrs.state_window > 0 else 4", + "shape": "[attrs.state_window, dim(shapes.query, 0), attrs.kv_num_heads, dim(shapes.query, 2) / attrs.q_num_heads, dim(shapes.value, 2) / attrs.kv_num_heads] if attrs.state_window > 0 else [dim(shapes.query, 0), attrs.kv_num_heads, dim(shapes.query, 2) / attrs.q_num_heads, dim(shapes.value, 2) / attrs.kv_num_heads]", + "optional": true, + "description": "Recurrent state from the previous step with shape `(B, H_kv, d_k, d_v)`, or `(W, B, H_kv, d_k, d_v)` when `state_window = W > 0`; defaults to zeros if absent." + }, + { + "role": "decay", + "dtype": "T", + "rank": 3, + "optional": true, + "description": "Exponential decay gate in log-space with shape `(B, T, H_kv * d_k)` or `(B, T, H_kv)`; required for `gated` and `gated_delta` modes." + }, + { + "role": "beta", + "dtype": "T", + "rank": 3, + "optional": true, + "description": "Update rate (sigmoid output) with shape `(B, T, H_kv)` or `(B, T, 1)`; required for `delta` and `gated_delta` modes." + } + ], + "outputs": [ + { + "role": "output", + "dtype": "T", + "rank": 3, + "shape": "[dim(shapes.query, 0), dim(shapes.query, 1), max(attrs.q_num_heads, attrs.kv_num_heads) * (dim(shapes.value, 2) / attrs.kv_num_heads)]", + "description": "Attention output with 3D packed shape `(B, T, max(H_q, H_kv) * d_v)`." + }, + { + "role": "present_state", + "dtype": "S", + "rank": "5 if attrs.state_window > 0 else 4", + "shape": "[attrs.state_window, dim(shapes.query, 0), attrs.kv_num_heads, dim(shapes.query, 2) / attrs.q_num_heads, dim(shapes.value, 2) / attrs.kv_num_heads] if attrs.state_window > 0 else [dim(shapes.query, 0), attrs.kv_num_heads, dim(shapes.query, 2) / attrs.q_num_heads, dim(shapes.value, 2) / attrs.kv_num_heads]", + "description": "Updated recurrent state with shape `(B, H_kv, d_k, d_v)`, or `(W, B, H_kv, d_k, d_v)` when `state_window = W > 0`." + } + ], + "attributes": { "chunk_size": 64, "scale": 0, "state_window": 0, "update_rule": "gated_delta" }, + "attributeDescriptions": { + "chunk_size": "Accepted for schema compatibility; does not affect the result.", + "kv_num_heads": "Number of key/value heads.", + "q_num_heads": "Number of query heads.", + "scale": "Scale applied to query-key products. Zero selects `1 / sqrt(d_k)`.", + "state_window": "Number of recent recurrent states retained in `present_state`, in the supported range 0 to 8; zero returns only the current state.", + "update_rule": "Recurrent update rule: `linear`, `gated`, `delta`, or `gated_delta`." + }, + "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"] }, + "args": { + "queryT": { "kind": "tensor", "semantic": "query", "role": "input" }, + "keyT": { "kind": "tensor", "semantic": "key", "role": "input" }, + "valueT": { "kind": "tensor", "semantic": "value", "role": "input" }, + "pastStateT": { "kind": "tensor", "semantic": "past_state", "role": "input", "required": false }, + "decayT": { "kind": "tensor", "semantic": "decay", "role": "input", "required": false }, + "betaT": { "kind": "tensor", "semantic": "beta", "role": "input", "required": false }, + "outputT": { "kind": "tensor", "semantic": "output", "role": "output" }, + "presentStateT": { "kind": "tensor", "semantic": "present_state", "role": "output" } + }, + "tunables": { + "tileV": 8, + "gatedTileV": 4, + "chunkSize": 16, + "chunkScanTokens": 8, + "chunkTileV": 32, + "chunkGroups": 8, + "chunkOutRows": 16, + "dvGroups": 4 + }, + "derive": { + "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)", + "effRule": "attrs.update_rule if attrs.update_rule else \"gated_delta\"", + "headDimK": "dim(shapes.query, 2) / max(1, attrs.q_num_heads)", + "headDimV": "dim(shapes.value, 2) / max(1, attrs.kv_num_heads)", + "nKeyHeads": "dim(shapes.key, 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.query, 2) > 0 and dim(shapes.value, 2) > 0 and dim(shapes.key, 2) > 0 and dim(shapes.query, 2) % attrs.q_num_heads == 0 and dim(shapes.value, 2) % attrs.kv_num_heads == 0 and dim(shapes.key, 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", + "batchSequenceOk": "dim(shapes.key, 0) == dim(shapes.query, 0) and dim(shapes.key, 1) == dim(shapes.query, 1) and dim(shapes.value, 0) == dim(shapes.query, 0) and dim(shapes.value, 1) == dim(shapes.query, 1)", + "scalarHeadDimFits": "headDimK <= deviceWorkgroupCap and headDimK <= 256", + "serialHeadDimFits": "headDimK > 0 and headDimK <= 16", + "stateWindow": "attrs.state_window if attrs.state_window is defined else 0", + "windowed": "stateWindow > 0", + "stateWindowOk": "stateWindow >= 0 and stateWindow <= 8", + "presentStateOk": "(dim(shapes.present_state, 0) == dim(shapes.query, 0) and dim(shapes.present_state, 1) == attrs.kv_num_heads and dim(shapes.present_state, 2) == headDimK and dim(shapes.present_state, 3) == (dim(shapes.value, 2) / attrs.kv_num_heads) and ranks.present_state == 4) if not windowed else (ranks.present_state == 5 and dim(shapes.present_state, 0) == stateWindow and dim(shapes.present_state, 1) == dim(shapes.query, 0) and dim(shapes.present_state, 2) == attrs.kv_num_heads and dim(shapes.present_state, 3) == headDimK and dim(shapes.present_state, 4) == (dim(shapes.value, 2) / attrs.kv_num_heads))", + "ioContractOk": "dim(shapes.outputT, 0) == dim(shapes.query, 0) and dim(shapes.outputT, 1) == dim(shapes.query, 1) and dim(shapes.outputT, 2) == outHeads * headDimV and stateWindowOk and presentStateOk", + "pastStateOk": "not present.pastStateT or ((ranks.past_state == 4 and dim(shapes.past_state, 0) == dim(shapes.query, 0) and dim(shapes.past_state, 1) == attrs.kv_num_heads and dim(shapes.past_state, 2) == headDimK and dim(shapes.past_state, 3) == headDimV) if not windowed else (ranks.past_state == 5 and dim(shapes.past_state, 0) == stateWindow and dim(shapes.past_state, 1) == dim(shapes.query, 0) and dim(shapes.past_state, 2) == attrs.kv_num_heads and dim(shapes.past_state, 3) == headDimK and dim(shapes.past_state, 4) == headDimV))", + "decayOk": "not present.decayT or (ranks.decay == 3 and dim(shapes.decay, 0) == dim(shapes.query, 0) and dim(shapes.decay, 1) == dim(shapes.query, 1) and (dim(shapes.decay, 2) == attrs.kv_num_heads or dim(shapes.decay, 2) == attrs.kv_num_heads * headDimK) and tensorDtypes.decay == tensorDtypes.query)", + "betaOk": "not present.betaT or (ranks.beta == 3 and dim(shapes.beta, 0) == dim(shapes.query, 0) and dim(shapes.beta, 1) == dim(shapes.query, 1) and (dim(shapes.beta, 2) == 1 or dim(shapes.beta, 2) == attrs.kv_num_heads) and tensorDtypes.beta == tensorDtypes.query)", + "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.query == \"float16\" or tensorDtypes.query == \"float32\") and tensorDtypes.key == tensorDtypes.query and tensorDtypes.value == tensorDtypes.query and tensorDtypes.outputT == tensorDtypes.query and (tensorDtypes.present_state == \"float16\" or tensorDtypes.present_state == \"float32\") and (not present.pastStateT or tensorDtypes.past_state == tensorDtypes.present_state) and f16Ok(tensorDtypes.query) and f16Ok(tensorDtypes.present_state)", + "commonContract": "headLayoutOk and batchSequenceOk and gateInputsOk and tensorDtypesOk and ioContractOk and pastStateOk", + "linearZeroContract": "effRule == \"linear\" and not present.pastStateT and commonContract", + "linearStateContract": "effRule == \"linear\" and present.pastStateT and commonContract", + "gatedZeroContract": "effRule == \"gated_delta\" and not present.pastStateT and commonContract", + "gatedStateContract": "effRule == \"gated_delta\" and present.pastStateT and commonContract", + "gatedOnlyZeroContract": "effRule == \"gated\" and not present.pastStateT and commonContract", + "gatedOnlyStateContract": "effRule == \"gated\" and present.pastStateT and commonContract", + "deltaZeroContract": "effRule == \"delta\" and not present.pastStateT and commonContract", + "deltaStateContract": "effRule == \"delta\" and present.pastStateT and commonContract", + "vec4DtypeOk": "tensorDtypes.key == tensorDtypes.query and tensorDtypes.value == tensorDtypes.query and tensorDtypes.outputT == tensorDtypes.query", + "gatedVec4DtypeOk": "vec4DtypeOk and tensorDtypes.decay == tensorDtypes.query and tensorDtypes.beta == tensorDtypes.query", + "decayVec4DtypeOk": "vec4DtypeOk and tensorDtypes.decay == tensorDtypes.query", + "betaVec4DtypeOk": "vec4DtypeOk and tensorDtypes.beta == tensorDtypes.query", + "vec4Lanes": "min(deviceWorkgroupCap, pow2ceil(ceil(headDimK / 4)))", + "vec4SubgroupsAvailable": "device.features.has(\"subgroups\") and has(device.adapterInfo, \"subgroupMinSize\") and vec4Lanes <= device.adapterInfo.subgroupMinSize", + "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", + "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.query, 0) > 0 and dim(shapes.query, 1) > 0", + "vec4UseSubgroups": "vec4SubgroupsAvailable and (vec4DvGroups == 1 or vec4SubgroupExact)", + "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, dim(shapes.value, 2) / attrs.kv_num_heads)) if vec4UseSubgroups else max(1, min(gatedVec4TileVBudget, dim(shapes.value, 2) / attrs.kv_num_heads))", + "vec4TileVPlain": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "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.query, 1), chunkSize)", + "kvPerKeyHead": "attrs.kv_num_heads / max(1, nKeyHeads)", + "decayPerElement": "present.decayT and dim(shapes.decay, 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.query, 0) * dim(shapes.query, 1) * (dim(shapes.query, 2) + dim(shapes.key, 2) + dim(shapes.value, 2)) * (2 if tensorDtypes.query == \"float16\" else 4)", + "chunkStatesBudget": "min(chunkActivationBytes, 134217728)", + "chunkStatesBytes": "dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV * 4", + "chunkWkBytes": "dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimK * 4", + "chunkUvecBytes": "dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV * 4", + "chunkGexpBytes": "dim(shapes.query, 0) * dim(shapes.query, 1) * dim(shapes.decay, 2) * 4 if present.decayT else 0", + "chunkScratchOk": "max(chunkStatesBytes, max(chunkWkBytes, max(chunkUvecBytes, chunkGexpBytes))) <= min(device.limits.maxStorageBufferBindingSize, device.limits.maxBufferSize) and chunkStatesBytes <= 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.query, 1) >= 1024 and chunkGeometryOk and chunkSharedOk and chunkScratchOk" + }, + "bindingSets": { + "baseStateIo": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + } + ], + "baseZeroIo": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + } + ], + "linearZeroIo": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + } + ], + "linearStateIo": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + } + ], + "gatedOnlyZeroIo": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + } + ], + "gatedOnlyStateIo": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + } + ], + "deltaZeroIo": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + } + ], + "deltaStateIo": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + } + ], + "baseState": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 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": "dim(shapes.query, 0) * attrs.kv_num_heads * headDimK * headDimV" + } + ] + } + } + ], + "baseZero": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 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": "dim(shapes.query, 0) * attrs.kv_num_heads * headDimK * headDimV" + } + ] + } + } + ], + "linearZero": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" }, + { "name": "stateWindow", "type": "u32", "value": "stateWindow" }, + { + "name": "stateSlotStride", + "type": "u32", + "value": "dim(shapes.query, 0) * attrs.kv_num_heads * headDimK * headDimV" + } + ] + } + } + ], + "linearState": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" }, + { "name": "stateWindow", "type": "u32", "value": "stateWindow" }, + { + "name": "stateSlotStride", + "type": "u32", + "value": "dim(shapes.query, 0) * attrs.kv_num_heads * headDimK * headDimV" + } + ] + } + } + ], + "gatedOnlyZero": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 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": "dim(shapes.query, 0) * attrs.kv_num_heads * headDimK * headDimV" + } + ] + } + } + ], + "gatedOnlyState": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 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": "dim(shapes.query, 0) * attrs.kv_num_heads * headDimK * headDimV" + } + ] + } + } + ], + "deltaZero": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 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": "dim(shapes.query, 0) * attrs.kv_num_heads * headDimK * headDimV" + } + ] + } + } + ], + "deltaState": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 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": "dim(shapes.query, 0) * attrs.kv_num_heads * headDimK * headDimV" + } + ] + } + } + ], + "chunkPrep_gated": [ + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkPrep_gated_delta": [ + { + "name": "decay", + "arg": "decayT", + "semantic": "decay", + "buffer": { "type": "read-only-storage" }, + "elementType": "$decayScalar" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkUt_delta": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { "name": "wk", "semantic": "wk", "buffer": { "type": "storage" }, "elementType": "f32" }, + { "name": "uvec", "semantic": "uvec", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkUt_gated_delta": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "beta", + "arg": "betaT", + "semantic": "beta", + "buffer": { "type": "read-only-storage" }, + "elementType": "$betaScalar" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "wk", "semantic": "wk", "buffer": { "type": "storage" }, "elementType": "f32" }, + { "name": "uvec", "semantic": "uvec", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkScan_linear_zero": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { "name": "states", "semantic": "states", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkScan_linear_state": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { "name": "states", "semantic": "states", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkOut_linear": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "states", "semantic": "states", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkScan_gated_zero": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { "name": "states", "semantic": "states", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkScan_gated_state": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { "name": "states", "semantic": "states", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkOut_gated": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "states", "semantic": "states", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { + "name": "value", + "arg": "valueT", + "semantic": "value", + "buffer": { "type": "read-only-storage" }, + "elementType": "$valueScalar" + }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkScan_delta_zero": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "wk", "semantic": "wk", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "uvec", "semantic": "uvec", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "states", "semantic": "states", "buffer": { "type": "storage" }, "elementType": "f32" }, + { "name": "deltas", "semantic": "deltas", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkScan_delta_state": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "wk", "semantic": "wk", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "uvec", "semantic": "uvec", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { "name": "states", "semantic": "states", "buffer": { "type": "storage" }, "elementType": "f32" }, + { "name": "deltas", "semantic": "deltas", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkOut_delta": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "states", "semantic": "states", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "deltas", "semantic": "deltas", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkScan_gated_delta_zero": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "wk", "semantic": "wk", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "uvec", "semantic": "uvec", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "states", "semantic": "states", "buffer": { "type": "storage" }, "elementType": "f32" }, + { "name": "deltas", "semantic": "deltas", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkScan_gated_delta_state": [ + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "wk", "semantic": "wk", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "uvec", "semantic": "uvec", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { + "name": "past_state", + "arg": "pastStateT", + "semantic": "past_state", + "buffer": { "type": "read-only-storage" }, + "elementType": "$stateScalar" + }, + { "name": "states", "semantic": "states", "buffer": { "type": "storage" }, "elementType": "f32" }, + { "name": "deltas", "semantic": "deltas", "buffer": { "type": "storage" }, "elementType": "f32" }, + { + "name": "present_state", + "arg": "presentStateT", + "semantic": "present_state", + "buffer": { "type": "storage" }, + "elementType": "$stateScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ], + "chunkOut_gated_delta": [ + { + "name": "query", + "arg": "queryT", + "semantic": "query", + "buffer": { "type": "read-only-storage" }, + "elementType": "$queryElem" + }, + { + "name": "key", + "arg": "keyT", + "semantic": "key", + "buffer": { "type": "read-only-storage" }, + "elementType": "$keyElem" + }, + { "name": "gexp", "semantic": "gexp", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "states", "semantic": "states", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { "name": "deltas", "semantic": "deltas", "buffer": { "type": "read-only-storage" }, "elementType": "f32" }, + { + "name": "output", + "arg": "outputT", + "semantic": "output", + "buffer": { "type": "storage" }, + "elementType": "$outputScalar" + }, + { + "name": "params", + "semantic": "kernel.params", + "buffer": { "type": "uniform" }, + "struct": { + "name": "Params", + "fields": [ + { "name": "batchSize", "type": "u32", "value": "dim(shapes.query, 0)" }, + { "name": "seqLength", "type": "u32", "value": "dim(shapes.query, 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.query, 2)" }, + { "name": "kPackedDim", "type": "u32", "value": "dim(shapes.key, 2)" }, + { "name": "vPackedDim", "type": "u32", "value": "dim(shapes.value, 2)" }, + { "name": "decayPackedDim", "type": "u32", "value": "dim(shapes.decay, 2) if present.decayT else 0" }, + { "name": "betaPackedDim", "type": "u32", "value": "dim(shapes.beta, 2) if present.betaT else 0" }, + { "name": "scale", "type": "f32", "value": "attrs.scale if attrs.scale else 0" } + ] + } + } + ] + }, + "variants": [ + { + "id": "linear_zero_chunked", + "priority": 30, + "when": ["linearZeroContract", "chunkedShapeOk"], + "constants": { + "updateRule": "\"linear\"", + "hasPastState": "false", + "usesDecay": "false", + "usesBeta": "false", + "decayPerElement": "false", + "headDimK": "headDimK", + "headDimV": "headDimV", + "chunkSize": "chunkSize", + "chunkTileK": "chunkTileK", + "chunkTileV": "chunkTileV", + "kvPerKeyHead": "kvPerKeyHead", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "intermediates": [ + { + "id": "states", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV]" + } + ], + "passes": [ + { + "id": "scan", + "name": "LinearAttention.ChunkScan", + "shader": "chunk-scan.wgsl.jinja", + "bindings": "chunkScan_linear_zero", + "constants": { "workgroupSize": "chunkScanWorkgroup", "chunkScanTokens": "chunkScanTokens" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (headDimV / chunkTileV)" } + }, + { + "id": "out", + "name": "LinearAttention.ChunkOutput", + "shader": "chunk-out.wgsl.jinja", + "bindings": "chunkOut_linear", + "constants": { "workgroupSize": "chunkOutWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * outHeads * chunkNumChunks" } + } + ] + }, + { + "id": "linear_state_chunked", + "priority": 30, + "when": ["linearStateContract", "chunkedShapeOk"], + "constants": { + "updateRule": "\"linear\"", + "hasPastState": "true", + "usesDecay": "false", + "usesBeta": "false", + "decayPerElement": "false", + "headDimK": "headDimK", + "headDimV": "headDimV", + "chunkSize": "chunkSize", + "chunkTileK": "chunkTileK", + "chunkTileV": "chunkTileV", + "kvPerKeyHead": "kvPerKeyHead", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "intermediates": [ + { + "id": "states", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV]" + } + ], + "passes": [ + { + "id": "scan", + "name": "LinearAttention.ChunkScan", + "shader": "chunk-scan.wgsl.jinja", + "bindings": "chunkScan_linear_state", + "constants": { "workgroupSize": "chunkScanWorkgroup", "chunkScanTokens": "chunkScanTokens" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (headDimV / chunkTileV)" } + }, + { + "id": "out", + "name": "LinearAttention.ChunkOutput", + "shader": "chunk-out.wgsl.jinja", + "bindings": "chunkOut_linear", + "constants": { "workgroupSize": "chunkOutWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * outHeads * chunkNumChunks" } + } + ] + }, + { + "id": "gated_zero_chunked", + "priority": 30, + "when": ["gatedOnlyZeroContract", "chunkedShapeOk"], + "constants": { + "updateRule": "\"gated\"", + "hasPastState": "false", + "usesDecay": "true", + "usesBeta": "false", + "decayPerElement": "decayPerElement", + "headDimK": "headDimK", + "headDimV": "headDimV", + "chunkSize": "chunkSize", + "chunkTileK": "chunkTileK", + "chunkTileV": "chunkTileV", + "kvPerKeyHead": "kvPerKeyHead", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "intermediates": [ + { + "id": "gexp", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * dim(shapes.query, 1) * dim(shapes.decay, 2)]" + }, + { + "id": "states", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV]" + } + ], + "passes": [ + { + "id": "prep", + "name": "LinearAttention.ChunkPrep", + "shader": "chunk-prep.wgsl.jinja", + "bindings": "chunkPrep_gated", + "constants": { "workgroupSize": "chunkFlatWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * chunkNumChunks" } + }, + { + "id": "scan", + "name": "LinearAttention.ChunkScan", + "shader": "chunk-scan.wgsl.jinja", + "bindings": "chunkScan_gated_zero", + "constants": { "workgroupSize": "chunkScanWorkgroup", "chunkScanTokens": "chunkScanTokens" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (headDimV / chunkTileV)" } + }, + { + "id": "out", + "name": "LinearAttention.ChunkOutput", + "shader": "chunk-out.wgsl.jinja", + "bindings": "chunkOut_gated", + "constants": { "workgroupSize": "chunkOutWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * outHeads * chunkNumChunks" } + } + ] + }, + { + "id": "gated_state_chunked", + "priority": 30, + "when": ["gatedOnlyStateContract", "chunkedShapeOk"], + "constants": { + "updateRule": "\"gated\"", + "hasPastState": "true", + "usesDecay": "true", + "usesBeta": "false", + "decayPerElement": "decayPerElement", + "headDimK": "headDimK", + "headDimV": "headDimV", + "chunkSize": "chunkSize", + "chunkTileK": "chunkTileK", + "chunkTileV": "chunkTileV", + "kvPerKeyHead": "kvPerKeyHead", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "intermediates": [ + { + "id": "gexp", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * dim(shapes.query, 1) * dim(shapes.decay, 2)]" + }, + { + "id": "states", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV]" + } + ], + "passes": [ + { + "id": "prep", + "name": "LinearAttention.ChunkPrep", + "shader": "chunk-prep.wgsl.jinja", + "bindings": "chunkPrep_gated", + "constants": { "workgroupSize": "chunkFlatWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * chunkNumChunks" } + }, + { + "id": "scan", + "name": "LinearAttention.ChunkScan", + "shader": "chunk-scan.wgsl.jinja", + "bindings": "chunkScan_gated_state", + "constants": { "workgroupSize": "chunkScanWorkgroup", "chunkScanTokens": "chunkScanTokens" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (headDimV / chunkTileV)" } + }, + { + "id": "out", + "name": "LinearAttention.ChunkOutput", + "shader": "chunk-out.wgsl.jinja", + "bindings": "chunkOut_gated", + "constants": { "workgroupSize": "chunkOutWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * outHeads * chunkNumChunks" } + } + ] + }, + { + "id": "delta_zero_chunked", + "priority": 30, + "when": ["deltaZeroContract", "chunkedShapeOk"], + "constants": { + "updateRule": "\"delta\"", + "hasPastState": "false", + "usesDecay": "false", + "usesBeta": "true", + "decayPerElement": "false", + "headDimK": "headDimK", + "headDimV": "headDimV", + "chunkSize": "chunkSize", + "chunkTileK": "chunkTileK", + "chunkTileV": "chunkTileV", + "kvPerKeyHead": "kvPerKeyHead", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "intermediates": [ + { + "id": "wk", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimK]" + }, + { + "id": "uvec", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV]" + }, + { + "id": "states", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV]" + }, + { + "id": "deltas", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV]" + } + ], + "passes": [ + { + "id": "ut", + "name": "LinearAttention.ChunkTransform", + "shader": "chunk-ut.wgsl.jinja", + "bindings": "chunkUt_delta", + "constants": { "workgroupSize": "chunkFlatWorkgroup", "chunkTileK": "chunkUtTileK" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks" } + }, + { + "id": "scan", + "name": "LinearAttention.ChunkScan", + "shader": "chunk-scan.wgsl.jinja", + "bindings": "chunkScan_delta_zero", + "constants": { "workgroupSize": "chunkScanWorkgroup", "chunkScanTokens": "chunkScanTokens" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (headDimV / chunkTileV)" } + }, + { + "id": "out", + "name": "LinearAttention.ChunkOutput", + "shader": "chunk-out.wgsl.jinja", + "bindings": "chunkOut_delta", + "constants": { "workgroupSize": "chunkOutWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * outHeads * chunkNumChunks" } + } + ] + }, + { + "id": "delta_state_chunked", + "priority": 30, + "when": ["deltaStateContract", "chunkedShapeOk"], + "constants": { + "updateRule": "\"delta\"", + "hasPastState": "true", + "usesDecay": "false", + "usesBeta": "true", + "decayPerElement": "false", + "headDimK": "headDimK", + "headDimV": "headDimV", + "chunkSize": "chunkSize", + "chunkTileK": "chunkTileK", + "chunkTileV": "chunkTileV", + "kvPerKeyHead": "kvPerKeyHead", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "intermediates": [ + { + "id": "wk", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimK]" + }, + { + "id": "uvec", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV]" + }, + { + "id": "states", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV]" + }, + { + "id": "deltas", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV]" + } + ], + "passes": [ + { + "id": "ut", + "name": "LinearAttention.ChunkTransform", + "shader": "chunk-ut.wgsl.jinja", + "bindings": "chunkUt_delta", + "constants": { "workgroupSize": "chunkFlatWorkgroup", "chunkTileK": "chunkUtTileK" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks" } + }, + { + "id": "scan", + "name": "LinearAttention.ChunkScan", + "shader": "chunk-scan.wgsl.jinja", + "bindings": "chunkScan_delta_state", + "constants": { "workgroupSize": "chunkScanWorkgroup", "chunkScanTokens": "chunkScanTokens" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (headDimV / chunkTileV)" } + }, + { + "id": "out", + "name": "LinearAttention.ChunkOutput", + "shader": "chunk-out.wgsl.jinja", + "bindings": "chunkOut_delta", + "constants": { "workgroupSize": "chunkOutWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * outHeads * chunkNumChunks" } + } + ] + }, + { + "id": "gated_delta_zero_chunked", + "priority": 30, + "when": ["gatedZeroContract", "chunkedShapeOk"], + "constants": { + "updateRule": "\"gated_delta\"", + "hasPastState": "false", + "usesDecay": "true", + "usesBeta": "true", + "decayPerElement": "decayPerElement", + "headDimK": "headDimK", + "headDimV": "headDimV", + "chunkSize": "chunkSize", + "chunkTileK": "chunkTileK", + "chunkTileV": "chunkTileV", + "kvPerKeyHead": "kvPerKeyHead", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "intermediates": [ + { + "id": "gexp", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * dim(shapes.query, 1) * dim(shapes.decay, 2)]" + }, + { + "id": "wk", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimK]" + }, + { + "id": "uvec", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV]" + }, + { + "id": "states", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV]" + }, + { + "id": "deltas", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV]" + } + ], + "passes": [ + { + "id": "prep", + "name": "LinearAttention.ChunkPrep", + "shader": "chunk-prep.wgsl.jinja", + "bindings": "chunkPrep_gated_delta", + "constants": { "workgroupSize": "chunkFlatWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * chunkNumChunks" } + }, + { + "id": "ut", + "name": "LinearAttention.ChunkTransform", + "shader": "chunk-ut.wgsl.jinja", + "bindings": "chunkUt_gated_delta", + "constants": { "workgroupSize": "chunkFlatWorkgroup", "chunkTileK": "chunkUtTileK" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks" } + }, + { + "id": "scan", + "name": "LinearAttention.ChunkScan", + "shader": "chunk-scan.wgsl.jinja", + "bindings": "chunkScan_gated_delta_zero", + "constants": { "workgroupSize": "chunkScanWorkgroup", "chunkScanTokens": "chunkScanTokens" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (headDimV / chunkTileV)" } + }, + { + "id": "out", + "name": "LinearAttention.ChunkOutput", + "shader": "chunk-out.wgsl.jinja", + "bindings": "chunkOut_gated_delta", + "constants": { "workgroupSize": "chunkOutWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * outHeads * chunkNumChunks" } + } + ] + }, + { + "id": "gated_delta_state_chunked", + "priority": 30, + "when": ["gatedStateContract", "chunkedShapeOk"], + "constants": { + "updateRule": "\"gated_delta\"", + "hasPastState": "true", + "usesDecay": "true", + "usesBeta": "true", + "decayPerElement": "decayPerElement", + "headDimK": "headDimK", + "headDimV": "headDimV", + "chunkSize": "chunkSize", + "chunkTileK": "chunkTileK", + "chunkTileV": "chunkTileV", + "kvPerKeyHead": "kvPerKeyHead", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "intermediates": [ + { + "id": "gexp", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * dim(shapes.query, 1) * dim(shapes.decay, 2)]" + }, + { + "id": "wk", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimK]" + }, + { + "id": "uvec", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV]" + }, + { + "id": "states", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * headDimK * headDimV]" + }, + { + "id": "deltas", + "dtype": "float32", + "shape": "[dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks * chunkSize * headDimV]" + } + ], + "passes": [ + { + "id": "prep", + "name": "LinearAttention.ChunkPrep", + "shader": "chunk-prep.wgsl.jinja", + "bindings": "chunkPrep_gated_delta", + "constants": { "workgroupSize": "chunkFlatWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * chunkNumChunks" } + }, + { + "id": "ut", + "name": "LinearAttention.ChunkTransform", + "shader": "chunk-ut.wgsl.jinja", + "bindings": "chunkUt_gated_delta", + "constants": { "workgroupSize": "chunkFlatWorkgroup", "chunkTileK": "chunkUtTileK" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * chunkNumChunks" } + }, + { + "id": "scan", + "name": "LinearAttention.ChunkScan", + "shader": "chunk-scan.wgsl.jinja", + "bindings": "chunkScan_gated_delta_state", + "constants": { "workgroupSize": "chunkScanWorkgroup", "chunkScanTokens": "chunkScanTokens" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (headDimV / chunkTileV)" } + }, + { + "id": "out", + "name": "LinearAttention.ChunkOutput", + "shader": "chunk-out.wgsl.jinja", + "bindings": "chunkOut_gated_delta", + "constants": { "workgroupSize": "chunkOutWorkgroup" }, + "dispatch": { "workgroups": "dim(shapes.query, 0) * outHeads * chunkNumChunks" } + } + ] + }, + { + "id": "linear_zero_serial_small_dk", + "priority": 20, + "when": ["linearZeroContract", "serialHeadDimFits"], + "constants": { + "updateRule": "\"linear\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "headDimK": "headDimK", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention.SerialSmallDk", + "shader": "linear-attention.serial.wgsl.jinja", + "bindings": "linearZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (dim(shapes.value, 2) / attrs.kv_num_heads)" + } + } + ] + }, + { + "id": "linear_state_serial_small_dk", + "priority": 20, + "when": ["linearStateContract", "serialHeadDimFits"], + "constants": { + "updateRule": "\"linear\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "headDimK": "headDimK", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention.SerialSmallDk", + "shader": "linear-attention.serial.wgsl.jinja", + "bindings": "linearState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (dim(shapes.value, 2) / attrs.kv_num_heads)" + } + } + ] + }, + { + "id": "gated_delta_zero_serial_small_dk", + "priority": 20, + "when": ["gatedZeroContract", "serialHeadDimFits"], + "constants": { + "updateRule": "\"gated_delta\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "headDimK": "headDimK", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention.SerialSmallDk", + "shader": "linear-attention.serial.wgsl.jinja", + "bindings": "baseZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (dim(shapes.value, 2) / attrs.kv_num_heads)" + } + } + ] + }, + { + "id": "gated_delta_state_serial_small_dk", + "priority": 20, + "when": ["gatedStateContract", "serialHeadDimFits", "tensorDtypes.present_state == tensorDtypes.past_state"], + "constants": { + "updateRule": "\"gated_delta\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "headDimK": "headDimK", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention.SerialSmallDk", + "shader": "linear-attention.serial.wgsl.jinja", + "bindings": "baseState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * (dim(shapes.value, 2) / attrs.kv_num_heads)" + } + } + ] + }, + { + "id": "linear_zero_scalar", + "priority": 0, + "when": ["linearZeroContract", "scalarHeadDimFits"], + "constants": { + "updateRule": "\"linear\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16", + "workgroupSize": "min(256, pow2ceil(dim(shapes.query, 2) / attrs.q_num_heads))", + "tileV": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.scalar.wgsl.jinja", + "bindings": "linearZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV)" + } + } + ] + }, + { + "id": "linear_state_scalar", + "priority": 0, + "when": ["linearStateContract", "scalarHeadDimFits"], + "constants": { + "updateRule": "\"linear\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16", + "workgroupSize": "min(256, pow2ceil(dim(shapes.query, 2) / attrs.q_num_heads))", + "tileV": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.scalar.wgsl.jinja", + "bindings": "linearState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV)" + } + } + ] + }, + { + "id": "gated_delta_zero_scalar", + "priority": 0, + "when": ["gatedZeroContract", "scalarHeadDimFits"], + "constants": { + "updateRule": "\"gated_delta\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16", + "workgroupSize": "min(256, pow2ceil(dim(shapes.query, 2) / attrs.q_num_heads))", + "tileV": "max(1, min(tunables.gatedTileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.scalar.wgsl.jinja", + "bindings": "baseZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV)" + } + } + ] + }, + { + "id": "gated_delta_state_scalar", + "priority": 0, + "when": ["gatedStateContract", "scalarHeadDimFits"], + "constants": { + "updateRule": "\"gated_delta\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16", + "workgroupSize": "min(256, pow2ceil(dim(shapes.query, 2) / attrs.q_num_heads))", + "tileV": "max(1, min(tunables.gatedTileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.scalar.wgsl.jinja", + "bindings": "baseState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV)" + } + } + ] + }, + { + "id": "gated_zero_scalar", + "priority": 0, + "when": ["gatedOnlyZeroContract", "scalarHeadDimFits"], + "constants": { + "updateRule": "\"gated\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16", + "workgroupSize": "min(256, pow2ceil(dim(shapes.query, 2) / attrs.q_num_heads))", + "tileV": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.scalar.wgsl.jinja", + "bindings": "gatedOnlyZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV)" + } + } + ] + }, + { + "id": "gated_state_scalar", + "priority": 0, + "when": ["gatedOnlyStateContract", "scalarHeadDimFits"], + "constants": { + "updateRule": "\"gated\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16", + "workgroupSize": "min(256, pow2ceil(dim(shapes.query, 2) / attrs.q_num_heads))", + "tileV": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.scalar.wgsl.jinja", + "bindings": "gatedOnlyState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV)" + } + } + ] + }, + { + "id": "delta_zero_scalar", + "priority": 0, + "when": ["deltaZeroContract", "scalarHeadDimFits"], + "constants": { + "updateRule": "\"delta\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16", + "workgroupSize": "min(256, pow2ceil(dim(shapes.query, 2) / attrs.q_num_heads))", + "tileV": "max(1, min(tunables.gatedTileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.scalar.wgsl.jinja", + "bindings": "deltaZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV)" + } + } + ] + }, + { + "id": "delta_state_scalar", + "priority": 0, + "when": ["deltaStateContract", "scalarHeadDimFits"], + "constants": { + "updateRule": "\"delta\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "useSubgroups": "device.features.has(\"subgroups\") and headDimK > 16", + "workgroupSize": "min(256, pow2ceil(dim(shapes.query, 2) / attrs.q_num_heads))", + "tileV": "max(1, min(tunables.gatedTileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "keyDtype": "tensorDtypes.query", + "valueDtype": "tensorDtypes.query", + "decayDtype": "tensorDtypes.query", + "betaDtype": "tensorDtypes.query", + "outputDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "keyElem": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.scalar.wgsl.jinja", + "bindings": "deltaState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV)" + } + } + ] + }, + { + "id": "linear_zero_vec4", + "priority": 10, + "when": ["linearZeroContract", "headDimK % 4 == 0", "vec4SharedFitsPlain", "vec4DtypeOk", "vec4HeadDimFits", "vec4BindingsNonEmpty"], + "constants": { + "updateRule": "\"linear\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "useSubgroups": "vec4UseSubgroups", + "tileV": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "keyElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows", + "vec4Lanes": "vec4Lanes", + "dvGroups": "vec4DvGroups" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.vec4.wgsl.jinja", + "bindings": "linearZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV), constants.dvGroups)" + } + } + ] + }, + { + "id": "linear_state_vec4", + "priority": 10, + "when": ["linearStateContract", "headDimK % 4 == 0", "vec4SharedFitsPlain", "vec4DtypeOk", "vec4HeadDimFits", "vec4BindingsNonEmpty"], + "constants": { + "updateRule": "\"linear\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "useSubgroups": "vec4UseSubgroups", + "tileV": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "keyElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows", + "vec4Lanes": "vec4Lanes", + "dvGroups": "vec4DvGroups" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.vec4.wgsl.jinja", + "bindings": "linearState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV), constants.dvGroups)" + } + } + ] + }, + { + "id": "gated_delta_zero_vec4", + "priority": 10, + "when": ["gatedZeroContract", "headDimK % 4 == 0", "vec4SharedFitsGatedBeta", "gatedVec4DtypeOk", "vec4HeadDimFits", "vec4BindingsNonEmpty"], + "constants": { + "updateRule": "\"gated_delta\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "useSubgroups": "vec4UseSubgroups", + "tileV": "gatedVec4TileV", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "keyElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows", + "vec4Lanes": "vec4Lanes", + "dvGroups": "vec4DvGroups" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.vec4.wgsl.jinja", + "bindings": "baseZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV), constants.dvGroups)" + } + } + ] + }, + { + "id": "gated_delta_state_vec4", + "priority": 10, + "when": ["gatedStateContract", "headDimK % 4 == 0", "tensorDtypes.present_state == tensorDtypes.past_state", "vec4SharedFitsGatedBeta", "gatedVec4DtypeOk", "vec4HeadDimFits", "vec4BindingsNonEmpty"], + "constants": { + "updateRule": "\"gated_delta\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "useSubgroups": "vec4UseSubgroups", + "tileV": "gatedVec4TileV", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "keyElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows", + "vec4Lanes": "vec4Lanes", + "dvGroups": "vec4DvGroups" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.vec4.wgsl.jinja", + "bindings": "baseState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV), constants.dvGroups)" + } + } + ] + }, + { + "id": "gated_zero_vec4", + "priority": 10, + "when": ["gatedOnlyZeroContract", "headDimK % 4 == 0", "vec4SharedFitsPlain", "decayVec4DtypeOk", "vec4HeadDimFits", "vec4BindingsNonEmpty"], + "constants": { + "updateRule": "\"gated\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "useSubgroups": "vec4UseSubgroups", + "tileV": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "keyElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows", + "vec4Lanes": "vec4Lanes", + "dvGroups": "vec4DvGroups" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.vec4.wgsl.jinja", + "bindings": "gatedOnlyZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV), constants.dvGroups)" + } + } + ] + }, + { + "id": "gated_state_vec4", + "priority": 10, + "when": ["gatedOnlyStateContract", "headDimK % 4 == 0", "tensorDtypes.present_state == tensorDtypes.past_state", "vec4SharedFitsPlain", "decayVec4DtypeOk", "vec4HeadDimFits", "vec4BindingsNonEmpty"], + "constants": { + "updateRule": "\"gated\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "useSubgroups": "vec4UseSubgroups", + "tileV": "max(1, min(tunables.tileV, dim(shapes.value, 2) / attrs.kv_num_heads))", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "keyElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "decayScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows", + "vec4Lanes": "vec4Lanes", + "dvGroups": "vec4DvGroups" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.vec4.wgsl.jinja", + "bindings": "gatedOnlyState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV), constants.dvGroups)" + } + } + ] + }, + { + "id": "delta_zero_vec4", + "priority": 10, + "when": ["deltaZeroContract", "headDimK % 4 == 0", "vec4SharedFitsGatedBeta", "betaVec4DtypeOk", "vec4HeadDimFits", "vec4BindingsNonEmpty"], + "constants": { + "updateRule": "\"delta\"", + "hasPastState": false, + "hasStateWindow": "windowed", + "useSubgroups": "vec4UseSubgroups", + "tileV": "gatedVec4TileV", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.present_state", + "queryElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "keyElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.present_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows", + "vec4Lanes": "vec4Lanes", + "dvGroups": "vec4DvGroups" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.vec4.wgsl.jinja", + "bindings": "deltaZero", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV), constants.dvGroups)" + } + } + ] + }, + { + "id": "delta_state_vec4", + "priority": 10, + "when": ["deltaStateContract", "headDimK % 4 == 0", "vec4SharedFitsGatedBeta", "betaVec4DtypeOk", "vec4HeadDimFits", "vec4BindingsNonEmpty"], + "constants": { + "updateRule": "\"delta\"", + "hasPastState": true, + "hasStateWindow": "windowed", + "useSubgroups": "vec4UseSubgroups", + "tileV": "gatedVec4TileV", + "queryDtype": "tensorDtypes.query", + "stateDtype": "tensorDtypes.past_state", + "queryElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "keyElem": "\"vec4\" if tensorDtypes.query == \"float16\" else \"vec4\"", + "valueScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "betaScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "outputScalar": "\"f16\" if tensorDtypes.query == \"float16\" else \"f32\"", + "stateScalar": "\"f16\" if tensorDtypes.past_state == \"float16\" else \"f32\"", + "chunkOutRows": "chunkOutRows", + "vec4Lanes": "vec4Lanes", + "dvGroups": "vec4DvGroups" + }, + "passes": [ + { + "id": "main", + "name": "LinearAttention", + "shader": "linear-attention.vec4.wgsl.jinja", + "bindings": "deltaState", + "dispatch": { + "workgroups": "dim(shapes.query, 0) * attrs.kv_num_heads * ceilDiv(ceilDiv(dim(shapes.value, 2) / attrs.kv_num_heads, constants.tileV), constants.dvGroups)" + } + } + ] + } + ] +}