使用Talos+Colab TPU调优Keras模型时出现tf.distribute.Strategy混合错误
Fixing "Mixing different tf.distribute.Strategy objects" Error with Talos and Colab TPU
Root Cause
The error occurs because your iris_model function initializes a new TPUStrategy instance every time Talos calls it during hyperparameter scans. Each TPU system initialization creates a distinct strategy object, and TensorFlow prohibits mixing multiple strategy instances in the same runtime session. This conflict triggers the RuntimeError you're seeing.
Solution
Move the TPU strategy initialization outside the model function so it runs only once before starting the Talos scan. Reuse this single strategy instance across all hyperparameter trials. Here's the corrected code with key improvements:
import os import tensorflow as tf import talos as ta from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense # Initialize TPU strategy ONCE before any model training resolver = tf.distribute.cluster_resolver.TPUClusterResolver(tpu='grpc://' + os.environ['COLAB_TPU_ADDR']) tf.config.experimental_connect_to_host(resolver.master()) tf.tpu.experimental.initialize_tpu_system(resolver) strategy = tf.distribute.experimental.TPUStrategy(resolver) # Disable eager execution for TF 2.0.0 compatibility (optional in newer TF versions) tf.compat.v1.disable_eager_execution() def iris_model(x_train, y_train, x_val, y_val, params): # Use the pre-initialized strategy scope to create the model with strategy.scope(): model = Sequential() model.add(Dense(32, input_dim=4, activation=params['activation'])) model.add(Dense(3, activation='softmax')) model.compile(optimizer=params['optimizer'], loss=params['losses']) # Prepare training dataset with batch size compatible with TPU (multiple of 8) train_dataset = tf.data.Dataset.from_tensor_slices((x_train, y_train)) train_dataset = train_dataset.cache().shuffle(1000, reshuffle_each_iteration=True).repeat().batch(params['batch_size'], drop_remainder=True) # Prepare validation dataset as tf.data.Dataset (required for TPU compatibility) val_dataset = tf.data.Dataset.from_tensor_slices((x_val, y_val)).batch(params['batch_size'], drop_remainder=True) # Fit the model: omit batch_size (dataset is already batched) steps_per_epoch = len(x_train) // params['batch_size'] out = model.fit(train_dataset, epochs=params['epochs'], validation_data=val_dataset, verbose=0, steps_per_epoch=steps_per_epoch) return out, model # Load dataset x, y = ta.templates.datasets.iris() # Hyperparameter grid: use batch sizes that are multiples of 8 (Colab TPU has 8 cores) p = {'activation': ['relu', 'elu'], 'optimizer': ['Nadam', 'Adam'], 'losses': ['logcosh'], 'batch_size': [24, 32, 40, 48], # Multiples of 8 to avoid TPU batch errors 'epochs': [10, 20]} # Run Talos hyperparameter scan scan_object = ta.Scan(x, y, model=iris_model, params=p, fraction_limit=0.1, experiment_name='first_test')
Key Improvements Explained
- Single Strategy Initialization: The TPU strategy is created once at the start, eliminating conflicts between multiple strategy instances.
- TPU-Compatible Batch Sizes: Colab TPUs have 8 cores, so batch sizes must be multiples of 8 to ensure even distribution across replicas (required when using
drop_remainder=True). - Validation as Dataset: TPUs perform best when both training and validation data are passed as
tf.data.Datasetobjects. - Redundant Parameter Removal: The
batch_sizeargument inmodel.fitis omitted because the dataset is already batched, preventing confusion.
Additional Notes
- If possible, upgrade to a newer TensorFlow version (2.8+) for better TPU support and eager execution compatibility.
- Ensure your Colab runtime is set to use a TPU (Runtime > Change runtime type > Hardware accelerator > TPU).
内容的提问来源于stack exchange,提问作者Sami Belkacem
相关产品推荐
相关产品推荐

