SetValueOperation 算子使用与实现解析:ascend-transformer-boost 原地张量切片赋值
【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost
导读
本文基于 ascend-transformer-boost 开源仓库中的 SetValueOperation 算子知识条目及其对应源码,完整讲解该算子的语义、参数定义、参数校验规则、Runner 执行链路与测试验证方式。SetValueOperation 是仓库中一个 tier S、单阶段(single-stage)、无独立输出张量的原地拷贝类算子,核心能力是把源张量 src 的内容写入目标张量 dst 的指定切片区域(dst[starts:ends] = src)。读完本文,你将掌握 SetValueParam 各字段的语义与取值范围、算子内部两层结构(Operation 层 + Ops Runner 层)的工作方式,以及如何通过仓库内的 CSV 测试用例快速验证算子行为。
1. 算子定位与知识库入口
SetValueOperation 在 ATB(Ascend Transformer Boost)算子知识库中被归类为other分类下的单算子条目,知识条目位于 .agent/knowledge/ops/other/set_value/index.md,其元数据为:
| 元数据项 | 值 | 说明 |
|---|---|---|
| op.name | set_value | 算子标识 |
| op.category | other | 功能分类 |
| op.tier | S | 复杂度分级(S 表示低复杂度单算子) |
| op.type | single | 单算子(非融合算子) |
| source.repo_path | src/ops/ops_infer/set_value/ | 算子源码目录 |
| Runner | ops_runner | 通过 OpsRunner 执行 |
| Pipeline | 单阶段 | 一次内核图(单节点)即可完成 |
知识条目本身是导航性质的:它将读者引向路由文件 .agent/knowledge/routing/set_value.md(其中包含文件清单、推荐阅读顺序与源码路径),并指向主索引 .agent/knowledge/README.md。路由文件标注了该算子关键信息:分类 infer、复杂度 S、文件数 4、Runner 类型 OpsRunner/Operation、ACLNN: no,即该算子不提供独立的 ACLNN 接口,而是走 ATB 原生的 Operation + OpsRunner 执行路径。
2. 源码文件结构与推荐阅读顺序
SetValueOperation 的实现非常精简,仅包含 4 个文件,全部位于 src/ops/ops_infer/set_value/:
| # | 文件 | 角色 | 关注重点 |
|---|---|---|---|
| 1 | set_value_operation.h | Operation 定义 | 输入输出数量、InferShape 签名、校验接口 |
| 2 | set_value_operation.cpp | Operation 实现 | 参数校验逻辑、CreateRunner() 决策 |
| 3 | set_value_ops_runner.h | Ops Runner 定义 | 原生 Ops 执行接口(SetupKernelGraph) |
| 4 | set_value_ops_runner.cpp | Ops Runner 实现 | 内核图构建、CopyOperation 节点、平台适配 |
路由文件给出的推荐阅读顺序是:先看set_value_operation.h了解输入输出数量与 InferShape 签名,再看set_value_operation.cpp理解 CreateRunner() 的决策逻辑,随后阅读set_value_ops_runner.h与set_value_ops_runner.cpp理解原生 Ops 的调用链与平台适配。参数头文件为 include/atb/infer_op_params.h。需要说明的是,从源码结构看,该算子并没有绑定独立的专属 kernel,而是通过 Runner 在内核图中复用CopyOperation(AsdOps 的 Copy 参数类型)完成实际拷贝,这一点在第 5 节详述。
3. SetValueParam 参数定义详解
算子的全部行为由 include/atb/infer_op_params.h 中定义的infer::SetValueParam结构体控制,其完整定义如下:
//! \struct SetValueParam //! \brief 将输入源张量中的内容拷贝到输入目标张量指定位置中. //! 该拷贝为原地拷贝,最终结果修改在输入目标张量中.<br> //! 输入目标张量 dst: [a,b,c], 输入源张量src: [d,e,f]. //! dst[starts[0]: ends[0], starts[1]: ends[1], starts[2]: ends[2]] = src.<br> //! 其中 ends[0]-starts[0]需为src第0维的维度大小,ends[1]-starts[1]需为src第1维的维度大小,ends[2]-starts[2]需为src第2维的维度大小。 struct SetValueParam { //! \brief 每一维拷贝起始位置 SVector<int64_t> starts; //! \brief 每一维拷贝结束位置后一个位置,拷贝到该位置前一个位置为止 SVector<int64_t> ends; //! \brief 每一维拷贝步长,当前仅支持strides为全1. SVector<int64_t> strides; //! \brief 预留参数 uint8_t rsv[8] = {0}; };参数语义与约束归纳如下:
| 字段 | 类型 | 语义 | 约束 |
|---|---|---|---|
| starts | SVector<int64_t> | 每一维拷贝起始位置(含) | 长度须等于张量维数;starts[i] >= 0;starts[i] < ends[i] |
| ends | SVector<int64_t> | 每一维拷贝结束位置的后一个位置(不含) | ends[i] <= dst 第 i 维大小 |
| strides | SVector<int64_t> | 每一维拷贝步长 | 当前仅支持全 1 |
| rsv | uint8_t[8] | 预留字段 | 置 0,CreateOperation 中通过 OP_PARAM_RSV_CHECK 校验 |
参数校验的关键公式(见源码 L129):每一维的拷贝元素个数(ends[i] - starts[i] - 1) / strides[i] + 1必须等于 src 对应维度大小;在 strides 全为 1 的前提下,该式退化为ends[i] - starts[i] == src.dims[i]。这意味着 dst 中被覆盖的切片形状必须与 src 完全一致,SetValue 本质上是一个"形状受限的原地张量切片拷贝"。
4. Operation 层:生命周期、InferShape 与参数校验
Operation 层由 set_value_operation.h 与 set_value_operation.cpp 实现。SetValueOperation继承自OperationBase,并在构造时从AtbOperationIrCfg单例加载名为"SetValueOperation"的算子 IR 配置。
4.1 输入输出数量
static const int32_t IN_TENSOR_NUM = 2; static const int32_t OUT_TENSOR_NUM = 0;该算子有2 个输入、0 个输出:输入 0 是目标张量 dst(会被原地修改),输入 1 是源张量 src。由于是原地拷贝,InferShapeImpl无需推导任何输出 shape,仅记录日志后直接返回NO_ERROR——这也解释了为什么算子没有独立的输出张量。
4.2 算子创建入口
CreateOperation模板特化是算子的工厂入口(set_value_operation.cpp):
template <> Status CreateOperation(const infer::SetValueParam &opParam, Operation **operation) { if (operation == nullptr) { return ERROR_INVALID_PARAM; } OP_PARAM_RSV_CHECK(opParam); *operation = new (std::nothrow) SetValueOperation(opParam); if (*operation == nullptr) { ATB_LOG(ERROR) << "failed to new operation"; return ERROR_OUT_OF_HOST_MEMORY; } return NO_ERROR; }它会先通过OP_PARAM_RSV_CHECK校验预留字段,再用new (std::nothrow)创建算子实例,内存分配失败时返回ERROR_OUT_OF_HOST_MEMORY。
4.3 参数校验规则(InferShapeCheck / SetupCheck)
InferShapeCheckImpl与SetupCheckImpl都调用两个私有方法完成校验:
DimNumCheckImpl(set_value_operation.cpp):
- dst 与 src 的维数必须相等,否则返回
ERROR_INVALID_TENSOR_DIM_NUM; - dst 维数必须等于
starts/ends/strides三个参数的长度,否则返回ERROR_INVALID_PARAM。
DimsCheckImpl(set_value_operation.cpp)逐条校验:
| 校验项 | 规则 | 错误码 |
|---|---|---|
| src 维度大小 | src.dims[i] <= dst.dims[i](每一维) | ERROR_INVALID_TENSOR_DIM |
| strides | 必须全为 1 | ERROR_INVALID_PARAM |
| starts/ends | starts[i] >= 0,ends[i] <= dst.dims[i],starts[i] < ends[i] | ERROR_INVALID_PARAM |
| 拷贝形状 | (ends[i]-starts[i]-1)/strides[i]+1 必须等于 src.dims[i] | ERROR_INVALID_PARAM |
| 差异维度数 | src 与 dst 至多两个维度不同,且其中一个必须是最高维(第 0 维) | ERROR_INVALID_TENSOR_DIM |
其中"差异维度数"的判定逻辑值得展开:代码遍历 dst 的第 1 维到最后一维,统计与 src 不同的维度个数 count。若 count > 1(即第 1 维之后至少两个维度不同)则报错;若 count == 0(即第 1 维之后全部相同)也报错——此时 src 与 dst 要么完全相同、要么只在第 0 维不同,而"只有一个维度不同时该维度不能是最高维"。这与参数头文件中的 warning 注释完全一致:
输入 src 和输入 dst 的各维度要求有一个或两个维度不相同:
- 如果有一个维度不相同,则这个维度不能是最高维(第 0 维);
- 如果有两个维度不相同,则其中一个不同的维度必须是最高维(第 0 维)。
换句话说,合法的形状组合是:src 相对 dst 在非最高维中恰好有一个维度变小(可选地第 0 维也变小),形成"切片区域"写入 dst。
4.4 Runner 决策与参数序列化
CreateRunner是 Operation 到执行层的桥接(set_value_operation.cpp):
std::shared_ptr<Runner> SetValueOperation::CreateRunner(Context &context) const { (void)context; return std::make_shared<SetValueOpsRunner>(param_); }它直接构造SetValueOpsRunner并传入参数副本。此外GetParamJson通过OpParamToJson(param_)将参数序列化为 JSON,供 IR 图 dump 与调试使用。整体调用链为:
CreateOperation<SetValueParam> → SetValueOperation → CreateRunner() → SetValueOpsRunner → SetupKernelGraph() → CopyOperation 内核图节点 → 执行拷贝5. Runner 层:SetValueOpsRunner 与 CopyOperation 内核图
执行层由 set_value_ops_runner.h 与 set_value_ops_runner.cpp 实现。SetValueOpsRunner继承自OpsRunner,唯一的核心工作是重写SetupKernelGraph,把算子参数翻译成一张单节点内核图。
5.1 输入张量编排
kernelGraph_.inTensors.resize(IN_TENSOR_COUNT); // 2 个输入 kernelGraph_.outTensors.resize(0); // 0 个输出 Mki::Tensor &dst = kernelGraph_.inTensors.at(inTensorNum++); Mki::Tensor &src = kernelGraph_.inTensors.at(inTensorNum++);dst 为输入 0,src 为输入 1,内核图不声明任何输出张量,与 Operation 层的GetOutputNum() == 0保持一致。
5.2 内核图节点构建
内核图仅含 1 个节点(kernelGraph_.nodes.resize(1)),节点类型为CopyOperation,参数类型为AsdOps::OpParam::Copy。构建逻辑(set_value_ops_runner.cpp)分三步:
- dstSize:直接把 src 各维大小填入
copyParam.dstSize,表示拷贝区域尺寸; - dstStride:先将 dst 各维自后向前做累积乘积得到
dimMatul[i](dst 第 i 维之后所有维度的总元素数),再令dstStride[i] = dimMatul[i+1] * param_.strides[i](最后一维直接取strides[last]),从而把 dst 视为行主序的线性缓冲,计算目标切片区域每一维的线性步长; - dstOffset:
dstOffset = starts[last] + Σ(starts[i] * dimMatul[i+1]),即把各维起始位置换算成 dst 线性内存中的起始偏移,并带有整型溢出保护——当(INT64_MAX - dstOffset) / dimMatul[i+1] < starts[i]时直接返回ERROR_INVALID_PARAM,避免偏移量上溢导致非法内存访问。
最终节点拓扑为:
copyNode.opDesc = {0, "CopyOperation", copyParam}; copyNode.inTensors = {&dst, &src}; copyNode.outTensors = {&dst};可见该内核图把src 拷贝到 dst 的指定线性偏移区域,输出仍指向 dst 本身,从内核图层面再次印证了"原地修改"语义。文件末尾的注册宏完成了算子与内核参数的绑定:
REG_RUNNER_TYPE(SetValueOpsRunner); REG_OP_PARAM(AsdOps::OpParam::Copy);5.3 单阶段流水线
整个执行路径只有一次内核图下发、一个 CopyOperation 节点,因此路由文件将其标记为"单阶段(single-stage)"流水线。相比多阶段融合算子,SetValue 的 host 侧开销极低:无需 tiling 决策、无需多节点编排,参数翻译完成后直接执行单节点拷贝。
6. 算子配置:dtype 与 format 支持
算子支持的数据类型与排布在 ops_configs/atb_ops_info.ini 的[SetValueOperation]段中声明:
[SetValueOperation] input0.name=x1 input0.dtype=float16,float,int32,int64,bf16 input0.format=nd,nd,nd,nd,nd input1.name=x2 input1.dtype=float16,float,int32,int64,bf16 input1.format=nd,nd,nd,nd,nd配置要点:
- 两个输入分别命名为 x1(dst)与 x2(src);
- 支持的 dtype 为float16、float、int32、int64、bf16共 5 种;
- 格式仅支持nd排布;
- 两个输入支持相同的 dtype/format 组合,要求 src 与 dst 类型一致。
7. 测试验证与用例剖析
仓库为该算子提供了多层次的测试资产,可用于验证语义与校验规则。
7.1 高层测试(CSV 用例)
高层测试目录 tests/high_level_test/SetValueOperation/ 按场景划分为 Smoke、Dtype_dataFormat、Boundary_value、Performance、Requirements 等多个子目录。以 Smoke 用例 为例,其结构与典型取值如下:
| CaseName | OpParam | InShape (dst;src) | InDType | ExpectedError |
|---|---|---|---|---|
| setValueInt64_smoke | {"starts":[0,0,0],"ends":[13,2,4096],"strides":[1,1,1]} | 28,10,4096;13,2,4096 | int64 | NO_ERROR |
| setValuefloat16_smoke | {"starts":[0,0,0],"ends":[13,2,4096],"strides":[1,1,1]} | 28,10,4096;13,2,4096 | float16 | NO_ERROR |
| setValueInt32_smoke | {"starts":[0,0,0],"ends":[13,2,4096],"strides":[1,1,1]} | 28,10,4096;13,2,4096 | int32 | NO_ERROR |
| setValuefloat_smoke | {"starts":[0,0,0],"ends":[13,2,4096],"strides":[1,1,1]} | 28,10,4096;13,2,4096 | float | NO_ERROR |
验证一下第 4 节的校验规则:dst=[28,10,4096],src=[13,2,4096],两者在第 0、1 维不同(且第 0 维为最高维)、第 2 维相同,符合"两个维度不同且其中之一是最高维"的要求;ends[0]-starts[0]=13==src.dims[0],ends[1]-starts[1]=2==src.dims[1],ends[2]-starts[2]=4096==src.dims[2],拷贝形状一致。同一 CSV 前半部分还包含大量 Generalization 随机形状用例(如 dst=[984,3680]、src=[13,1024],或 dst=[15,12,2304]、src=[13,12,1024]),覆盖 int32/int64/float16/float 等类型,ExpectedError 均为 NO_ERROR。
此外 Dtype_dataFormat 与 Boundary_value 目录分别覆盖类型/格式组合与边界形状,Requirements 验证 bf16 精度支持,Performance 提供性能基线(BaseLine)对照。
7.2 API 测试
在 API 测试侧,仓库提供了 test_set_value.py 与对应的 set_value.csv 用例数据,可从 Python 侧驱动算子执行并做数值对比。结合 CSV 测试框架(tests/apitest/opstest/csv/)与高层测试框架(tests/high_level_test/operation_test.py)即可将上述用例一键跑通,用于回归验证与精度检查。
8. 使用要点与限制总结
综合参数头文件注释、源码校验逻辑与测试用例,使用 SetValueOperation 时需要重点注意以下几点:
- 原地语义:算子没有输出张量,结果直接写回输入 dst,调用方需保证 dst 可写且允许被就地修改;
- 形状强约束:dst 与 src 必须同维数;src 各维均不能大于 dst 对应维;两者恰好有 1 或 2 个维度不同,且唯一的差异维不能是第 0 维、两个差异维中必须包含第 0 维;
- strides 仅支持全 1:strides 数组长度必须等于维数,且所有元素必须为 1,否则返回
ERROR_INVALID_PARAM; - 切片区间语义:
ends[i]是开区间上界(拷贝到 ends[i]-1 为止),且拷贝形状(ends[i]-starts[i]-1)/strides[i]+1必须与 src 第 i 维大小严格一致; - 类型与格式:支持 float16/float/int32/int64/bf16 五种 dtype,格式仅支持 nd;
- 执行路径:不提供 ACLNN 接口,统一走
SetValueOperation → SetValueOpsRunner → CopyOperation 内核图的单阶段执行链路。
该算子以极简的"一张内核图 + 一个 CopyOperation 节点"实现了通用的原地切片赋值能力,适合在需要将某张量内容写入另一张量指定区域的场景中作为基础构件使用。若需深入了解其知识条目组织方式,可继续阅读 .agent/knowledge/ops/other/set_value/index.md 与路由文件 .agent/knowledge/routing/set_value.md。
【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考