使用BERT模型训练时遇ValueError及Colab会话崩溃问题求助
BERT训练报错与内存问题修复方案
一、解决ValueError: too many values to unpack (expected 2)
错误原因
你的代码对训练/验证数据集执行了重复的批量处理操作:
- 先对
train_data调用batch(32)完成批量处理,后续在model_auto.fit()中又再次调用batch(32),导致输入张量被二次打包,形状变为三维(None, None, None),而BERT模型期望的输入是二维张量(batch_size, seq_length),因此触发维度解包错误。
修复步骤
修改model_auto.fit()中的数据集调用,移除重复的shuffle()和batch()操作:
model_auto.fit( train_data, # 已提前完成shuffle、batch、repeat,无需再次处理 validation_data=validation_data, # 同理,已提前完成batch操作 epochs=2 )
确认数据集预处理代码仅做一次批量处理:
# 训练集预处理(已正确完成所有必要操作) train_data = convert_examples_to_tf_dataset(list(train_InputExamples), tokenizer) train_data = train_data.shuffle(25).batch(32).repeat(2) # 验证集预处理(已正确完成批量处理) validation_data = convert_examples_to_tf_dataset(list(validation_InputExamples), tokenizer) validation_data = validation_data.batch(32)
二、解决未耗尽RAM却提示内存用尽的问题
常见原因与修复方案
- 显存不足(而非RAM不足)
Colab的GPU显存通常远小于RAM,即使RAM有剩余,显存耗尽也会导致会话崩溃:
- 减小批量大小:将
batch(32)改为batch(16)或batch(8),降低单次迭代的显存占用。 - 替换轻量模型:使用
distilbert-base-uncased替代标准BERT,参数仅为BERT-base的40%,性能损失极小:from transformers import TFDistilBertForSequenceClassification model_auto = TFDistilBertForSequenceClassification.from_pretrained('distilbert-base-uncased', num_labels=你的类别数量)
- 优化数据集加载
给tf.data.Dataset添加缓存与预取操作,提升数据加载效率并减少内存波动:
train_data = train_data.shuffle(25).batch(32).repeat(2).cache().prefetch(tf.data.AUTOTUNE) validation_data = validation_data.batch(32).cache().prefetch(tf.data.AUTOTUNE)
- 手动清理残留内存
训练前清理TensorFlow会话的残留内存,释放无用资源:
import gc import tensorflow as tf gc.collect() tf.keras.backend.clear_session()
内容的提问来源于stack exchange,提问作者Nandhini Palanikumar
相关产品推荐
相关产品推荐

