【发布时间】:2019-04-05 21:56:47
【问题描述】:
我在 Keras 中设置了一个自动编码器。我希望能够根据预定的“精度”向量对输入向量的特征进行加权。这个连续值向量与输入的长度相同,每个元素都在[0, 1]范围内,对应输入元素的置信度,其中1表示完全置信,0表示不置信。
我为每个示例都有一个精度向量。
我已经定义了一个考虑到这个精度向量的损失。在这里,低置信度特征的重建被降低权重。
def MAEpw_wrapper(y_prec):
def MAEpw(y_true, y_pred):
return K.mean(K.square(y_prec * (y_pred - y_true)))
return MAEpw
我的问题是精度张量 y_prec 取决于批次。我希望能够根据当前批次更新y_prec,以便每个精度向量与其观察正确关联。
我做了以下事情:
global y_prec
y_prec = K.variable(P[:32])
这里的P 是一个 numpy 数组,其中包含所有精度向量,其索引对应于示例。我将y_prec 初始化为具有32 批大小的正确形状。然后我定义以下DataGenerator:
class DataGenerator(Sequence):
def __init__(self, batch_size, y, shuffle=True):
self.batch_size = batch_size
self.y = y
self.shuffle = shuffle
self.on_epoch_end()
def on_epoch_end(self):
self.indexes = np.arange(len(self.y))
if self.shuffle == True:
np.random.shuffle(self.indexes)
def __len__(self):
return int(np.floor(len(self.y) / self.batch_size))
def __getitem__(self, index):
indexes = self.indexes[index * self.batch_size: (index+1) * self.batch_size]
# Set precision vector.
global y_prec
new_y_prec = K.variable(P[indexes])
y_prec = K.update(y_prec, new_y_prec)
# Get training examples.
y = self.y[indexes]
return y, y
我的目标是在生成批处理的同一函数中更新y_prec。这似乎正在按预期更新y_prec。然后我定义我的模型架构:
dims = [40, 20, 2]
model2 = Sequential()
model2.add(Dense(dims[0], input_dim=64, activation='relu'))
model2.add(Dense(dims[1], input_dim=dims[0], activation='relu'))
model2.add(Dense(dims[2], input_dim=dims[1], activation='relu', name='bottleneck'))
model2.add(Dense(dims[1], input_dim=dims[2], activation='relu'))
model2.add(Dense(dims[0], input_dim=dims[1], activation='relu'))
model2.add(Dense(64, input_dim=dims[0], activation='linear'))
最后,我编译并运行:
model2.compile(optimizer='adam', loss=MAEpw_wrapper(y_prec))
model2.fit_generator(DataGenerator(32, digits.data), epochs=100)
digits.data 是一个 numpy 观察数组。
但是,这最终会定义单独的图表:
StopIteration: Tensor("Variable:0", shape=(32, 64), dtype=float32_ref) must be from the same graph as Tensor("Variable_4:0", shape=(32, 64), dtype=float32_ref).
我已经搜索了 SO 来解决我的问题,但我没有找到任何工作。任何有关如何正确执行此操作的帮助表示赞赏。
【问题讨论】:
标签: python tensorflow keras loss-function