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

3D U-Net训练正常但model.predict()报GpuLaunchKernel错误求助

问题描述

训练4类多标签3D U-Net时,model.fit()无报错且模型有学习效果,但调用model.predict()时触发GPU内核报错:

85/85 - 56s
2022-12-22 18:26:24.265485: F tensorflow/core/kernels/concat_lib_gpu_impl.cu.cc:165] Non-OK-status: GpuLaunchKernel( concat_variable_kernel<T, IntType, true>, config.block_count, config.thread_per_block, smem_usage, gpu_device.stream(), input_ptrs, output_scan, static_cast(output->dimension(0)), static_cast(output->dimension(1)), output->data()) status: Internal: invalid configuration argument
/cm/local/apps/slurm/var/spool/job5510720/slurm_script: line 14: 1945 Aborted

简化后的代码如下:

import tensorflow as tf
from keras.models import Model
from keras.models import load_model
from tensorflow.keras.optimizers import Adam, SGD
from keras.layers import Conv3D, MaxPooling3D, Conv3DTranspose, UpSampling3D, Concatenate


def unet(input_shape,filters,kernel,model_name):

    strides_1 = (1,1,1)
    strides_2 = (2,2,2)
    ins = Input(shape=input_shape,name='input_1')

    encode1a = Conv3D(filters=filters, kernel_size=kernel, activation='relu', padding='same', name='encode1a', strides=strides_1)(x)
    encode1b = Conv3D(filters=filters, kernel_size=kernel, activation='relu', padding='same', name='encode1b', strides=strides_1)(encode1a)
    pool1 = MaxPooling3D(pool_size=(2, 2, 2), padding='same', name='pool1')(encode1b)

    encode2a = Conv3D(filters=2*filters, kernel_size=kernel, activation='relu', padding='same', name='encode2a', strides=strides_1)(pool1)
    encode2b = Conv3D(filters=2*filters, kernel_size=kernel, activation='relu', padding='same', name='encode2b', strides=strides_1)(encode2a)
    pool2 = MaxPooling3D(pool_size=(2, 2, 2), padding='same', name='pool2')(encode2b)

    encode3a = Conv3D(filters=4*filters, kernel_size=kernel, activation='relu', padding='same', name='encode3a', strides=strides_1)(pool2)
    encode3b = Conv3D(filters=4*filters, kernel_size=kernel, activation='relu', padding='same', name='encode3b', strides=strides_1)(encode3a)
    pool3 = MaxPooling3D(pool_size=(2, 2, 2), padding='same', name='pool3')(encode3b)

    # Bottleneck
    #--------------------------
    bottom_a = Conv3D(filters=8*filters, kernel_size=kernel, activation='relu', padding='same')(pool3)
    bottom_b = Conv3D(filters=8*filters, kernel_size=kernel, activation='relu', padding='same')(bottom_a)

    # Decoding 
    #--------------------------
    up2   = Concatenate(axis=4)([Conv3DTranspose(filters=4*filters, kernel_size=(2,2,2), strides=strides_2, padding='same')(bottom_b), encode3b])
    decode2a = Conv3D(filters=4*filters, kernel_size=kernel, activation='relu', padding='same',name='decode1a')(up2)
    decode2b = Conv3D(filters=4*filters, kernel_size=kernel, activation='relu', padding='same',name='decode1b')(decode2a)

    up3   = Concatenate(axis=4)([Conv3DTranspose(filters=2*filters, kernel_size=(2,2,2), strides=strides_2, padding='same')(decode2b), encode2b])
    decode1a = Conv3D(filters=2*filters, kernel_size=kernel, activation='relu', padding='same',name='decode2a')(up3)
    decode1b = Conv3D(filters=2*filters, kernel_size=kernel, activation='relu', padding='same',name='decode2b')(decode1a)

    up4   = Concatenate(axis=4)([Conv3DTranspose(filters=filters, kernel_size=(2,2,2), strides=strides_2, padding='same')(decode1b), encode1b])
    decode0a = Conv3D(filters=filters, kernel_size=kernel, activation='relu', padding='same',name='decode3a')(up4)
    decode0b = Conv3D(filters=filters, kernel_size=kernel, activation='relu', padding='same',name='decode3b')(decode0a)

    # Output
    flatten = Convolution3D(filters=4, kernel_size=(1,1,1), activation='softmax')(decode0b)
    model = Model(inputs=ins, outputs=flatten, name=model_name)
    return model


FILTERS = 32
KERNEL = (3,3,3)
MODEL_NAME = 'multi-unet-test'
LR = 3e-3

strategy = tf.distribute.MirroredStrategy()
print('Number of devices: {}'.format(strategy.num_replicas_in_sync))
with strategy.scope():
    model = nets.unet((None,None,None,1),FILTERS,KERNEL,model_name=MODEL_NAME)
    model.compile(optimizer=nets.Adam(lr=LR),loss=tf.keras.losses.SparseCategoricalCrossentropy(),metrics=['accuracy'])
model.summary()

X_train, Y_train = load_dataset_all(FILE_DEN,FILE_MSK,SUBGRID)
# this is a function for loading input and mask fields
# outputs shapes of [256,128,128,128,4]

history = model.fit(X_train, Y_train, batch_size = 4, epochs = 50, verbose = 2, shuffle = True, validation_split = 0.2)
model.save(MODEL_NAME)

# Load and predict 
# this is actually in another script but I'm putting this all in one go:
model = load_model(MODEL_NAME)
model.compile(loss=model.loss,optimizer=model.optimizer,metrics=['accuracy'])
# load test data:
X_test = load_dataset()
Y_test = model.predict(X_test, batch_size = 4, verbose = 2)

已尝试的解决方法:

  • 调整测试集样本数适配batch size(从[343,128,128,128,4]改为[340,128,128,128,4])
  • 更换不同版本的TF/CUDA(TF2.4.1+CUDA11.6、TF2.9.2+CUDA11.2)
解决方案

针对GPU concat内核报错,给出以下针对性修复建议:

1. 修复模型输入层引用错误

U-Net定义中,输入层变量为ins,但第一个Conv3D层错误引用了未定义的x:

encode1a = Conv3D(...)()(x)  # 此处应替换为ins

修正为:

encode1a = Conv3D(filters=filters, kernel_size=kernel, activation='relu', padding='same', name='encode1a', strides=strides_1)(ins)

这个错误可能导致训练时张量结构隐性不匹配,加载模型后预测触发底层GPU错误。

2. 替换动态输入维度为固定尺寸

训练时使用(None,None,None,1)动态输入形状,在多GPU分布式训练+模型加载预测的场景下,易引发GPU内核配置异常。建议改为与训练数据匹配的固定尺寸:

model = nets.unet((128,128,128,1),FILTERS,KERNEL,model_name=MODEL_NAME)

确保测试集X_test的形状为[N,128,128,128,1],与训练集一致。

3. 加载模型时禁用分布式策略

训练时用MirroredStrategy,但加载模型时若环境策略配置不一致,会导致张量拼接错误。加载模型时直接调用,无需包裹策略:

# 无需strategy.scope()包裹
model = load_model(MODEL_NAME)

若需多GPU预测,需重新在策略范围内编译并适配测试数据分布。

4. 调整预测参数与GPU内存分配

报错源于GPU内核配置异常,大概率是显存占用过高导致:

  • 减小预测batch size(如改为2或1)
  • 预测前开启GPU内存增长模式,避免显存瞬间占满:
gpus = tf.config.list_physical_devices('GPU')
if gpus:
    try:
        for gpu in gpus:
            tf.config.experimental.set_memory_growth(gpu, True)
    except RuntimeError as e:
        print(e)

5. 验证Concatenate层的轴与张量尺寸

模型中使用Concatenate(axis=4),3D张量通道轴默认是axis=-1(对应索引3,形状为[batch, d, h, w, channels]),需确认:

  • 所有待拼接的张量(如Conv3DTranspose输出与对应编码层输出)在通道轴外的尺寸完全匹配
  • padding='same'在所有卷积、池化、转置卷积层中正确应用,避免尺寸偏差

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.07 05:55:15