Download build/webgpu/manifest.json from webgpu-kernels/com.microsoft.QMoE: direct link, hf CLI and curl.
- Browser
- Download file 25.5 kB
-
https://huggingface.co/kernels/webgpu-kernels/com.microsoft.QMoE/resolve/v1/build/webgpu/manifest.json
- Command line
-
hf download hf://webgpu-kernels/com.microsoft.QMoE@v1/build/webgpu/manifest.json
-
curl -L -o manifest.json https://huggingface.co/kernels/webgpu-kernels/com.microsoft.QMoE/resolve/v1/build/webgpu/manifest.json
25.5 kB
| { | |
| "domain": "com.microsoft", | |
| "name": "QMoE", | |
| "sinceVersion": 1, | |
| "inputs": { | |
| "inputT": { "onnx": "input", "dtype": "T" }, | |
| "routerT": { "onnx": "router_probs", "dtype": "T", "rank": 2 }, | |
| "fc1T": { "onnx": "fc1_experts_weights", "dtype": "T1", "rank": 3 }, | |
| "fc1ScalesT": { "onnx": "fc1_scales", "dtype": "T2" }, | |
| "fc2T": { "onnx": "fc2_experts_weights", "dtype": "T1", "rank": 3 }, | |
| "fc2ScalesT": { "onnx": "fc2_scales", "dtype": "T2" } | |
| }, | |
| "outputs": { "outputT": { "onnx": "output", "dtype": "T", "shape": "shapes.inputT" } }, | |
| "attributes": { | |
| "activation_alpha": { "default": 1 }, | |
| "activation_beta": { "default": 0 }, | |
| "activation_type": { "default": "relu" }, | |
| "expert_weight_bits": { "default": 4 }, | |
| "k": { "default": 1 }, | |
| "normalize_routing_weights": { "default": 0 }, | |
| "quant_type": { "default": "int" }, | |
| "swiglu_fusion": { "default": 0 }, | |
| "use_sparse_mixer": { "default": 0 }, | |
| "weights_prepacked": { "default": -1 }, | |
| "block_size": {}, | |
| "swiglu_limit": {} | |
| }, | |
| "attributeConstraints": { | |
| "activation_type": { "values": ["relu", "swiglu"] }, | |
| "expert_weight_bits": { "values": [4, 8] }, | |
| "normalize_routing_weights": { "values": [0, 1] }, | |
| "quant_type": { "values": ["int"] }, | |
| "swiglu_fusion": { "values": [0, 1] }, | |
| "use_sparse_mixer": { "values": [0] }, | |
| "weights_prepacked": { "values": [-1, 0] } | |
| }, | |
| "typeConstraints": { "T": ["float32"], "T1": ["uint8"], "T2": ["float32"] }, | |
| "tunables": { | |
| "workgroupSize": { "default": 64 }, | |
| "decodeLanes": { "default": 32 }, | |
| "decodeBlockTarget": { "default": 1024 }, | |
| "decodeMinLaneTrips": { "default": 4 }, | |
| "groupThreads": { "default": 8 }, | |
| "groupRegM": { "default": 4 }, | |
| "groupRegN": { "default": 4 }, | |
| "groupTileK": { "default": 16 }, | |
| "groupRouteWorkgroup": { "default": 256 } | |
| }, | |
| "derive": { | |
| "deviceWorkgroupCap": "min(device.limits.maxComputeInvocationsPerWorkgroup, device.limits.maxComputeWorkgroupSizeX)", | |
| "wave32Adapter": "has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize == 32 and device.adapterInfo.subgroupMaxSize == 32", | |
| "canPinSubgroupSize32": "device.features.has(\"subgroups\") and device.features.has(\"subgroup-size-control\") and has(device.adapterInfo, \"subgroupMinSize\") and has(device.adapterInfo, \"subgroupMaxSize\") and device.adapterInfo.subgroupMinSize <= 32 and device.adapterInfo.subgroupMaxSize >= 32", | |
| "pinSubgroupSize32": "canPinSubgroupSize32 and not wave32Adapter", | |
| "wave32Effective": "wave32Adapter or pinSubgroupSize32", | |
| "groupTileK": "tunables.groupTileK", | |
| "groupThreads": "tunables.groupThreads", | |
| "regM": "tunables.groupRegM", | |
| "regN": "tunables.groupRegN", | |
| "groupRouteWorkgroup": "tunables.groupRouteWorkgroup", | |
| "groupTileM": "tunables.groupThreads * tunables.groupRegM", | |
| "groupTileN": "tunables.groupThreads * tunables.groupRegN", | |
| "groupTileKVec": "ceilDiv(tunables.groupTileK, 4)", | |
| "groupThreadCount": "tunables.groupThreads * tunables.groupThreads", | |
| "groupSharedBytes": "(groupTileM * tunables.groupTileK + 2 * groupTileN * tunables.groupTileK + groupTileM) * 4", | |
| "topK": "attrs.k", | |
| "activationType": "attrs.activation_type", | |
| "workgroupSizeOk": "tunables.workgroupSize >= 1 and tunables.workgroupSize <= deviceWorkgroupCap", | |
| "decodeLanesOk": "tunables.decodeLanes >= 1 and tunables.decodeLanes <= tunables.workgroupSize and tunables.workgroupSize % tunables.decodeLanes == 0", | |
| "decodeRows": "max(1, tunables.workgroupSize / max(1, tunables.decodeLanes))", | |
| "decodeDeviceOk": "workgroupSizeOk and decodeLanesOk and tunables.decodeLanes <= device.limits.maxComputeWorkgroupSizeX and decodeRows <= device.limits.maxComputeWorkgroupSizeY and tunables.workgroupSize * 8 <= device.limits.maxComputeWorkgroupStorageSize", | |
| "fc1Rows": "dim(shapes.fc1T, 1)", | |
| "hasSwigluLimit": "has(attrs, \"swiglu_limit\")", | |
| "swigluLimit": "attrs.swiglu_limit if has(attrs, \"swiglu_limit\") else 0", | |
| "workgroupSize": "tunables.workgroupSize", | |
| "hiddenSize": "dim(shapes.inputT, ranks.inputT - 1)", | |
| "numTokens": "numel(shapes.inputT) / max(1, hiddenSize)", | |
| "weightBits": "attrs.expert_weight_bits", | |
| "quantBlockSize": "attrs.block_size if has(attrs, \"block_size\") else 0", | |
| "packSize": "2 if weightBits == 4 else 1", | |
| "quantMidpoint": "8 if weightBits == 4 else 128", | |
| "fusionSize": "2 if attrs.activation_type == \"swiglu\" and attrs.swiglu_fusion == 1 else 1", | |
| "interSize": "dim(shapes.fc1T, 1) / fusionSize", | |
| "fc1PackedCols": "dim(shapes.fc1T, 2)", | |
| "fc2PackedCols": "dim(shapes.fc2T, 2)", | |
| "colWiseScales": "quantBlockSize == 0", | |
| "fc1ScaleBlocks": "1 if colWiseScales else hiddenSize / max(1, quantBlockSize)", | |
| "fc2ScaleBlocks": "1 if colWiseScales else interSize / max(1, quantBlockSize)", | |
| "activationSupported": "(attrs.activation_type == \"relu\" and attrs.swiglu_fusion == 0) or (attrs.activation_type == \"swiglu\" and attrs.swiglu_fusion == 1)", | |
| "routingModeSupported": "attrs.normalize_routing_weights == 0 or attrs.normalize_routing_weights == 1", | |
| "rawWeightLayout": "attrs.weights_prepacked == -1 or attrs.weights_prepacked == 0", | |
| "inputOutputShapeOk": "((ranks.inputT == 2 and ranks.outputT == 2 and dim(shapes.outputT, 0) == dim(shapes.inputT, 0) and dim(shapes.outputT, 1) == dim(shapes.inputT, 1)) or (ranks.inputT == 3 and ranks.outputT == 3 and dim(shapes.outputT, 0) == dim(shapes.inputT, 0) and dim(shapes.outputT, 1) == dim(shapes.inputT, 1) and dim(shapes.outputT, 2) == dim(shapes.inputT, 2))) and hiddenSize > 0", | |
| "quantBlockSizeOk": "colWiseScales or (quantBlockSize >= 16 and pow2ceil(quantBlockSize) == quantBlockSize and hiddenSize % quantBlockSize == 0 and interSize % quantBlockSize == 0)", | |
| "quantScalesOk": "tensorDtypes.fc1ScalesT == \"float32\" and tensorDtypes.fc2ScalesT == \"float32\" and dim(shapes.fc1ScalesT, 0) == dim(shapes.routerT, 1) and dim(shapes.fc2ScalesT, 0) == dim(shapes.routerT, 1) and dim(shapes.fc1ScalesT, 1) == dim(shapes.fc1T, 1) and dim(shapes.fc2ScalesT, 1) == hiddenSize and ((ranks.fc1ScalesT == 2 and ranks.fc2ScalesT == 2) if colWiseScales else (ranks.fc1ScalesT == 3 and ranks.fc2ScalesT == 3 and dim(shapes.fc1ScalesT, 2) == fc1ScaleBlocks and dim(shapes.fc2ScalesT, 2) == fc2ScaleBlocks))", | |
| "quantShapeOk": "inputOutputShapeOk and ranks.routerT == 2 and ranks.fc1T == 3 and ranks.fc2T == 3 and dim(shapes.routerT, 0) == numTokens and dim(shapes.fc1T, 0) == dim(shapes.routerT, 1) and dim(shapes.fc2T, 0) == dim(shapes.routerT, 1) and dim(shapes.fc1T, 1) % fusionSize == 0 and dim(shapes.fc2T, 1) == hiddenSize and dim(shapes.fc1T, 2) * packSize == hiddenSize and dim(shapes.fc2T, 2) * packSize == interSize and quantBlockSizeOk and quantScalesOk", | |
| "quantContract": "activationSupported and routingModeSupported and rawWeightLayout and quantShapeOk and topK >= 1 and topK <= dim(shapes.routerT, 1)", | |
| "hiddenChunkFits": "topK * interSize * 4 <= device.limits.maxStorageBufferBindingSize and topK * interSize * 4 <= device.limits.maxBufferSize", | |
| "hiddenChunkTokens": "numTokens if interSize == 0 else min(numTokens, max(1, floor(min(device.limits.maxStorageBufferBindingSize, device.limits.maxBufferSize) / (topK * interSize * 4))))", | |
| "hiddenChunkCount": "max(1, ceilDiv(numTokens, max(1, hiddenChunkTokens)))", | |
| "routeScratchBytes": "numTokens * topK * 4", | |
| "routedScratchFits": "routeScratchBytes <= device.limits.maxStorageBufferBindingSize and routeScratchBytes <= device.limits.maxBufferSize", | |
| "groupSlots": "hiddenChunkTokens * topK", | |
| "groupMaxTiles": "ceilDiv(groupSlots, max(1, groupTileM)) + dim(shapes.routerT, 1)", | |
| "groupSlotOutBytes": "hiddenChunkTokens * topK * hiddenSize * 4", | |
| "groupedDeviceOk": "groupThreadCount <= device.limits.maxComputeInvocationsPerWorkgroup and tunables.groupThreads <= device.limits.maxComputeWorkgroupSizeX and tunables.groupThreads <= device.limits.maxComputeWorkgroupSizeY and tunables.groupRouteWorkgroup <= deviceWorkgroupCap and groupSharedBytes <= device.limits.maxComputeWorkgroupStorageSize and dim(shapes.routerT, 1) * 8 <= device.limits.maxComputeWorkgroupStorageSize and tunables.groupTileK % 4 == 0", | |
| "groupedDispatchOk": "groupMaxTiles <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(interSize, max(1, groupTileN)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(hiddenSize, max(1, groupTileN)) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "groupSlotOutFits": "groupSlotOutBytes <= device.limits.maxStorageBufferBindingSize and groupSlotOutBytes <= device.limits.maxBufferSize", | |
| "groupedShapeOk": "numTokens * topK * 4 >= groupTileM * dim(shapes.routerT, 1)", | |
| "groupedContract": "quantContract and workgroupSizeOk and hiddenChunkFits and routedScratchFits and interSize > 0 and groupedDeviceOk and groupedDispatchOk and groupSlotOutFits and groupedShapeOk", | |
| "sgmatWorkgroup": "128", | |
| "sgmatSubgroups": "4", | |
| "sgmatRowSubtiles": "2", | |
| "sgmatTileCols": "64", | |
| "groupedSgmatStageElements": "max(groupTileM * 32, sgmatSubgroups * 4 * 64)", | |
| "groupedSgmatSharedBytes": "(groupedSgmatStageElements + sgmatTileCols * 32) * 4", | |
| "groupedSgmatOk": "groupedContract and groupTileM == 32 and interSize % 32 == 0 and groupedSgmatSharedBytes <= device.limits.maxComputeWorkgroupStorageSize and ceilDiv(hiddenSize, sgmatTileCols) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and hiddenSize % 32 == 0", | |
| "decodeDispatchOk": "ceilDiv(interSize, decodeRows) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and ceilDiv(hiddenSize, decodeRows) <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535) and hiddenChunkTokens * topK <= min(device.limits.maxComputeWorkgroupsPerDimension, 65535)", | |
| "decodeLaneDepth": "min(hiddenSize, interSize) / max(1, tunables.decodeLanes)", | |
| "decodeDepthOk": "decodeLaneDepth >= tunables.decodeMinLaneTrips", | |
| "splitFc1Blocks": "ceilDiv(numTokens * topK * interSize, max(1, tunables.workgroupSize))", | |
| "decodeOccupancyOk": "splitFc1Blocks <= tunables.decodeBlockTarget", | |
| "decodeContract": "quantContract and decodeDeviceOk and decodeDispatchOk and hiddenChunkFits and routedScratchFits and interSize > 0 and decodeDepthOk and (decodeOccupancyOk or not groupedContract)", | |
| "hidden": "hiddenSize", | |
| "experts": "dim(shapes.routerT, 1)", | |
| "inter": "interSize" | |
| }, | |
| "bindings": { | |
| "output": { "arg": "outputT", "elementType": "f32" }, | |
| "router_probs": { "arg": "routerT", "elementType": "f32" }, | |
| "route_expert": { "scratch": "routeExpert", "elementType": "u32" }, | |
| "route_mix": { "scratch": "routeMix", "elementType": "f32" }, | |
| "route_expert_u32": { | |
| "scratch": "routeExpert", | |
| "name": "route_expert", | |
| "buffer": "read-only-storage", | |
| "elementType": "u32" | |
| }, | |
| "slot_list": { "scratch": "slotList", "elementType": "u32" }, | |
| "tile_meta": { "scratch": "tileMeta", "elementType": "u32" }, | |
| "params_route": { "name": "params", "struct": [{ "name": "tokenCount", "type": "u32", "value": "numTokens" }] }, | |
| "params": { | |
| "struct": [ | |
| { "name": "tokenOffset", "type": "u32", "value": "repeat.chunk * hiddenChunkTokens" }, | |
| { | |
| "name": "tokenCount", | |
| "type": "u32", | |
| "value": "min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens)" | |
| } | |
| ] | |
| }, | |
| "input": { "arg": "inputT", "elementType": "f32" }, | |
| "slot_list_u32": { "scratch": "slotList", "name": "slot_list", "buffer": "read-only-storage", "elementType": "u32" }, | |
| "tile_meta_u32": { "scratch": "tileMeta", "name": "tile_meta", "buffer": "read-only-storage", "elementType": "u32" }, | |
| "fc1_experts_weights": { "arg": "fc1T", "elementType": "u32" }, | |
| "fc1_scales": { "arg": "fc1ScalesT", "elementType": "f32" }, | |
| "hidden_act": { "scratch": "hiddenAct", "elementType": "f32" }, | |
| "params__uniform": { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "activationAlpha", "type": "f32", "value": "attrs.activation_alpha" }, | |
| { "name": "activationBeta", "type": "f32", "value": "attrs.activation_beta" }, | |
| { "name": "tokenOffset", "type": "u32", "value": "repeat.chunk * hiddenChunkTokens" } | |
| ] | |
| }, | |
| "hidden_act_f32": { | |
| "scratch": "hiddenAct", | |
| "name": "hidden_act", | |
| "buffer": "read-only-storage", | |
| "elementType": "f32" | |
| }, | |
| "fc2_experts_weights": { "arg": "fc2T", "elementType": "u32" }, | |
| "fc2_scales": { "arg": "fc2ScalesT", "elementType": "f32" }, | |
| "slot_out": { "scratch": "slotOut", "elementType": "f32" }, | |
| "slot_out_f32": { "scratch": "slotOut", "name": "slot_out", "buffer": "read-only-storage", "elementType": "f32" }, | |
| "route_mix_f32": { "scratch": "routeMix", "name": "route_mix", "buffer": "read-only-storage", "elementType": "f32" } | |
| }, | |
| "variants": [ | |
| { | |
| "id": "quant_zero_inter", | |
| "priority": 20, | |
| "when": ["quantContract", "workgroupSizeOk", "interSize == 0"], | |
| "passes": [ | |
| { | |
| "id": "output_stage", | |
| "name": "QMoE.OutputStageZeroInter", | |
| "shader": "qmoe-output-zero-inter.wgsl.jinja", | |
| "derive": { "outputElementCount": "numel(shapes.outputT)" }, | |
| "bindings": ["output"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numel(shapes.outputT)), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numel(shapes.outputT)), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "quant_grouped_sgmat_routed", | |
| "priority": 32, | |
| "when": ["groupedSgmatOk", "wave32Effective"], | |
| "requires": { | |
| "features": ["subgroups", "chromium-experimental-subgroup-matrix"], | |
| "subgroupMatrixConfigs": [{ "componentType": "f32", "resultComponentType": "f32", "M": 8, "N": 8, "K": 8 }] | |
| }, | |
| "intermediates": [ | |
| { "id": "routeExpert", "dtype": "uint32", "shape": "[numTokens * topK]" }, | |
| { "id": "routeMix", "dtype": "float32", "shape": "[numTokens * topK]" }, | |
| { "id": "hiddenAct", "dtype": "float32", "shape": "[hiddenChunkTokens * topK * interSize]" }, | |
| { "id": "slotList", "dtype": "uint32", "shape": "[max(1, groupSlots)]" }, | |
| { "id": "tileMeta", "dtype": "uint32", "shape": "[1 + 3 * groupMaxTiles]" }, | |
| { "id": "slotOut", "dtype": "float32", "shape": "[max(1, hiddenChunkTokens * topK * hiddenSize)]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "route_stage", | |
| "name": "QMoE.RouteStage", | |
| "shader": "qmoe-route-stage.wgsl.jinja", | |
| "bindings": ["router_probs", "route_expert", "route_mix", "params_route"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numTokens), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numTokens), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "token_chunks", | |
| "repeat": { "count": "hiddenChunkCount", "index": "chunk" }, | |
| "passes": [ | |
| { | |
| "id": "group_stage", | |
| "name": "QMoE.GroupStage", | |
| "shader": "expert-group-slots.wgsl.jinja", | |
| "bindings": ["route_expert_u32", "slot_list", "tile_meta", "params"], | |
| "dispatch": { "x": 1 } | |
| }, | |
| { | |
| "id": "fc1_activation_stage", | |
| "name": "QMoE.FC1ActivationStageGroupedSgmat", | |
| "shader": "qmoe-fc1-activation-grouped-sgmat.wgsl.jinja", | |
| "bindings": ["input", "slot_list_u32", "tile_meta_u32", "fc1_experts_weights", "fc1_scales", "hidden_act", "params__uniform"], | |
| "dispatch": { "x": "groupMaxTiles", "y": "ceilDiv(interSize * fusionSize, sgmatTileCols)" } | |
| }, | |
| { | |
| "id": "output_stage", | |
| "name": "QMoE.OutputStageGroupedSgmat", | |
| "shader": "qmoe-output-grouped-sgmat.wgsl.jinja", | |
| "bindings": ["hidden_act_f32", "slot_list_u32", "tile_meta_u32", "fc2_experts_weights", "fc2_scales", "slot_out"], | |
| "dispatch": { "x": "groupMaxTiles", "y": "ceilDiv(hiddenSize, sgmatTileCols)" } | |
| }, | |
| { | |
| "id": "mix_stage", | |
| "name": "QMoE.MixStage", | |
| "shader": "expert-slot-mix.wgsl.jinja", | |
| "bindings": ["slot_out_f32", "route_mix_f32", "output", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * hiddenSize), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * hiddenSize), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "quant_grouped_routed", | |
| "priority": 30, | |
| "when": ["groupedContract"], | |
| "intermediates": [ | |
| { "id": "routeExpert", "dtype": "uint32", "shape": "[numTokens * topK]" }, | |
| { "id": "routeMix", "dtype": "float32", "shape": "[numTokens * topK]" }, | |
| { "id": "hiddenAct", "dtype": "float32", "shape": "[hiddenChunkTokens * topK * interSize]" }, | |
| { "id": "slotList", "dtype": "uint32", "shape": "[max(1, groupSlots)]" }, | |
| { "id": "tileMeta", "dtype": "uint32", "shape": "[1 + 3 * groupMaxTiles]" }, | |
| { "id": "slotOut", "dtype": "float32", "shape": "[max(1, hiddenChunkTokens * topK * hiddenSize)]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "route_stage", | |
| "name": "QMoE.RouteStage", | |
| "shader": "qmoe-route-stage.wgsl.jinja", | |
| "bindings": ["router_probs", "route_expert", "route_mix", "params_route"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numTokens), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numTokens), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "token_chunks", | |
| "repeat": { "count": "hiddenChunkCount", "index": "chunk" }, | |
| "passes": [ | |
| { | |
| "id": "group_stage", | |
| "name": "QMoE.GroupStage", | |
| "shader": "expert-group-slots.wgsl.jinja", | |
| "bindings": ["route_expert_u32", "slot_list", "tile_meta", "params"], | |
| "dispatch": { "x": 1 } | |
| }, | |
| { | |
| "id": "fc1_activation_stage", | |
| "name": "QMoE.FC1ActivationStageGrouped", | |
| "shader": "qmoe-fc1-activation-grouped.wgsl.jinja", | |
| "bindings": ["input", "slot_list_u32", "tile_meta_u32", "fc1_experts_weights", "fc1_scales", "hidden_act", "params__uniform"], | |
| "dispatch": { "x": "groupMaxTiles", "y": "ceilDiv(interSize, groupTileN)" } | |
| }, | |
| { | |
| "id": "output_stage", | |
| "name": "QMoE.OutputStageGrouped", | |
| "shader": "qmoe-output-grouped.wgsl.jinja", | |
| "bindings": ["hidden_act_f32", "slot_list_u32", "tile_meta_u32", "fc2_experts_weights", "fc2_scales", "slot_out"], | |
| "dispatch": { "x": "groupMaxTiles", "y": "ceilDiv(hiddenSize, groupTileN)" } | |
| }, | |
| { | |
| "id": "mix_stage", | |
| "name": "QMoE.MixStage", | |
| "shader": "expert-slot-mix.wgsl.jinja", | |
| "bindings": ["slot_out_f32", "route_mix_f32", "output", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * hiddenSize), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * hiddenSize), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "quant_gemv_routed", | |
| "priority": 20, | |
| "when": ["decodeContract"], | |
| "derive": { "decodeLanes": "tunables.decodeLanes" }, | |
| "intermediates": [ | |
| { "id": "routeExpert", "dtype": "uint32", "shape": "[numTokens * topK]" }, | |
| { "id": "routeMix", "dtype": "float32", "shape": "[numTokens * topK]" }, | |
| { "id": "hiddenAct", "dtype": "float32", "shape": "[hiddenChunkTokens * topK * interSize]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "route_stage", | |
| "name": "QMoE.RouteStage", | |
| "shader": "qmoe-route-stage.wgsl.jinja", | |
| "bindings": ["router_probs", "route_expert", "route_mix", "params_route"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numTokens), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numTokens), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "token_chunks", | |
| "repeat": { "count": "hiddenChunkCount", "index": "chunk" }, | |
| "passes": [ | |
| { | |
| "id": "fc1_activation_stage", | |
| "name": "QMoE.FC1ActivationStageGemv", | |
| "shader": "qmoe-fc1-activation-gemv.wgsl.jinja", | |
| "bindings": ["input", "route_expert_u32", "fc1_experts_weights", "fc1_scales", "hidden_act", "params__uniform"], | |
| "dispatch": { | |
| "x": "ceilDiv(interSize, decodeRows)", | |
| "y": "min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * topK" | |
| } | |
| }, | |
| { | |
| "id": "output_stage", | |
| "name": "QMoE.OutputStageGemv", | |
| "shader": "qmoe-output-gemv.wgsl.jinja", | |
| "bindings": [ | |
| "hidden_act_f32", | |
| "route_expert_u32", | |
| "route_mix_f32", | |
| "fc2_experts_weights", | |
| "fc2_scales", | |
| "output", | |
| { | |
| "name": "params", | |
| "struct": [{ "name": "tokenOffset", "type": "u32", "value": "repeat.chunk * hiddenChunkTokens" }] | |
| } | |
| ], | |
| "dispatch": { | |
| "x": "ceilDiv(hiddenSize, decodeRows)", | |
| "y": "min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens)" | |
| } | |
| } | |
| ] | |
| } | |
| ] | |
| }, | |
| { | |
| "id": "quant_split_routed", | |
| "priority": 10, | |
| "when": ["quantContract", "workgroupSizeOk", "hiddenChunkFits", "routedScratchFits", "interSize > 0"], | |
| "intermediates": [ | |
| { "id": "routeExpert", "dtype": "uint32", "shape": "[numTokens * topK]" }, | |
| { "id": "routeMix", "dtype": "float32", "shape": "[numTokens * topK]" }, | |
| { "id": "hiddenAct", "dtype": "float32", "shape": "[hiddenChunkTokens * topK * interSize]" } | |
| ], | |
| "passes": [ | |
| { | |
| "id": "route_stage", | |
| "name": "QMoE.RouteStage", | |
| "shader": "qmoe-route-stage.wgsl.jinja", | |
| "bindings": ["router_probs", "route_expert", "route_mix", "params_route"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((numTokens), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((numTokens), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "token_chunks", | |
| "repeat": { "count": "hiddenChunkCount", "index": "chunk" }, | |
| "passes": [ | |
| { | |
| "id": "fc1_activation_stage", | |
| "name": "QMoE.FC1ActivationStage", | |
| "shader": "qmoe-fc1-activation-stage.wgsl.jinja", | |
| "bindings": [ | |
| "input", | |
| "route_expert_u32", | |
| "fc1_experts_weights", | |
| "fc1_scales", | |
| "hidden_act", | |
| { | |
| "name": "params", | |
| "struct": [ | |
| { "name": "activationAlpha", "type": "f32", "value": "attrs.activation_alpha" }, | |
| { "name": "activationBeta", "type": "f32", "value": "attrs.activation_beta" }, | |
| { "name": "tokenOffset", "type": "u32", "value": "repeat.chunk * hiddenChunkTokens" }, | |
| { | |
| "name": "tokenCount", | |
| "type": "u32", | |
| "value": "min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens)" | |
| } | |
| ] | |
| } | |
| ], | |
| "dispatch": { | |
| "x": "min(ceilDiv((min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * topK * interSize), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * topK * interSize), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| }, | |
| { | |
| "id": "output_stage", | |
| "name": "QMoE.OutputStage", | |
| "shader": "qmoe-output-stage.wgsl.jinja", | |
| "bindings": ["hidden_act_f32", "route_expert_u32", "route_mix_f32", "fc2_experts_weights", "fc2_scales", "output", "params"], | |
| "dispatch": { | |
| "x": "min(ceilDiv((min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * hiddenSize), (workgroupSize)), 65535)", | |
| "y": "ceilDiv(ceilDiv((min(hiddenChunkTokens, numTokens - repeat.chunk * hiddenChunkTokens) * hiddenSize), (workgroupSize)), 65535)", | |
| "z": 1 | |
| } | |
| } | |
| ] | |
| } | |
| ] | |
| } | |
| ] | |
| } | |