半难批处理孪生神经网络训练时RAM线性增长问题求助
问题描述
在Windows系统通过Docker运行Ubuntu虚拟机,使用Jupyter Notebook训练带半难批处理的孪生神经网络时,RAM占用随epochs线性增长,峰值达60GB后训练从GPU切换至CPU,导致速度骤降。即使在RTX 4090上减小batch size,Docker环境下的RAM问题仍未解决;改用siamese_net.fit_generator()替代siamese_net.fit()也无改善。
超参数
batch_size = 16 epochs = 10 steps_per_epoch = int(x_train.shape[0]/batch_size) val_steps = int(x_test.shape[0]/batch_size) alpha = 0.2 num_hard = int(batch_size * 0.5) # Number of semi-hard triplet examples in the batch lr = 0.00006 optimiser = 'Adam' emb_size = 10
难批处理生成函数
def create_hard_batch(batch_size, num_hard, split="train"): # Adjust for color images: width x height x depth x_anchors = np.zeros((batch_size, x_train_w, x_train_h, x_train_d)) x_positives = np.zeros((batch_size, x_train_w, x_train_h, x_train_d)) x_negatives = np.zeros((batch_size, x_train_w, x_train_h, x_train_d)) if split == "train": data = x_train data_y = y_train else: data = x_test data_y = y_test # Generate num_hard number of hard examples: hard_batches = [] batch_losses = [] rand_batches = [] # Get some random batches for i in range(0, batch_size): hard_batches.append(create_batch(1, split)) # Adjusted create_batch function for color images A_emb = embedding_model.predict(hard_batches[i][0]) P_emb = embedding_model.predict(hard_batches[i][1]) N_emb = embedding_model.predict(hard_batches[i][2]) # Compute d(A, P) - d(A, N) for each selected batch batch_losses.append(np.sum(np.square(A_emb-P_emb),axis=1) - np.sum(np.square(A_emb-N_emb),axis=1)) # Sort batch_loss by distance, highest first, and keep num_hard of them hard_batch_selections = [x for _, x in sorted(zip(batch_losses,hard_batches), key=lambda x: x[0])] hard_batches = hard_batch_selections[:num_hard] # Get batch_size - num_hard number of random examples num_rand = batch_size - num_hard for i in range(0, num_rand): rand_batch = create_batch(1, split) # Adjusted create_batch function for color images rand_batches.append(rand_batch) selections = hard_batches + rand_batches for i in range(0, len(selections)): x_anchors[i] = selections[i][0] x_positives[i] = selections[i][1] x_negatives[i] = selections[i][2] return [x_anchors, x_positives, x_negatives]
数据生成器
def data_generator(batch_size=16, num_hard=50, split="train"): while True: x = create_hard_batch(batch_size, num_hard, split) y = np.zeros((batch_size, 3*emb_size)) yield x, y
训练函数
siamese_history = siamese_net.fit_generator( data_generator(batch_size, num_hard), steps_per_epoch=steps_per_epoch, epochs=epochs, verbose=1, callbacks=callbacks, workers=0, validation_data=data_generator(batch_size, num_hard, split="test"), validation_steps=val_steps);
解决方案
1. 清理生成过程中的内存泄漏点
- 手动清理临时变量:在
create_hard_batch函数末尾,使用完hard_batches、batch_losses、rand_batches后,添加del hard_batches, batch_losses, rand_batches,并调用import gc; gc.collect()强制触发垃圾回收,避免临时张量持续占用内存。 - 优化
predict调用:每次调用embedding_model.predict会产生新的张量,可将多次predict合并为一次批量计算,或者在每次predict后调用tf.keras.backend.clear_session()清理Keras会话中的残留资源。
2. 修正生成器参数错误
data_generator的默认参数num_hard=50与超参数中num_hard=8不符,会导致生成远超预期的难样本,额外占用内存。需将生成器调用的num_hard统一为超参数定义的值,避免冗余数据生成。
3. 优化批处理生成逻辑
- 减少临时batch生成量:当前为筛选
num_hard个难样本,生成了batch_size个临时batch,可改为仅生成num_hard*2个临时batch,筛选出最符合条件的样本,减少不必要的内存占用。 - 复用数组内存:预先分配
x_anchors、x_positives、x_negatives数组,每次生成batch时直接覆盖数据,而非重复创建新数组。
4. Docker环境内存管控
- 限制容器内存配额:启动Docker容器时添加
--memory=40g --memory-swap=40g参数,强制限制容器内存使用,避免系统内存耗尽触发GPU转CPU。 - 关闭内存交换:开启
--oom-kill-disable=false,当内存不足时终止进程(需配合内存限制使用),防止系统陷入低效的内存交换状态。
5. Keras训练流程优化
- 改用
fit替代fit_generator:fit_generator已被弃用,fit原生支持生成器输入,且内存管理更高效,修改训练代码为:
siamese_history = siamese_net.fit( data_generator(batch_size, num_hard), steps_per_epoch=steps_per_epoch, epochs=epochs, verbose=1, callbacks=callbacks, validation_data=data_generator(batch_size, num_hard, split="test"), validation_steps=val_steps )
- 添加内存清理回调:自定义Keras回调,在每个batch训练完成后清理内存:
import tensorflow as tf import gc class MemoryCleanupCallback(tf.keras.callbacks.Callback): def on_batch_end(self, batch, logs=None): tf.keras.backend.clear_session() gc.collect() callbacks.append(MemoryCleanupCallback())
内容的提问来源于stack exchange,提问作者Ahmed Altunkaya
相关产品推荐
相关产品推荐

