You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何对接收字典输入的tf.keras组合模型使用GradientTape求导

错误根因

你遇到的报错是两个原因共同导致的:

  1. GradientTape仅能跟踪tf.Tensor类型的对象,你构造的data字典中所有值都是numpy数组,不会被梯度带记录,因此触发属性不存在的报错
  2. 你之前尝试的predict()方法返回的是numpy数组,会直接切断梯度传播链,必须直接调用模型实例获取张量类型的预测结果

正确实现代码

不需要修改super_model的任何结构,多输入Keras模型原生支持值为张量的字典输入,调整输入构造和梯度带逻辑即可:

import tensorflow as tf
import numpy as np

angles = [0] * 21
# 把字典中每个输入都转换为tf.Tensor类型
data = {
    'x1_model_input': tf.convert_to_tensor([angles[0:3]], dtype=tf.float32),
    'x2_model_input': tf.convert_to_tensor([angles[3:6]], dtype=tf.float32),
    'x3_model_input': tf.convert_to_tensor([[angles[6]]], dtype=tf.float32),
    'x4_model_input': tf.convert_to_tensor([angles[7:13]], dtype=tf.float32),
    'x5_model_input': tf.convert_to_tensor([angles[13:15]], dtype=tf.float32),
    'x6_model_input': tf.convert_to_tensor([angles[15:21]], dtype=tf.float32)
}

with tf.GradientTape() as tape:
    # 显式监听所有输入张量
    for tensor in data.values():
        tape.watch(tensor)
    pred = super_model(data)

# 得到的grads是和data同结构的字典,每个key对应输入的梯度张量
grads = tape.gradient(pred, data)

# 如果需要转成numpy格式,可单独提取处理
# x1_grad = grads['x1_model_input'].numpy()

补充说明

如果需要对原始的angles变量直接求梯度,可以把angles声明为tf.Variable后再拆分构造输入字典,梯度会直接回传到angles变量。

内容的提问来源于stack exchange,提问作者OmG

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 02:45:04