lite_boost.layers.rope_apply

查看源文件
lite_boost.layers.rope_apply(x, grid_sizes, freqs)[源代码]

x 应用旋转位置编码(RoPE)。

x 的交错复对 \((x[..., 2k], x[..., 2k+1])\)freqs 的旋转角旋转,cos/sin频率表按每样本的 grid_sizes 展开并进程内缓存;本rank切片遵循序列并行(SP)切分,当 seq_len % sp_size != 0 时末rank切片在表外零填充(填充位置输入为零,输出保持为零)。调用前需初始化 torch.distributed (单进程可调用 dist.init_process_group(backend="hccl", world_size=1, rank=0))。本接口仅支持A2,不支持300I Duo。

cos/sin 表以 LRU 方式进程内缓存,条目数与估算字节数均有上限 (_ROPE_CACHE_MAX_ENTRIES=32_ROPE_CACHE_MAX_BYTES=2 GiB),缓存内存有界;被淘汰的表再次命中 时按 key 重建,结果不变。长时间运行显存持续增长时,可在业务空闲点定期调用 torch.npu.empty_cache()

符号说明:Tfreqs 频率表长度,B 为batch大小,F/H/W 为每样本的(帧数、高、宽)网格(原始序列长度 seq_len = F*H*W),D 为head_dim,N 为头数,s 为按SP切分后的每rank序列长度(padded_seq_len / sp_sizepadded_seq_len 为各样本SP切分前统一补齐后的长度)。sF*H*W 相关但不完全相同。单卡且无需补齐时 s == F*H*W;一般情况下 s >= ceil(seq_len / sp_size) (按补齐后的长度均分),末rank切片可能超出短样本的表范围,超出位置补零(该处输入亦为零,输出保持为零)。

参数:
  • x (Tensor) - shape为 \((B, s, N, D)\) 的输入张量。B 为batch大小,s 为按SP切分后的序列长度(padded_seq_len / sp_size),N 为头数,D 为head_dim(偶数,按交错复对旋转)。支持float16、float32和bfloat16。

  • grid_sizes (Tensor) - shape为 \((B, 3)\) 的整型张量,每样本的 \((F, H, W)\) 网格(seq_len = F*H*W)。

  • freqs (Tensor) - shape为 \((T, D//2)\) 的复数张量(极坐标 \(e^{i\theta}\)),须与 x 在同一设备上。

返回:

Tensor, 与 x 同shape,输出固定为float32。

异常:
  • RuntimeError - x 非4维、D 非偶数、freqsx 不在同一设备时抛出。

  • ValueError - 调用前未初始化 torch.distributed 时抛出。

  • TypeError - grid_sizes 非2维 \((B, 3)\) 时抛出。

支持平台:

Ascend

样例:

>>> import os
>>> import torch
>>> import torch_npu
>>> import torch.distributed as dist
>>> from lite_boost.layers import rope_apply
>>> os.environ["MASTER_ADDR"] = "127.0.0.1"
>>> os.environ["MASTER_PORT"] = "29500"
>>> torch.npu.set_device(0)
>>> torch.npu.set_compile_mode(jit_compile=False)  # complex freqs on some platforms
>>> if not dist.is_initialized():
...     dist.init_process_group(backend="hccl", world_size=1, rank=0)
>>> x = torch.randn(1, 16, 8, 32, device="npu")
>>> grid_sizes = torch.tensor([[4, 4, 4]], device="npu")
>>> freqs = torch.exp(1j * torch.rand(1024, 16)).to("npu")
>>> y = rope_apply(x, grid_sizes, freqs)
>>> print(y.shape)
torch.Size([1, 16, 8, 32])
>>> print(y.dtype)
torch.float32
>>> dist.destroy_process_group()