自定义Keras Layer中用for循环触发InaccessibleTensorError如何解决?
问题原因
- TensorFlow 图执行模式下,
for batch in tf.range(self.batch_size)会被自动转换为tf.while_loop结构,循环体内生成的张量属于独立的函数作用域,无法直接存入Python列表后在循环外访问,直接触发InaccessibleTensorError。 - 原
compute_output_shape方法依赖只有在call运行时才会赋值的self.batch_size,静态推导输出形状时也会出现异常。 - 额外性能提醒:你当前的输出维度是
3^16 = 43046721,单样本输出就有4300多万个浮点值,batch_size=10的话单批次输出就需要1.7G左右显存(按float32计算),实际运行大概率会触发显存不足错误,建议先确认你的计算逻辑是否正确,是否真的需要这么大的输出维度。
修复方案
使用tf.map_fn代替显式for循环遍历批次维度,它会自动处理循环内张量的作用域问题,同时修正输出形状推导逻辑:
import tensorflow as tf from tensorflow.keras.layers import Layer, Input from tensorflow.keras import models import numpy as np class WeightedLayer(Layer): def __init__(self, n_input, n_memb, **kwargs): super(WeightedLayer, self).__init__(**kwargs) self.n = n_input # 16 features self.m = n_memb # 3 self.output_dim = self.m ** self.n # 提前计算输出维度 def build(self, batch_input_shape): super(WeightedLayer, self).build(batch_input_shape) def call(self, input_): # 定义单样本的处理逻辑 def process_single_sample(x): # x 形状为 (16, 3) cp = x[0, :] for d in range(1, self.n): # 动态调整形状 xd = tf.reshape(x[d, :], (1,) * d + (self.m,)) c = tf.reshape(cp, (self.m,) * d + (1,)) cp = tf.matmul(c, xd) return tf.reshape(cp, (self.output_dim,)) # 对批次内所有样本并行应用处理逻辑 return tf.map_fn(process_single_sample, input_, dtype=tf.float32) def compute_output_shape(self, batch_input_shape): # 不需要依赖运行时的batch_size,直接用输入的批次维度 return tf.TensorShape((batch_input_shape[0], self.output_dim))
原有的模型调用代码不需要修改即可适配修正后的层。
你当前的输出维度
3^16超过4300万,即使修正了语法错误,实际运行也会因为显存不足无法训练。建议先把n_input调整为较小值(比如4,对应输出维度3^4=81)验证逻辑正确性后,再根据需求调整参数或者对输出做降维处理。
内容的提问来源于stack exchange,提问作者zahra moradi
相关产品推荐
相关产品推荐

