CANN opbase tensor_view_utils 预留接口深度解析:CanPickViewAsContiguous 与 Validate 的实现原理与使用指引
【免费下载链接】opbase本项目是CANN算子库的基础框架库,为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase
导读
本文聚焦 CANN 算子库基础框架库(opbase)中tensor_view_utils模块的预留接口——CanPickViewAsContiguous(两个重载)与Validate,结合头文件声明、C++ 实现源码与单元测试,深入讲解"转置连续存储"判定、view 元数据合法性校验的底层算法。读完本文,你将理解这些预留接口的能力边界、算法细节与适用限制,能够在算子开发中正确判断何时可以安全引用它们(或选择绕开它们,改用稳定接口IsContiguous)。
预留接口警示:依据 预留接口文档,本章接口为预留接口,后续有可能变更或废弃,不建议开发者使用,开发者无需关注。本文所述行为以当前仓库代码为准,若后续版本发生变更,请以新版文档与源码为准。
一、接口全景:三组预留接口的定位
tensor_view_utils是 opbase 为算子侧提供的 tensor 视图工具模块,其公开头文件位于 include/nnopbase/opdev/tensor_view_utils.h,实现位于 src/nnopbase/composite_op/utils/tensor_view_utils.cpp。模块共提供四个接口,其中IsContiguous为稳定接口(见 IsContiguous 文档),其余三个为预留接口:
| 接口定义 | 功能说明 |
|---|---|
CanPickViewAsContiguous(std::initializer_list<const aclTensor *> tensorList) | 判断给定 tensorList 是否连续存储或者转置连续存储。 |
CanPickViewAsContiguous(const aclTensor *tensor) | 判断给定 tensor 是否连续存储或者转置连续存储。 |
Validate(const aclTensor *tensor) | 判断给定 tensor 的 view shape、view stride、view offset 是否合法。 |
从命名与语义可以看出,这三组接口服务于一条核心链路:判断一个(或一组)张量的视图(view)能否被"当成"连续张量来对待。这在算子实现中具有实际价值——连续张量可以走最简化的访存路径与 tiling 策略,而"转置连续"(如transpose/permute产生的张量)在重新排布后同样可以视为连续,从而复用连续内核。
二、CanPickViewAsContiguous(const aclTensor *tensor):单张量转置连续判定
2.1 头文件契约
头文件中的注释明确了该接口的语义(include/nnopbase/opdev/tensor_view_utils.h):
/** * @brief Check whether the input tensor can be regarded as a view of a contiguous tensor. * @param tensor The input tensor * @return True/false */ bool CanPickViewAsContiguous(const aclTensor* tensor);即:判断输入 tensor能否被视为某个连续张量的视图。注意这与IsContiguous不同——IsContiguous只回答"是否连续",而CanPickViewAsContiguous额外回答了"虽然不是标准连续,但经过转置(permute)后是否连续"。
2.2 源码实现:三步判定算法
核心实现在 src/nnopbase/composite_op/utils/tensor_view_utils.cpp:
bool CanPickViewAsContiguous(const aclTensor* tensor) { const auto& viewShape = tensor->GetViewShape(); const auto& viewStrides = tensor->GetViewStrides(); if (IsContiguous(viewShape, viewStrides)) { return true; } bool mayTranspose = false; bool mayBroadcast = false; auto strideShapePairs = BuildStrideShapePairs(viewShape, viewStrides, mayTranspose, mayBroadcast); if (mayBroadcast) { return false; } if (!mayTranspose) { return false; } std::sort(strideShapePairs.rbegin(), strideShapePairs.rend()); return IsContiguous(strideShapePairs); }算法分为三步:
- 直接连续判定:先按标准连续规则检查
viewShape与viewStrides,若已连续则直接返回true。 - 广播检测:调用
BuildStrideShapePairs扫描所有维度,若发现存在stride[i] == 0 && shape[i] != 1的维度,说明该视图由广播产生(stride 为 0 意味着该维度所有元素共享同一存储位置),无法视为连续视图,返回false。 - 转置连续判定:若不存在广播,且维度间 stride 出现"非单调递减"(即
lastStride < viewStrides[i],mayTranspose = true),说明存在维度交换(转置)。此时将所有(stride, shape)对按 stride 降序排列,再执行一次连续校验——如果排布后各维度 stride 恰好等于其外层维度尺寸的累积乘积,则说明该张量是对某个连续张量做转置得到的视图,可以"被当作连续视图"。
2.3 辅助函数细节:BuildStrideShapePairs
BuildStrideShapePairs 的扫描逻辑如下:
inline StrideShapePairs BuildStrideShapePairs(const op::Shape& viewShape, const op::Strides& viewStrides, bool& mayTranspose, bool& mayBroadcast) { StrideShapePairs strideShapePairs; strideShapePairs.reserve(viewStrides.size()); int64_t lastStride = INT64_MAX; for (size_t i = 0; i < viewStrides.size(); i++) { if (viewStrides[i] == 0 && viewShape[i] != 1) { mayBroadcast = true; return strideShapePairs; } if (viewStrides[i] != 0 && viewShape[i] != 1) { strideShapePairs.emplace_back(std::make_pair(viewStrides[i], viewShape[i])); if (lastStride < viewStrides[i]) { mayTranspose = true; } lastStride = viewStrides[i]; } } return strideShapePairs; }需要特别留意的两点:
- size-1 维度被剔除:
shape[i] == 1的维度不参与 stride 排序与连续校验。因为 size 为 1 的维度其 stride 无实际意义(只有一个元素,无所谓跨度),剔除后不影响"能否视为连续"的结论。 - stride 为 0 且 shape 不为 1 即判为广播:这是最容易误用的场景。例如
{4, 1, 6, 7}存储形状上的 broadcast 视图{4, 5, 6, 7}(中间维 stride 为 0),虽然viewShape看起来是"规则的",但物理上每个 5 元素都复用同一份数据,绝不能被当作连续存储处理,接口返回false。
2.4 单元测试验证
test_tensor_view_utils.cpp(tests/nnopbase/st/composite_op/test_tensor_view_utils.cpp下存在同构用例)给出了四类典型场景:
// 广播视图:stride 含 0 且 shape != 1 → false tensor = CreateAclTensor({4, 5, 6, 7}, {42, 0, 0, 1}, 0, {4, 1, 6, 7}); EXPECT_FALSE(op::CanPickViewAsContiguous({tensor, tensor, tensor})); // 转置+广播混合 → false tensor = CreateAclTensor({4, 6, 5, 7}, {42, 7, 0, 1}, 0, {4, 1, 6, 7}); EXPECT_FALSE(op::CanPickViewAsContiguous({tensor, tensor, tensor})); // 标准连续 → true tensor = CreateAclTensor({4, 5, 6, 7}, {210, 42, 7, 1}, 0, {4, 5, 6, 7}); EXPECT_TRUE(op::CanPickViewAsContiguous({tensor, tensor, tensor})); // 转置连续(前两维互换)→ true tensor = CreateAclTensor({5, 4, 6, 7}, {42, 210, 7, 1}, 0, {4, 5, 6, 7}); EXPECT_TRUE(op::CanPickViewAsContiguous({tensor, tensor, tensor})); // 仅中间维被"掐掉"(stride 0 且 shape=1)的切片 → false tensor = CreateAclTensor({4, 5, 6, 7}, {42, 0, 7, 1}, 0, {4, 1, 6, 7}); EXPECT_FALSE(op::CanPickViewAsContiguous({tensor, tensor, tensor}));其中{5, 4, 6, 7}+ strides{42, 210, 7, 1}是转置连续的精髓示例:存储形状{4, 5, 6, 7}的连续 strides 为{210, 42, 7, 1},将第 0、1 维互换后得到{5, 4, 6, 7}的视图 strides{42, 210, 7, 1},排序后仍满足连续条件,因此判定为 true——这就是算子常说的"transpose 后的张量可以当作连续张量处理"的判定依据。
三、CanPickViewAsContiguous(tensorList):多张量一致性判定
3.1 头文件契约
/** * @brief Check whether all input tensors can be regarded as a view of contiguous tensors * and all tensors have the same view feature. * @param tensorList The input tensors. */ bool CanPickViewAsContiguous(std::initializer_list<const aclTensor*> tensorList);头文件强调了两层含义:所有输入张量都能被视为连续张量的视图,并且所有张量具有相同的视图特征(same view feature)。
3.2 源码实现:先一致、后单判
实现在 src/nnopbase/composite_op/utils/tensor_view_utils.cpp:
bool CanPickViewAsContiguous(std::initializer_list<const aclTensor*> tensorList) { if (tensorList.size() == 0) { return true; } auto firstTensor = *(tensorList.begin()); for (auto tensor = tensorList.begin() + 1; tensor != tensorList.end(); tensor++) { if ((*tensor)->GetViewShape() != firstTensor->GetViewShape() || (*tensor)->GetViewStrides() != firstTensor->GetViewStrides()) { return false; } } return CanPickViewAsContiguous(firstTensor); }算法要点:
- 空列表返回 true:不构成任何约束,符合数学上"全称命题对空集成立"的约定,也便于上层循环统一调用。
- 逐张量比对 view shape 与 view stride:任一 tensor 的
viewShape或viewStrides与首个张量不同,立即返回false。这里只比较 view 层元数据,不比较存储形状与 offset——因为后续真正关心的是"这些张量能否以同一套连续视图的访存方式统一处理"。 - 一致后复用单张量判定:以第一个张量为代表调用单参版本。这一设计避免了重复计算,也保证了两组重载的判定口径完全一致。
测试中的典型反例(test_tensor_view_utils.cpp):
auto tensor = CreateAclTensor({4, 5, 5, 7}, {42, 0, 0, 1}, 0, {4, 1, 6, 7}); auto tensor2 = CreateAclTensor({3, 5, 5, 7}, {42, 0, 0, 1}, 0, {4, 1, 6, 7}); // 两个张量 viewShape 不同({4,5,5,7} vs {3,5,5,7})→ false EXPECT_FALSE(op::CanPickViewAsContiguous({tensor, tensor2, tensor}));这一接口在算子中的典型用途,是在批处理场景(例如对多个输入 tensor 执行同一算子)中,预先确认所有输入可以统一走"连续视图"路径,从而省去逐张量分支判断。
四、Validate:view 元数据合法性校验
4.1 头文件契约
/** * @brief Check whether the input tensor is valid. * @param tensor The input tensor * @return bool True/false */ bool Validate(const aclTensor* tensor);功能为判断给定 tensor 的view shape、view stride、view offset 是否合法。这里的"合法"包含两层含义:维度结构自洽、且 view 在物理存储范围内不越界。
4.2 源码实现:维度匹配 + 越界检测
实现在 src/nnopbase/composite_op/utils/tensor_view_utils.cpp:
bool Validate(const aclTensor* tensor) { auto viewShape = tensor->GetViewShape(); auto viewStrides = tensor->GetViewStrides(); auto viewOffset = tensor->GetViewOffset(); if (viewShape.GetDimNum() != viewStrides.size()) { OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ViewShape and ViewStride mismatch."); return false; } auto storageSize = tensor->GetStorageShape().GetShapeSize(); int64_t maxViewOffset = viewOffset; int64_t minViewOffset = viewOffset; for (size_t i = 0; i < viewStrides.size(); i++) { maxViewOffset += std::max(static_cast<int64_t>(0), (viewStrides[i] * (viewShape[i] - 1))); minViewOffset += std::min(static_cast<int64_t>(0), (viewStrides[i] * (viewShape[i] - 1))); } if (maxViewOffset + 1 > storageSize || minViewOffset < 0) { OP_LOGE(ACLNN_ERR_PARAM_INVALID, "ViewShape overlap."); return false; } return true; }算法要点:
- 维度结构自洽:
viewShape.GetDimNum() != viewStrides.size()时直接报错返回false。view 的每个维度必须有对应的 stride,维度数不匹配说明视图元数据被破坏或构造错误。注意GetViewStrides()返回的容器大小即 view 维数,这与IsContiguous实现中"按 strides 大小反向遍历 shape"的假设保持一致。 - 越界上界检测:对所有维度累加
stride[i] * (shape[i] - 1)中大于 0 的部分得到maxViewOffset,即视图最后一个元素相对存储起点的最大偏移。若maxViewOffset + 1 > storageSize(+1是因为偏移是从 0 计数的元素下标,最后一个元素的字节位置还要再占一个元素),说明视图超出了物理存储范围,非法。 - 越界下界检测:同样累加负 stride 贡献得到
minViewOffset,若minViewOffset < 0,说明负 stride 使视图起点"反向越界"(跑到存储之前),非法。这覆盖了 stride 为负的视图(如反向切片[::-1])中 offset 校准不当的场景。
4.3 单元测试验证
test_tensor_view_utils.cpp:
// 合法视图:viewShape {4,5,6,7},strides {210,42,7,1},offset 0,存储 {4,5,6,7}(size=840) tensor = CreateAclTensor({4, 5, 6, 7}, {210, 42, 7, 1}, 0, {4, 5, 6, 7}); EXPECT_TRUE(op::Validate(tensor)); // 非法视图:首维 stride 从 210 改为 211 → maxViewOffset = 211*3+42*4+7*5+1*6 = 842 // 842 + 1 = 843 > 840(存储 size)→ 越界 → false tensor = CreateAclTensor({4, 5, 6, 7}, {211, 42, 7, 1}, 0, {4, 5, 6, 7}); EXPECT_FALSE(op::Validate(tensor));第一个用例中,210*3 + 42*4 + 7*5 + 1*6 = 840,840 + 1 = 841 > 840?注意这里精确计算:maxViewOffset = 0 + 210*(4-1) + 42*(5-1) + 7*(6-1) + 1*(7-1) = 630 + 168 + 35 + 6 = 839,839 + 1 = 840 == storageSize,恰好等于存储大小,判定合法。而将首维 stride 增大到 211 后maxViewOffset = 842,843 > 840,立即触发"ViewShape overlap"错误。这个对比精确展示了"视图最后一个元素必须落在存储内"的边界语义——合法的连续视图其最大偏移恰好卡在存储末尾,多一个字节都不行。
五、与 IsContiguous 的关系与选择建议
IsContiguous(稳定接口)与CanPickViewAsContiguous(预留接口)共享同一个底层连续判定函数,但语义边界不同:
| 对比维度 | IsContiguous | CanPickViewAsContiguous |
|---|---|---|
| 接口状态 | 稳定接口 | 预留接口(可能变更或废弃) |
| 连续判定 | 严格标准连续 | 标准连续 + 转置连续(排除广播) |
| 空指针行为 | 返回 true 并打印 ERROR(见 IsContiguous 文档) | 由内部连续判定兜底返回 true |
| 典型用途 | 判断是否可走连续访存路径 | 判断视图是否可"重塑为"连续视图统一处理 |
从 IsContiguous 的源码实现 可以看到它与CanPickViewAsContiguous在前提假设上的差异:IsContiguous对私有格式(private format,如IsPrivateFormat判定的特殊存储格式)张量直接返回 true,而CanPickViewAsContiguous完全基于 view shape/stride 的数学特征判定,不感知格式。因此:
- 常规算子开发中,判定"能否走连续内核"应优先使用稳定接口
IsContiguous; - 只有当确实需要利用"转置可重排为连续"这一特性(例如对 permute 后的张量统一做连续化 tiling)时,才考虑预留接口
CanPickViewAsContiguous,并做好接口随版本变更的兼容预案; Validate作为视图元数据完整性校验工具,可用于入参校验阶段提前拦截越界或结构不一致的视图,减少后续访存阶段的非法内存访问风险。
六、扩展阅读与验证路径
若希望进一步验证接口行为,可关注以下仓库路径:
- 头文件声明与注释:include/nnopbase/opdev/tensor_view_utils.h,其中包含了各接口的契约说明(如
IsContiguous的四类连续条件); - 核心实现:src/nnopbase/composite_op/utils/tensor_view_utils.cpp,本文所有算法分析均对应其中的具体行;
- 单元测试:tests/nnopbase/ut/composite_op/test_tensor_view_utils.cpp,覆盖连续、转置连续、广播、越界等全部判定分支;集成测试见 tests/nnopbase/st/composite_op/test_tensor_view_utils.cpp;
- 模块总览文档:tensor_view_utils 模块文档,列出稳定接口与预留接口的完整清单。
总结
tensor_view_utils的三组预留接口围绕"视图能否视为连续"这一算子访存优化核心问题展开:CanPickViewAsContiguous(const aclTensor*)通过"剔除 size-1 维度、检测广播、stride 降序重排后再判连续"三步算法,识别标准连续与转置连续两种可连续化场景;其列表重载在保证所有输入 view shape/stride 完全一致后复用单张量判定;Validate则通过"维数匹配 + 正负方向越界检测"保障视图元数据安全。基于当前仓库源码与测试用例,它们的语义清晰、边界明确,但仍需牢记其预留接口身份——在正式算子代码中使用前,请评估接口变更风险,并优先考虑稳定接口IsContiguous。
【免费下载链接】opbase本项目是CANN算子库的基础框架库,为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考