我在字节做后端,为什么还要学TensorFlow 2.0
最近组里要搞一个内部工具,自动解析运维工单截图,提取关键信息分类归档。以前靠正则加规则硬堆,截图字体或布局一变就崩。领导拍板上深度学习,我这个写Go和Java的后端被迫啃起TensorFlow 2.0。
一开始很抗拒,但2.0把Keras收编为官方高级API,默认Eager Execution,写起来像普通Python,能直接print张量值,不用再开Session,门槛降了不少。
几个基础概念:张量可理解为多维数组,类似Numpy的ndarray,但能上GPU。变量是训练中要更新的参数。自动微分GradientTape最舒服,前向计算写在上下文里,调tape.gradient(loss, variables)即可出梯度。
import tensorflow as tf
x = tf.Variable(3.0)
with tf.GradientTape() as tape:
y = x ** 2 + 2 * x + 1
dy_dx = tape.gradient(y, x)
print(dy_dx.numpy()) # 8.0
跑通时感觉比写反射还简单。
需求数据集约六千张截图,十二个类别。选型上PyTorch社区强,但公司推理服务多用TF Serving部署,对齐成本更低。tf.data.Dataset管道配合map、batch、prefetch很方便。坑:prefetch必须放管道最后,否则异步加载失效,我因顺序写反训练慢近一倍,排查一晚上才发现是官方示例缩进误导。
模型用预训练MobileNetV2做特征提取,上面接两层全连接,冻结base model只训练新增层。验证集准确率约91%,内部工具够用。最大感受是学习率和batch size很敏感:batch size从32改到64,同样学习率下loss直接震荡发散,学习率调小到四分之一才稳住。
最近注意到Computer Use,让模型像人一样操作电脑界面。如果工单工具能进化到这种程度,直接让模型自己截图、识别、填结果,前景很大。国内大模型也在卷多模态和工具调用,后续可试试用工单分类的文本部分做增强,看准确率能否再拉高。
总的来说,TensorFlow 2.0对后端很友好,不懂底层图优化也能快速搭出能跑的模型。真正调好部署路还长,但至少深夜写代码时不会再想砸电脑了。

评论 0