com.microsoft.EngramGate

com.microsoft · ONNX Runtime contrib operator · contrib since_version 1

Description

Fuses the Engram gate. For each (batch, sequence, hc_mult) row it RMS-normalizes key and query with their per-stream scales, forms dot = sum(RMSNorm(key) * RMSNorm(query)) / sqrt(hidden_size), and writes sigmoid(sign(dot) * sqrt(max(abs(dot), 1e-6))) * value, broadcasting the value row shared by every hyper-connection. That 1e-6 floor is fixed and is not epsilon; a zero dot product gives a gate of exactly 0.5. Both sums of squares and the scaled cross term accumulate in one float32 pass, and only the store narrows. Bfloat16 is not implemented.

See the ONNX Runtime EngramGate contrib-operator spec for the reference semantics.

Inputs

Name Upstream name Logical dtype Rank Shape Description Presence
keyT key T 4 — Projected Engram keys with shape (batch_size, sequence_length, hc_mult, hidden_size). required
queryT query T 4 — Hidden-state queries, exactly the same shape as key. The upstream kernels require equality rather than broadcasting, and so does this one. required
valueT value T 3 — Projected Engram value shared by every hyper-connection, with shape (batch_size, sequence_length, hidden_size). Each row of it is gated hc_mult times, once per stream. required
keyNormScaleT key_norm_scale T 2 — RMSNorm scale for keys with shape (hc_mult, hidden_size): one weight row per hyper-connection stream, selected by the row's stream index. required
queryNormScaleT query_norm_scale T 2 — RMSNorm scale for queries with shape (hc_mult, hidden_size), indexed by the same stream index as the key scale. required

Outputs

Name Upstream name Logical dtype Rank Shape Description Presence
outputT output T 4 same as keyT The gated value tensor, with the same shape and dtype as key. Every element of a row carries the same gate, so the row's stream index is only observable through the scale lookup. required

Attributes

Default values (overridable per request):

Attribute Default Description
epsilon 0.00001 Constant added to both mean-of-squares denominators before the reciprocal square root. It does not reach the gate: the 1e-6 floor under the square root is a separate fixed constant.

Type constraints

Variable Allowed dtypes
T float32, float16

Implementation variants

One implementation is selected per call from the device capabilities, the request shapes and the dtypes; these notes say what each one covers.

  • subgroup_rows_vec4 — One aligned lane group per gated row inside a subgroup, with the row's value words held in registers across the fold. The three accumulators travel as one vector through a segmented butterfly over the low lane bits, so several rows share a subgroup with no workgroup memory and no barrier.
  • packed_rows_vec4 — Portable vec4 route: the workgroup is split into a power-of-two lane group per row and each group folds its own contiguous slice of one shared vector array, so a short row does not idle the workgroup and no subgroup support is required. It also carries rows too wide to stage in registers, which it walks twice.
  • subgroup_rows_scalar — The barrier-free lane-group schedule for a hidden size that is not a multiple of four: the same segmented butterfly over scalar loads.
  • packed_rows_scalar — Scalar fallback for a hidden size that is not a multiple of four, and the route every device can select. It keeps the lane-group-per-row schedule so a narrow row still fills the workgroup.

Device requirements

Some implementation variants require subgroups. These are route-specific capabilities, not package-wide requirements; availability also depends on the request shape and dtype.

Files

Use with @huggingface/kernels

npm install --save-exact @huggingface/[email protected]

Required output shapes and logical data types are inferred from the supplied inputs and attributes; result tensors are allocated automatically.

The version: 1 option selects the published kernel contract; it is independent of any operator opset, contrib since_version, or model version. It follows the v1 branch as fixes land. To pin exact artifact bytes, pass a 40-character commit revision instead of version.

Replace each *Data placeholder with a typed array containing the corresponding input data.

import { getKernel } from "@huggingface/kernels";

const kernel = await getKernel("webgpu-kernels/com.microsoft.EngramGate", { version: 1 });
const { outputT } = await kernel({
  keyT: { data: keyTData, shape: [1, 1, 1, 2] },
  queryT: { data: queryTData, shape: [1, 1, 1, 2] },
  valueT: { data: valueTData, shape: [1, 1, 2] },
  keyNormScaleT: { data: keyNormScaleTData, shape: [1, 2] },
  queryNormScaleT: { data: queryNormScaleTData, shape: [1, 2] },
});
Downloads last month
-
kernel
webgpu
wgsl
apache-2.0
WebGPU

Requires WebGPU support. See the compatibility table.