IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Python class

RotaryEmbedding

RotaryEmbedding​

class max.experimental.nn.rope.RotaryEmbedding(weight)

source

Bases: Module

Applies Rotary Positional Embeddings (RoPE) to input tensors.

RoPE encodes positional information using complex-valued rotations applied to query and key vectors in attention. This encoding is relative, allowing models to better generalize to sequences longer than those seen during training.

See “RoFormer: Enhanced Transformer with Rotary Position Embedding” (https://arxiv.org/abs/2104.09864)

from max.experimental import random
from max.experimental.nn.rope import RotaryEmbedding
from max.experimental.tensor import Tensor

# (max_sequence_length, head_dim // 2, 2).
rope = RotaryEmbedding(weight=Tensor.zeros([2048, 64, 2]))

# Apply to query or key tensors in attention.
# Shape: (batch, seq_len, num_heads, head_dim)
random.set_seed(0)
query = random.normal([4, 128, 12, 128])
query_with_rope = rope(query, start_pos=0)

print(query_with_rope.shape)  # [4, 128, 12, 128]

Parameters:

weight (Tensor)

dim​

property dim: int

source

Returns the embedding dimension.

forward()​

forward(x, start_pos=0)

source

Applies rotary positional embeddings (RoPE) to x.

seq_len is inferred from the shape of x.

Parameters:

  • x (Tensor) – Activation tensor with shape (batch, seq_len, n_kv_heads, head_dim). x is interpreted as a complex number valued tensor where the head_dim dimension is alternating pairs of (real, imaginary) parts.
  • start_pos (int | str | Dim | integer | TypedAttr) – starting position of input tensor, defaults to 0 if None

Returns:

Input activation tensor with rotary positional embeddings applied and the same shape as x.

Return type:

Tensor

max_sequence_length​

property max_sequence_length: int

source

Returns the maximum sequence length.

weight​

weight: Tensor

source