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
相关产品推荐
相关产品推荐

