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

使用多进程+TensorFlow Dataset预测时遇轴越界及CUDA内存错误

问题:TensorFlow Dataset结合多进程预测的系列错误解决

问题背景

延续Keras模型多进程预测的需求,改用TensorFlow Dataset处理数据后,出现轴越界、CUDA初始化失败、GPU内存不足三类错误,以下是完整复现代码及错误详情:

复现代码

import tensorflow as tf
import numpy as np
from multiprocessing import Pool
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Conv2D, MaxPooling2D,\
Dense, Flatten

# GPU配置
gpus = tf.config.experimental.list_physical_devices('GPU')
for gpu in gpus:
    tf.config.experimental.set_memory_growth(gpu, True)
    
def model_arch():
    models = Sequential()
    models.add(Conv2D(64, (5, 5), padding="same", activation="relu", input_shape=(28, 28, 1)))
    models.add(MaxPooling2D(pool_size=(2, 2)))
    models.add(Conv2D(128, (5, 5), padding="same", activation="relu"))
    models.add(MaxPooling2D(pool_size=(2, 2)))
    models.add(Conv2D(256, (5, 5), padding="same", activation="relu"))
    models.add(MaxPooling2D(pool_size=(2, 2)))
    models.add(Flatten())
    models.add(Dense(256, activation="relu"))
    models.add(Dense(10, activation="softmax"))
    return models

def _apply_df(data):
    model = model_arch()
    model.load_weights("/home/ggous/model_mnist.h5")
    return model.predict(data)

def apply_by_multiprocessing(data, workers):
    pool = Pool(processes=workers)
    result = pool.map(_apply_df, np.array_split(data, workers))
    pool.close()
    return list(result)

def resize_and_rescale(data):
    data = tf.cast(data, tf.float32)
    data /= 255.0
    return data
    
def prepare(ds):
    ds = ds.map(resize_and_rescale)
    return ds.batch(1)

def after_prepare(data):
    tens_data = tf.data.Dataset.from_tensor_slices(data)
    tens_data = prepare(tens_data)
    return tens_data

def main():
    fashion_mnist = tf.keras.datasets.fashion_mnist
    _, (test_images, test_labels) = fashion_mnist.load_data()

    test_images = after_prepare(test_images)
    results = apply_by_multiprocessing(test_images, workers=3)
    print(test_images.shape)           
    print(len(results))                
    print([x.shape for x in results])  

if __name__ == "__main__":
    main()

错误现象

  1. 轴越界错误:
    axis1: axis 0 is out of bounds for array of dimension 0
    
  2. CUDA初始化错误:
    F tensorflow/stream_executor/cuda/cuda_driver.cc:146] Failed setting context: CUDA_ERROR_NOT_INITIALIZED: initialization error
    
  3. 添加spawn启动方式后的GPU内存不足错误:
    出现大量CUDA_ERROR_OUT_OF_MEMORY日志,最终提示卷积算子无可用算法,核心原因是显存耗尽。

解决方案

一、解决轴越界错误

原因:np.array_split无法直接拆分TensorFlow Dataset对象,会将其视为单元素数组,拆分后子元素维度为空,导致model.predict报错。

解决方法:
将Dataset转换为numpy数组后再拆分,修改数据预处理流程:

def preprocess_data(data):
    # 直接返回带通道维度的numpy数组,替代Dataset
    data = tf.cast(data, tf.float32) / 255.0
    return data[..., tf.newaxis].numpy()

# 在main函数中替换原after_prepare调用
test_data = preprocess_data(test_images)
results = apply_by_multiprocessing(test_data, workers=3)

二、解决CUDA初始化错误

原因:多进程默认的fork启动方式会继承父进程的GPU上下文,导致子进程GPU资源冲突。

解决方法:
在主入口开头设置spawn启动方式,强制子进程重新初始化GPU上下文:

if __name__ == "__main__":
    multiprocessing.set_start_method('spawn', force=True)
    main()

三、解决GPU内存不足错误

原因:每个子进程加载独立模型实例,多进程同时占用显存;batch=1导致GPU利用率低,加剧内存碎片化。

解决方法:

  1. 限制子进程显存占用:在子进程函数中单独配置显存限制
    def _apply_df(data):
        gpus = tf.config.experimental.list_physical_devices('GPU')
        if gpus:
            try:
                # 根据显存大小调整限制值,比如1024MB
                tf.config.experimental.set_virtual_device_configuration(
                    gpus[0],
                    [tf.config.experimental.VirtualDeviceConfiguration(memory_limit=1024)]
                )
                tf.config.experimental.set_memory_growth(gpus[0], True)
            except RuntimeError as e:
                print(e)
        model = model_arch()
        model.load_weights("/home/ggous/model_mnist.h5")
        return model.predict(data)
    
  2. 减少工作进程数:根据GPU显存大小调整,比如从3改为1-2
  3. 增大batch size:将prepare函数中的batch(1)改为batch(32)或batch(64),提升GPU利用率
  4. 父进程不加载模型:主进程仅负责数据拆分和进程管理,不初始化模型,减少显存占用

修正后完整代码

import tensorflow as tf
import numpy as np
import multiprocessing
from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Conv2D, MaxPooling2D, Dense, Flatten

def model_arch():
    models = Sequential()
    models.add(Conv2D(64, (5, 5), padding="same", activation="relu", input_shape=(28, 28, 1)))
    models.add(MaxPooling2D(pool_size=(2, 2)))
    models.add(Conv2D(128, (5, 5), padding="same", activation="relu"))
    models.add(MaxPooling2D(pool_size=(2, 2)))
    models.add(Conv2D(256, (5, 5), padding="same", activation="relu"))
    models.add(MaxPooling2D(pool_size=(2, 2)))
    models.add(Flatten())
    models.add(Dense(256, activation="relu"))
    models.add(Dense(10, activation="softmax"))
    return models

def _apply_df(data):
    # 子进程单独配置GPU显存限制
    gpus = tf.config.experimental.list_physical_devices('GPU')
    if gpus:
        try:
            tf.config.experimental.set_virtual_device_configuration(
                gpus[0],
                [tf.config.experimental.VirtualDeviceConfiguration(memory_limit=1024)]
            )
            tf.config.experimental.set_memory_growth(gpus[0], True)
        except RuntimeError as e:
            print(e)
    model = model_arch()
    model.load_weights("/home/ggous/model_mnist.h5")
    return model.predict(data)

def apply_by_multiprocessing(data, workers):
    pool = multiprocessing.Pool(processes=workers)
    result = pool.map(_apply_df, np.array_split(data, workers))
    pool.close()
    pool.join()
    return np.concatenate(result, axis=0)

def preprocess_data(data):
    # 统一预处理为带通道维度的numpy数组
    data = tf.cast(data, tf.float32) / 255.0
    return data[..., tf.newaxis].numpy()

def main():
    fashion_mnist = tf.keras.datasets.fashion_mnist
    _, (test_images, test_labels) = fashion_mnist.load_data()
    
    # 预处理为模型需要的输入格式
    test_data = preprocess_data(test_images)
    
    # 多进程预测
    results = apply_by_multiprocessing(test_data, workers=2)
    
    print(f"原始数据形状: {test_images.shape}")
    print(f"预测结果总形状: {results.shape}")

if __name__ == "__main__":
    multiprocessing.set_start_method('spawn', force=True)
    main()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 19:25:43