You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

  1. Single Strategy Initialization: The TPU strategy is created once at the start, eliminating conflicts between multiple strategy instances.
  2. 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).
  3. Validation as Dataset: TPUs perform best when both training and validation data are passed as tf.data.Dataset objects.
  4. Redundant Parameter Removal: The batch_size argument in model.fit is 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.14 09:01:35