如何在单GPU上并行运行单个tf.keras.Model的model.fit以提升算力利用率
单GPU并行运行多个Keras模型训练的实现方案
当单GPU训练时利用率仅25%,通过并行运行多个model.fit实例是提升算力利用率的有效方式,核心是解决显存共享、资源隔离和数据集分片问题,以下是两种可行的实现方法:
一、显存预配置(必做步骤)
TensorFlow默认会占用全部GPU显存,必须先调整显存策略,让多个模型可以共享显存资源:
import tensorflow as tf gpus = tf.config.list_physical_devices('GPU') if gpus: try: # 开启显存动态增长,按需分配显存 for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True) # 或者为每个进程设置固定显存上限(比如单GPU16G的话,分给4个模型各4G) # tf.config.set_logical_device_configuration(gpus[0], [tf.config.LogicalDeviceConfiguration(memory_limit=4096)]) except RuntimeError as e: print(e)
二、方法一:Python多进程实现(推荐)
Python多线程受GIL限制,计算密集型任务更适合用多进程。每个进程独立训练一个模型,处理数据集的不同分片,完全隔离资源。
实现代码
import multiprocessing as mp import tensorflow as tf from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense # 全局配置显存动态增长 gpus = tf.config.list_physical_devices('GPU') if gpus: for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True) def train_single_model(process_id, batch_size, epochs, x_train, y_train, num_processes): # 每个进程加载专属数据集分片 dataset = tf.data.Dataset.from_tensor_slices( (x_train[process_id::num_processes], y_train[process_id::num_processes]) ) dataset = dataset.batch(batch_size).prefetch(tf.data.AUTOTUNE) # 每个进程独立构建模型 model = Sequential([ Dense(64, activation='relu', input_shape=(10,)), Dense(32, activation='relu'), Dense(1, activation='sigmoid') ]) model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy']) print(f"进程 {process_id} 启动训练") model.fit(dataset, epochs=epochs, verbose=1) model.save(f"trained_model_{process_id}.h5") if __name__ == '__main__': # 根据GPU性能调整并行数量(先从2-4开始试) num_parallel = 4 batch_size = 32 epochs = 10 # 加载超大规模数据集(实际场景可替换为从文件/TFRecord加载) x_train = tf.random.normal((100000, 10)) y_train = tf.random.uniform((100000, 1), 0, 2, dtype=tf.int32) # 创建进程池并启动训练任务 pool = mp.Pool(num_parallel) for idx in range(num_parallel): pool.apply_async( train_single_model, args=(idx, batch_size, epochs, x_train, y_train, num_parallel) ) pool.close() pool.join() print("所有训练任务完成")
注意事项
- 每个进程必须独立构建模型、加载数据集,禁止跨进程共享TensorFlow张量或模型实例
- 数据集分片要均匀,保证每个进程的计算量相当
- 通过
nvidia-smi监控GPU利用率,逐步调整num_parallel的数值,直到利用率稳定在80%-90%左右
三、方法二:单进程内多模型并行训练
如果不想用多进程,可以在单个TensorFlow进程内手动实现多模型并行训练,通过自定义训练循环替代model.fit,将多个模型的训练步骤打包为并行操作。
实现代码
import tensorflow as tf from tensorflow.keras.models import Sequential from tensorflow.keras.layers import Dense # 配置显存动态增长 gpus = tf.config.list_physical_devices('GPU') if gpus: tf.config.experimental.set_memory_growth(gpus[0], True) def build_base_model(): # 定义基础模型结构 model = Sequential([ Dense(64, activation='relu', input_shape=(10,)), Dense(32, activation='relu'), Dense(1, activation='sigmoid') ]) model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy']) return model # 准备分片数据集 num_models = 4 x_train = tf.random.normal((100000, 10)) y_train = tf.random.uniform((100000, 1), 0, 2, dtype=tf.int32) datasets = [] for i in range(num_models): ds = tf.data.Dataset.from_tensor_slices((x_train[i::num_models], y_train[i::num_models])) ds = ds.batch(32).prefetch(tf.data.AUTOTUNE) datasets.append(ds) # 创建多个独立模型实例 models = [build_base_model() for _ in range(num_models)] # 自定义并行训练步骤 @tf.function def parallel_train_step(models, datasets): losses = [] for model, ds in zip(models, datasets): for x, y in ds: with tf.GradientTape() as tape: y_pred = model(x, training=True) loss = model.compiled_loss(y, y_pred, regularization_losses=model.losses) # 计算并应用梯度 gradients = tape.gradient(loss, model.trainable_variables) model.optimizer.apply_gradients(zip(gradients, model.trainable_variables)) # 更新指标 model.compiled_metrics.update_state(y, y_pred) losses.append(loss) return losses # 执行训练循环 epochs = 10 for epoch in range(epochs): print(f"Epoch {epoch+1}/{epochs}") losses = parallel_train_step(models, datasets) # 打印每个模型的训练结果 for idx, model in enumerate(models): metric_results = {m.name: m.result().numpy() for m in model.metrics} print(f"模型 {idx}: 损失={losses[idx].numpy():.4f}, 准确率={metric_results['accuracy']:.4f}") model.reset_metrics()
注意事项
- 手动训练循环需要自己处理梯度计算、指标更新等细节,灵活性高但代码复杂度更高
- 需确保显存足够容纳多个模型的参数和中间计算结果
内容的提问来源于stack exchange,提问作者Arsen Zahray
相关产品推荐
相关产品推荐

