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

Google Colab中TensorFlow模型TPU训练连接问题求助

TPU训练TensorFlow模型时出现GRPC连接拒绝错误

我已在CPU和GPU上成功运行TensorFlow神经网络模型,现因数据集较大,尝试在TPU上训练模型。按常规方式初始化TPU策略:

tpu_resolver = tf.distribute.cluster_resolver.TPUClusterResolver()  # 自动检测TPU
tf.config.experimental_connect_to_cluster(tpu_resolver)  # 连接到TPU集群
tf.tpu.experimental.initialize_tpu_system(tpu_resolver)  # 初始化TPU系统
strategy = tf.distribute.TPUStrategy(tpu_resolver)
tpu_device = tpu_resolver.master()  # 获取TPU设备URI
print("Running on TPU:", tpu_device)

运行后输出:

Running on TPU: grpc://10.74.203.82:8470

但在strategy.scope()下训练模型时,出现连接拒绝的GRPC错误导致训练终止,错误信息如下:

err: File "/content/SeniorHonoursProject/BaCoN-II/train.py", line 169, in my_train
err: new_history = model.fit(train_dataset.dataset, epochs=epochs,
err: File "/usr/local/lib/python3.10/dist-packages/keras/src/utils/traceback_utils.py", line 70, in error_handler
err: raise e.with_traceback(filtered_tb) from None
err: File "/usr/local/lib/python3.10/dist-packages/tensorflow/python/framework/ops.py", line 362, in _numpy
err: raise core._status_to_exception(e) from None  # pylint: disable=protected-access
err: tensorflow.python.framework.errors_impl.InternalError: 8 root error(s) found.
err: (0) INTERNAL: {{function_node __inference_train_function_11526}} failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused
err: Additional GRPC error information from remote target /job:localhost/replica:0/task:0/device:CPU:0:
err: :UNKNOWN:failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused {grpc_status:14, created_time:"2024-03-22T23:36:04.934268533+00:00"}
err: [[{{node MultiDeviceIteratorGetNextFromShard}}]]
err: Executing non-communication op <MultiDeviceIteratorGetNextFromShard> originally returned UnavailableError, and was replaced by InternalError to avoid invoking TF network error handling logic.
err: [[RemoteCall]]
err: [[IteratorGetNextAsOptional]]
err: [[TPUReplicate/_compile/_9902494219978988908/_4/_384]]
err: (1) INTERNAL: {{function_node __inference_train_function_11526}} failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused
err: Additional GRPC error information from remote target /job:localhost/replica:0/task:0/device:CPU:0:
err: :UNKNOWN:failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused {grpc_status:14, created_time:"2024-03-22T23:36:04.934268533+00:00"}
err: [[{{node MultiDeviceIteratorGetNextFromShard}}]]
err: Executing non-communication op <MultiDeviceIteratorGetNextFromShard> originally returned UnavailableError, and was replaced by InternalError to avoid invoking TF network error handling logic.
err: [[RemoteCall]]
err: [[IteratorGetNextAsOptional]]
err: [[Pad_3/_250]]
err: (2) INTERNAL: {{function_node __inference_train_function_11526}} failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused
err: Additional GRPC error information from remote target /job:localhost/replica:0/task:0/device:CPU:0:
err: :UNKNOWN:failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused {grpc_status:14, created_time:"2024-03-22T23:36:04.934268533+00:00"}
err: [[{{node MultiDeviceIteratorGetNextFromShard}}]]
err: Executing non-communication op <MultiDeviceIteratorGetNextFromShard> originally returned UnavailableError, and was replaced by InternalError to avoid invoking TF network error handling logic.
err: [[RemoteCall]]
err: [[IteratorGetNextAsOptional]]
err: [[Pad_15/_466]]
err: (3) INTERNAL: {{function_node __inference_train_function_11526}} failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused
err: Additional GRPC error information from remote target /job:localhost/replica:0/task:0/device:CPU:0:
err: :UNKNOWN:failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused {grpc_status:14, created_time:"2024-03-22T23:36:04.934268533+00:00"}
err: [[{{node MultiDeviceIteratorGetNextFromShard}}] ... [truncated]
err: Exception ignored in atexit callback: <function async_wait at 0x7ec0f4a74790>
err: Traceback (most recent call last):
err: File "/usr/local/lib/python3.10/dist-packages/tensorflow/python/eager/context.py", line 2833, in async_wait
err: context().sync_executors()
err: File "/usr/local/lib/python3.10/dist-packages/tensorflow/python/eager/context.py", line 749, in sync_executors
err: pywrap_tfe.TFE_ContextSyncExecutors(self._context_handle)
err: tensorflow.python.framework.errors_impl.InternalError: 8 root error(s) found.
err: (0) INTERNAL: {{function_node __inference_train_function_11526}} failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused
err: Additional GRPC error information from remote target /job:localhost/replica:0/task:0/device:CPU:0:
err: :UNKNOWN:failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused {grpc_status:14, created_time:"2024-03-22T23:36:04.934268533+00:00"}
err: [[{{node MultiDeviceIteratorGetNextFromShard}}]]
err: Executing non-communication op <MultiDeviceIteratorGetNextFromShard> originally returned UnavailableError, and was replaced by InternalError to avoid invoking TF network error handling logic.
err: [[RemoteCall]]
err: [[IteratorGetNextAsOptional]]
err: [[TPUReplicate/_compile/_9902494219978988908/_4/_384]]
err: (1) INTERNAL: {{function_node __inference_train_function_11526}} failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused
err: Additional GRPC error information from remote target /job:localhost/replica:0/task:0/device:CPU:0:
err: :UNKNOWN:failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused {grpc_status:14, created_time:"2024-03-22T23:36:04.934268533+00:00"}
err: [[{{node MultiDeviceIteratorGetNextFromShard}}]]
err: Executing non-communication op <MultiDeviceIteratorGetNextFromShard> originally returned UnavailableError, and was replaced by InternalError to avoid invoking TF network error handling logic.
err: [[RemoteCall]]
err: [[IteratorGetNextAsOptional]]
err: [[Pad_3/_250]]
err: (2) INTERNAL: {{function_node __inference_train_function_11526}} failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused
err: Additional GRPC error information from remote target /job:localhost/replica:0/task:0/device:CPU:0:
err: :UNKNOWN:failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused {grpc_status:14, created_time:"2024-03-22T23:36:04.934268533+00:00"}
err: [[{{node MultiDeviceIteratorGetNextFromShard}}]]
err: Executing non-communication op <MultiDeviceIteratorGetNextFromShard> originally returned UnavailableError, and was replaced by InternalError to avoid invoking TF network error handling logic.
err: [[RemoteCall]]
err: [[IteratorGetNextAsOptional]]
err: [[Pad_15/_466]]
err: (3) INTERNAL: {{function_node __inference_train_function_11526}} failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused
err: Additional GRPC error information from remote target /job:localhost/replica:0/task:0/device:CPU:0:
err: :UNKNOWN:failed to connect to all addresses; last error: UNKNOWN: ipv4:127.0.0.1:50188: Failed to connect to remote host: Connection refused {grpc_status:14, created_time:"2024-03-22T23:36:04.934268533+00:00"}
err: [[{{node MultiDeviceIteratorGetNextFromShard}}] ... [truncated]
out: g1D)

以下是初始化模型、运行训练及数据管道的代码:

模型初始化与训练代码

with strategy.scope():
  n_batches_eff = training_dataset.n_batches // strategy.num_replicas_in_sync
  lr_fn = tf.optimizers.schedules.ExponentialDecay(FLAGS.lr, n_batches_eff, FLAGS.decay)
  optimizer = tf.keras.optimizers.Adam(lr_fn)

with strategy.scope():
            model=make_model(#自定义模型构建函数)
            if FLAGS.bayesian:
                 loss=BayesianLoss(n_train_examples=training_dataset.n_batches*training_dataset.batch_size, n_val_examples=validation_dataset.n_batches*validation_dataset.batch_size, TPU=FLAGS.TPU)
                loss.set_model(model)
            else:
                if FLAGS.TPU:
                    loss = tf.keras.losses.CategoricalCrossentropy(from_logits=True, reduction=tf.keras.losses.Reduction.NONE)
                else:
                    loss=tf.keras.losses.CategoricalCrossentropy(from_logits=True)
            model.compile(optimizer=optimizer, loss=loss, metrics=['accuracy'])


 with strategy.scope():
            val_steps_per_epoch = val_dataset.n_batches // strategy.num_replicas_in_sync
            train_steps_per_epoch = train_dataset.n_batches // strategy.num_replicas_in_sync
            new_history = model.fit(train_dataset.dataset, epochs=epochs,
                                validation_data=val_dataset.dataset,
                                callbacks=[callback], steps_per_epoch=train_steps_per_epoch, validation_steps=val_steps_per_epoch, initial_epoch=last_epoch)

数据管道代码

with self.strategy.scope():
                if self.shuffle:
                    dataset = dataset.shuffle(buffer_size=len(list_IDs))
                dataset.cache()
                global_batchsize = self.batch_size * self.strategy.num_replicas_in_sync
                global_batchsize = tf.cast(global_batchsize, dtype=tf.int64)
                dataset = dataset.batch(global_batchsize)
                dataset = dataset.map(self.normalize_and_onehot, num_parallel_calls=tf.data.experimental.AUTOTUNE)
                dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)
                dataset = self.strategy.experimental_distribute_dataset(dataset)

请问如何解决这个问题?


解决建议
  • 移除strategy.experimental_distribute_dataset调用:TPU策略下TensorFlow会自动处理数据集分发,手动调用可能导致迭代器跨设备连接冲突。
  • 统一strategy.scope()上下文:将优化器定义、模型构建、数据集处理的所有代码放在同一个strategy.scope()内,避免跨上下文的资源不兼容。
  • 调整数据集处理顺序:先执行map和prefetch,再进行batch操作;同时移除dataset.cache(),改用TFRecord格式存储数据直接供TPU读取,避免CPU缓存数据的跨设备访问问题。
  • 重置TPU连接:初始化TPU系统后添加tf.keras.backend.clear_session()重置默认图形,再创建策略和模型。
  • 检查自定义损失函数:确保BayesianLoss完全使用TensorFlow原生操作,无Python逻辑嵌入,且TPU模式下损失的reduction设置为NONE,由策略手动聚合损失。
  • 验证TensorFlow版本:使用TF 2.15+版本,旧版本可能存在TPU迭代器的GRPC连接bug。

内容的提问来源于stack exchange,提问作者Arnav Agarwal

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 04:29:50