Keras GPU训练遇OOM错误,如何添加report_tensor_allocations_upon_oom到RunOptions?
解决Keras GPU训练OOM时启用张量分配报告的方法
嘿,这个问题我之前帮不少开发者排查过,其实按照错误提示的要求添加report_tensor_allocations_upon_oom的方法很直接,分两种常见的Keras训练场景来操作:
场景1:使用Keras高层API(model.fit())
如果你用的是model.fit()、model.evaluate()这类封装好的训练方法,只需要导入TensorFlow的RunOptions和RunMetadata,然后在训练时传入对应的参数即可:
import tensorflow as tf from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense # 示例模型(替换成你自己的模型) model = Sequential([ Dense(64, activation='relu', input_shape=(10,)), Dense(1) ]) model.compile(optimizer='adam', loss='mse') # 配置OOM时的张量分配报告选项 run_options = tf.compat.v1.RunOptions(report_tensor_allocations_upon_oom=True) run_metadata = tf.compat.v1.RunMetadata() # 在fit中传入配置好的参数 model.fit( x_train, y_train, epochs=10, batch_size=32, options=run_options, run_metadata=run_metadata )
这里用tf.compat.v1是因为Keras高层API对TensorFlow 2.x原生RunOptions的支持还不够完善,用兼容模块能确保稳定生效。
场景2:使用自定义训练循环(tf.GradientTape)
如果你的训练逻辑是基于tf.GradientTape的自定义循环,需要在调用模型或者被@tf.function装饰的训练步骤中传入run_options参数:
import tensorflow as tf # 自定义训练步骤 @tf.function def train_step(inputs, labels, model, loss_fn, optimizer, run_options): with tf.GradientTape() as tape: predictions = model(inputs, training=True) loss = loss_fn(labels, predictions) gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables)) return loss # 初始化OOM报告选项 run_options = tf.compat.v1.RunOptions(report_tensor_allocations_upon_oom=True) # 启动训练循环 for epoch in range(10): for x_batch, y_batch in train_dataset: # 传入run_options到训练步骤中 current_loss = train_step( x_batch, y_batch, model, loss_fn, optimizer, run_options=run_options )
生效后效果
当OOM错误再次触发时,控制台会输出所有已分配张量的详细清单,包括每个张量的形状、内存占用量、分配的代码位置等信息。你可以通过这些内容快速定位到那个占用了大量显存的“元凶”张量,进而针对性地优化(比如调整batch size、清理无用张量、改用更高效的层结构等)。
额外小提示
如果你的显存经常被一次性占满,还可以先开启TensorFlow的动态显存分配,避免程序启动就吞掉全部显存:
gpus = tf.config.list_physical_devices('GPU') if gpus: try: # 动态分配显存 for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True) logical_gpus = tf.config.list_logical_devices('GPU') print(len(gpus), "Physical GPUs,", len(logical_gpus), "Logical GPUs") except RuntimeError as e: # 动态显存分配必须在程序启动时设置 print(e)
内容的提问来源于stack exchange,提问作者dspeyer
相关产品推荐
相关产品推荐

