如何对接收字典输入的tf.keras组合模型使用GradientTape求导
错误根因
你遇到的报错是两个原因共同导致的:
GradientTape仅能跟踪tf.Tensor类型的对象,你构造的data字典中所有值都是numpy数组,不会被梯度带记录,因此触发属性不存在的报错- 你之前尝试的
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
相关产品推荐
相关产品推荐

