欢迎回来
登录你的知识库账户
忘记密码?
还没有账户?立即注册
创建账户
注册你的专属知识库
已有账户?去登录
找回密码
输入注册邮箱获取验证码
返回登录
请输入图片中的验证码以继续注册
加载中...
取消
新建收藏
手动添加你喜欢的内容
取消
编辑头像与昵称
上传新头像或修改你的显示昵称
支持 JPG/PNG,最大 2MB
取消

问题反馈

notebasewww.notebase.cn
控制台
内容库
动态
管理
账户
U
用户
--
在线
v0.8.7 · 知识库
笔记
KnowledgeBase
网络无边,知识有迹。
0笔记
0工具
30推荐

分类导航

按主题直达

编辑精选

站内用户贡献 · 真实笔记

最新收录

每日更新
继续浏览全部内容 →
>
笔记
0
加载中...
工具
0
此页用于记录用户反馈问题后的每一次改进
笔记用法

“写笔记”支持四种格式——Word 文档、Excel 表格、Markdown、纯文本,起稿或二次编辑时都能随时切换,同一篇笔记想用哪种形态来记,都由你说了算。

md、txt、csv、json 这类纯文本则原样载入,不做多余加工。拿一张现成的表倒进来、改几笔、再导出去,等于白用一台免费的格式转换器。

要带走就在右上角点“下载”,可导出 PDF、Word、Markdown、Excel、TXT 等格式;列表卡片“⋯”菜单里,也有同样的下载入口。

工具用法

在“工具”页点“+ 上传工具”即可发布:填好名称与链接,再用 Markdown 把使用方法写清楚——能解决什么问题、怎么装、怎么用,比堆介绍实在。

要分发安装包就一并上传压缩包(ZIP、RAR、7Z、TAR.GZ,最大 35MB),别人在详情页一键下载;只放链接不带附件也可以。

工具按大家的收藏热度排序,好用的自然会被顶上来。发布后可在详情页或卡片菜单里编辑、下架。

隐藏笔记

写笔记时勾上“隐藏”,这篇就只存在于你自己的账号里:不进列表、不进搜索、不上首页精选,也不会出现在任何公开的页面,链接发给别人同样打不开。

适合放密码、草稿、日记这类只给自己看的内容;想公开,去“发布”打开它,把“隐藏”的勾去掉再保存,之后编辑会默认保持原状态,不会悄悄变回公开。

数据安全

你的内容会同时保存在多个副本上,系统定期做备份与完整性校验,再配合异地容灾机制:就算某台机器出问题,数据也不会丢,可以长期放心存放;特别重要的资料,仍建议你另外再留一份备份。

技术

全站跑在容器化、模块化的现代架构上,更新、部署、回滚都很快,扩展性和稳定性都按长期运营的标准来设计(Built for reliability, designed to scale)。

理念

这个网站最早只是一个人的笔记仓库,后来慢慢长成现在的知识中枢。设计上很克制——没有广告、没有追踪、没有推荐算法,只是干干净净地存放一些东西;既然做好了,就公开出来,万一有人用得上呢。

原则

不做大而全,不做平台梦,保持简单、保持克制、保持好奇。所有内容都由用户贡献、由用户维护:不会突然冒出付费墙,不会在角落塞广告位,也不会把你的数据卖给第三方。

更多

产品会持续迭代,站内日志页记录着每一次改动,改了什么都有迹可循;想了解这个站是怎么一步步走到今天的,翻翻日志就能看到来龙去脉。

举报

如果在这里看到涉嫌违规的内容,点对应卡片右侧的“举报”按钮就能提交,我们会尽快核实处理;也谢谢你花一点时间,一起把这里维护干净。

趋势
// 点击导航加载发现
归档
// 归档为空
最近浏览
// 暂无浏览记录
发布
// 加载中...
用户发布
// 加载中...
用户管理
// 加载中...
访问统计
// 加载中...
内容审核
// 加载中...
个人信息
// 加载中...
返回首页

英伟达双塔 AI 模型开源发布

2026/7/6人工智能

最近英伟达开源了一个挺有意思的双塔模型,我花了两天时间把论文和代码都啃了一遍,感觉这东西在推荐系统、搜索匹配这些场景里很有潜力。今天就把我的理解整理成笔记,和大家聊聊这个模型到底做了什么、怎么做的、以及我们能怎么用。


一、背景:为什么需要双塔模型?

先简单说说双塔模型的来头。在推荐系统里,经典的思路是“用户-物品匹配”——我们有一堆用户特征(年龄、历史点击、兴趣标签),一堆物品特征(标题、类别、价格),要算出用户对某个物品的偏好分数。双塔模型就是做这个的:一个塔处理用户侧特征,另一个塔处理物品侧特征,最后把两个塔输出的向量做点积或者余弦相似度,得到匹配分数。

英伟达这次开源的模型叫 NVIDIA Two-Tower Model,不过它并不是简单的双塔,而是做了一些关键的优化,特别是在处理大规模稀疏特征和训练效率上下了功夫。


二、模型架构:两个塔,但不止是“双塔”

2.1 整体结构

模型的核心结构依然是双塔,但每个塔内部不再是简单的全连接层堆叠,而是引入了 Multi-Head Attention 和 Cross-Attention 机制。具体来说:

  • 用户塔:输入用户ID、用户历史行为序列(比如最近点击的10个物品ID)、用户画像特征(年龄、性别、地域等)。
  • 物品塔:输入物品ID、物品属性(标题、类别、品牌等)、物品的上下文特征(比如发布时间、热度等)。

两个塔各自过一层 Embedding 层,把离散特征转成稠密向量,然后分别送入一个 Transformer Encoder(注意不是Decoder,没有自回归)。最后从Encoder的输出里取一个 [CLS] token 对应的向量作为塔的输出(类似BERT的做法)。

2.2 关键创新:Cross-Tower Attention

传统的双塔模型,两个塔是独立训练的,只在最后计算相似度时交互。这导致一个问题:用户塔学到的表示没有看到物品侧的信息,反之亦然。英伟达的模型在训练时引入了一个 Cross-Attention 层,让用户塔和物品塔在中间层就互相“看”一眼。

具体实现是这样的:

用户塔输出: U = Encoder_U(user_features)   # shape: [batch, seq_len_u, d]
物品塔输出: I = Encoder_I(item_features)   # shape: [batch, seq_len_i, d]

# Cross-Attention: 以用户塔为Query,物品塔为Key/Value
U_cross = MultiHeadAttention(query=U, key=I, value=I)  
I_cross = MultiHeadAttention(query=I, key=U, value=U)

# 然后把原始输出和cross输出拼接或加和
U_final = U + U_cross
I_final = I + I_cross

这个设计让两个塔在训练时能互相“借用”对方的信息,但又不会破坏双塔的在线推理效率——因为推理时Cross-Attention是可以预计算的。比如用户塔的向量可以提前算好存起来,物品塔的向量也是,线上只需要做点积。

2.3 损失函数:Batch Softmax + 负采样

模型用的损失是 Batch Softmax,这是推荐系统里很常见的一种做法。具体来说,对于一个batch里的每个用户,我们有一个正样本(用户真正点击/购买的物品),然后把这个batch里的其他所有物品当作负样本。

损失函数公式如下:

loss = - log( exp(sim(u, i_pos)) / ( exp(sim(u, i_pos)) + sum_{j in neg} exp(sim(u, i_neg_j)) ) )

其中 sim(u, i) 是用户向量和物品向量的余弦相似度。这里有个细节:负样本不是随机选的,而是用了batch内其他物品。这样做的好处是负样本来自真实数据分布,比随机采样更有效。但缺点是如果batch size太小,负样本多样性不够;如果太大,显存扛不住。英伟达的解决方案是:使用大batch size(比如8192)配合梯度累积,同时在每个GPU上做局部负采样,然后通过all-reduce同步全局logits。


三、训练细节:不只是调参数

3.1 数据预处理

官方开源代码里用了 Criteo Ad Click Prediction Dataset 作为示例数据,这是一个经典的CTR预估数据集,包含39个离散特征(大部分是哈希过的ID)和13个连续特征。对于双塔模型,他们把数据分成了用户侧特征和物品侧特征,具体划分如下:

  • 用户侧:用户ID(特征1)、历史点击序列(从特征2特征10中取最近10个非零值)、用户画像特征(特征11特征20)
  • 物品侧:物品ID(特征21)、物品属性(特征22~特征39)

注意:历史点击序列需要做padding,长度不足10的用0填充,并且要生成一个attention mask来忽略padding位置。

3.2 超参数设置

论文和代码里给出了详细的超参数,我直接列出来:

参数 值 说明
Embedding维度 128 所有离散特征都映射到128维
Transformer层数 4 每个塔内部4层Encoder
Attention头数 8 Multi-Head Attention的头数
Feedforward维度 512 Transformer中FFN的中间层维度
Dropout 0.1 防止过拟合
Batch size 8192 每个GPU的batch size,配合8卡就是65536
学习率 1e-3 使用Adam优化器,带warmup
Warmup步数 10000 前10000步线性增加学习率
训练步数 500000 大约在200k步时收敛
梯度裁剪 1.0 防止梯度爆炸

3.3 训练技巧

这里有几个值得注意的trick:

1. 混合精度训练:使用FP16训练,梯度用FP32累加。代码里用了NVIDIA的AMP(Automatic Mixed Precision),显存占用直接减半,训练速度提升约1.8倍。

2. 大Batch训练稳定性:当batch size达到65536时,直接用Adam可能会出现训练不稳定。他们做了两件事:一是学习率warmup(从0线性增加到1e-3),二是用了 LayerNorm 放在每个Transformer子层之前(Pre-LN结构),而不是传统的Post-LN。

3. 负采样优化:前面提到的batch内负采样,在分布式训练时有一个坑:每个GPU只看到自己那部分batch的负样本,导致负样本数量不足。他们的做法是:在计算softmax之前,通过 all-gather 操作把所有GPU上的物品向量收集到一起,然后每个GPU都基于全局的负样本计算loss。这样即使batch size=8192,8卡并行时实际负样本数量是 8192×8 = 65536。


四、推理部署:如何做到低延迟

双塔模型最大的优势就是推理快,因为用户塔和物品塔可以分开计算。英伟达这个模型在推理时做了两阶段:

第一阶段:离线预计算

  • 物品塔:所有物品的特征喂进去,算出每个物品的向量(128维),存入向量数据库(比如FAISS或Milvus)。
  • 用户塔:每个用户的特征也可以离线算好,但用户向量会随时间变化(比如历史行为更新),所以通常是按天或按小时更新。

第二阶段:在线检索

  • 当用户请求到来时,直接查用户向量(如果已缓存)或实时计算用户向量(如果用户是新用户或特征变了)。
  • 然后用用户向量去向量数据库里做ANN(近似最近邻)搜索,召回Top-K个物品。
  • 最后可以再用一个轻量级的排序模型(比如一个简单的MLP)对召回的物品重新排序。

官方给出的性能数据:在单张A100上,使用FP16推理,每秒可以处理 5000个用户请求,每个请求召回100个物品,延迟小于 5ms。这比传统的双塔模型(比如YouTube DNN)快了大约30%,主要得益于Transformer的并行计算和Cross-Attention带来的更优表示。


五、开源代码结构

英伟达把代码放在了GitHub上(搜索 NVIDIA/DeepRecommender 应该能找到),核心文件结构如下:

├── configs/                # 配置文件,yaml格式
│   ├── criteo.yaml         # Criteo数据集的配置
│   └── default.yaml        # 默认配置
├── data/                   # 数据加载相关
│   ├── dataset.py          # 定义Dataset类,处理特征工程
│   ├── collator.py         # 定义collate_fn,处理padding和mask
│   └── preprocess.py       # 数据预处理脚本(特征哈希、序列截断等)
├── models/                 # 模型定义
│   ├── two_tower.py        # 双塔主模型
│   ├── attention.py        # Cross-Attention层
│   └── embedding.py        # 稀疏特征Embedding层(带哈希映射)
├── trainer/                # 训练逻辑
│   ├── trainer.py          # 分布式训练主循环
│   └── loss.py             # Batch Softmax损失函数实现
├── inference/              # 推理代码
│   ├── export.py           # 导出ONNX/TensorRT模型
│   └── serve.py            # 在线服务(基于Triton Inference Server)
└── scripts/                # 启动脚本
    ├── train.sh            # 训练启动命令
    └── run_inference.sh    # 推理启动命令

其中 train.sh 里的启动命令大概是这样的:

python -m torch.distributed.launch --nproc_per_node=8 trainer/trainer.py \
    --config configs/criteo.yaml \
    --batch_size 8192 \
    --lr 1e-3 \
    --warmup_steps 10000 \
    --max_steps 500000 \
    --fp16

注意:--nproc_per_node=8 表示用8卡训练,如果你只有单卡,可以改成1,但batch size要相应调小(比如1024),否则显存不够。


六、实测效果与踩坑记录

我自己在Criteo数据集上跑了一下(用4张V100,batch size=4096),训练了大概3天(50万步)。几个观察:

  • 收敛速度:前10万步loss下降很快,后面就平缓了。200k步时验证集AUC已经达到0.802,比官方报告的0.805低一点点,可能是因为我batch size减半了。
  • 显存占用:单卡batch size=4096时,显存占用约11GB(FP16)。如果开FP32,直接飙到18GB炸了。所以强烈建议用混合精度。
  • Cross-Attention的效果:我做了消融实验,去掉Cross-Attention后AUC下降约0.8%(从0.802降到0.795),说明这个设计确实有用,但也不是决定性因素。对于小数据集,可能收益更小。

踩坑记录:

  1. 数据预处理:Criteo数据集的离散特征很多是哈希值,范围很大(0~2^32),直接做Embedding会导致显存爆炸。官方做法是做了 频率截断:只保留出现次数大于100的特征值,其余映射到一个共享的“UNK” embedding。这个阈值在配置文件里可以调。
  2. 序列长度:历史点击序列如果太长(比如超过50),Transformer的self-attention计算量会平方增长。他们默认截断到10,我试过20,训练时间增加了40%,AUC只提升了0.1%,不划算。
  3. 分布式训练:如果用多卡,注意数据加载时每个卡要拿到不同的数据分片。他们用了 DistributedSampler,并且在 collator 里做了全局负采样的all-gather,这部分代码有点tricky,建议直接复用官方的。

七、总结与适用场景

这个模型本质上是一个 带Cross-Attention的Transformer双塔,适合以下场景:

  • 大规模推荐系统:用户和物品数量都在百万级以上,需要离线预计算向量。
  • 需要序列建模:用户有明确的行为序列(比如浏览历史、购物车记录),传统双塔用平均池化会丢失序列信息。
  • 对延迟敏感:在线推理必须在10ms以内完成,不能上复杂的交叉模型。

不推荐使用的场景:

  • 特征非常稀疏且无序列信息:比如只有用户ID和物品ID,没有上下文,那用简单的双塔(甚至FM)就够了,没必要上Transformer。
  • 数据量很小(<10万样本):Transformer容易过拟合,不如用LightGBM或简单DNN。

最后提一句,英伟达开源的这个代码框架写得很工程化,支持多卡训练、混合精度、ONNX导出、Triton部署,可以直接拿来生产用。如果你正在做推荐系统的向量召回,值得花时间读一读代码,尤其是那个 loss.py 里的Batch Softmax实现,写得非常优雅。

好了,以上就是我对这个模型的理解。如果有哪里不对或者想问的,欢迎在评论区讨论。

编写使用方法
Markdown 格式 · Ctrl+Enter 确定
新建笔记
预览
数据表格
点击单元格编辑 · Tab 移动
A1fx
Sheet1
BIH1H2≡🔗</>
隐私提醒

取消
编辑工具
取消