news 2026/9/14 17:40:25

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
第46课:TensorFlow|多输入多输出复杂模型设计【业务多维度数据融合建模】

文章目录

    • 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 定义模型(函数式API)
      • 4.4 编译模型(多损失、多指标)
      • 4.5 数据加载(tf.data)
      • 4.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,LabelEncoder

4. 代码实战教学

4.1 生成模拟电商数据

为简化演示,生成合成数据:用户画像(年龄、性别)、商品图片(32x32灰度)、商品描述(10个词的序列),标签:点击率(0/1)、转化概率(0~1)。

np.random.seed(42)num_samples=5000# 用户特征age=np.random.randint(18,70,size=num_samples)gender=np.random.choice([0,1],size=num_samples)# 0女1男user_features=np.column_stack([age,gender])# 商品图片(模拟随机噪声)images=np.random.rand(num_samples,32,32,1).astype(np.float32)# 商品描述(模拟整数序列,长度10,词汇表大小100)text_seq=np.random.randint(1,100,size=(num_samples,10))# 标签click=np.random.binomial(1,0.3,size=num_samples)# 点击率30%conversion=np.random.uniform(0,1,size=num_samples)*click# 只有点击后才可能转化# 划分split=int(0.8*num_samples)train_user=user_features[:split]train_img=images[:split]train_text=text_seq[:split]train_click=click[:split]train_conv=conversion[:split]test_user=user_features[split:]test_img=images[split:]test_text=text_seq[split:]test_click=click[split:]test_conv=conversion[split:]

4.2 数据标准化

scaler=StandardScaler()train_user=scaler.fit_transform(train_user)test_user=scaler.transform(test_user)

4.3 定义模型(函数式API)

我们将构建三个输入分支:

  • 用户特征分支:全连接网络
  • 图像分支:简单CNN
  • 文本分支:Embedding + LSTM

然后融合特征,输出两个任务。

# 输入层user_input=layers.Input(shape=(2,),name='user_input')image_input=layers.Input(shape=(32,32,1),name='image_input')text_input=layers.Input(shape=(10,),name='text_input',dtype=tf.int32)# 用户分支user_dense=layers.Dense(32,activation='relu')(user_input)user_dense=layers.Dropout(0.2)(user_dense)# 图像分支x=layers.Conv2D(32,3,activation='relu')(image_input)x=layers.MaxPooling2D(2)(x)x=layers.Conv2D(64,3,activation='relu')(x)x=layers.GlobalAveragePooling2D()(x)image_features=layers.Dropout(0.2)(x)# 文本分支embedding_layer=layers.Embedding(input_dim=101,output_dim=32,input_length=10)text_emb=embedding_layer(text_input)text_lstm=layers.LSTM(32,return_sequences=False)(text_emb)text_features=layers.Dropout(0.2)(text_lstm)# 融合:拼接所有分支的特征concat_features=layers.concatenate([user_dense,image_features,text_features],name='fusion')# 共享层shared=layers.Dense(64,activation='relu')(concat_features)shared=layers.Dropout(0.3)(shared)# 输出分支click_output=layers.Dense(1,activation='sigmoid',name='click_output')(shared)conv_output=layers.Dense(1,activation='linear',name='conv_output')(shared)# 构建模型model=tf.keras.Model(inputs=[user_input,image_input,text_input],outputs=[click_output,conv_output])model.summary()

4.4 编译模型(多损失、多指标)

model.compile(optimizer=optimizers.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.data)

batch_size=64train_dataset=tf.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_dataset=tf.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 训练与评估

history=model.fit(train_dataset,epochs=30,validation_data=test_dataset,verbose=1,callbacks=[tf.keras.callbacks.EarlyStopping(patience=3)])# 评估results=model.evaluate(test_dataset,verbose=0)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)# 简单实现:通过一个小网络学习权重concat=layers.concatenate(features_list)# (batch, sum_dim)attention_weights=layers.Dense(len(features_list),activation='softmax')(concat)# 加权求和weighted_sum=tf.zeros_like(features_list[0])fori,finenumerate(features_list):weighted_sum+=attention_weights[:,i:i+1]*freturnweighted_sum

5. 案例实操演练

案例:电商多模态点击率和转化率联合预测

5.1 使用真实数据集(模拟实际预处理)

假设我们已有用户行为日志、商品图片URL、商品标题文本。下面演示完整的数据加载与预处理流水线。

# 模拟读取数据importpandasaspd df=pd.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_gender=LabelEncoder()df['gender_code']=le_gender.fit_transform(df['gender'])# 数值特征标准化scaler=StandardScaler()df[['age']]=scaler.fit_transform(df[['age']])# 模拟图片和文本预处理(此处省略,假设已转换为张量)

5.2 定义更完善的输入分支(使用预训练图像特征)

# 图像特征提取(使用MobileNetV2,冻结)fromtensorflow.keras.applicationsimportMobileNetV2 img_input=layers.Input(shape=(224,224,3),name='image_raw')base_model=MobileNetV2(include_top=False,weights='imagenet',pooling='avg')base_model.trainable=Falseimage_embedding=base_model(img_input)image_features=layers.Dense(128,activation='relu')(image_embedding)# 文本分支(使用预训练词向量简化)text_input=layers.Input(shape=(20,),dtype=tf.int32,name='text_seq')embedding=layers.Embedding(5000,64)(text_input)text_lstm=layers.LSTM(64)(embedding)text_features=layers.Dense(128,activation='relu')(text_lstm)# 用户特征分支user_id_input=layers.Input(shape=(1,),name='user_id')user_embed=layers.Embedding(1000,32)(user_id_input)user_embed=layers.Flatten()(user_embed)user_demo_input=layers.Input(shape=(2,),name='user_demo')# age,genderuser_dense=layers.Dense(32,activation='relu')(user_demo_input)user_concat=layers.concatenate([user_embed,user_dense])user_features=layers.Dense(64,activation='relu')(user_concat)# 融合所有特征fusion=layers.concatenate([user_features,image_features,text_features])shared=layers.Dense(128,activation='relu')(fusion)shared=layers.Dropout(0.3)(shared)click_out=layers.Dense(1,activation='sigmoid',name='click')(shared)conv_out=layers.Dense(1,activation='sigmoid',name='conversion')(shared)# 转换概率multi_model=tf.keras.Model(inputs=[user_id_input,user_demo_input,img_input,text_input],outputs=[click_out,conv_out])multi_model.compile(optimizer='adam',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_var=0.0,**kwargs):super().__init__(**kwargs)self.num_tasks=num_tasks self.log_vars=tf.Variable(initial_log_var*tf.ones(num_tasks),trainable=True,name='log_vars')defcall(self,y_true,y_pred):# 假设y_pred是列表,包含每个任务的预测loss=0.0foriinrange(self.num_tasks):task_loss=tf.reduce_mean(tf.keras.losses.binary_crossentropy(y_true[i],y_pred[i]))loss+=tf.exp(-self.log_vars[i])*task_loss+self.log_vars[i]returnloss# 注意:上述简化,实际使用需适配模型输出结构,或使用自定义训练循环。

5.4 多输入模型的可视化

tf.keras.utils.plot_model(multi_model,to_file='multimodal_model.png',show_shapes=True,show_layer_names=True)

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_prob=multi_model.predict(sample)print(f"Click 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 基础作业

  1. 修改电商案例,增加一个辅助任务(预测商品点击后是否加入购物车),构建三输出模型。
  2. 尝试使用注意力机制融合三个模态特征,对比与简单拼接的性能差异。
  3. 实现不确定性加权损失,并在训练中观察可学习参数log_vars的变化。

7.3 进阶实操作业

任务:多模态情感分析(文本+语音)

  • 使用CMU-MOSI数据集(文本和语音特征,输出情感得分[-3,3]回归)。
  • 构建双输入模型:文本分支(BERT或LSTM)、语音分支(1D CNN)。
  • 输出为情感得分(回归)和情感类别(分类)多任务。
  • 使用不同的融合策略(拼接、门控)对比性能。

7.4 思考拓展题

  1. 在多任务学习中,如果某个任务的数据量远小于其他任务,应该如何调整损失权重以避免该任务被忽略?

  2. 为什么有时在多输入模型中需要为不同输入使用不同的归一化参数?在多模态模型中如何维护这些参数?

  3. 梯度冲突问题具体表现是什么?请查阅PCGrad(Projecting Conflicting Gradients)方法并简述其原理。


下一课预告:深度学习项目全流程规范——我们将系统讲解从需求分析、数据标注、模型开发、测试到上线的完整项目生命周期管理。


🔗《TensorFlow2.x: 深度学习入门到高阶实战教程》系列课程导航

去订阅

第一部分:基础入门(1-10 课)
第二部分:神经网络核心(11-25 课)
第三部分:进阶网络与框架高阶(26-40 课)
第四部分:企业实战与项目落地(41-50 课)

🌟 感谢您耐心阅读到这里!
💡 如果本文对您有所启发欢迎:
👍 点赞📌 收藏 📤 分享给更多需要的伙伴。
🗣️ 期待在评论区看到您的想法, 共同进步。
🔔 关注我,持续获取更多干货内容~
🤗 我们下篇文章见~

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/14 17:39:22

CGA评估体系在老年照护分级中的应用与实践

1. 项目背景与核心价值老年人能力评估与照护分级是养老服务体系中的关键环节。CGA(Comprehensive Geriatric Assessment)作为一套系统的老年综合评估工具,在临床和养老机构中发挥着越来越重要的作用。去年参与某地级市养老机构标准化建设时&a…

作者头像 李华
网站建设 2026/9/14 17:37:15

沈阳高精度绿地数据处理:解压、投影拼接与拓扑质检全攻略

简介:沈阳高精度绿地数据是一套基于WGS1984坐标系、以栅格图像为核心的GIS地理信息资料,适用于城市规划、环境监测、生态研究和GIS教学等需要精细分析绿地覆盖的从业者与学生。压缩包共7个文件,涵盖tif栅格主数据、tfw世界文件、xml元数据与辅…

作者头像 李华
网站建设 2026/9/14 17:36:09

LangGraph与AutoGen多智能体框架选型指南

1. 多智能体框架的技术选型困境在构建基于大语言模型(LLM)的复杂应用时,开发团队常常面临一个关键决策:如何在LangGraph和AutoGen这两个主流多智能体框架之间做出选择?这个问题看似简单,实则涉及技术架构、…

作者头像 李华
网站建设 2026/9/14 17:36:09

Buzz 离线语音转文字:10 分钟跑通本地 Whisper 部署与字幕制作

Buzz 离线语音转文字:10 分钟跑通本地 Whisper 部署与字幕制作 【免费下载链接】buzz Buzz transcribes and translates audio offline on your personal computer. Powered by OpenAIs Whisper. 项目地址: https://gitcode.com/GitHub_Trending/buz/buzz Bu…

作者头像 李华
网站建设 2026/9/14 17:35:59

Flutter鸿蒙应用负载与功耗问题定位实战指南

做鸿蒙上的Flutter性能问题,最头疼的不是代码本身,而是“不知道去哪看数据”。同样一个App,在Android上跑得好好的,换到鸿蒙上就出现发热、掉帧、后台耗电异常,而且排查工具链跟以前完全不一样,adb那套指令…

作者头像 李华