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