从测试转开发三年,我终于搞懂了 TensorFlow 2.0 的 Embedding

程序员App
2026-05-01 18:36
阅读 1796

去年双十一前两周,我们成都这边的天气刚入秋,办公室空调还开着冷风,产品那边突然甩过来一个需求:“能不能做个商品推荐模型?用户看了 A 就推荐 B,类似淘宝那种。”

我当时差点一口老血喷出来——我是从测试岗转开发的,虽然也啃过《动手学深度学习》,但真刀真枪上模型还是头一回。更要命的是,Deadline 是双十一下周一上线!运维小哥还在旁边悠悠地说:“别慌,反正你们 AI 部分挂了也不算 P0 级事故……”(我信你个鬼!)

不过吐槽归吐槽,活儿还得干。翻了翻公司历史代码库,发现之前用的是 TensorFlow 1.x,写得跟天书一样,session、placeholder 满天飞。我果断决定:上 TensorFlow 2.0。一来它默认 eager execution,调试友好;二来 Keras 集成进来了,API 清晰不少——对我们这种“半路出家”的开发者太友好了。


为什么是 TF 2.0?因为我不想再写 session.run()

记得刚转开发那会儿,看同事跑 TF 1.x 的代码,光初始化 session 就写了八行。每次 debug 都得把整个图跑一遍,改一行参数就得等三分钟。现在 TF 2.0 直接像写 Python 一样运行:

import tensorflow as tf

a = tf.constant([1, 2, 3])
b = tf.constant([4, 5, 6])
c = a + b  # 直接执行!不用 sess.run()
print(c)  # 输出: tf.Tensor([5 7 9], shape=(3,), dtype=int32)

这种即时反馈对非科班出身的我来说简直是救命稻草。特别是半夜改模型时,再也不用祈祷“这次应该没问题吧”,直接 print 就行。


Embedding:不是嵌入,是“翻译”用户行为

回到推荐系统的需求。我们的原始数据是用户点击日志:user_id, item_id, timestamp。但神经网络不能直接吃 ID 啊,总不能让模型认为 item_1001item_1000 “大”吧?

这时候 Embedding 就派上用场了。简单说,它把离散的 ID 映射成连续的向量。比如:

  • item_1000[0.2, -0.5, 0.8]
  • item_1001[-0.1, 0.3, 0.6]

这些向量在高维空间里能捕捉语义相似性——如果两个商品经常被同一群人点击,它们的 embedding 向量就会靠得很近。

TF 2.0 里用起来超简单:

# 假设我们有 10000 个商品
embedding_layer = tf.keras.layers.Embedding(
    input_dim=10000,   # 词表大小(这里就是 item 数量)
    output_dim=64      # 向量维度
)

# 输入一批 item_id,比如 [100, 205, 888]
item_ids = tf.constant([[100], [205], [888]])
embeddings = embedding_layer(item_ids)
print(embeddings.shape)  # (3, 1, 64)

注意:input_dim 必须大于等于所有可能的 ID 最大值+1,不然会报 Index out of bounds。我就在这栽过一次,凌晨两点才发现数据里混了个测试 ID 99999,而我的 input_dim 只设了 10000……


模型搭起来:用 Keras Functional API 玩点花的

我们的目标是预测用户下一个会点什么商品。典型的序列推荐问题。我参考了 YouTube DNN 的思路,但简化了结构。

输入有两个:

  • 用户最近点击的 N 个商品 ID(序列)
  • 用户静态特征(比如城市、性别)

输出:对所有商品打分,取 top-K 推荐。

def build_recommend_model(num_items, max_seq_len=10, user_feat_dim=5):
    # 输入层
    seq_input = tf.keras.Input(shape=(max_seq_len,), name='item_seq')
    user_input = tf.keras.Input(shape=(user_feat_dim,), name='user_feat')
    
    # 商品序列 embedding
    item_emb = tf.keras.layers.Embedding(num_items, 64)(seq_input)  # (batch, seq_len, 64)
    item_emb = tf.keras.layers.GlobalAveragePooling1D()(item_emb)   # (batch, 64)
    
    # 用户特征全连接
    user_dense = tf.keras.layers.Dense(32, activation='relu')(user_input)
    
    # 拼接
    concat = tf.keras.layers.concatenate([item_emb, user_dense])
    dense = tf.keras.layers.Dense(128, activation='relu')(concat)
    dense = tf.keras.layers.Dense(64, activation='relu')(dense)
    
    # 输出层:对每个商品打分
    output = tf.keras.layers.Dense(num_items, activation='softmax', name='item_probs')(dense)
    
    model = tf.keras.Model(inputs=[seq_input, user_input], outputs=output)
    return model

训练时用 sparse_categorical_crossentropy,因为标签是单个商品 ID(不是 one-hot)。


调参那些事儿:Claude 和 Cursor 救我狗命

说实话,调参过程一度让我怀疑人生。loss 下不去,AUC 在 0.55 打转(随机猜是 0.5)。正准备放弃时,我想起最近团队在试用几个 AI 编程助手。

  • Cursor:我让它分析 loss 曲线,它建议我降低 learning rate 并加 dropout。果然有效!
  • Claude:问它“为什么 embedding 维度设 64 而不是 128”,它解释说:对于中小规模数据,高维 embedding 容易过拟合,64 是经验值。

还有一次,我把 GlobalAveragePooling1D 错写成了 GlobalMaxPooling1D,模型效果暴跌。Claude 一眼就看出来:“平均池化更适合捕捉整体偏好,最大池化会过度关注个别热门商品。”

至于 Gemini?我们还没接入,听说 Google 内部用得挺猛,但国内访问不太稳。等哪天它能本地部署了,一定要试试。

说真的,这些 AI 助手不能替代思考,但能极大减少“低级错误排查时间”。对于我们这种非算法背景的开发者,简直是外挂。


效果对比:从 0.55 到 0.72 的跨越

经过几轮迭代,最终线上 A/B 测试结果如下:

模型版本 AUC 点击率提升 训练耗时(epoch)
初始版(无用户特征) 0.55 +0% 8min
加入用户特征 0.63 +12% 9min
加 dropout & 调 lr 0.72 +28% 10min

最惊喜的是,embedding 层自己学会了商品聚类!我把商品 embedding 用 t-SNE 降维可视化,发现:

  • 手机壳和贴膜靠得很近
  • 奶粉和尿不湿扎堆
  • 游戏耳机和机械键盘居然也在一块(看来我们用户都是硬核玩家)

产品经理看到图后直呼“这比规则引擎智能多了”,还主动给我们加了两天 buffer 做压测——破天荒没催进度!


给 fellow 半路转开发者的建议

  1. 别怕从简单模型开始。我一开始就想上 Transformer,结果连数据预处理都搞不定。先跑通一个能 train 的 baseline,再慢慢加 complexity。
  2. 善用 Keras 的 callbacksModelCheckpointEarlyStopping 这些能省下大量手动监控的时间。
  3. Embedding 初始化很重要。默认是随机初始化,但如果你有预训练好的向量(比如 Word2Vec),可以用 embeddings_initializer 注入。
  4. TF 2.0 的 @tf.function 是双刃剑。它能加速,但调试困难。开发阶段建议关掉,上线前再开。

写在最后

从测试转开发这三年,我越来越觉得:工程能力 + 领域理解 > 纯算法炫技。TensorFlow 2.0 的设计哲学也印证了这一点——它把复杂的图计算隐藏起来,让开发者专注业务逻辑。

现在回头看那个双十一流水线,虽然模型不算 SOTA,但它稳定跑了半年多,每天服务几十万用户。上周五下班前,运维小哥还发消息说:“你那个推荐模块,内存占用比隔壁 Go 服务还低,牛啊!”

我笑了笑,关掉电脑,骑共享单车去吃火锅了。成都的秋天,适合慢下来,也适合慢慢把代码写好。

评论 0

最热最新
暂无评论
程序员AppLv.1
0
影响力
0
文章
0
粉丝