本目录算子仅支持 NPU 调用。
目录结构
layout_transform_ops
├── layout_transform_ops.cpp
├── README.md
└── recat_copy_async
└── atlasA5
├── op_host
│ ├── recat_copy_async.cpp
│ └── recat_copy_async_tiling.h
├── op_kernel
│ ├── recat_copy_async.cpp
│ └── recat_copy_async_kernel.h
├── recat_copy_async.json
└── run.sh
硬件支持情况
| 实现目录 |
典型硬件 |
atlasA5/ |
Atlas A5 训练系列 |
接口定义
torch.ops.fbgemm.recat_embedding_grad_output(
Tensor grad_output,
int[] num_features_per_rank,
) -> Tensor
torch.ops.fbgemm.recat_embedding_grad_output_mixed_D(
Tensor grad_output,
int[] dim_sum_per_rank,
) -> Tensor
torch.ops.fbgemm.recat_embedding_grad_output_mixed_D_batch(
Tensor grad_output,
Tensor dim_sum_per_rank,
Tensor cumsum_dim_sum_per_rank,
) -> Tensor
功能说明
将模型并行各 rank 的 embedding 梯度按 rank 重新拼接(recat),从全局布局重排为分片(sharded)布局,用于分布式训练的反向聚合。
recat_embedding_grad_output:固定维度D,输入 [B, T_global, D],按 num_features_per_rank 切分 T_global 维度。
recat_embedding_grad_output_mixed_D:混合维度,输入 [B, Sum(D)],按 dim_sum_per_rank 切分 D 维度;切分参数以 Python list 传入。
recat_embedding_grad_output_mixed_D_batch:与 mixed_D 功能相同,切分参数以 Tensor 形式传入。
数据布局转换
输入 (按维度连续存储): 输出 (按 rank 连续拼接):
[B, D0 + D1 + ... + DT-1] --> [B*D0, B*D1, ..., B*DT-1] (1D)
每个 rank t 的输出段长度为 B * dim_current,其中 dim_current = dim_sum_per_rank[t]。
仿真/伪代码
def recat_mixed_D(grad_output, dim_sum_per_rank):
B, dim_sum = grad_output.shape
sharded = torch.empty(B * dim_sum, dtype=grad_output.dtype)
cum = 0
for t, dim_current in enumerate(dim_sum_per_rank):
if dim_current == 0:
continue
sharded[B*cum : B*cum + B*dim_current] = grad_output[:, cum:cum+dim_current].reshape(-1)
cum += dim_current
return sharded
参数说明
recat_embedding_grad_output
| 名称 |
输入/输出 |
类型 |
数据格式/形状 |
说明 |
| grad_output |
输入 |
Tensor[float32/float16] |
[B, T_global, D] |
局部批次梯度,必须连续 |
| num_features_per_rank |
输入 |
int[] |
[T] |
每个 rank 的特征数量,sum == T_global |
| sharded_grad_output |
输出 |
Tensor |
[B * T_global * D] |
1D sharded 梯度,类型同 grad_output |
recat_embedding_grad_output_mixed_D / _mixed_D_batch
| 名称 |
输入/输出 |
类型 |
数据格式/形状 |
说明 |
| grad_output |
输入 |
Tensor[float32/float16] |
[B, Sum(D)] |
局部批次梯度,必须连续 |
| dim_sum_per_rank |
输入 |
int[] |
[T] |
每个 rank 的维度和,sum == Sum(D) |
| sharded_grad_output |
输出 |
Tensor |
[B * Sum(D)] |
1D sharded 梯度,类型同 grad_output |
recat_embedding_grad_output_mixed_D_batch
| 名称 |
输入/输出 |
类型 |
数据格式/形状 |
说明 |
| grad_output |
输入 |
Tensor[float32/float16] |
[B, Sum(D)] |
局部批次梯度,必须连续 |
| dim_sum_per_rank |
输入 |
Tensor[int64] |
[T] |
每个 rank 的维度和,sum == Sum(D) |
| cumsum_dim_sum_per_rank |
输入 |
Tensor[int64] |
[T] |
dim_sum_per_rank 的前缀和 |
| sharded_grad_output |
输出 |
Tensor |
[B * Sum(D)] |
1D sharded 梯度,类型同 grad_output |
参数约束
grad_output 必须连续
grad_output 与 dim_sum_per_rank / cumsum_dim_sum_per_rank 必须在同一 NPU 设备上
sum(dim_sum_per_rank) == grad_output.size(1)(mixed_D 系列)
sum(num_features_per_rank) == grad_output.size(1)(固定 D,size(1) == T_global)
cumsum_dim_sum_per_rank 为 dim_sum_per_rank 的排他前缀和, cumsum_dim_sum_per_rank[t]为每张卡的特征起始索引。
- 支持
dim_current == 0 的空 rank(跳过)
调用示例
import torch
import fbgemm_gpu
import fbgemm_ascend
torch.npu.set_device("npu:0")
B, T_global, D = 2, 5, 2
num_features_per_rank = [2, 3]
grad_output = torch.tensor([
[[1, 2], [3, 4], [5, 6], [7, 8], [9, 10]],
[[11, 12], [13, 14], [15, 16], [17, 18], [19, 20]]
], dtype=torch.float32, device="npu:0")
out = torch.ops.fbgemm.recat_embedding_grad_output(grad_output, num_features_per_rank)
B, dim_sum = 2, 5
dim_sum_per_rank = [2, 3]
grad_output = torch.tensor([
[1, 2, 3, 4, 5],
[6, 7, 8, 9, 10]
], dtype=torch.float32, device="npu:0")
out = torch.ops.fbgemm.recat_embedding_grad_output_mixed_D(grad_output, dim_sum_per_rank)
dim_sum_per_rank_t = torch.tensor([2, 3], dtype=torch.int64, device="npu:0")
cumsum_dim_sum_per_rank_t = torch.tensor([0, 2], dtype=torch.int64, device="npu:0")
out = torch.ops.fbgemm.recat_embedding_grad_output_mixed_D_batch(
grad_output, dim_sum_per_rank_t, cumsum_dim_sum_per_rank_t
)
编译与测试