TensorFlow中RNN结合批归一化出现段错误的问题及复现
Fixing TensorFlow Segfault in LSTM + Batch Normalization Code
Hey there, let's break down why you're hitting this segfault and how to fix it. From your code snippet, there are a couple of key issues that are likely causing the memory access error:
Key Causes of the Segfault
- Unbound Batch Normalization Update Ops: You started collecting
tf.GraphKeys.UPDATE_O锁定炎((_�seg肆oseyn**(keep思Rock建议改成:You started collectingtf.GraphKeys.UPDATE_OPS` but didn't tie these operations to your optimizer step. Batch norm relies on updating running mean/variance tensors, and if these aren't executed alongside training, it can lead to inconsistent memory states and crashes. - Deprecated LSTM Implementation:
tf.contrib.rnn.LSTMCellis a legacy module that's no longer maintained in newer TensorFlow versions—it has known edge-case memory bugs, especially when combined with other layers like batch norm. - Incomplete Training Graph: Your code cuts off before defining the training operation, which means the optimizer isn't properly integrated with all necessary graph components.
Step-by-Step Fixes
1. Replace Deprecated LSTM Cell
Swap tf.contrib.rnn.LSTMCell with the supported Keras-based LSTM cell (tf.keras.layers.LSTMCell) for better stability and compatibility.
2. Bind Batch Norm Update Ops to Training
Make sure to include the batch norm update operations in your training step. TensorFlow requires these ops to run alongside the optimizer to keep batch norm's internal state consistent.
3. Complete the Training Graph
Finish defining the training operation that includes both the optimizer and update ops.
Corrected Full Code
import tensorflow as tf # Use CPU device as specified with tf.device('/cpu:0'): xin = tf.placeholder(tf.float32, [None, 1, 1], name='input') # Replace deprecated contrib LSTM with Keras LSTMCell rnn_cell = tf.keras.layers.LSTMCell(1) out, _ = tf.nn.dynamic_rnn(rnn_cell, xin, dtype=tf.float32) # Batch normalization with training flag out = tf.layers.batch_normalization(out, training=True) out = tf.identity(out, name='output') # Define loss (add your actual loss function here; example provided) y_true = tf.placeholder(tf.float32, [None, 1, 1], name='target') loss = tf.reduce_mean(tf.square(out - y_true)) optimiser = tf.train.AdamOptimizer(0.0001) # Collect batch norm update operations update_ops = tf.get_collection(tf.GraphKeys.UPDATE_OPS) # Wrap training step to include update ops with tf.control_dependencies(update_ops): train_op = optimiser.minimize(loss) # Test the graph to verify no segfault with tf.Session() as sess: sess.run(tf.global_variables_initializer()) # Run a dummy training step dummy_input = [[[0.5]]] dummy_target = [[[0.7]]] _, loss_val = sess.run([train_op, loss], feed_dict={xin: dummy_input, y_true: dummy_target}) print(f"Training step completed, loss: {loss_val}")
Additional Debugging Tips
- Update TensorFlow: If you're using an older version (pre-2.x), upgrade to a stable TensorFlow 1.x release (like 1.15) or migrate to TensorFlow 2.x with compatibility mode—many memory bugs were fixed in later versions.
- Check Batch Size: Avoid extremely small or variable batch sizes that might trigger edge-case memory handling issues in batch norm.
- Validate Device Placement: Ensure all operations are properly placed on
/cpu:0(your code does this, but double-check if any ops are accidentally placed on reasoningBo.on implementing sanity分
精神牛考虑 /*D More同学的GPU if available).
内容的提问来源于stack exchange,提问作者Thomas Bastiani
相关产品推荐
相关产品推荐

