Keras图像二进制哈希模型内存占用过高问题求助
我正在构建一个输出图像二进制哈希的网络,采用两个VGG-19并行训练,传入图像对让相似图像哈希更接近、相异图像哈希差异更大,硬件是Geforce GTX 1080(12GB显存)。
目前遇到的问题是:处理100/3200张图像后内存耗尽,htop显示32GB内存及交换空间全被占用,停用_calculate_binary函数则问题消失,改用numpy数组也存在同样问题。
相关代码与变量说明
训练代码片段
# positive images prim_model.fit(data[index][0], temp_label, epochs=1, verbose=0) sec_model.fit(data[index][i], temp_label, epochs=1, verbose=0) model_vars._calculate_binary([prim_model, sec_model], [index, 0, index, i]) # negative images prim_model.fit(data[index][0], temp_label, epochs=1, verbose=0) sec_model.fit(data[index][i], temp_label, epochs=1, verbose=0) model_vars._calculate_binary([prim_model, sec_model], [index, 0, index, i])
model_vars存储的关键变量
U = a tensor of shape (64, 3200) where 64 is binary bits of output and 3200 is number of images and U represents the output of all the images from prim_model(first model) V = a tensor of same shape which holds output of sec_model B = a tensor of shape(16, 3200) storing the final binary values of output
_calculate_binary函数核心逻辑
for index in xrange(3200): Q = some_calculations (a 2-d tenor of shape(16, 3200) Q_star_c = tf.reshape(tf.transpose(Q)[:, (index)], [self.kbit, 1] ) #extracting a column from Q U_star_c = #A column extracted from U V_star_c = #A column extracted from V self.U_1 = tf.concat( [ self.U[:, 0:index], self.U[:, index+1: self.total_images]] , axis=1) #Removing the column extracted above from the original now the size of U_1 is (16, 3199) self.V_1 = #same as above self.B = #slicing the original B tensor #Now doing some calcultion to calculate B_star_c (binary value of index'th image B_star_c = tf.scalar_mul(-1, \ tf.sign(tf.add(tf.matmul(tf.scalar_mul(2, self.B), \ tf.add(tf.matmul(self.U_1, U_star_c, transpose_a=True), tf.matmul(self.V_1, V_star_c, transpose_a=True)) ) , Q_star_c)) ) #Now combining the final generated binary column to the original Binary tensor making the size of B to be (16, 3200) again self.B = tf.concat( [ self.B[:, 0:index], tf.concat( [B_star_c, self.B[:, index:self.total_images]], axis=1)], axis=1)
从你的代码和现象来看,内存持续增长的核心原因几乎可以锁定在_calculate_binary函数里的张量操作链式累积和不必要的全量张量复制上,尤其是在循环里反复对self.U、self.V、self.B做切片拼接操作,再加上TensorFlow默认的计算图保留机制,会导致内存里堆积大量中间张量和历史计算节点。下面给你几个针对性的解决办法:
1. 避免循环内反复拼接张量,改用索引更新而非全量复制
你现在每次循环都通过tf.concat来移除某一列再重新拼接,这会创建大量新的张量副本(TensorFlow的张量默认是不可变的)。对于B张量的更新,完全可以用tf.tensor_scatter_nd_update直接更新指定位置的列,不需要反复拼接:
# 替换原来的拼接更新逻辑 # 先准备更新的索引:index对应的列位置 update_indices = [[bit_idx, index] for bit_idx in range(self.kbit)] # 将B_star_c转为匹配的形状(self.kbit,) B_star_c_flat = tf.reshape(B_star_c, (-1,)) # 直接更新B张量的指定列 self.B = tf.tensor_scatter_nd_update(self.B, update_indices, B_star_c_flat)
同样,self.U_1和self.V_1不需要每次都创建新的拼接张量,而是在计算时直接排除指定索引的列,比如用矩阵全量乘积减去对应列的贡献,避免复制大张量:
# 替代原来创建U_1的逻辑,避免全量拼接 full_u_product = tf.matmul(self.U, U_star_c, transpose_a=True) index_u_product = tf.matmul(tf.expand_dims(self.U[:, index], axis=1), U_star_c, transpose_a=True) u_product = full_u_product - index_u_product
2. 用tf.function封装函数,避免计算图节点累积
如果没有用tf.function装饰_calculate_binary,每次循环都会在计算图里添加新节点,导致计算图越来越大。把函数用tf.function装饰,让TensorFlow编译成高效的图操作,减少内存开销:
@tf.function(reduce_retracing=True) def _calculate_binary(self, models, indices): # 先把全局张量读入局部变量,操作后再一次性赋值回去 U_local = self.U V_local = self.V B_local = self.B for index in tf.range(3200): # ... 你的计算逻辑 ... # 用tensor_scatter_nd_update更新B_local # ... # 最后一次性更新全局变量 self.U = U_local self.V = V_local self.B = B_local
3. 分批处理图像,不要一次性加载全量数据
一次性存储3200张图像的特征张量已经占用不少内存,再加上循环中间张量很容易爆内存。可以把数据分成小批次处理,每批处理完就释放对应内存:
# 把3200张图像分成50批,每批64张 batch_size = 64 for batch_start in range(0, 3200, batch_size): batch_end = batch_start + batch_size # 加载当前批次的图像特征 U_batch = prim_model.predict(data[batch_start:batch_end]) V_batch = sec_model.predict(data[batch_start:batch_end]) # 计算当前批次的B值 B_batch = self._calculate_binary_batch(U_batch, V_batch) # 更新全局B张量的对应位置 update_indices = [[bit_idx, batch_start + idx] for bit_idx in range(16) for idx in range(batch_size)] self.B = tf.tensor_scatter_nd_update(self.B, update_indices, tf.reshape(B_batch, (-1,))) # 手动释放当前批次内存 del U_batch, V_batch, B_batch tf.keras.backend.clear_session()
4. 优化Q张量的计算逻辑
如果每次循环都重新计算全量(16,3200)的Q张量,会带来巨大内存开销。可以检查Q的计算逻辑:如果Q是全局统计量,就提前计算一次复用;如果Q只和当前index相关,就只计算对应列,不需要生成全量张量。
5. 用tf.Variable存储核心张量
如果self.U、self.V、self.B是普通张量,建议改用tf.Variable存储,因为Variable在更新时会复用内存空间,而普通张量每次赋值都会创建新副本:
# 初始化时替换为Variable self.U = tf.Variable(initial_U_tensor, trainable=False) self.V = tf.Variable(initial_V_tensor, trainable=False) self.B = tf.Variable(initial_B_tensor, trainable=False)
内容的提问来源于stack exchange,提问作者Deepak Sharma

