Skip to contents

Applies scaling to queries: \(q_{\text{scaled}} = q \cdot \text{mlp}(\log n)\), where a small MLP learns to map sequence length to scaling factors.

Usage

SSMaxMLP(num_heads, n_hidden = 64L, elementwise = FALSE, head_dim = NULL)

Arguments

num_heads

integer(1) Number of attention heads.

n_hidden

integer(1) Number of hidden units in the MLP. Default: 64L.

elementwise

logical(1) If TRUE, apply elementwise scaling per head dimension so that each element in the head dimension gets its own scaling factor. Default: FALSE.

head_dim

integer(1) or NULL. Dimension of each attention head. Required when elementwise = TRUE.

Value

An nn_module object (class generator).