Google Colab Pro训练CNN时步数未知导致无限循环问题求助
Hey there, let's work through why your CNN training is showing Unknown steps and stuck in an infinite loop. This is a common issue with TensorFlow's ImageDataGenerator, and it's easy to fix once you know the root cause!
Why This Happens
TensorFlow displays Unknown steps when it can't automatically calculate how many batches make up one epoch (i.e., when it doesn't know when to end an epoch). Without this number, the training loop has no stopping point for each epoch, leading to that infinite run.
Step-by-Step Solutions
1. Manually Specify steps_per_epoch and validation_steps
This is the most reliable fix. You calculate these values based on your dataset size and batch size:
steps_per_epoch = total_training_images // batch_sizevalidation_steps = total_validation_images // batch_size
Using your dataset numbers (10018 training images, 1336 validation images), here's how to implement this in your model.fit() call:
# Define your batch size (make sure this matches what you used in flow_from_directory) batch_size = 32 # Calculate steps steps_per_epoch = 10018 // batch_size validation_steps = 1336 // batch_size # Update your fit call model.fit( train_generator, epochs=300, steps_per_epoch=steps_per_epoch, validation_data=val_generator, validation_steps=validation_steps )
Even if your dataset size isn't perfectly divisible by the batch size, using integer division here is fine—you'll just skip a few leftover samples per epoch (or you can add 1 to include them, but it's usually not necessary).
2. Double-Check Your flow_from_directory Setup
Make sure you explicitly set the batch_size parameter when creating your generators. If this is missing or incorrect, TensorFlow might fail to infer the epoch steps correctly:
train_generator = ImageDataGenerator().flow_from_directory( '/content/drive/path/to/train', target_size=(224, 224), # Adjust to your image size batch_size=batch_size, # Critical: match this to your steps calculation class_mode='binary' # Since you have 2 classes ) val_generator = ImageDataGenerator().flow_from_directory( '/content/drive/path/to/validation', target_size=(224, 224), batch_size=batch_size, class_mode='binary' )
3. Fix the Deprecation Warning (Bonus)
While this doesn't cause the infinite loop, it's good to clean up that period warning. Replace the period argument in your ModelCheckpoint with save_freq:
- If you want to save every N epochs: use
save_freq='epoch' - Example:
from tensorflow.keras.callbacks import ModelCheckpoint checkpoint = ModelCheckpoint( 'model_checkpoint.h5', save_freq='epoch', # Replaces deprecated `period` save_best_only=True )
Final Note
Your TensorFlow version (2.2.0-rc3) is a bit outdated, which might contribute to the auto-inference issue. If you can, updating to a newer stable version (e.g., 2.x latest) could prevent this from happening in the future, but the steps above should fix it immediately.
内容的提问来源于stack exchange,提问作者Alechandro

