Skip to main content

//clika-runtime/io.clika.runtime/Ops/qkRmsNorm

qkRmsNorm

[common]
fun qkRmsNorm(query: Tensor, key: Tensor? = null, value: Tensor? = null, queryWeight: Tensor? = null, keyWeight: Tensor? = null, headDim: Long = 0, eps: Double? = null): List<Tensor>

qkRmsNorm(query: Tensor, key: Tensor? = null, value: Tensor? = null, queryWeight: Tensor? = null, keyWeight: Tensor? = null, headDim: Long = 0L, eps: Double? = null): the qk_rms_norm operator. Per-head RMS norm over packed attention projections: query, key and value in ONE call, no reshapes. Each contiguous head_dim run of query (and key, when present) is its own normalization group: out = x / sqrt(mean(x^2) + eps) * w. value passes through untouched (it rides along so one call serves the projection triplet). Equivalent to reshaping [S, heads*head_dim] to [S, heads, head_dim], applying rms_norm, and reshaping back, with none of those steps.