如何在MacBook Pro M2 Max上用GPU加速Llama2模型训练?
解决M2 Max上Llama-2-7b训练GPU未利用及内存不足问题
核心问题分析
- 训练循环逻辑错误:原代码嵌套两层
for循环,将每个batch重复处理整个数据集次数,导致计算量暴增、速度极慢,且内存持续累积。 - 序列长度不合理:
max_length=5000远超Llama-2的最优上下文效率,单序列占用内存过高。 - 模型加载无内存优化:直接加载FP16精度的7B模型,M2 Max的GPU内存难以承载长序列训练。
- 语法错误导致设备迁移失败:
input_ids.to(device)行的换行逗号导致张量未正确传到MPS设备。
优化步骤
1. 修复训练循环逻辑
移除内层多余的for i in range(len(train_loader1))循环,确保每个batch只被处理一次。
2. 降低序列长度
将max_length调整为512或1024(根据任务需求,最大不超过4096,Llama-2的官方上下文窗口),大幅减少单序列内存占用。
3. 启用模型量化加载
使用bitsandbytes库将模型量化为4位/8位精度,内存占用直接降低75%/50%,同时基本不损失任务性能。
4. 梯度累积模拟大batch
保持小batch_size的同时,通过梯度累积实现大batch的训练效果,提升GPU利用率,减少内存波动。
5. 内存清理与效率优化
- 训练循环中定期调用
torch.mps.empty_cache()清理无用内存 - 减少循环内打印操作,降低IO耗时
- 确保所有张量正确迁移到MPS设备
完整优化后代码
from transformers import AutoModelForSequenceClassification, AutoTokenizer from peft import LoraConfig, get_peft_model # 可选,用LoRA进一步减少训练内存 from torch.utils.data import DataLoader, TensorDataset import torch import time # 模型与tokenizer加载 model_name = "meta-llama/Llama-2-7b-hf" tokenizer = AutoTokenizer.from_pretrained(model_name, use_auth_token=True) tokenizer.pad_token = tokenizer.eos_token # 启用4位量化加载模型,大幅降低内存占用 model = AutoModelForSequenceClassification.from_pretrained( model_name, num_labels=2, load_in_4bit=True, device_map="auto", # 自动分配模型到可用设备 bnb_4bit_use_double_quant=True, bnb_4bit_quant_type="nf4", bnb_4bit_compute_dtype=torch.bfloat16 ) # 可选:启用LoRA微调,仅训练部分参数,进一步减少内存需求 # lora_config = LoraConfig( # r=8, # lora_alpha=32, # target_modules=["q_proj", "v_proj"], # lora_dropout=0.05, # bias="none", # task_type="SEQ_CLS" # ) # model = get_peft_model(model, lora_config) # model.print_trainable_parameters() device = torch.device("mps") # 量化模型已自动分配,无需手动to(device) # 分词:降低序列长度到512 train_encodings1 = tokenizer(list(X_train), truncation=True, padding=True, max_length=512) test_encodings1 = tokenizer(list(X_test), truncation=True, padding=True, max_length=512) # 构建Dataset与DataLoader train_dataset1 = TensorDataset( torch.tensor(train_encodings1['input_ids']), torch.tensor(train_encodings1['attention_mask']), torch.tensor(y_train) ) test_dataset1 = TensorDataset( torch.tensor(test_encodings1['input_ids']), torch.tensor(test_encodings1['attention_mask']), torch.tensor(y_test) ) # 适当提升batch_size,结合梯度累积 batch_size = 4 gradient_accumulation_steps = 4 # 等效于batch_size=16 train_loader1 = DataLoader(train_dataset1, batch_size=batch_size, shuffle=True) test_loader1 = DataLoader(test_dataset1, batch_size=batch_size, shuffle=False) # 优化器与损失函数 optimizer = torch.optim.AdamW(model.parameters(), lr=2e-5) criterion = torch.nn.CrossEntropyLoss() # 训练流程 num_epochs = 3 for epoch in range(num_epochs): model.train() print(time.ctime()) start = time.time() total_loss = 0.0 step_count = 0 for batch_idx, (input_ids, attention_mask, labels) in enumerate(train_loader1): # 迁移张量到MPS设备 input_ids = input_ids.to(device) attention_mask = attention_mask.to(device) labels = labels.to(device) # 前向传播 outputs = model(input_ids, attention_mask=attention_mask) loss = criterion(outputs.logits, labels) # 梯度累积 loss = loss / gradient_accumulation_steps loss.backward() total_loss += loss.item() * gradient_accumulation_steps step_count += 1 # 累积到指定步数再更新参数 if (batch_idx + 1) % gradient_accumulation_steps == 0: optimizer.step() optimizer.zero_grad() # 清理内存 torch.mps.empty_cache() # 计算平均损失 average_loss = total_loss / len(train_loader1) print(f"Epoch {epoch + 1}, Average Loss: {average_loss:.4f}") stop = time.time() print(f"Training time: {stop - start:.2f}s") # 可选:测试集评估 # model.eval() # test_loss = 0.0 # with torch.no_grad(): # for input_ids, attention_mask, labels in test_loader1: # input_ids = input_ids.to(device) # attention_mask = attention_mask.to(device) # labels = labels.to(device) # outputs = model(input_ids, attention_mask=attention_mask) # loss = criterion(outputs.logits, labels) # test_loss += loss.item() # print(f"Test Loss: {test_loss / len(test_loader1):.4f}")
额外注意事项
- 安装依赖:确保已安装
bitsandbytes、peft库,执行pip install bitsandbytes peft - 环境变量设置:若仍存在内存问题,可临时设置
export PYTORCH_MPS_HIGH_WATERMARK_RATIO=0.9(不建议设为0.0,避免系统崩溃) - LoRA微调:若内存仍紧张,启用代码中注释的LoRA配置,仅训练模型的小部分参数,内存占用可进一步降低
内容的提问来源于stack exchange,提问作者Debalina
相关产品推荐
相关产品推荐

