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

Tensorflow/Keras训练卷积神经网络内存错误:能否内存中训练?

问题描述

训练一个大型鸟类图像数据集,训练集包含58388张224×224的3通道图像,对应张量尺寸为(58388, 224, 224, 3),占用32.74GB内存,使用的是24GB显存的Nvidia Geforce RTX 3090显卡。

模型构建、编译及训练代码如下:

num_classes = len(np.unique(y_train))
img_height = X_train.shape[1]
img_width = X_train.shape[2]

model = Sequential([
  layers.Rescaling(1./255, input_shape=(img_height, img_width, 3)),
  layers.Conv2D(16, 3, padding='same', activation='relu'),
  layers.MaxPooling2D(),
  layers.Conv2D(32, 3, padding='same', activation='relu'),
  layers.MaxPooling2D(),
  layers.Conv2D(64, 3, padding='same', activation='relu'),
  layers.MaxPooling2D(),
  layers.Flatten(),
  layers.Dense(128, activation='relu'),
  layers.Dense(num_classes)
])

model.compile(optimizer='adam',
              loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
              metrics=['accuracy'])

history = model.fit(
    X_train, 
    y_train,
    batch_size=32,
    epochs=25,
    validation_split=0.03,
)

运行时出现错误:

InternalError: Failed copying input tensor from /job:localhost/replica:0/task:0/device:CPU:0 to /job:localhost/replica:0/task:0/device:GPU:0 in order to run _EagerConst: Dst tensor is not initialized.

该错误提示GPU内存不足,即使将batch_size调小至1仍无法解决。按理解训练时仅单个批次会传输至GPU显存,而非整个数据集,因此无法理解问题原因。

已通过tf.keras.utils.image_dataset_from_directory创建数据集生成器完成训练,现咨询:是否存在无需使用生成器、直接将数据集保存在内存中训练模型的方法?

解决方案
  • 转换为TensorFlow张量并启用内存增长
    将X_train和y_train转换为tf.Tensor,同时开启GPU内存增长模式,让TensorFlow按需分配显存,避免一次性占用过多:

    import tensorflow as tf
    
    # 开启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)
    
    # 转换为tf.Tensor
    X_train_tensor = tf.convert_to_tensor(X_train, dtype=tf.float32)
    y_train_tensor = tf.convert_to_tensor(y_train, dtype=tf.int32)
    
    # 后续模型编译和训练不变,直接传入张量
    history = model.fit(
        X_train_tensor, 
        y_train_tensor,
        batch_size=32,
        epochs=25,
        validation_split=0.03,
    )
    
  • 降低数据精度
    把数据从float32转为float16或bfloat16,大幅减少内存占用,同时启用混合精度训练:

    # 转换数据为float16精度
    X_train_low_precision = X_train.astype('float16')
    
    # 开启混合精度训练
    from tensorflow.keras.mixed_precision import set_global_policy
    set_global_policy('mixed_float16')
    
    # 模型编译逻辑不变
    model.compile(optimizer='adam',
                  loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
                  metrics=['accuracy'])
    
    # 传入低精度数据训练
    history = model.fit(
        X_train_low_precision, 
        y_train,
        batch_size=32,
        epochs=25,
        validation_split=0.03,
    )
    

    混合精度训练会自动处理模型参数和计算的精度转换,在不明显损失精度的前提下降低显存占用。

  • 手动分批次训练
    自行实现分批次逻辑,每次仅将单个批次的数据传入GPU,避免整个数据集张量被意外复制到GPU:

    import numpy as np
    
    batch_size = 32
    epochs = 25
    validation_split = 0.03
    val_size = int(len(X_train) * validation_split)
    
    # 拆分验证集和训练集
    X_val, y_val = X_train[:val_size], y_train[:val_size]
    X_train_subset = X_train[val_size:]
    y_train_subset = y_train[val_size:]
    
    for epoch in range(epochs):
        print(f"Epoch {epoch+1}/{epochs}")
        # 打乱训练集顺序
        idx = np.random.permutation(len(X_train_subset))
        X_train_shuffled = X_train_subset[idx]
        y_train_shuffled = y_train_subset[idx]
        
        # 逐批次训练
        for i in range(0, len(X_train_shuffled), batch_size):
            X_batch = X_train_shuffled[i:i+batch_size]
            y_batch = y_train_shuffled[i:i+batch_size]
            # 手动更新模型权重
            loss, acc = model.train_on_batch(X_batch, y_batch)
            print(f"Batch {i//batch_size+1}, Loss: {loss:.4f}, Accuracy: {acc:.4f}", end='\r')
        
        # 验证模型性能
        val_loss, val_acc = model.evaluate(X_val, y_val, verbose=0)
        print(f"\nValidation Loss: {val_loss:.4f}, Validation Accuracy: {val_acc:.4f}")
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 20:50:31