如何微调Meta Llama-3 8B
来源:互联网
时间:2026-08-22 14:44:20
Meta这次发布的Llama 3系列,包含8B和70B两种参数规模的预训练和指令调优模型,确实让开源社区眼前一亮。指令调优版专门为对话场景做了优化,在多项行业基准测试中,表现优于不少开源聊天模型。开发团队在实用性和安全性上下了不少功夫,这一点从最终效果上就能看出来。
目录概览:
微调
开始微调Llama-3 8B
第 1 步:安装库
pip install huggingface_hub ipython:安装两个基础库,huggingface_hub用于从 Hugging Face Hub 访问模型,ipython用于交互式编码。"unsloth[colab] @ git+https://github.com/unslothai/unsloth.git" "unsloth[conda] @git+https://github.com/unslothai/unsloth.git":从 GitHub 安装 Unsloth 库,分别指定了 Google Colab ([colab]) 和 conda 环境 ([conda]) 的选项。export HF_TOKEN=xxxxxxxxxxxxx:设置 Hugging Face Hub 的身份验证令牌,实际值出于安全考虑被隐藏。
pip install huggingface_hub ipython "unsloth[colab] @ git+https://github.com/unslothai/unsloth.git" "unsloth[conda] @git+https://github.com/unslothai/unsloth.git" export HF_TOKEN=xxxxxxxxxxxxx
安装 Wandb
安装 Wandb 库:
pip install wandb- 登录 Wandb:
wandb login,会提示输入 API 密钥,用于跟踪训练过程。
pip install wandb wandb logio
导入库
import os from unsloth import FastLanguageModel import torch from trl import SFTTrainer from transformers import TrainingArguments from datasets import load_dataset
加载数据集
设置最大序列长度:
max_seq_length = 2048,每个训练样本允许的最大 token 数,用于控制内存和计算资源。定义数据 URL:
url指向 JSONL 格式的数据集地址。- 加载数据集:
dataset = load_dataset("json", data_files = {"train" : url}, split = "train"),使用datasets库从 URL 加载训练数据。
load_dataset("json")指定数据格式为 JSON。data_files字典用键 "train" 和 URL 定义训练数据位置。split="train"表示加载训练集。
max_seq_length = 2048
url = "https://huggingface.co/datasets/laion/OIG/resolve/main/unified_chip2.jsonl"
dataset = load_dataset("json", data_files = {"train" : url}, split = "train")
加载 Llama-3-8B
# 2. Load Llama3 model
model, tokenizer = FastLanguageModel.from_pretrained(
model_name = "unsloth/llama-3-8b-bnb-4bit",
max_seq_length = max_seq_length,
dtype = None,
load_in_4bit = True
)
生成文本函数
def generate_text(text):
inputs = tokenizer([text], return_tensors="pt").to("cuda")
outputs = model.generate(**inputs, max_new_tokens=20, use_cache=True)
tokenizer.batch_decode(outputs)
print("Before training
")
设置 LoRA 参数并开始训练
model = FastLanguageModel.get_peft_model(
model,
r = 16,
target_modules = ["q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj",],
lora_alpha = 16,
lora_dropout = 0,
bias = "none",
use_gradient_checkpointing = True,
random_state = 3407,
max_seq_length = max_seq_length,
use_rslora = False,
loftq_config = None,
)
开始训练
trainer = SFTTrainer(
model = model,
train_dataset = dataset,
dataset_text_field = "text",
max_seq_length = max_seq_length,
tokenizer = tokenizer,
args = TrainingArguments(
per_device_train_batch_size = 2,
gradient_accumulation_steps = 4,
warmup_steps = 10,
max_steps = 60,
fp16 = not torch.cuda.is_bf16_supported(),
bf16 = torch.cuda.is_bf16_supported(),
logging_steps = 1,
output_dir = "outputs",
optim = "adamw_8bit",
weight_decay = 0.01,
lr_scheduler_type = "linear",
seed = 3407,
),
)
trainer.train()
测试模型
print("
########
After training
")
generate_text(": List the top 5 most popular movies of all time.
: ")
保存模型
model.sa ve_pretrained("lora_model")
model.sa ve_pretrained_merged("outputs", tokenizer, sa ve_method = "merged_16bit",)
model.push_to_hub_merged("YOURUSERNAME/llama3-8b-oig-unsloth-merged", tokenizer, sa ve_method = "merged_16bit", token = os.environ.get("HF_TOKEN"))
model.push_to_hub("YOURUSERNAME/llama3-8b-oig-unsloth", tokenizer, sa ve_method = "lora", token = os.environ.get("HF_TOKEN")) -
- 关于宇宙的好的网名有哪些
- 角色扮演 | 1
- 网名