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

在Google Colab训练Alexnet(猫狗数据集)时内存耗尽求助

问题分析与解决方案

一、内存耗尽的核心原因

  1. 全量数据集Shuffle导致内存过载
    你设置了buffer_size=DATASET_SIZE,这会让TensorFlow把整个猫狗数据集(约23k张图片)全部加载到内存中执行shuffle操作,直接占满Colab的RAM,这是会话崩溃的最主要原因。

  2. TensorFlow与PyTorch混合使用的内存冗余
    用TensorFlow加载预处理数据后,又转成PyTorch张量,这个过程会产生两份内存拷贝(TF张量+Torch张量),进一步加剧内存消耗。同时两个深度学习框架同时运行,也会占用额外的系统资源。

  3. 标签处理的冗余操作
    先用to_categorical生成独热标签,后续又通过torch.argmax转回类别索引,完全没必要,徒增内存占用。

二、代码实现的修正建议

1. 优化数据集Shuffle逻辑

将shuffle的buffer_size改为合理值(比如1000),既保证数据打乱效果,又避免占用过多内存:

dataset = dataset.shuffle(buffer_size=1000, seed=42)  # 替换原buffer_size=DATASET_SIZE

2. 统一深度学习框架(二选一)

方案A:全用TensorFlow(推荐,匹配现有数据加载流程)

如果你的AlexNet是用TensorFlow实现的,直接用TF原生训练流程,避免跨框架转换:

loss_fn = tf.keras.losses.CategoricalCrossentropy()
optimizer = tf.keras.optimizers.Adam()

for step, (images, labels) in enumerate(train_ds.take(1)):
    with tf.GradientTape() as tape:
        outputs = model(images)
        loss = loss_fn(labels, outputs)
    
    grads = tape.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(grads, model.trainable_variables))
    
    print(f"Simulated single training step: Loss = {loss.numpy():.4f}")

方案B:全用PyTorch

如果坚持用PyTorch,改用PyTorch原生数据集加载方式,消除TensorFlow的内存开销:

from torchvision import datasets, transforms
from torch.utils.data import random_split, DataLoader

# 定义预处理流程
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),  # 自动将像素值归一化到[0,1]
])

# 加载猫狗数据集(需确保数据集路径正确,可使用torchvision内置数据集或本地路径)
dataset = datasets.ImageFolder(root="./cats_vs_dogs", transform=transform)
train_size = int(0.75 * len(dataset))
val_size = int(0.1 * len(dataset))
test_size = len(dataset) - train_size - val_size

train_ds, val_ds, test_ds = random_split(dataset, [train_size, val_size, test_size])

# 创建数据加载器
train_loader = DataLoader(train_ds, batch_size=4, shuffle=True, num_workers=2)
val_loader = DataLoader(val_ds, batch_size=4, num_workers=2)
test_loader = DataLoader(test_ds, batch_size=4, num_workers=2)

# 纯PyTorch训练步骤
for inputs, labels in train_loader:
    inputs = inputs.to(device)
    labels = labels.to(device)
    
    optimizer.zero_grad()
    outputs = model(inputs)
    loss = loss_fn(outputs, labels)
    loss.backward()
    optimizer.step()
    
    print(f"Simulated single training step: Loss = {loss.item():.4f}")
    break  # 仅执行一轮示例

3. 简化标签处理

  • 用TensorFlow时,保留to_categorical生成的独热标签,搭配CategoricalCrossentropy损失函数即可。
  • 用PyTorch时,不需要独热编码,直接使用类别索引和CrossEntropyLoss损失函数,删除多余的torch.argmax步骤。

4. 额外内存优化技巧

  • 开启Colab GPU加速(Runtime > Change runtime type > 选择GPU),利用显存分担内存压力。
  • 适当调大DataLoader的num_workers参数,让数据加载在后台异步执行,减少内存阻塞。
  • 检查AlexNet模型结构,避免设置过大的全连接层神经元数量,减少模型参数占用的内存。

三、总结

你的代码核心问题是全量数据集shuffle和跨框架内存冗余,修正这两点后即可解决内存耗尽导致的会话崩溃问题。同时统一深度学习框架、简化标签处理,能让训练流程更高效稳定。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.03 19:04:49