TensorFlow 2.0 入门:从婚庆推荐算法到生产部署的实战手记

爬虫不想爬
2025-12-24 04:05
阅读 1663

上周五晚上十点半,我一边啃着冷掉的麦当劳,一边在 VS Code 里调试一个 TensorFlow 模型——没错,这就是备婚+刷题+加班三重奏的日常。上个月刚跟未婚夫敲定婚礼日期,结果领导又甩来一个新需求:“能不能用 AI 给用户推荐个性化的婚礼主题?” 我翻了个白眼,但转念一想:反正最近在准备跳槽,不如借机深入学下 TensorFlow 2.0,顺便把算法、模型部署这些硬核技能补一补。

说干就干。作为一个重度依赖 ChatGPT 和 Claude 的程序媛(别笑,谁还没个“AI搭子”?),我很快搭起了一个原型。今天这篇笔记,就记录下从零入门 TF 2.0 的过程,尤其是如何把一个看似“高大上”的机器学习任务,落地到我们团队那套老旧的 Spring Boot 后端 + 前端 JavaScript 生态里。


为什么是 TensorFlow 2.0?

先说背景:我们公司是个中型 SaaS 平台,主要做婚庆全流程管理。前端用 React,后端是 Spring Boot,数据库是 MySQL + Redis。产品经理想要一个“智能灵感推荐”功能——比如用户上传了婚纱照,系统能自动推荐匹配的场地风格、花艺搭配等。

这本质上是个多模态推荐问题,但 MVP 阶段我们先简化成:基于用户历史行为(点击、收藏、浏览时长)预测其对某类婚礼主题的兴趣得分。典型的监督学习 + 排序算法场景。

选 TensorFlow 2.0 而不是 PyTorch?原因很现实:

  • 团队运维熟悉 TF Serving
  • 公司已有 TF 1.x 遗留模型,升级比迁移成本低
  • Keras 集成后开发体验更友好(对我这种非算法岗太友好了)

核心概念快速过一遍

TF 2.0 最大的变化就是 Eager Execution 默认开启,再也不用手动建图、跑 session 了!写代码像写 NumPy 一样自然:

import tensorflow as tf

# 看,直接打印张量值!
x = tf.constant([1, 2, 3])
print(x)  # 直接输出 [1 2 3],不用 sess.run()

几个关键概念必须搞清:

概念 说明 类比
tf.Tensor 多维数组,数据的基本单位 类似 NumPy array
tf.Variable 可训练参数(如权重) 模型的“记忆”
tf.function 将 Python 函数编译为图,加速执行 类似 JIT 编译
tf.data 高效数据管道 数据加载的瑞士军刀

我踩的第一个坑就是数据加载。一开始直接用 pandas.read_csv() 塞进模型,结果内存爆炸。后来改用 tf.data.Dataset,配合 cache()prefetch(),训练速度直接起飞:

def make_dataset(csv_path, batch_size=32):
    dataset = tf.data.experimental.make_csv_dataset(
        csv_path,
        batch_size=batch_size,
        label_name='interest_score',
        num_epochs=1,
        shuffle=True
    )
    return dataset.cache().prefetch(tf.data.AUTOTUNE)

构建你的第一个推荐模型

我们的数据很简单:每行是一个 (user_id, theme_id, click_count, view_duration, ...),标签是 interest_score(0~1 之间的浮点数)。

模型结构我选了经典的 Wide & Deep —— Wide 部分处理特征交叉(比如“25岁女性 + 海边婚礼”),Deep 部分捕捉非线性关系。TF 2.0 用 Keras Functional API 写起来超清爽:

def build_wide_and_deep_model(user_vocab, theme_vocab):
    # 输入层
    user_input = tf.keras.Input(shape=(), name='user_id', dtype=tf.string)
    theme_input = tf.keras.Input(shape=(), name='theme_id', dtype=tf.string)
    
    # Wide 部分:交叉特征
    crossed = tf.keras.layers.experimental.preprocessing.HashedCrossing(
        num_bins=10000, output_mode='one_hot'
    )([user_input, theme_input])
    wide_output = tf.keras.layers.Dense(1)(crossed)
    
    # Deep 部分:嵌入 + MLP
    user_embedding = tf.keras.layers.StringLookup(vocabulary=user_vocab)(user_input)
    user_embedded = tf.keras.layers.Embedding(len(user_vocab), 16)(user_embedding)
    
    theme_embedding = tf.keras.layers.StringLookup(vocabulary=theme_vocab)(theme_input)
    theme_embedded = tf.keras.layers.Embedding(len(theme_vocab), 16)(theme_embedding)
    
    deep = tf.keras.layers.concatenate([user_embedded, theme_embedded])
    deep = tf.keras.layers.Dense(128, activation='relu')(deep)
    deep = tf.keras.layers.Dense(64, activation='relu')(deep)
    deep_output = tf.keras.layers.Dense(1)(deep)
    
    # 合并输出
    output = tf.keras.layers.Add()([wide_output, deep_output])
    output = tf.keras.layers.Activation('sigmoid')(output)  # 输出 0~1
    
    model = tf.keras.Model(inputs=[user_input, theme_input], outputs=output)
    return model

注意这里用了 StringLookup 做字符串到索引的映射——告别手动 One-Hot!这也是 TF 2.0 的预处理层优势:模型自带特征工程,导出后无需额外处理。


训练与调优:那些让我想砸键盘的时刻

训练初期 loss 死活不降,一度怀疑人生。后来发现是标签没归一化(兴趣得分原始范围 0~100)。加上 tf.keras.utils.normalize() 后,一个 epoch 就收敛了。

另一个坑是过拟合。我们的数据量不大(约 50 万条),模型稍微复杂就 memorize。解决方案:

  • 加 Dropout(0.3 起步)
  • Early Stopping(监控 val_loss)
  • Label Smoothing(把 1 改成 0.95,0 改成 0.05)
model.compile(
    optimizer='adam',
    loss=tf.keras.losses.BinaryCrossentropy(label_smoothing=0.1),
    metrics=['AUC']
)

callbacks = [
    tf.keras.callbacks.EarlyStopping(patience=3, restore_best_weights=True),
    tf.keras.callbacks.ReduceLROnPlateau(factor=0.5, patience=2)
]

history = model.fit(
    train_dataset,
    validation_data=val_dataset,
    epochs=20,
    callbacks=callbacks
)

最终 AUC 达到 0.87,在业务可接受范围内(产品经理居然没挑刺,感动哭了)。


导出模型:让 Spring Boot 能调用

模型训练完只是开始,如何集成到现有系统才是难点。我们的后端是 Java Spring Boot,总不能让 Java 直接跑 Python 吧?

方案一:TF Serving
方案二:导出 SavedModel,用 TensorFlow Java API 调用

我选了方案二——少维护一个服务。导出超简单:

model.save('wedding_recommend_model')

然后在 Spring Boot 项目里加依赖:

<dependency>
    <groupId>org.tensorflow</groupId>
    <artifactId>tensorflow-core-platform</artifactId>
    <version>0.4.0</version>
</dependency>

Java 调用代码(简化版):

try (Graph graph = Graph.importGraphDef(Files.readAllBytes(modelPath))) {
    try (Session session = new Session(graph)) {
        Tensor<?> userIdTensor = Tensors.create(new String[]{"user_123"});
        Tensor<?> themeIdTensor = Tensors.create(new String[]{"beach_theme"});
        
        Tensor<?> result = session.runner()
            .feed("user_id", userIdTensor)
            .feed("theme_id", themeIdTensor)
            .fetch("StatefulPartitionedCall") // 注意:TF 2.x 默认输出节点名
            .run().get(0);
            
        float score = result.floatValue();
        return score;
    }
}

⚠️ 注意:TF 2.x 导出的 SavedModel 默认输出节点名是 StatefulPartitionedCall,别被坑了!


前端怎么用?JavaScript 也能玩!

虽然主要推理在后端,但前端偶尔也需要轻量级预测(比如实时反馈)。TF.js 完美支持加载 SavedModel:

// 在 React 组件里
import * as tf from '@tensorflow/tfjs';

useEffect(() => {
  const loadModel = async () => {
    const model = await tf.loadGraphModel('/models/wedding_recommend_model/model.json');
    const score = model.predict({
      user_id: tf.tensor(['user_123']),
      theme_id: tf.tensor(['beach_theme'])
    }).dataSync()[0];
    setRecommendScore(score);
  };
  loadModel();
}, []);

不过要注意:TF.js 对 SavedModel 的兼容性有限,复杂模型建议用后端。


工具链:我的效率倍增器

整个过程中,这些工具救了我狗命:

  • TensorBoard:可视化训练曲线,比 print loss 高级一万倍
  • Weights & Biases:自动记录超参和指标,方便对比实验
  • Docker:封装 TF 环境,避免“在我机器上能跑”
  • Claude:帮我解释报错信息,比如 InvalidArgumentError: indices[0] = 0 is not in [0, 0) 这种反人类错误

特别是 TensorBoard,一行代码搞定:

tensorboard_callback = tf.keras.callbacks.TensorBoard(log_dir="./logs")
model.fit(..., callbacks=[tensorboard_callback])

然后 tensorboard --logdir=./logs,浏览器打开就能看 loss/AUC 曲线、计算图、甚至 embedding 可视化!


心得:算法工程师 ≠ 炼丹师

很多人觉得搞 AI 就是调参炼丹,但实际工作中,数据质量 > 模型复杂度 > 超参优化。我们 80% 的时间花在:

  • 清洗脏数据(比如用户误点导致的异常高分)
  • 构建合理的负样本(随机采样 vs. hard negative mining)
  • 设计评估指标(线上 AB 测试比离线 AUC 更重要)

另外,工程化能力决定模型上限。一个再牛的算法,如果无法部署、无法监控、无法迭代,就是废铁。这也是为什么我坚持要把模型塞进 Spring Boot —— 落地才有价值。


结语:边备婚边学 AI,痛并快乐着

现在这个推荐模型已经上线两周,CTR 提升了 18%,产品经理终于请我喝了杯喜茶(虽然还是没加薪)。下周就要去试婚纱了,但晚上还得刷 LeetCode 准备跳槽面试……程序员的生活啊,永远在 deadline 和梦想之间反复横跳。

如果你也在用 TensorFlow 2.0 做业务落地,欢迎留言交流!特别是怎么把 TF 模型塞进 Java 生态的坑,我真的踩过太多……(以及,有没有人知道婚庆行业的推荐系统还有什么骚操作?求分享!)

最后送大家一句我贴在显示器上的话:“模型会过拟合,但热爱不会。”

—— 一个在 bug 和捧花之间努力平衡的程序媛

评论 0

最热最新
暂无评论
爬虫不想爬Lv.1
0
影响力
0
文章
0
粉丝