在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都支持)或float16dtype,和模型参数保持一致,避免类型转换的额外内存开销; - 预处理时直接在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
相关产品推荐
相关产品推荐

