多输入多输出模型必须用functional api(即tf.keras.model),因sequential仅支持单输入单输出线性堆叠;需分别声明input、显式调用层、用字典组织输入输出,并在compile中按名配置损失与指标。

多输入多输出模型必须用 tf.keras.Model,不能用 Sequential
因为 Sequential 只支持单输入单输出的线性堆叠,遇到多个输入张量(比如图像 + 文本特征)或多个预测目标(比如同时回归价格 + 分类类别),它会直接报错 ValueError: Input tensors must be of the same type 或更模糊的图构建失败。只有 Functional API 能显式定义分支结构和合并逻辑。
实操要点:
- 每个输入需独立声明
tf.keras.Input,带明确的shape和可选name - 所有中间层(如
Dense、Conv2D)必须显式调用,不能靠顺序隐含 - 输出端用字典组织:键名将作为训练时
y的 key,也影响model.compile中的损失函数配置 - 模型实例化必须传入
inputs=[...]和outputs={...}两个参数,缺一不可
输入数据格式必须是字典或列表,不能直接拼 NumPy 数组
喂数据时如果把多个输入强行 np.concatenate 或堆成高维数组,fit() 会提示 ValueError: Failed to convert a NumPy array to a Tensor —— 因为模型在构建时已按命名输入注册了张量结构,运行时必须严格匹配。
正确做法:
- 训练时
x传字典:{'img_input': x_img, 'text_input': x_text},key 必须与Input(name=...)一致 - 或传列表:
[x_img, x_text],但要求输入声明顺序与列表索引严格对应(不推荐,易错) - 输出
y同理:用字典{'price_pred': y_price, 'class_pred': y_class},否则损失函数找不到对应 target - 验证集
validation_data格式必须和训练集完全一致
compile() 里损失函数要按输出名对齐,权重可选但别漏掉
如果输出字典是 {'reg_out': ..., 'cls_out': ...},而 compile(loss='mse'),TensorFlow 会尝试把同一个损失套到两个输出上,大概率报 TypeError: Expected float32, got None(尤其当分类输出用了 softmax 而回归用了 linear)。
必须显式指定:
- 损失函数:用字典
{'reg_out': 'mse', 'cls_out': 'sparse_categorical_crossentropy'} - 损失权重(可选):
loss_weights={'reg_out': 1.0, 'cls_out': 0.5},用于调节多任务梯度幅度 - 监控指标也建议用字典形式,比如
metrics={'reg_out': 'mae', 'cls_out': 'accuracy'}
注意:如果某个输出不需要参与反向传播(比如只做推理用的辅助 head),可在该输出层加 trainable=False,但更稳妥的做法是在 compile 中将其损失设为 None。
保存和加载模型要用 save_weights_only=False,否则结构丢失
用 model.save('path') 是最安全的;但如果手动调用 model.save_weights() 再试图用 load_weights() 恢复,会报 ValueError: You are trying to load a weight file containing 12 layers into a model with 8 layers —— 因为多输入多输出模型的计算图包含多个入口/出口节点,仅存权重无法重建连接关系。
务必确认:
- 保存时用
tf.keras.models.save_model(model, 'my_mimo_model')或直接model.save(...) - 加载时用
tf.keras.models.load_model('my_mimo_model'),不要用load_weights - 若需跨环境部署(如转 TFLite),先用
tf.keras.models.clone_model复制结构,再载入权重,避免依赖原始 Python 定义
真正麻烦的不是写法,而是调试时发现某一路输入没进图、某路输出梯度为零——这种问题不会报错,只会让 loss 不下降。建议从单输入单输出 baseline 开始,逐步加分支,并用 model.summary() 对照每层 input/output shape 是否符合预期。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











