十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

CANN ops-nn EmbeddingHashTableLookupOrInsert 算子详解:NPU Hash 表按 key 查询与插入的完整实现指南

CANN ops-nn EmbeddingHashTableLookupOrInsert 算子详解:NPU Hash 表按 key 查询与插入的完整实现指南 CANN ops-nn EmbeddingHashTableLookupOrInsert 算子详解NPU Hash 表按 key 查询与插入的完整实现指南【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn本文是 CANN ops-nn 神经网络算子库中EmbeddingHashTableLookupOrInsert算子的技术指南。该算子面向大规模稀疏特征场景在 NPU 上实现按 key 查询 Hash 表命中即返回 value、未命中即插入的原子语义支撑推荐系统 Embedding 训练/推理。读完本文你将掌握该算子的功能语义、全部输入输出与属性参数、产品支持矩阵以及从 shape 推导、tiling 分片到 SIMT 内核的源码级实现原理。算子概述与功能语义EmbeddingHashTableLookupOrInsert是 CANN ops-nn 仓库hash目录下的一组 Hash 表算子之一同目录还包含 init_embedding_hash_table、embedding_hash_table_export、embedding_hash_table_import、embedding_hash_table_apply_adam_w 等共同构成 NPU 上的 Hash 表算子族。算子的核心功能语义见 README.md根据 key 值查看 table 中是否存在 key如果存在则不插入 value 值并且导出 key 当前位置上的值如果不存在则对 key 进行 hash找到位置后插入 value。即这是一个lookup-or-insert查询或插入原子操作key 已存在于 Hash 表中不重复插入把该 key 所在桶bucket的 value 序列导出到输出key 不存在对 key 做哈希定位将 value 插入到 Hash 表同时从语义上返回该位置的 value。这类语义正是稀疏 Embedding 场景第一次见到某个 id 就把随机初始化的向量写进表、之后直接读表的标准需求避免了先查后插两步操作之间的并发竞态。产品支持情况该算子目前仅支持 Ascend 950 系列产品其余产品线均不支持详见 README.md产品是否支持Ascend 950PR / Ascend 950DT√Atlas A3 训练系列产品 / Atlas A3 推理系列产品✗Atlas A2 训练系列产品 / Atlas A2 推理系列产品✗Atlas 200I/500 A2 推理产品✗Atlas 推理系列产品✗Atlas 训练系列产品✗这一支持范围与源码中的算子注册一致在 embedding_hash_table_lookup_or_insert_def.cpp 中算子仅为ascend950、ascend960dt、ascend350三个 AICore 平台添加了配置对应地op_host/config/目录下也只提供了 ascend950、ascend960dt、ascend350 三个平台的二进制配置。README 表格中的 Ascend 950PR/Ascend 950DT 与源码中ascend950/ascend960dt一一对应。输入输出与属性参数说明输入/输出参数名输入/输出描述数据类型数据格式table_handle输入输入 Hash 表 handle 句柄里面包含了 Hash 表的表头地址等INT64NDkeys输入查询 key 序列INT64NDvalues输出查询 key 如果已存在返回的对应位置上的 value 序列FLOATND属性参数名输入/输出/属性描述数据类型数据格式bucket_size输入属性Hash 表桶数量INT64-embedding_dim输入属性Hash 表桶深度INT64-filter_mode输入属性准入过滤模式默认no_filter不过滤STRING-filter_freq输入属性准入频次阈值filter_mode生效时使用默认 0INT64-default_key_or_value输入属性是否使用默认 key 或 value默认falseBOOL-default_key输入属性默认 key 值默认 0INT64-default_value输入属性默认 value 值默认 0.0FLOAT-filter_key_flag输入属性是否启用 filter_key 过滤默认falseBOOL-filter_key输入属性需要过滤的 key 值默认 -1INT64-源码层面的参数佐证上述参数定义在算子原型 embedding_hash_table_lookup_or_insert_proto.h 与注册实现 embedding_hash_table_lookup_or_insert_def.cpp 中均有完整声明table_handle、keys均为REQUIRED输入数据类型固定DT_INT64、格式FORMAT_NDvalues为REQUIRED输出固定DT_FLOAT即 float32、FORMAT_NDbucket_size、embedding_dim为REQUIRED属性必填无默认值其余属性均为OPTIONAL并有默认值filter_modeno_filter、filter_freq0、default_key_or_valuefalse、default_key0、default_value0.0、filter_key_flagfalse、filter_key-1。关于属性的语义原型注释给出了更精确的说明见 proto.hfilter_mode取值no_filter或counter。counter启用基于计数器counter的准入过滤模式no_filter关闭过滤功能filter_freq过滤器阈值filter thresholdfilter_mode生效时使用default_key_or_value为true时返回default_key对应的值为false时返回default_valuedefault_key / default_value用户设置的默认 key / 默认 valuefilter_key_flag / filter_keyfilter_key_flagtrue时启用filter_key过滤被过滤的输入 key 会返回default_value当default_key_or_valuefalse时或改写为default_key当default_key_or_valuetrue时。输入输出 Shape 与数据类型推导算子的输出 shape 由 InferShape 逻辑决定见 embedding_hash_table_lookup_or_insert_infershape.cpp读取输入keys的 shape将各维度乘积得到 key 总数key_size对任意 N 维 keys 输入做了展平处理输出values被设置为二维shape[key_size, embedding_dim]——第一维是 key 个数第二维是属性embedding_dim桶深度。数据类型推导InferDtype则严格校验table_handle与keys必须是DT_INT64values必须是DT_FLOAT任一不匹配即返回失败见 infershape.cpp。关于table_handle的内容算子原型注释说明其 shape 为[5]见 proto.h在内核侧table_handle的首个 64 位整数存的是 Hash 表本身的内存地址内核先取地址再解引用得到表头与表体见 lookup_or_insert_base.h。内核实现原理从哈希定位到并发安全插入算子的计算内核位于 op_kernel/embedding_hash_table_lookup_or_insert.cpp入口函数按 TilingKey 分发到两条 SIMT 内核路径EMBEDDING_HASH_TABLE_LOOKUP_OR_INSERT_TILING_KEY_GENERAL1001通用路径 kernel_lookup_or_insert_general.h支持任意embedding_dimEMBEDDING_HASH_TABLE_LOOKUP_OR_INSERT_TILING_KEY_OPT_DIM1002特化路径 kernel_lookup_or_insert_opt_dim.h仅当embedding_dim ∈ {1, 2, 4, 8, 16, 32}时启用对常见小维度做编译期展开优化。两条路径共享同一套核心查找/插入算法与桶内存布局公共逻辑在 lookup_or_insert_base.h 中定义。桶bucket内存布局从lookup_or_insert_base.h中的常量可以还原每个桶的布局偏移均按字节计桶首 8 字节存放当前桶的keyint64COUNTER_OFFSET 88 字节 int64 计数器counter用于filter_modecounter的准入统计TABLE_STATE_OFFSET 164 字节状态字段1表示该桶已写入有效 keyTABLE_FLAG_OFFSET_FOR_B32 204 字节 flag 字段与 evict淘汰算子的标记EVICTED_FLAG_MASK 1 3配合VALUES_OFFSET 24桶内 value 序列起始位置长度为embedding_dim * sizeof(float)。单桶总大小按 8 字节向上取整对齐bucketSize_ RoundUpTo8(VALUES_OFFSET embeddingDim * sizeof(float))见 lookup_or_insert_base.h保证相邻桶之间内存对齐。查找与插入算法流程ComputeLookupOrInsert通用路径见 kernel_lookup_or_insert_general.h的处理流程如下key 过滤可选若启用filter_key_flag且当前 key 等于filter_key则根据default_key_or_value选择为0时直接向输出行写default_value并跳过后续处理否则把 keys 输入中的该位置改写为default_key并以默认 key 参与后续查找哈希定位控制线程threadX 中的第一条线程用MurmurHash3(keys[i], sizeof(int64_t), 0) % tableSize计算初始桶号tableSize即属性bucket_size线性探测 CAS 抢占从初始桶开始最多探测tableSize次。使用asc_atomic_cas在桶的 flag 位置20~23 字节做 0→BIG_ENDIAN_ONE的原子比较交换CAS 成功原值为 0说明桶空闲当前线程写 key、__threadfence()内存屏障后置状态为 1插入成功并insertCountsCAS 失败自旋等待状态字段变为 1等占用线程完成写入随后读取桶内 key与当前 key 相等则查询命中不等则探测下一个桶(currIdx 1) % tableSizeevict 标记清理查询命中时若发现桶 flag 带EVICTED_FLAG_MASK该桶曾被淘汰算子标记则用 CAS 清除该标记并累加插入计数与 evict 算子的逻辑相照应见 kernel_lookup_or_insert_general.h计数与导出命中的桶由控制线程对其 counter 做asc_atomic_add(1)随后通过__shfl把成功标志与桶地址广播给同组 X 线程由多线程协作把桶内 value 序列拷贝到输出values[i * embeddingDim ...]。线程协作与访存合并优化每个 AICore 内按 (x, y) 二维组织 SIMT 线程threadXNum个 X 线程协作搬运一行 value多个 Y 线程各自按步长遍历 key。为提升访存带宽内核按embedding_dim的对齐情况选择合并档位merge见 kernel_lookup_or_insert_general.hdim % 4 0→merge 4读侧用 2×float2B64、写侧用 float4B128dim % 2 0→merge 2float2B64否则 →merge 1标量 floatB32。由于桶内 value 区起始偏移 24B、桶 stride 仅 8B 对齐读侧最高只能使用 B64 短向量而输出行i*dim*4在dim%40时是 16B 对齐写侧可以升级为 B128float4代码注释对此做了明确说明。OPT_DIM特化路径进一步在编译期展开这些档位消除运行时分支见 kernel_lookup_or_insert_opt_dim.h。插入计数回写每个 Y 线程把自己insertCounts写入 UB 中对应的槽位随后用 SIMD 归约ComputeInplaceReduceSumB64求和再通过AtomicAdd把新增数量累加到 table_handle 的HANDLE_SIZE_ALL_OFFSET下标 2与HANDLE_SIZE_ALL_NOEXPORT_OFFSET下标 4两个统计字段上见 kernel_lookup_or_insert_general.h供导出export算子统计表容量时使用。Tiling 分片策略与 WorkspaceTiling 逻辑在 embedding_hash_table_lookup_or_insert_tiling_arch35.cpp 中实现关键决策如下读取与校验从输入 shape 计算keyNumkeys 的 shape 乘积并校验 keys 的 dtype 必须为DT_INT64线程分片threadXNum按next_pow2(ceil(embeddingDim / merge))计算且不超过 WARP_SIZE32threadYNum MAX_THREAD_NUM / threadXNum其中MAX_THREAD_NUM在非 FPGA 平台为 512TilingKey 选择embeddingDim属于{1, 2, 4, 8, 16, 32}时走OPT_DIM特化路径否则走GENERAL通用路径见 tiling_arch35.cpp 与 L206-L210核数分配启动核数SetBlockDim取ceil(keyNum / threadYNum)与平台 AIV 核数coreNumAiv的较小值Workspace申请固定 16MB 系统 WorkspaceASCENDC_TOOLS_WORKSPACE 16777216内核启动时先校验 workspace 非空再取用户 workspace见 embedding_hash_table_lookup_or_insert.cppUB 预留为每个 Y 线程预留一个 B64 槽位存放各自插入计数按 UB block 对齐。Tiling 相关的单元测试位于 test_embedding_hash_table_lookup_or_insert_tiling.cppInferShape 测试位于 test_embedding_hash_table_lookup_or_insert_infershape.cpp可作为验证行为与回归的参考。约束与调用说明约束原文档 README.md 声明本算子无额外约束约束说明无。需要留意的前提条件主要来自实现层面仅支持 Ascend 950PR/950DT 平台见上文产品支持矩阵、table_handle与keys必须为 INT64、values必须为 FLOAT且使用前需要先通过同族的初始化算子建立 Hash 表。调用说明原文档中调用方式一栏为无/无/无即当前仓库未提供可直接运行的样例代码。若需在图中接入该算子可依据算子原型 embedding_hash_table_lookup_or_insert_proto.h 中REG_OP声明的接口契约构造算子节点EmbeddingHashTableLookupOrInsert( table_handle : Tensor(INT64, ND) # 输入shape [5] 的句柄 keys : Tensor(INT64, ND) # 输入待查询/插入的 key 序列 ) - values : Tensor(FLOAT32, ND) # 输出shape [N, embedding_dim] attr: bucket_size : int # 必填桶数量 embedding_dim : int # 必填桶深度 filter_mode : str # 默认 no_filter可选 counter filter_freq : int # 默认 0 default_key_or_value : bool # 默认 false default_key : int # 默认 0 default_value : float # 默认 0.0 filter_key_flag : bool # 默认 false filter_key : int # 默认 -1典型使用场景结合算子语义与同族算子init_embedding_hash_table 初始化表、embedding_hash_table_export / import 导入导出、embedding_hash_table_apply_adam_w 参数更新该算子的典型落地路径为训练启动时用init_embedding_hash_table按bucket_size、embedding_dim初始化 Hash 表并取得table_handle每个 step 的 Embedding 前向过程调用EmbeddingHashTableLookupOrInsert已训练过的 id 直接取回向量新出现的 id 自动插入配合default_key_or_value/default_value控制冷启动初始值用filter_modecounterfilter_freq控制稀疏特征的准入用filter_key_flag/filter_key屏蔽填充位等特殊 key训练过程中由embedding_hash_table_apply_adam_w对表中向量做优化器更新训练结束用export/import做表的持久化导出与恢复。总结EmbeddingHashTableLookupOrInsert是 CANN ops-nn 面向 NPU 稀疏场景提供的高性能 Hash 表查询/插入算子其价值在于把查表 缺省插入合并为一次原子操作并通过 SIMT 线程协作、访存合并、CAS 并发控制、维度特化与分层 tiling 等手段在 Ascend 950 系列产品上获得可扩展的吞吐。理解本文涉及的算子原型proto.h、内核算法op_kernel与 tiling 策略tiling_arch35.cpp即可在此基础上进行二次开发、性能调优或集成到自定义的训练/推理链路中。【免费下载链接】ops-nn本项目是CANN提供的神经网络类计算算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-nn创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表