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

TensorFlow2.6多GPU推理:如何利用全部GPU并保持数据集顺序

TensorFlow 2.6 单机多GPU推理(保持数据集顺序)

1. 初始化多GPU分布式策略

使用MirroredStrategy实现单机多卡并行,它会自动在所有目标GPU上复制模型并分配计算任务,需在加载模型前完成初始化。

import tensorflow as tf
import numpy as np

# 默认使用所有可见GPU,也可显式指定10块GPU
strategy = tf.distribute.MirroredStrategy()
# 显式指定GPU的写法:
# strategy = tf.distribute.MirroredStrategy(devices=[f"/GPU:{i}" for i in range(10)])

2. 在策略范围内加载模型

必须在策略的scope()上下文内加载或构建模型,确保TensorFlow为多GPU环境适配模型参数。

with strategy.scope():
    # 加载预训练模型,替换为你的模型路径
    model = tf.keras.models.load_model("your_trained_model.h5")
    # 若是自定义模型,直接在此处构建
    # model = tf.keras.Sequential([...])

3. 构建有序推理数据集

要保证输出顺序与输入一致,数据集处理需严格避免打乱,同时合理设置batch:

  • 禁用shuffle操作
  • batch_size建议设为GPU数量的整数倍,均衡各卡负载
  • 保留最后一个不足额batch(drop_remainder=False),避免丢失样本
# 示例:以numpy数组作为输入,替换为你的实际数据源
input_data = np.random.rand(10000, 224, 224, 3)  # 模拟10000个样本
dataset = tf.data.Dataset.from_tensor_slices(input_data)

# 设置batch size,10块GPU每批各处理10个样本,总batch=100
dataset = dataset.batch(batch_size=100, drop_remainder=False)
# 预取数据提升推理效率
dataset = dataset.prefetch(tf.data.AUTOTUNE)

4. 执行分布式推理并保持顺序

方法1:直接使用model.predict()(推荐)

TensorFlow原生predict方法在分布式策略下会自动处理多卡并行,且默认严格保持输入与输出的顺序对应,无需手动合并结果。

# 自动多GPU推理,结果顺序与输入dataset完全一致
predictions = model.predict(dataset, verbose=1)

方法2:自定义推理逻辑(适合复杂场景)

若需自定义推理步骤,可通过strategy.run()执行并行计算,再用strategy.gather()合并各卡结果,保证顺序对齐:

@tf.function
def inference_step(inputs):
    # 自定义推理逻辑,这里直接调用模型
    return model(inputs, training=False)

predictions = []
for batch in dataset:
    # 在所有GPU上并行执行推理
    per_replica_preds = strategy.run(inference_step, args=(batch,))
    # 合并各GPU的结果,顺序与输入batch一致
    batch_preds = strategy.gather(per_replica_preds, axis=0)
    predictions.append(batch_preds.numpy())

# 合并所有batch的结果,最终顺序与原输入数据完全匹配
final_predictions = np.concatenate(predictions, axis=0)

5. 关键注意事项

  • 绝对禁止在推理阶段对数据集做shuffle操作,否则会破坏顺序
  • 若模型含自定义层,需确保层内操作兼容分布式策略,可通过tf.distribute.get_replica_context()处理特殊逻辑
  • 大内存压力下,可使用dataset.cache()将数据缓存到内存(或磁盘),减少IO等待

内容的提问来源于stack exchange,提问作者haoran.li

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 13:30:45