ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

使用 Spark Scala 调用 TensorFlow 2.0 训练好的模型进行分布式推断

使用 Spark Scala 调用 TensorFlow 2.0 训练好的模型进行分布式推断 使用 Spark Scala 调用 TensorFlow 2.0 训练好的模型进行分布式推断【免费下载链接】eat_tensorflow2_in_30_daysTensorflow2.0 is delicious, just eat it! 项目地址: https://gitcode.com/gh_mirrors/ea/eat_tensorflow2_in_30_days本篇技术指南聚焦于如何利用 TensorFlow for Java 官方接口在 SparkScala集群中加载训练好的 TensorFlow 模型并完成分布式模型推断。文章完整覆盖了从 protobuf 模型文件准备、Maven 依赖配置到 Driver 端单机调试、RDD 与 DataFrame 两种分布式调用方式的全流程读者学完后能够在实际工程项目中让训练好的 Keras 模型在成百上千台机器上并行执行预测。〇背景与总体思路TensorFlow 训练好的模型以原生方式保存成 protobufSavedModel文件后可以有多种部署形态通过 tensorflow-js 在浏览器中运行、通过 tensorflow-lite 在移动端运行、通过 tensorflow-serving 提供 HTTP API 服务参见 6-6,使用tensorflow-serving部署模型.md以及通过TensorFlow for Java在 Java 或 Spark(Scala) 中直接加载模型进行预测。如果使用 PySpark实现起来相对简单——每个 executor 上用 Python 加载模型分别预测即可但工程上为了性能考虑通常使用的是 Scala 版本的 Spark。因此本文采用TensorFlow for Java这一官方 Java 接口在 Spark 中调用训练好的模型借助 Spark 的分布式计算能力让模型在集群的成百上千个 executor 上分布式并行执行推断。在 Spark(Scala) 中调用 TensorFlow 模型进行预测需要完成以下五个步骤准备 protobuf 模型文件用 tf.keras 训练模型并导出为 SavedModel 格式创建 Spark(Scala) 项目在项目中添加 Java 版本的 TensorFlow 对应 jar 包依赖Driver 端加载模型调试在 Driver 上先加载模型并单机验证预测正确通过 RDD 在 executor 上加载模型利用广播机制把模型分发到各 executor 上分布式推断通过 DataFrame 在 executor 上加载模型将推断方法注册为 SparkSQL UDF 后分布式推断。一准备 protobuf 模型文件我们使用 tf.keras 训练一个简单的线性回归模型并保存成 protobuf 文件。该模型的导出结果也保留在仓库中可查看 data/linear_model 目录下的saved_model.pb与variables权重目录。import tensorflow as tf from tensorflow.keras import models,layers,optimizers ## 样本数量 n 800 ## 生成测试用数据集 X tf.random.uniform([n,2],minval-10,maxval10) w0 tf.constant([[2.0],[-1.0]]) b0 tf.constant(3.0) Y Xw0 b0 tf.random.normal([n,1],mean 0.0,stddev 2.0) # 表示矩阵乘法,增加正态扰动 ## 建立模型 tf.keras.backend.clear_session() inputs layers.Input(shape (2,),name inputs) #设置输入名字为inputs outputs layers.Dense(1, name outputs)(inputs) #设置输出名字为outputs linear models.Model(inputs inputs,outputs outputs) linear.summary() ## 使用fit方法进行训练 linear.compile(optimizerrmsprop,lossmse,metrics[mae]) linear.fit(X,Y,batch_size 8,epochs 100) tf.print(w ,linear.layers[1].kernel) tf.print(b ,linear.layers[1].bias) ## 将模型保存成pb格式文件 export_path ./data/linear_model/ version 1 #后续可以通过版本号进行模型版本迭代与管理 linear.save(export_pathversion, save_formattf)上述代码中值得注意的要点输入层通过name inputs显式命名输出层通过name outputs命名这两个名字在导出后会体现在 SavedModel 的签名SignatureDef中是后续在 Scala 端feed/fetch张量的关键依据训练完成后linear.layers[1].kernel与linear.layers[1].bias即 Dense 层的权重与偏置理想情况下分别接近[[2.0],[-1.0]]与3.0linear.save(export_pathversion, save_formattf)以 SavedModel 格式导出version目录用于模型版本迭代与管理这也是后续 Spark 端加载路径的一部分。导出成功后可以用如下命令查看模型目录结构!ls {export_pathversion}结果包含saved_model.pb、variables权重文件目录与可能的assets目录与仓库中 data/linear_model 的目录结构一致。接着使用saved_model_cli查看模型的签名信息# 查看模型文件相关信息 !saved_model_cli show --dir {export_pathstr(version)} --all输出中会包含类似如下的关键信息也是后续 Scala 代码中feed和fetch的张量名来源MetaGraphDef with tag-set: serve contains the following SignatureDefs: signature_def[serving_default]: The given SavedModel SignatureDef contains the following input(s): inputs[inputs] tensor_info: dtype: DT_FLOAT shape: (-1, 2) name: serving_default_inputs:0 The given SavedModel SignatureDef contains the following output(s): outputs[outputs] tensor_info: dtype: DT_FLOAT shape: (-1, 1) name: StatefulPartitionedCall:0 Method name is: tensorflow/serving/predict模型文件信息中这些标红的部分tag-set、serving_default_inputs:0、StatefulPartitionedCall:0等都是后面 Scala 调用时会用到的serveSavedModelBundle.load 时的 tag 参数一般固定传serveserving_default_inputs:0sess.runner().feed(...)时指定的输入张量名StatefulPartitionedCall:0fetch(...)时指定的输出张量名。二创建 Spark(Scala) 项目并添加 jar 包依赖如果使用 Maven 管理项目需要添加如下 jar 包依赖!-- https://mvnrepository.com/artifact/org.tensorflow/tensorflow -- dependency groupIdorg.tensorflow/groupId artifactIdtensorflow/artifactId version1.15.0/version /dependency由于 Java 版 TensorFlow 的 Maven 坐标体系该 jar 包还会传递依赖org.tensorflow:libtensorflow和org.tensorflow:libtensorflow_jni两个 jar 包分别负责 Java API 与本地 JNI 实现。如果构建环境无法从 Maven 中央仓库拉取也可以直接从 Maven 仓库页面下载org.tensorflow.tensorflow的 jar 包以及其依赖的org.tensorflow.libtensorflow和org.tensorflow.libtensorflow_jni的 jar 包放到项目中对应版本号为 1.15.0。需要说明的是Java 版 TensorFlow 的 API 形态仍与 TensorFlow 1.x 的静态计算图模式一致需要通过SavedModelBundle加载模型、建立Session然后指定 feed 与 fetch 的 tensor 再run()这一点与 TensorFlow 2.x 的 Eager 模式 API 有差异使用时要格外注意。三在 Driver 端加载 TensorFlow 模型调试本节示范代码在 Jupyter Notebook 中演示需要安装toreeApache Toree以支持 Spark(Scala) 内核。先在 Driver 端加载模型做单机验证确保模型文件、张量名、输入输出格式全部正确后再进入分布式阶段可以大幅降低排查问题的成本。import scala.collection.mutable.WrappedArray import org.{tensorflowtf} //注load函数的第二个参数一般都是“serve”可以从模型文件相关信息中找到 val bundle tf.SavedModelBundle .load(/Users/liangyun/CodeFiles/eat_tensorflow2_in_30_days/data/linear_model/1,serve) //注在java版本的tensorflow中还是类似tensorflow1.0中静态计算图的模式需要建立Session, 指定feed的数据和fetch的结果, 然后 run. //注如果有多个数据需要喂入可以连续使用多个feed方法 //注输入必须是float类型 val sess bundle.session() val x tf.Tensor.create(Array(Array(1.0f,2.0f),Array(2.0f,3.0f))) val y sess.runner().feed(serving_default_inputs:0, x) .fetch(StatefulPartitionedCall:0).run().get(0) val result Array.ofDimFloat(0).toInt,y.shape()(1).toInt) y.copyTo(result) if(x ! null) x.close() if(y ! null) y.close() if(sess ! null) sess.close() if(bundle ! null) bundle.close() result输出如下Array(Array(3.019596), Array(3.9878292))代码中的关键点SavedModelBundle.load(path, serve)的第二个参数serve与模型导出时的 tag-set 对应可从saved_model_cli show --all的输出中查到tf.Tensor.create(Array(Array(1.0f,2.0f),Array(2.0f,3.0f)))构造输入张量注意输入必须是 float 类型与模型导出时DT_FLOAT对应feed(serving_default_inputs:0, x)与fetch(StatefulPartitionedCall:0)的张量名都来自模型签名信息如果有多个输入需要喂入可以连续使用多个feed方法预测结果通过y.copyTo(result)拷贝到与输出 shape 一致的 Float 二维数组中x.close()、y.close()、sess.close()、bundle.close()依次释放资源防止内存泄漏这在长期运行的 Spark 作业中尤为重要。四通过 RDD 在 executor 上加载 TensorFlow 模型Driver 端调试通过后下一步是把模型分发到集群的各个 executor 上。这里采用 Spark 的广播机制broadcast先在 Driver 端加载SavedModelBundle再通过sc.broadcast将其发送到各个 executor之后在mapPartitions中取出广播值批量构造张量并调用模型完成推断。import org.apache.spark.sql.SparkSession import scala.collection.mutable.WrappedArray import org.{tensorflowtf} val spark SparkSession .builder() .appName(TfRDD) .enableHiveSupport() .getOrCreate() val sc spark.sparkContext //在Driver端加载模型 val bundle tf.SavedModelBundle .load(/Users/liangyun/CodeFiles/master_tensorflow2_in_20_hours/data/linear_model/1,serve) //利用广播将模型发送到executor上 val broads sc.broadcast(bundle) //构造数据集 val rdd_data sc.makeRDD(List(Array(1.0f,2.0f),Array(3.0f,5.0f),Array(6.0f,7.0f),Array(8.0f,3.0f))) //通过mapPartitions调用模型进行批量推断 val rdd_result rdd_data.mapPartitions(iter { val arr iter.toArray val model broads.value val sess model.session() val x tf.Tensor.create(arr) val y sess.runner().feed(serving_default_inputs:0, x) .fetch(StatefulPartitionedCall:0).run().get(0) //将预测结果拷贝到相同shape的Float类型的Array中 val result Array.ofDimFloat(0).toInt,y.shape()(1).toInt) y.copyTo(result) result.iterator }) rdd_result.take(5) bundle.close输出如下Array(Array(3.019596), Array(3.9264367), Array(7.8607616), Array(15.974984))这里有两个值得注意的工程细节使用mapPartitions而不是map在每个分区内先取出该分区的所有数据iter.toArray只建立一次Session、只创建一次输入张量进行批量推断避免每条记录都重复创建 Session 带来的巨大开销这也是分布式并行推断性能的关键广播模型而非直接序列化SavedModelBundle本身是不可序列化的大对象通过sc.broadcast在每个 executor 上只保存一份拷贝供该 executor 上的所有任务共享这与在每个 executor 上加载模型的语义一致。五通过 DataFrame 在 executor 上加载 TensorFlow 模型除了在 RDD 上调用模型进行分布式推断我们也可以在 DataFrame 数据上调用 TensorFlow 模型。主要思路是将推断方法注册成为一个 SparkSQL 函数UDF然后就可以像使用普通 SQL 函数一样在 DataFrame 上完成预测。import org.apache.spark.sql.SparkSession import scala.collection.mutable.WrappedArray import org.{tensorflowtf} object TfDataFrame extends Serializable{ def main(args:Array[String]):Unit { val spark SparkSession .builder() .appName(TfDataFrame) .enableHiveSupport() .getOrCreate() val sc spark.sparkContext import spark.implicits._ val bundle tf.SavedModelBundle .load(/Users/liangyun/CodeFiles/master_tensorflow2_in_20_hours/data/linear_model/1,serve) val broads sc.broadcast(bundle) //构造预测函数并将其注册成sparkSQL的udf val tfpredict (features:WrappedArray[Float]) { val bund broads.value val sess bund.session() val x tf.Tensor.create(Array(features.toArray)) val y sess.runner().feed(serving_default_inputs:0, x) .fetch(StatefulPartitionedCall:0).run().get(0) val result Array.ofDimFloat(0).toInt,y.shape()(1).toInt) y.copyTo(result) val y_pred result(0)(0) y_pred } spark.udf.register(tfpredict,tfpredict) //构造DataFrame数据集将features放到一列中 val dfdata sc.parallelize(List(Array(1.0f,2.0f),Array(3.0f,5.0f),Array(7.0f,8.0f))).toDF(features) dfdata.show //调用sparkSQL预测函数增加一个新的列作为y_preds val dfresult dfdata.selectExpr(features,tfpredict(features) as y_preds) dfresult.show bundle.close } }运行主函数TfDataFrame.main(Array())输出如下---------- | features| ---------- |[1.0, 2.0]| |[3.0, 5.0]| |[7.0, 8.0]| ---------- ------------------- | features| y_preds| ------------------- |[1.0, 2.0]| 3.019596| |[3.0, 5.0]|3.9264367| |[7.0, 8.0]| 8.828995| -------------------DataFrame 方式的几个关键设计点预测函数tfpredict接收WrappedArray[Float]DataFrame 中 ArrayType 列在 Scala 端的对应类型内部从广播变量取回模型、构造 shape 为(1, 2)的输入张量完成单样本预测返回标量Floatspark.udf.register(tfpredict, tfpredict)将函数注册为 SparkSQL UDF之后即可在selectExpr中直接以tfpredict(features)的形式调用selectExpr(features,tfpredict(features) as y_preds)在保留原特征列的同时追加一列预测结果完全沿用 SparkSQL 生态便于与既有的数据仓库、Hive 表分析流程无缝衔接这也是enableHiveSupport()存在的意义。六总结与扩展思路以上我们分别在 Spark 的RDD数据结构和DataFrame数据结构上实现了调用一个 tf.keras 实现的线性回归模型进行分布式模型推断。在实际使用时只需把加载路径、输入输出张量名换成你自己的模型信息即可流程完全一致。在此基础上稍作修改就可以用 Spark 调用训练好的各种复杂的神经网络模型CNN、RNN、Transformer 等进行分布式推断。事实上TensorFlow 并不仅仅适合实现神经网络其底层的计算图语言可以表达各种数值计算过程利用其丰富的低阶 API可以在 TensorFlow 2.0 上实现任意机器学习模型如广义线性模型、树模型、因子分解机等结合 4-5,AutoGraph和tf.Module.md 中介绍的tf.Module便捷封装功能可以将训练好的任意机器学习模型导出成模型文件并在 Spark 上分布式调用执行——这为工程应用提供了巨大的想象空间。几点重要的使用前提与注意事项供读者参考模型格式本文方案要求模型以 SavedModelprotobuf格式导出仓库中的 data/linear_model 即为可直接用于 Spark 端加载的完整示例模型目录版本匹配Java 版 TensorFlow 1.15.0 的 APISavedModelBundle、Session、runner与 2.x 的 Eager 模式 API 不同代码中所有 feed/fetch 的张量名都必须以saved_model_cli show --all输出的实际签名为准分布式语义通过sc.broadcast广播模型、使用mapPartitions批量推断、以及把推断注册为 SparkSQL UDF是保证大规模分布式推断性能与易用性的三个核心手段。【免费下载链接】eat_tensorflow2_in_30_daysTensorflow2.0 is delicious, just eat it! 项目地址: https://gitcode.com/gh_mirrors/ea/eat_tensorflow2_in_30_days创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表