TorchRec大量使用Jagged Tensor

简介: Jagged Tensor(锯齿张量)是专为变长序列设计的紧凑存储格式,用values+lengths/offsets替代padding,显著节省内存与计算。广泛应用于推荐系统中用户行为、多值标签等不等长特征处理,如HSTU模型中的拼接、拆分与矩阵乘法操作。

什么是 Jagged Tensor

Jagged Tensor(锯齿张量)是一种变长嵌套张量,用于表示每个样本长度不同的特征。与普通的 padded 2D tensor 不同,它用 values(拼接的值)+ offsets/lengths(每个样本的长度)来紧凑存储,避免大量 padding 浪费。


例子 1:用户历史点击序列(长度不一致)

假设一个 batch 有 3 个用户,他们的历史点击 item_id 分别为:

plaintext

用户A: 点击了 [101, 202, 303]         → 长度 3
用户B: 点击了 [404]                    → 长度 1  
用户C: 点击了 [505, 606, 707, 808]    → 长度 4

Padded 方式(浪费内存):

python

# 需要 pad 到最长 4,大量 0 浪费
tensor([[101, 202, 303,   0],
        [404,   0,   0,   0],
        [505, 606, 707, 808]])   # shape: (3, 4)

Jagged Tensor 方式(紧凑):

python

values  = [101, 202, 303, 404, 505, 606, 707, 808]   # 所有值拼接,长度 8
lengths = [3, 1, 4]                                     # 每个用户的序列长度
offsets = [0, 3, 4, 8]                                  # 累加偏移量


例子 2:多值标签特征(tag)

用户打的标签数量不同:

plaintext

用户A: tags = ["科技", "体育"]           → 长度 2
用户B: tags = ["美食", "旅行", "摄影"]   → 长度 3
用户C: tags = ["音乐"]                   → 长度 1

python

values  = [科技, 体育, 美食, 旅行, 摄影, 音乐]   # 拼接
lengths = [2, 3, 1]


例子 3:jagged_tensors.py 中三个 op 的实际场景

这三个 op 来源于 generative-recommenders(HSTU 模型),用于处理用户行为序列 + 候选 item 拼接的场景:

concat_2D_jagged — 拼接用户历史和候选 item

plaintext

用户A历史: [e1, e2, e3](3个embedding)    候选item: [c1, c2](2个embedding)
用户B历史: [e4](1个embedding)             候选item: [c3, c4, c5](3个embedding)

拼接后:

plaintext

用户A: [e1, e2, e3, c1, c2]       → 长度 3+2=5
用户B: [e4, c3, c4, c5]           → 长度 1+3=4

用 padded tensor 需要 shape (2, 5, dim),用户B 有 1 行 padding。

用 jagged tensor 只存 values: (9, dim) + offsets,零浪费。

split_2D_jagged — 从拼接结果中拆回用户/候选

HSTU 的 Transformer 对拼接后的序列做完 self-attention 后,需要把用户部分和候选部分拆回来:

plaintext

拼接的attention输出: [a1, a2, a3, a4, a5, a6, a7, a8, a9]
offsets_left  = [0, 3, 4]      ← 用户历史的偏移
offsets_right = [0, 2, 5]      ← 候选item的偏移
→ left:  用户表征
→ right: 候选打分

jagged_dense_bmm_broadcast_add — 变长序列的矩阵乘法

对每个用户的变长序列做投影:out = jagged × dense + bias

plaintext

jagged: (sum_B(M_i), K)  ← 所有用户序列拼接,总共 sum 个 token
dense:  (B, K, N)         ← 每个用户一个投影矩阵
bias:   (B, N)            ← 每个用户一个偏置
out:    (sum_B(M_i), N)   ← 变长输出

传统做法需要 pad → bmm → unpad,jagged 版本直接在紧凑格式上计算,省内存且更快。


为什么推荐系统特别需要 Jagged Tensor?

推荐系统中到处都是变长特征:

  • 用户点击历史(有人点了 5 个,有人点了 500 个)
  • 多值 ID 特征(用户标签、商品类目)
  • 序列特征中的多值子特征(每个行为关联多个属性)

如果全部 pad 到最大长度,内存和计算浪费巨大。Jagged Tensor 是推荐系统处理这类数据的标准做法,torchrec 的 KeyedJaggedTensor 也是基于同样的理念

相关实践学习
机器学习概览及常见算法
机器学习(Machine Learning, ML)是人工智能的核心,专门研究计算机怎样模拟或实现人类的学习行为,以获取新的知识或技能,重新组织已有的知识结构使之不断改善自身的性能,它是使计算机具有智能的根本途径,其应用遍及人工智能的各个领域。 本课程将带你入门机器学习,掌握机器学习的概念和常用的算法。
相关文章
|
5月前
|
存储 搜索推荐 PyTorch
为什么使用 TorchRec 训练和推理更快
本文结合TorchEasyRec实践,从四大维度解析推荐系统加速:1)KeyedJaggedTensor统一变长特征,实现Embedding批量融合查找;2)自动分布式分片突破单卡显存瓶颈;3)TrainPipelineSparseDist流水线并行,重叠通信与计算;4)fbgemm-gpu融合优化器,减少显存访问。端到端提升训练效率与扩展性。
596 9
|
4月前
|
缓存 关系型数据库 数据库
【赵渝强老师】PostgreSQL的数据预热扩展pg_prewarm
本文详解PostgreSQL扩展pg_prewarm,支持手工与自动两种数据预热方式:手工调用pg_prewarm()函数将表数据预加载至缓冲区;自动模式通过后台进程周期记录并重启时恢复热点数据,显著提升查询性能。
297 1
【赵渝强老师】PostgreSQL的数据预热扩展pg_prewarm
|
9月前
|
存储 机器学习/深度学习 搜索推荐
08_昇腾推荐系统加速算子:FBGEMM算子库
FBGEMM算子库适配昇腾平台,支持Torchrec模型在DCNV2和GR等推荐模型中的高效运行。已完成JaggedToPaddedDense、DenseToJagged、HstuDenseForward/Backward等核心算子的移植与优化,并引入自定义算子提升生成式推荐性能,助力推荐系统训练加速。
|
9月前
|
存储 缓存 搜索推荐
02_昇腾推荐系统架构解析:嵌入表存储到多级缓存的全链路设计
昇腾推荐系统采用多级缓存架构,基于达芬奇架构NPU实现HBM与DDR协同的Embedding存储。通过FastHashMap与动态Swap机制,结合LRU/LFU准入淘汰策略,支持大规模稀疏特征高效训练。软件层面深度适配TorchRec,提供统一接口,实现计算与通信重叠,提升端到端性能,适用于电商、短视频等大模型推荐场景。
02_昇腾推荐系统架构解析:嵌入表存储到多级缓存的全链路设计
|
3月前
|
存储 运维 监控
RFID技术赋能变电站守护万家灯火
变电站仓库是电网安全运行的关键枢纽。RFID技术通过抗金属标签、智能读写与平台联动,实现物资全生命周期闭环管理——入库秒级建档、存储动态监控、盘点高效精准、领用快速核验,显著提升调配效率与运维可靠性,筑牢电力物资保障防线,守护万家灯火。(239字)
|
5月前
|
存储 人工智能 JavaScript
Prompt、Context、Harness:AI Agent 工程的三层架构解析
2023年重“Prompt”(如何说),2025年重“Context”(看到什么),2026年跃升至“Harness”(系统级约束与验证)。三者非替代而是分层:Prompt优化表达,Context管理信息环境,Harness构建可信执行系统——模型是马,Harness才是缰绳、马鞍与路。
1435 10
Prompt、Context、Harness:AI Agent 工程的三层架构解析
|
4月前
|
人工智能 监控 安全
基于 Vercel 生成式 AI 的规模化钓鱼攻击机理与防御体系研究
本文剖析生成式AI(如v0.dev)与Vercel云平台被滥用于工业化钓鱼攻击的新威胁:攻击者零代码生成高仿真品牌登录页,依托免费托管快速上线、弹性重建,并通过Telegram实时窃密。传统检测全面失效。研究提出覆盖平台治理、内容检测、行为分析、身份加固、运营响应的五层闭环防御体系,强调多维特征融合与主动防护。(239字)
244 2
|
4月前
|
人工智能 安全 API
网络威胁频发,如何用IP离线库提升风险IP识别与实时响应能力?
面对攻击成本骤降、防御被动失衡的困局,IP离线库成为扭转战局的关键:本地部署、日更数据、毫秒查询,支持归属地/ASN/网络类型等20+维风险识别。断网、限流、隔离网络下仍可精准溯源、快速封禁,让威胁判断真正自主可控。(239字)
200 3
|
5月前
|
人工智能 JSON 自然语言处理
Vibe Coding 老翻车?可能是你的 AI 根本读不懂产品文档
Vibe Coding(氛围感编程)火爆Reddit/HN,主打自然语言驱动开发。但PDF需求文档常因排版复杂(表格错位、流程图丢失、多栏混乱)导致AI生成代码出错。解法:接入上海AI实验室开源的MinerU MCP Server,自动无损解析PDF为结构化Markdown,让AI真正“读懂”PRD,提升代码准确率。(239字)
628 6

热门文章

最新文章