CANN pypto 的 one_hot 算子详解:整数索引到 one-hot 编码的张量转换与 Tile 切分实践 CANN pypto 的 one_hot 算子详解整数索引到 one-hot 编码的张量转换与 Tile 切分实践【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto导读pypto.one_hot是 CANN pypto 张量/算子编程框架中用于将整数索引 Tensor 转换为 one-hot 编码 Tensor 的向量类接口输入中的每个整数被映射为一条仅对应位置为 1、其余全为 0 的向量常用于分类标签编码、MoE 专家路由等场景。本文以官方 API 文档 pypto-one_hot.md 为主体完整覆盖其产品支持情况、函数原型、参数/返回值约束与 TileShape 切分配置并结合仓库源码与测试用例深入说明其实现校验、代码生成与切片组装的实际行为帮助读者正确编写可运行、可验证的 one-hot 内核。产品支持情况当前仓库中pypto.one_hot已在以下昇腾产品系列上得到支持Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持功能说明pypto.one_hot将整数 Tensor 转换为对应的 one-hot 编码。核心语义是输入张量中的每个整数元素v被转换为一条长度为num_classes的向量其中第v个位置从 0 开始计数取值为 1其余位置取值为 0。这一转换在分类任务的标签展开、索引表构建例如 MoE 专家选择矩阵等场景中被广泛使用仓库中的分布式系统测试 test_moe_distributed_dispatch_combine.py 即用pypto.one_hot(expert_ids_flat, moe_expert_num)构造专家 ID 的 one-hot 路由表再配合cumsum完成专家分发计数。函数原型one_hot(input: Tensor, num_classes: int) - TensorPython 层的接口定义位于 python/pypto/op/other.py通过op_wrapper装饰后对外暴露其 docstring 明确给出输入输出语义op_wrapper def one_hot(input: Tensor, num_classes: int) - Tensor: ...参数说明参数名输入/输出说明input输入源操作数。支持的类型为Tensor。Tensor 支持的数据类型为DT_INT8DT_INT16DT_INT32DT_INT64。支持 1-3 维。内部元素需为非负数。不支持空 TensorShape Size 不大于 2147483647即 INT32_MAX。num_classes输入one-hot 编码长度。需大于 input 中最大元素。除了文档给出的约束外Python 层实现 python/pypto/op/other.py 还做了更细的运行时校验input必须是pypto_impl.Tensor实例否则抛出PyptoError(0xF00001, TypeError(input must be aTensor))num_classes必须是 Python 内置int否则抛出TypeErrornum_classes -1时抛出RuntimeError(num_classes must be specified)用于表达“未指定类别数”的语义num_classes 0时抛出RuntimeError(num_classes must be a positive integer)保证编码长度合法。最终调用底层实现pypto_impl.OneHot(input, num_classes)完成算子的构建。返回值说明返回一个 Shape 为(input, num_classes)、数据类型为DT_INT64的 Tensor。即输出在输入形状的末尾追加一维长度等于num_classes输出元素类型固定为 64 位整型。这一约定与 PyTorch 的torch.nn.functional.one_hot行为一致仓库系统测试 test_onehot.py 正是使用torch.nn.functional.one_hot(inputs_cpu[0].long(), config.num_classes)生成期望值来对拍验证。约束说明TileShape 约束TileShape对输出进行切分其维度必须与输出维度一致且尾轴最后一维必须等于num_classes。这是因为 one-hot 的最后一维本质上是“类别展开轴”展开长度由类别数决定不可被切分必须整轴全载全量加载到 tile 中。格式约束Tensor 类型输入不支持TileOpFormat.TILEOP_NZ格式即输入需使用非 NZ 布局如 ND参与计算测试用例 onehot_test_case.py 中的输入格式也统一标注为format: ND。调用示例TileShape 设置示例调用该 operation 接口前应通过pypto.set_vec_tile_shapes设置 TileShape。TileShape 的维度应与输出一致且最后一个维度必须等于num_classest 轴不可切分。示例 1输入 input shape 为[m, n]输出为[m, n, t]其中t num_classes。TileShape 设置为[m1, n1, t1]则m1、n1分别用于切分 m、n 轴t1必须等于num_classest 轴不可切必须保证 t 轴全载。pypto.set_vec_tile_shapes(4, 16, 32)示例 2输入为一维[m]输出为[m, t]TileShape 设置为[m1, t1]其中t1 num_classes。仓库单元测试 test_dynamic_shape_error.py 中的内核即采用这种配置pypto.frontend.jit(runtime_optionsSIM_RUNTIME_OPTIONS) def one_hot_kernel( a: pypto.Tensor([pypto.DYNAMIC, ...], pypto.DT_INT32), out: pypto.Tensor([], pypto.DT_INT32), ): pypto.set_vec_tile_shapes(4, 5, 32) # num_classes 5t 轴必须等于 5 out[:] pypto.one_hot(a, 5)接口调用示例x pypto.tensor([3], pypto.DT_INT32) y pypto.one_hot(x, 5)结果示例如下输入数据x: [0, 2, 4] 输出数据y: [[1, 0, 0, 0, 0], [0, 0, 1, 0, 0], [0, 0, 0, 0, 1]]即输入元素 0 映射到[1, 0, 0, 0, 0]元素 2 映射到[0, 0, 1, 0, 0]元素 4 映射到[0, 0, 0, 0, 1]验证了“第 v 位为 1、其余为 0”的语义。该示例同时出现在接口 docstring 的 Examples 一节中可作为复制运行的参考。结合源码深入实现校验与动态 Shape 限制从 python/pypto/op/other.py 可以看到one_hot对参数做严格类型与取值范围校验后才下沉到pypto_impl.OneHot。这意味着num_classes必须显式指定且为正整数-1 会被解释为“未指定”并直接报错输入必须是 Tensor且内部元素为非负整数文档约束因为 one-hot 需要以元素值作为索引位负值或超过类别数的值在语义上非法。动态 Shape 限制one-hot 的输出尾轴长度由num_classes决定属于典型的“计算相关形状”因此输入张量不允许携带动态维度。单元测试 test_dynamic_shape_error.py 专门验证了这一行为当输入以pypto.Tensor([pypto.DYNAMIC, ...], pypto.DT_INT32)声明动态 shape 时调用pypto.one_hot(a, 5)会触发CheckTensorDynamicShape异常DYNAMIC_SHAPE_ERROR。这提示开发者one-hot 的输入 shape 必须在编译期确定不能使用pypto.DYNAMIC占位。结合源码深入代码生成与 Tile 切分原理在向量代码生成侧one_hot的底层指令发射实现在 framework/src/codegen/npu/codegen_vector_unary.cppPrintOneHotLayout()走PrintTileOpWithFullParamsInOrder()路径面向支持 Tile Tensor 的新架构如 Ascend 950 系列PrintOneHot()面向传统向量路径将原始 shape 归一化为 4 维NormalizeShape并从算子属性numClasses中取出类别数按BLOCK_SIZE / sizeof(int64_t)的对齐规则把num_classes向上对齐(numClasses align - 1) / align * align生成带模板参数的底层 tile op 调用。输出缓冲区固定以int64_t*类型参与计算与文档“返回 DT_INT64”的约定吻合。因此“尾轴必须等于 num_classes”的 TileShape 约束本质上是底层实现的要求t 轴作为类别展开轴需要一次性全载才能完成按位展开切分会导致类别索引跨 tile 而无法正确生成编码。结合源码深入系统测试如何验证正确性仓库系统测试 test_onehot.py 提供了完整的验证范式内核通过pypto.frontend.jit编译内核内部先pypto.set_vec_tile_shapes(*config.tile_shape)设置切分再用pypto.loop按执行视图逐块处理最后用pypto.assemble(result, output_offset, output0)把每块结果组装回输出 Tensor期望值由torch.nn.functional.one_hot(inputs_cpu[0].long(), config.num_classes)生成输入输出均在npu:{device_id}上执行最后assert_outputs对拍测试用例配置 onehot_test_case.py 覆盖三种典型场景OneHot_test_7输入(32, 32)int32输出(32, 32, 32)int64num_classes 32TileShape(16, 8, 32)t 轴全载为 32OneHot_test_14输入(120,)输出(120, 64)num_classes 64TileShape(32, 64)OneHot_test_16输入(128,)元素范围 451~975输出(128, 1007)num_classes 1007TileShape(8, 1007)覆盖大类别数场景。这些用例同时印证了文档的两个要点输出尾轴恒等于num_classes且各用例的 TileShape 尾轴均等于类别数32 / 64 / 1007t 轴从未被切分。常见错误与排查建议结合文档约束与源码校验逻辑使用pypto.one_hot时最容易出现的四类错误及对策如下错误场景触发原因处理建议num_classes must be specified传入num_classes -1显式指定正整数类别数num_classes must be a positive integer传入num_classes 0保证编码长度为正整数input must be a Tensor传入非 Tensor 对象先通过pypto.tensor(...)或既有 Tensor 视图构造输入CheckTensorDynamicShape异常输入声明了pypto.DYNAMIC动态维度将输入 shape 固定为编译期可知的静态形状此外还需牢记输入元素必须为非负且小于num_classes否则对应位索引越界TileShape 的尾轴必须等于num_classes且不可切分Tensor 输入不能使用TileOpFormat.TILEOP_NZ格式。小结pypto.one_hot提供了一条从整数索引 Tensor 到 one-hot 编码 Tensor 的简洁通路调用前需通过set_vec_tile_shapes设置与输出维度一致、尾轴等于num_classes的 TileShape输入仅支持 1-3 维整型DT_INT8/16/32/64静态 shape输出固定为DT_INT64。开发者可参照 test_onehot.py 的内核写法将one_hot与pypto.loop、pypto.assemble组合在昇腾 NPU 上完成大规模索引的 one-hot 展开与切片组装。【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考