LayoutTransformOps

本目录算子仅支持 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 训练系列

接口定义

# 1. 固定 D 维度
torch.ops.fbgemm.recat_embedding_grad_output(
    Tensor grad_output,            # [B, T_global, D]
    int[] num_features_per_rank,   # [T]
) -> Tensor                        # [B * T_global * D]

# 2. 混合 D 维度 (Python list 参数)
torch.ops.fbgemm.recat_embedding_grad_output_mixed_D(
    Tensor grad_output,            # [B, Sum(D)]
    int[] dim_sum_per_rank,        # [T]
) -> Tensor                        # [B * Sum(D)]

# 3. 混合 D 维度 (Tensor 参数)
torch.ops.fbgemm.recat_embedding_grad_output_mixed_D_batch(
    Tensor grad_output,            # [B, Sum(D)]
    Tensor dim_sum_per_rank,       # [T], int64
    Tensor cumsum_dim_sum_per_rank,# [T], int64
) -> Tensor                        # [B * Sum(D)]

功能说明

将模型并行各 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
        # src: grad_output[:, cum : cum + dim_current]  -> [B, dim_current]
        # dst: sharded[B*cum : B*cum + B*dim_current]   -> 展平
        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_outputdim_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_rankdim_sum_per_rank 的排他前缀和, cumsum_dim_sum_per_rank[t]为每张卡的特征起始索引。
  • 支持 dim_current == 0 的空 rank(跳过)

调用示例

import torch
import fbgemm_gpu  # noqa:F401
import fbgemm_ascend  # noqa:F401

torch.npu.set_device("npu:0")

# ---- 1. recat_embedding_grad_output (固定 D) ----
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)
# out.shape == (20,)
# rank0: [1,2,3,4, 11,12,13,14]  rank1: [5,6,7,8,9,10, 15,16,17,18,19,20]

# ---- 2. recat_embedding_grad_output_mixed_D (Python list) ----
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)
# out == [1,2,6,7, 3,4,5,8,9,10]

# ---- 3. recat_embedding_grad_output_mixed_D_batch (Tensor) ----
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
)
# out == [1,2,6,7, 3,4,5,8,9,10]

编译与测试