资讯动态

第46课:TensorFlow|多输入多输出复杂模型设计【业务多维度数据融合建模】

发布时间:2026/9/14 17:40:41 来源:尧图企业网站定制
文章目录1. 课前导读1.1 本节课学习目标1.2 知识重难点1.3 学习前置条件1.4 学完可掌握能力1.5 行业应用场景2. 核心理论精讲2.1 多输入多输出模型的定义2.2 数据融合策略2.3 多任务损失加权2.4 处理异构输入2.5 联合训练与任务间信息共享2.6 评估与推理3. 环境搭建与工具配置4. 代码实战教学4.1 生成模拟电商数据4.2 数据标准化4.3 定义模型函数式API4.4 编译模型多损失、多指标4.5 数据加载tf.data4.6 训练与评估4.7 特征融合的改进注意力融合5. 案例实操演练5.1 使用真实数据集模拟实际预处理5.2 定义更完善的输入分支使用预训练图像特征5.3 使用不确定性加权可学习损失权重5.4 多输入模型的可视化5.5 推理与部署6. 常见坑点与排错总结6.1 数据对齐坑点6.2 模型定义坑点6.3 多损失编译坑点6.4 训练坑点7. 知识点总结 课后作业7.1 核心知识点梳理7.2 基础作业7.3 进阶实操作业7.4 思考拓展题《TensorFlow2.x: 深度学习入门到高阶实战教程》系列课程导航1. 课前导读1.1 本节课学习目标理解多输入多输出模型的业务需求数据异构、任务关联。掌握Keras函数式API定义多输入多输出模型的方法。学习数据融合策略早期融合输入层拼接、中期融合特征层拼接/加权、晚期融合决策层融合。掌握多任务学习中损失函数的加权策略固定权重、动态调节、不确定性加权。能够处理不同输入类型数值、类别、图像、文本的预处理和编码。完成电商推荐案例输入用户画像数值类别、商品图像CNN特征、商品描述文本嵌入输出点击率二分类和转化概率回归。1.2 知识重难点类别内容重点函数式API的多输入多输出建模特征融合层的设计多损失编译与联合训练处理不同尺寸输入难点图像与文本特征的时间与空间对齐多任务损失的权重动态调节梯度冲突的缓解易混淆点多输入模型的数据输入格式列表或字典多输出中不同任务使用不同激活函数sigmoid vs linear损失权重的设置位置1.3 学习前置条件已掌握TensorFlow基础模型构建Sequential和函数式API。了解CNN、Embedding层的基本使用第21、29课。熟悉数据预处理和tf.data第16、37课。1.4 学完可掌握能力独立构建处理异构数据的复杂模型满足工业多模态需求。实现多任务学习提高模型泛化能力。能够针对不同任务设计合适的损失函数和评估指标。1.5 行业应用场景电商推荐融合用户、商品、交互特征预测点击率和转化率。自动驾驶多传感器输入相机、激光雷达、毫米波雷达多输出目标检测、车道线分割。医疗诊断结合影像、病历、基因组数据预测疾病类型和严重程度。内容推荐图文视频多模态预测观看时长和互动率。2. 核心理论精讲2.1 多输入多输出模型的定义多输入模型接收多个张量作为输入例如用户特征ID、年龄、性别等结构化数据商品图像像素矩阵商品描述文本序列多输出模型产生多个预测输出例如点击率二分类转化概率回归评分多分类Keras函数式API允许定义有向无环图每个输入和输出都可以独立命名在编译时可以为每个输出指定不同的损失函数和权重。2.2 数据融合策略早期融合输入级融合将原始特征预处理后直接拼接成一个长向量然后输入共享网络。适用于特征维度相差不大且语义对齐的情况。中期融合特征级融合不同模态的数据先分别提取中间特征如CNN、RNN然后通过拼接、加权求和、注意力机制等方式融合再输入后续网络。这是最常见的方法。晚期融合决策级融合每个模态单独建立模型进行预测最终对结果进行投票或加权平均。用于模型差异大且可独立训练的场景。2.3 多任务损失加权总损失[\mathcal{L} \sum_{i} w_i \mathcal{L}_i]固定权重根据任务重要性手动设置。动态权重方法Uncertainty Weighting通过可学习的噪声参数 (\sigma_i)损失为 (\mathcal{L} \sum_i \frac{1}{2\sigma_i^2} \mathcal{L}_i \log \sigma_i)。噪声大则自动降低权重。GradNorm调整权重使各任务梯度范数相近。2.4 处理异构输入数值特征标准化后直接输入。类别特征Embedding层或One-hot编码。图像特征使用预训练CNN如EfficientNet提取特征向量或端到端训练。文本特征用预训练词嵌入RNN/Transformer或直接用BERT提取句向量。2.5 联合训练与任务间信息共享多任务学习的优势辅助任务可提供归纳偏置提升主任务泛化能力。但需注意任务之间可能存在冲突导致负迁移。可通过调整损失权重、梯度滤波等方法缓解。2.6 评估与推理多输出模型评估时需分别报告每个任务的指标。推理时可以仅计算所需任务的输出。3. 环境搭建与工具配置沿用第45课环境额外安装pillow用于图像处理若未安装。conda activate tf213 pipinstallpillow项目结构multimodal/ ├── data/ # 模拟数据生成脚本 ├── models/ # 保存模型 ├── train.py └── predict.py导入importtensorflowastfimportnumpyasnpimportpandasaspdimportmatplotlib.pyplotaspltfromtensorflow.kerasimportlayers,models,optimizers,losses,metricsfromsklearn.preprocessingimportStandardScaler,LabelEncoder4. 代码实战教学4.1 生成模拟电商数据为简化演示生成合成数据用户画像年龄、性别、商品图片32x32灰度、商品描述10个词的序列标签点击率0/1、转化概率0~1。np.random.seed(42)num_samples5000# 用户特征agenp.random.randint(18,70,sizenum_samples)gendernp.random.choice([0,1],sizenum_samples)# 0女1男user_featuresnp.column_stack([age,gender])# 商品图片模拟随机噪声imagesnp.random.rand(num_samples,32,32,1).astype(np.float32)# 商品描述模拟整数序列长度10词汇表大小100text_seqnp.random.randint(1,100,size(num_samples,10))# 标签clicknp.random.binomial(1,0.3,sizenum_samples)# 点击率30%conversionnp.random.uniform(0,1,sizenum_samples)*click# 只有点击后才可能转化# 划分splitint(0.8*num_samples)train_useruser_features[:split]train_imgimages[:split]train_texttext_seq[:split]train_clickclick[:split]train_convconversion[:split]test_useruser_features[split:]test_imgimages[split:]test_texttext_seq[split:]test_clickclick[split:]test_convconversion[split:]4.2 数据标准化scalerStandardScaler()train_userscaler.fit_transform(train_user)test_userscaler.transform(test_user)4.3 定义模型函数式API我们将构建三个输入分支用户特征分支全连接网络图像分支简单CNN文本分支Embedding LSTM然后融合特征输出两个任务。# 输入层user_inputlayers.Input(shape(2,),nameuser_input)image_inputlayers.Input(shape(32,32,1),nameimage_input)text_inputlayers.Input(shape(10,),nametext_input,dtypetf.int32)# 用户分支user_denselayers.Dense(32,activationrelu)(user_input)user_denselayers.Dropout(0.2)(user_dense)# 图像分支xlayers.Conv2D(32,3,activationrelu)(image_input)xlayers.MaxPooling2D(2)(x)xlayers.Conv2D(64,3,activationrelu)(x)xlayers.GlobalAveragePooling2D()(x)image_featureslayers.Dropout(0.2)(x)# 文本分支embedding_layerlayers.Embedding(input_dim101,output_dim32,input_length10)text_embembedding_layer(text_input)text_lstmlayers.LSTM(32,return_sequencesFalse)(text_emb)text_featureslayers.Dropout(0.2)(text_lstm)# 融合拼接所有分支的特征concat_featureslayers.concatenate([user_dense,image_features,text_features],namefusion)# 共享层sharedlayers.Dense(64,activationrelu)(concat_features)sharedlayers.Dropout(0.3)(shared)# 输出分支click_outputlayers.Dense(1,activationsigmoid,nameclick_output)(shared)conv_outputlayers.Dense(1,activationlinear,nameconv_output)(shared)# 构建模型modeltf.keras.Model(inputs[user_input,image_input,text_input],outputs[click_output,conv_output])model.summary()4.4 编译模型多损失、多指标model.compile(optimizeroptimizers.Adam(0.001),loss{click_output:losses.BinaryCrossentropy(),conv_output:losses.MeanSquaredError()},loss_weights{click_output:1.0,conv_output:0.5},# 转化任务权重较低metrics{click_output:[metrics.BinaryAccuracy(),metrics.AUC()],conv_output:[metrics.MeanAbsoluteError()]})4.5 数据加载tf.databatch_size64train_datasettf.data.Dataset.from_tensor_slices(({user_input:train_user,image_input:train_img,text_input:train_text},{click_output:train_click,conv_output:train_conv})).batch(batch_size).shuffle(1000).prefetch(tf.data.AUTOTUNE)test_datasettf.data.Dataset.from_tensor_slices(({user_input:test_user,image_input:test_img,text_input:test_text},{click_output:test_click,conv_output:test_conv})).batch(batch_size).prefetch(tf.data.AUTOTUNE)4.6 训练与评估historymodel.fit(train_dataset,epochs30,validation_datatest_dataset,verbose1,callbacks[tf.keras.callbacks.EarlyStopping(patience3)])# 评估resultsmodel.evaluate(test_dataset,verbose0)print(Test results:)forname,valinzip(model.metrics_names,results):print(f{name}:{val:.4f})4.7 特征融合的改进注意力融合替代简单的拼接可以使用门控注意力机制动态加权各模态特征。defattention_fusion(features_list,num_features):# features_list: list of tensors each shape (batch, feat_dim)# 简单实现通过一个小网络学习权重concatlayers.concatenate(features_list)# (batch, sum_dim)attention_weightslayers.Dense(len(features_list),activationsoftmax)(concat)# 加权求和weighted_sumtf.zeros_like(features_list[0])fori,finenumerate(features_list):weighted_sumattention_weights[:,i:i1]*freturnweighted_sum5. 案例实操演练案例电商多模态点击率和转化率联合预测5.1 使用真实数据集模拟实际预处理假设我们已有用户行为日志、商品图片URL、商品标题文本。下面演示完整的数据加载与预处理流水线。# 模拟读取数据importpandasaspd dfpd.DataFrame({user_id:np.random.randint(1,1000,10000),age:np.random.randint(18,70,10000),gender:np.random.choice([M,F],10000),item_id:np.random.randint(1,500,10000),click:np.random.binomial(1,0.3,10000),conversion:np.random.uniform(0,1,10000)})# 对类别特征做标签编码le_genderLabelEncoder()df[gender_code]le_gender.fit_transform(df[gender])# 数值特征标准化scalerStandardScaler()df[[age]]scaler.fit_transform(df[[age]])# 模拟图片和文本预处理此处省略假设已转换为张量5.2 定义更完善的输入分支使用预训练图像特征# 图像特征提取使用MobileNetV2冻结fromtensorflow.keras.applicationsimportMobileNetV2 img_inputlayers.Input(shape(224,224,3),nameimage_raw)base_modelMobileNetV2(include_topFalse,weightsimagenet,poolingavg)base_model.trainableFalseimage_embeddingbase_model(img_input)image_featureslayers.Dense(128,activationrelu)(image_embedding)# 文本分支使用预训练词向量简化text_inputlayers.Input(shape(20,),dtypetf.int32,nametext_seq)embeddinglayers.Embedding(5000,64)(text_input)text_lstmlayers.LSTM(64)(embedding)text_featureslayers.Dense(128,activationrelu)(text_lstm)# 用户特征分支user_id_inputlayers.Input(shape(1,),nameuser_id)user_embedlayers.Embedding(1000,32)(user_id_input)user_embedlayers.Flatten()(user_embed)user_demo_inputlayers.Input(shape(2,),nameuser_demo)# age,genderuser_denselayers.Dense(32,activationrelu)(user_demo_input)user_concatlayers.concatenate([user_embed,user_dense])user_featureslayers.Dense(64,activationrelu)(user_concat)# 融合所有特征fusionlayers.concatenate([user_features,image_features,text_features])sharedlayers.Dense(128,activationrelu)(fusion)sharedlayers.Dropout(0.3)(shared)click_outlayers.Dense(1,activationsigmoid,nameclick)(shared)conv_outlayers.Dense(1,activationsigmoid,nameconversion)(shared)# 转换概率multi_modeltf.keras.Model(inputs[user_id_input,user_demo_input,img_input,text_input],outputs[click_out,conv_out])multi_model.compile(optimizeradam,loss{click:binary_crossentropy,conversion:binary_crossentropy},metrics{click:accuracy,conversion:accuracy})5.3 使用不确定性加权可学习损失权重classUncertaintyWeightedLoss(tf.keras.losses.Loss):def__init__(self,num_tasks,initial_log_var0.0,**kwargs):super().__init__(**kwargs)self.num_tasksnum_tasks self.log_varstf.Variable(initial_log_var*tf.ones(num_tasks),trainableTrue,namelog_vars)defcall(self,y_true,y_pred):# 假设y_pred是列表包含每个任务的预测loss0.0foriinrange(self.num_tasks):task_losstf.reduce_mean(tf.keras.losses.binary_crossentropy(y_true[i],y_pred[i]))losstf.exp(-self.log_vars[i])*task_lossself.log_vars[i]returnloss# 注意上述简化实际使用需适配模型输出结构或使用自定义训练循环。5.4 多输入模型的可视化tf.keras.utils.plot_model(multi_model,to_filemultimodal_model.png,show_shapesTrue,show_layer_namesTrue)5.5 推理与部署# 预测单个样本sample{user_id:np.array([[123]]),user_demo:np.array([[0.5,0.2]]),image_raw:np.random.rand(1,224,224,3).astype(np.float32),text_seq:np.random.randint(1,5000,size(1,20))}click_prob,conv_probmulti_model.predict(sample)print(fClick probability:{click_prob[0][0]:.4f}, Conversion probability:{conv_prob[0][0]:.4f})6. 常见坑点与排错总结6.1 数据对齐坑点坑1多输入数据的样本顺序不一致导致模型训练时错位。解决使用tf.data.Dataset时确保所有输入和标签来自同一切片。坑2不同输入批次大小不匹配如图像和文本长度批次内自动补齐需确保批次维度一致。6.2 模型定义坑点坑3使用concatenate时各分支特征维度必须匹配除了最后一维。坑4函数式API中层重用需注意是否共享权重。若要共享Embedding应定义一次并多次调用。6.3 多损失编译坑点坑5编译时loss字典的键必须与输出层名称完全一致。若未指定Keras会为未指定的输出自动分配默认损失可能错误。坑6loss_weights中的权重不会自动归一化需根据任务重要性手动设置。6.4 训练坑点坑7不同任务收敛速度不同可能导致训练不稳定。可先训练主任务一段时间再开启辅助任务。坑8梯度冲突多个任务的梯度方向不一致导致优化困难。可采用梯度归一化或PCGrad。7. 知识点总结 课后作业7.1 核心知识点梳理函数式API定义多输入多输出模型的标准方法。数据融合早期、中期、晚期融合的适用场景。多任务损失固定权重、不确定性加权。异构输入处理数值、类别、图像、文本的特征提取。7.2 基础作业修改电商案例增加一个辅助任务预测商品点击后是否加入购物车构建三输出模型。尝试使用注意力机制融合三个模态特征对比与简单拼接的性能差异。实现不确定性加权损失并在训练中观察可学习参数log_vars的变化。7.3 进阶实操作业任务多模态情感分析文本语音使用CMU-MOSI数据集文本和语音特征输出情感得分[-3,3]回归。构建双输入模型文本分支BERT或LSTM、语音分支1D CNN。输出为情感得分回归和情感类别分类多任务。使用不同的融合策略拼接、门控对比性能。7.4 思考拓展题在多任务学习中如果某个任务的数据量远小于其他任务应该如何调整损失权重以避免该任务被忽略为什么有时在多输入模型中需要为不同输入使用不同的归一化参数在多模态模型中如何维护这些参数梯度冲突问题具体表现是什么请查阅PCGradProjecting Conflicting Gradients方法并简述其原理。下一课预告深度学习项目全流程规范——我们将系统讲解从需求分析、数据标注、模型开发、测试到上线的完整项目生命周期管理。《TensorFlow2.x: 深度学习入门到高阶实战教程》系列课程导航去订阅第一部分基础入门1-10 课第二部分神经网络核心11-25 课第三部分进阶网络与框架高阶26-40 课第四部分企业实战与项目落地41-50 课 感谢您耐心阅读到这里 如果本文对您有所启发欢迎 点赞 收藏 分享给更多需要的伙伴。️ 期待在评论区看到您的想法, 共同进步。 关注我持续获取更多干货内容 我们下篇文章见

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价