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

使用Unsloth在Colab微调模型时遇PicklingError问题求助

解决Unsloth训练时PicklingError:无法序列化SFTConfig类的问题

问题描述

在Google Colab Pro Plus环境中使用Unsloth训练模型,执行trainer.train()步骤时触发PicklingError,错误信息为:

PicklingError: Can't pickle <class 'trl.trainer.sft_config.SFTConfig'>: it's not the same object as trl.trainer.sft_config.SFTConfig

已尝试更换H100、A100、L4、T4及高内存GPU,使用Google示例JSON数据和Hugging Face的yahma/alpaca-cleaned数据集,问题均未解决。

错误栈

PicklingError                             Traceback (most recent call last)
/tmp/ipykernel_22154/2279315892.py in <cell line: 0>()
----> 1 trainer_stats = trainer.train()

10 frames
/usr/local/lib/python3.12/dist-packages/torch/serialization.py in _save(obj, zip_file, pickle_module, pickle_protocol, _disable_byteorder_record)
   1225 
   1226     pickler = PyTorchPickler(data_buf, protocol=pickle_protocol)
-> 1227     pickler.dump(obj)
   1228 
   1229     # The class def keeps the persistent_id closure alive, leaking memory.

PicklingError: Can't pickle <class 'trl.trainer.sft_config.SFTConfig'>: it's not the same object as trl.trainer.sft_config.SFTConfig

原训练配置代码

#Train the model using HuggingFace TRLs wait for the trainer variable to be created
import sys
import importlib
import torch
from datasets import load_dataset
# Force reload TRL components to sync memory references
if "trl" in sys.modules:
    importlib.reload(sys.modules["trl"])

from transformers import TrainingArguments
from unsloth import is_bfloat16_supported
from trl import SFTConfig, SFTTrainer

trainer = SFTTrainer(
    output_dir = "/content/drive/MyDrive/outputDir",
    model = model,
    tokenizer = tokenizer,
    train_dataset = dataset,
    dataset_text_field = "text",
    max_seq_length = max_seq_length,
    dataset_num_proc = 2,
    packing = False, # Can make training 5x faster for short sequences.
    args = TrainingArguments(
        per_device_train_batch_size = 1,#it makes no difference when it is 2
        gradient_accumulation_steps = 1,#when i set the gradient_accumulation_steps to 1or23o4 the loss decreasa up to steps 7 and8, then it starts to increse again
        warmup_steps = 1,
         num_train_epochs = 1, # Set this for 1 full training run.
        gradient_checkpointing = True,
        max_steps = 60,
        learning_rate = 2e-4,
        fp16 = not is_bfloat16_supported(),
        bf16 = is_bfloat16_supported(),
        logging_steps = 1,
        optim = "adamw_8bit",
        weight_decay = 0.01,
        lr_scheduler_type = "linear",
        seed = 3407,
        
        report_to = "none", # Use this for WandB etc
    ),
)

解决方案

这个错误的核心原因是手动重载TRL模块导致类引用不一致:importlib.reload(sys.modules["trl"])会重新加载TRL模块,但之前导入的SFTConfig、SFTTrainer等类来自旧模块实例,与重载后的模块类不是同一个对象,pickle序列化时无法匹配。

修复步骤

  1. 移除TRL模块重载代码:删除if "trl" in sys.modules: importlib.reload(sys.modules["trl"])部分,避免类引用冲突。
  2. 确保依赖版本兼容:确认Colab中Unsloth和TRL版本为官方推荐的兼容组合,可通过重新安装最新版解决潜在版本问题:
    !pip install --upgrade unsloth trl transformers
    
  3. 修改后的训练配置代码:
    # 训练模型,创建trainer实例
    import torch
    from datasets import load_dataset
    from transformers import TrainingArguments
    from unsloth import is_bfloat16_supported
    from trl import SFTConfig, SFTTrainer
    
    trainer = SFTTrainer(
        output_dir = "/content/drive/MyDrive/outputDir",
        model = model,
        tokenizer = tokenizer,
        train_dataset = dataset,
        dataset_text_field = "text",
        max_seq_length = max_seq_length,
        dataset_num_proc = 2,
        packing = False, # 短序列训练时开启可提速5倍
        args = TrainingArguments(
            per_device_train_batch_size = 1,
            gradient_accumulation_steps = 1,
            warmup_steps = 1,
            num_train_epochs = 1, # 设为1表示完整训练一轮
            gradient_checkpointing = True,
            max_steps = 60,
            learning_rate = 2e-4,
            fp16 = not is_bfloat16_supported(),
            bf16 = is_bfloat16_supported(),
            logging_steps = 1,
            optim = "adamw_8bit",
            weight_decay = 0.01,
            lr_scheduler_type = "linear",
            seed = 3407,
            report_to = "none", # 如需WandB等工具可修改此处
        ),
    )
    

验证执行

运行修改后的代码,执行trainer_stats = trainer.train()即可正常启动训练。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.10 19:13:10