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

在SFTTrainer中实现加权损失函数时遭遇CUDA内存溢出问题

问题排查与修复方案

一、为啥会慢还OOM?

1. 损失计算写得太低效

要是你手动循环每个样本的pieces、逐token计算加权损失,那基本等于放弃了PyTorch的张量并行优化,还会触发大量CPU-GPU数据来回传,速度直接拉垮,显存也会因为频繁分配小张量爆掉。

2. 权重张量没优化

如果权重是在CPU上生成的,训练时才临时移到GPU,或者用了float32这种占空间的 dtype,再或者每个token单独存权重(重复值也不复用),显存浪费会特别严重。

3. 没跟上Unsloth的优化逻辑

Unsloth本身做了LoRA高效实现、梯度检查点这些内存优化,要是你自定义损失时绕开了它的模型包装,或者没开自动混合精度,等于白瞎了这些优化,显存扛不住很正常。

4. 数据集处理太冗余

要是没把每个样本的pieces合并成完整序列,训练时动态拼接,会导致每个batch的张量形状乱飘,触发额外内存分配;再加上没过滤过长序列,单个batch的总token数超标,直接爆显存没商量。

二、具体怎么修?

1. 重写损失计算,用批量张量运算

别搞逐样本循环,把所有token和对应的权重整成大张量,用PyTorch广播机制批量算:

def compute_weighted_loss(model, inputs, return_outputs=False):
    outputs = model(**inputs)
    logits = outputs.logits
    weights = inputs["weights"]
    labels = inputs["labels"]

    # 用交叉熵损失,忽略padding(假设padding id是0)
    loss_fct = torch.nn.CrossEntropyLoss(reduction="none")
    # 对齐logits和labels的维度(shift掉最后一个token的logits,第一个token的labels)
    shift_logits = logits[..., :-1, :].contiguous()
    shift_labels = labels[..., 1:].contiguous()
    shift_weights = weights[..., 1:].contiguous()

    # 批量算逐token损失,乘权重后取平均
    per_token_loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1))
    weighted_loss = (per_token_loss * shift_weights.view(-1)).mean()

    return (weighted_loss, outputs) if return_outputs else weighted_loss

关键前提:预处理时要把每个样本的pieces按顺序拼成完整text,同时生成对应的权重序列(每个token对应一个权重,padding部分设为0),并且把权重张量和labels、input_ids一起移到GPU。

2. 优化权重张量的内存占用

  • 权重用bfloat16(Colab的T4/A100都支持)或float16 dtype,和模型参数保持一致,避免类型转换的额外内存开销;
  • 预处理时直接在GPU生成权重(显存够的话),或者给DataLoader开pin_memory=True,加速CPU到GPU的传输;
  • 重复的权重模式(比如某段text全用2.0权重),直接用repeat/expand复用,别存一堆重复值。

3. 适配Unsloth的训练优化

  • 必须用Unsloth的FastLanguageModel.get_peft_model包装模型,开LoRA和梯度检查点:
model = FastLanguageModel.get_peft_model(
    model,
    r=16,
    lora_alpha=32,
    lora_dropout=0.05,
    bias="none",
    use_gradient_checkpointing=True,  # 砍内存必备
    random_state=3407,
    max_seq_length=2048,
)
  • 把自定义损失函数传给SFTTrainer的compute_loss参数,别瞎改模型前向逻辑;
  • 开自动混合精度,在training_args里设fp16=True或bf16=True,Colab里直接用就行。

4. 优化数据集预处理和加载

  • 预处理时把每个样本的pieces拼成完整字符串,同时生成权重序列:
def preprocess_function(examples):
    full_texts = []
    weight_sequences = []
    for pieces in examples["pieces"]:
        text_parts = []
        weights = []
        for piece in pieces:
            # 先转token,避免重复编码
            tokens = tokenizer(piece["text"], add_special_tokens=False)["input_ids"]
            text_parts.append(piece["text"])
            # 给这段text的每个token分配权重
            weights.extend([piece["weight"]] * len(tokens))
        # 加bos/eos特殊token,同时给它们分配权重(比如设为1.0)
        full_text = tokenizer.bos_token + " ".join(text_parts) + tokenizer.eos_token
        full_texts.append(full_text)
        weight_sequences.append([1.0] + weights + [1.0])
    
    # 批量编码,固定max_length
    tokenized = tokenizer(full_texts, padding="max_length", truncation=True, max_length=2048)
    # 把权重序列填充到max_length,padding部分设为0
    padded_weights = [ws + [0.0]*(2048 - len(ws)) if len(ws) <2048 else ws[:2048] for ws in weight_sequences]
    tokenized["weights"] = padded_weights
    
    return tokenized
  • 用with_format指定设备和dtype,避免训练时动态转换:
tokenized_dataset = tokenized_dataset.with_format("torch", device="cuda", dtype={"weights": torch.bfloat16})
  • 减小batch size,或者开梯度累积(training_args.gradient_accumulation_steps=4),降低单次显存占用。

5. 其他小技巧

  • 训练前跑torch.cuda.empty_cache()清显存;
  • 要是用T4 GPU,把DataLoader的persistent_workers关掉,避免额外显存占用;
  • 关闭不必要的日志和检查点存储,减少IO和内存消耗。

内容的提问来源于stack exchange,提问作者Jelle De Loecker

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 05:35:07