跳到主要內容
AI News HubLIVE
來源內容 · 翻譯待補全6 分鐘閱讀

待翻譯:Build Your Own Post-training Pipeline

文章摘要

AI 服務暫時不可用,以下為來源摘要,待恢復後補全翻譯:This is the final post in a four-part series about post-training. If you missed them, check out part 1, part 2, and part 3. Time to get your hands dirty! I’ll take you through implementing the key pieces of the classic ChatGPT pipeline: SFT, then reward model training, then PPO. The goal isn’t to reproduce […]

來源O'Reilly AI & ML Radar作者: Sharon Zhou
待翻譯:Build Your Own Post-training Pipeline
回報錯誤

更正管道尚未開通,可先複製下方文章資訊留存。

查看更正說明
直接讀正文

AI 服務暫時不可用,以下為來源正文,待恢復後補全翻譯。

This is the final post in a four-part series about post-training. If you missed them, check out part 1, part 2, and part 3. Time to get your hands dirty! I’ll take you through implementing the key pieces of the classic ChatGPT pipeline: SFT, then reward model training, then PPO. The goal isn’t to reproduce InstructGPT, because that took a large team and thousands of GPU-hours. But getting hands-on will help make the concepts from this series concrete, so that when you look at each stage’s code, you understand what it’s doing to the model and why. We’ll use Qwen2.5-1.5B as the base model. It’s small enough to train on a single node with a few GPUs but large enough that you can observe real behavioral changes from post-training. For SFT, you’ll use torchtune, Meta’s PyTorch-native fine-tuning library. For RL, you’ll use verl (from ByteDance’s team), a production-grade RL post-training framework built on Ray. verl was designed from the ground up for distributed RL training with proper rollout orchestration, and it’s what you’d actually reach for if you were doing this at scale. You’ll need a machine with at least 2 GPUs (4 or 8 is better) and around 80 GB of total GPU memory. You’ll likely use an LLM to flesh out the boilerplate and get this running end to end. I used LLMs pointed to the recent docs to generate these. What matters as you read the code in this section is that you understand the config choices and data formats well enough to direct that process and debug when something goes wrong. Stage 1: SFT on demonstrations SFT is the easiest stage to implement. You’re taking the pretrained model and training it on {prompt, response} pairs using the same next-token prediction loss from pretraining, except you only compute loss on the response tokens. Your SFT data should be in a conversational format. Here’s what a single example looks like: { "messages": [ {"role": "user", "content": "Why do people like golden retrievers?"}, {"role": "assistant", "content": "Golden retrievers are one of the most popular dog breeds for good reason. They have a gentle, friendly temperament that makes them great family dogs. They're patient with kids and tend to get along well with other animals. They're also highly trainable and eager to please, which is why they're commonly used as service dogs and therapy dogs. And honestly, the goofy enthusiasm doesn't hurt either. They never really grow out of that puppy energy."} ] } Remember what a pretrained model does with “Why do people like golden retrievers?” It might spit out “Why do people like labrador retrievers?” because it’s just completing patterns. Your SFT data is what teaches it to actually answer the question instead. In practice, you’d have thousands of these examples covering a range of tasks: question answering, summarization, creative writing, coding help, multiturn dialogue, and refusals for harmful requests. For this walkthrough, assume you have a JSONL file of these conversations. torchtune uses YAML configs and built-in recipes. Here’s what a config might look like: # sft_config.yaml # Model model: _component_: torchtune.models.qwen2_5.lora_qwen2_5_1_5b lora_attn_modules: ['q_proj', 'k_proj', 'v_proj', 'output_proj'] lora_rank: 64 lora_alpha: 128 # Tokenizer tokenizer: _component_: torchtune.models.qwen2_5.qwen2_5_tokenizer path: /path/to/Qwen2.5-1.5B/vocab.json merges_file: /path/to/Qwen2.5-1.5B/merges.txt max_seq_len: 2048 # Checkpointer — loads the pretrained weights checkpointer: _component_: torchtune.training.FullModelHFCheckpointer checkpoint_dir: /path/to/Qwen2.5-1.5B output_dir: ./sft_checkpoint model_type: QWEN2 # Dataset dataset: _component_: torchtune.datasets.chat_dataset source: sft_data.jsonl conversation_style: sharegpt max_seq_len: 2048 train_on_input: false # only compute loss on the assistant's response tokens # Training seed: 42 epochs: 3 batch_size: 4 gradient_accumulation_steps: 4 # effective batch size of 16 optimizer: _component_: torch.optim.AdamW lr: 2e-5 weight_decay: 0.01 lr_scheduler: _component_: torchtune.training.lr_schedulers.get_cosine_schedule_with_warmup num_warmup_steps: 50 dtype: bf16 compile: false Then launch it: tune run lora_finetune_distributed --config sft_config.yaml What actually matters in this config: train_on_input: false masks the loss on prompt tokens so the model only learns to generate good responses, not to mimic user messages. If you accidentally set this to true, the model wastes capacity learning to produce prompts. A learning rate of 2e-5 is the standard starting point for SFT. That’s aggressive enough to shift behavior in a few epochs but not so large that you destroy the pretrained weights. Use LoRA instead of full fine-tuning. For a 1.5B model you could do either, but LoRA is the practical default because it’s faster, uses a lot less GPU memory, and at this model size the quality gap is negligible. At larger model sizes, LoRA becomes even more useful in saving compute. 3 epochs because the dataset is small. InstructGPT trained for 16 epochs on ~13K examples. Smaller datasets need more passes, but watch validation loss for overfitting. This is something you should tune to see different results across different-sized datasets. After this stage, try chatting with the model and try a bunch of different comparisons to the base model. Ask it “Why do people like golden retrievers?” and you should get a real answer now, not a list of related questions. The model should be able to hold a conversation and follow basic instructions. It’s already dramatically more useful than the base model, even if the responses aren’t always great. Realistically, you should set up a good evaluation (“evals”) to assess the quality of the model, hyperparameter tune, and determine your data mix. In this section, you’ll just focus on looking at the code. Stage 2: Training the reward model For your reward model, start with your data of preference pairs, which will look like this: { "prompt": "What's the capital of that country that celebrates with a lot of colored powders?", "chosen": "You're probably thinking of Holi, the festival where people throw bright colored powders. That celebration is most famously associated with India. The capital of India is New Delhi.", "rejected": "The capital of that country that celebrates with a lot of colored flowers is Amsterdam, the Netherlands, known for its tulip festivals." } The “rejected” response isn’t just wrong about the festival. It also misread “powders” as “flowers” and jumped to an incorrect answer. For reward model training, you’ll use TRL’s RewardTrainer. Reward model training is a straightforward classification task. from transformers import AutoModelForSequenceClassification, AutoTokenizer from trl import RewardTrainer, RewardConfig from datasets import load_dataset # Start from the SFT checkpoint — it already understands the response distribution model = AutoModelForSequenceClassification.from_pretrained( "./sft_checkpoint", num_labels=1, torch_dtype="bfloat16", ) tokenizer = AutoTokenizer.from_pretrained("./sft_checkpoint") tokenizer.pad_token = tokenizer.eos_token dataset = load_dataset("json", data_files="preference_data.jsonl", split="train") training_args = RewardConfig( output_dir="./reward_model", per_device_train_batch_size=8, num_train_epochs=1, # just 1 epoch — reward models overfit fast learning_rate=1e-5, # lower than SFT, be gentle bf16=True, max_length=2048, logging_steps=10, save_strategy="epoch", ) trainer = RewardTrainer( model=model, args=training_args, train_dataset=dataset, tokenizer=tokenizer, ) trainer.train() trainer.save_model("./reward_model/final") As you skim the code, take note of these three things: Start from an SFT checkpoint, not a base model. The reward model needs to understand the distribution of responses it’ll be scoring, and the SFT model is closer to that distribution. 1 epoch only. Reward models overfit quickly. The InstructGPT team found only 1 epoch was needed. Learning rate of 1e-5, lower than SFT. Smaller updates. After training, check the reward model by scoring a few responses manually. Give it a clearly good response and a clearly bad one for the same prompt and make sure the good one gets a higher score. If this basic test fails, something is wrong with your data or training. Don’t skip this step: It’s better to catch problems sooner than many GPU hours later! Stage 3: PPO with verl Now you can take the SFT model and optimize it against the reward model using PPO. verl handles the hard parts: coordinating rollout generation across workers, managing the four models that need to be in memory simultaneously (policy, reference, reward, critic), and orchestrating the update loop. First, prepare your prompts. verl expects parquet format. Each row needs a prompt field containing the tokenized and chat-template-formatted prompt. Here’s an example of preparing it: import pandas as pd from transformers import AutoTokenizer from datasets import load_dataset tokenizer = AutoTokenizer.from_pretrained("./sft_checkpoint") raw_prompts = load_dataset("json", data_files="rl_prompts.jsonl", split="train") def format_prompt(example): messages = [{"role": "user", "content": example["prompt"]}] formatted = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=True ) return {"prompt": formatted} formatted = raw_prompts.map(format_prompt) df = pd.DataFrame(formatted) df.to_parquet("data/rl_prompts.parquet") The prompts for RL don’t need labels. The reward model provides the signal. You just need a diverse set of prompts that covers the types of requests your model expects to see when you deploy it. Next, define the reward function. verl lets you wrap your reward model in a function that takes a batch of prompts and responses and returns rewards: # reward_fn.py — verl will call this during training import torch from transformers import AutoModelForSequenceClassification, AutoTokenizer class RewardFunction: def init(self, reward_model_path="./reward_model/final"): self.model = AutoModelForSequenceClassification.from_pretrained( reward_model_path, torch_dtype=torch.bfloat16 ).cuda().eval() self.tokenizer = AutoTokenizer.from_pretrained(reward_model_path) self.tokenizer.pad_token = self.tokenizer.eos_token def call(self, prompts, responses): """ prompts: list of prompt strings responses: list of response strings Returns: list of scalar rewards """ texts = [p + r for p, r in zip(prompts, responses)] inputs = self.tokenizer( texts, return_tensors="pt", padding=True, truncation=True, max_length=2048 ).to(self.model.device) with torch.no_grad(): rewards = self.model(**inputs).logits.squeeze(-1) return rewards.tolist() Now configure the PPO training run. verl uses YAML configuration files that specify the models, hyperparameters, and infrastructure: # ppo_config.yaml data: train_files: data/rl_prompts.parquet prompt_key: prompt max_prompt_length: 1024 max_response_length: 1024 actor_rollout_ref: model: path: ./sft_checkpoint actor: optim: lr: 1e-6 # very low — RL updates should be gentle lr_warmup_steps: 10 ppo_mini_batch_size: 64 ppo_micro_batch_size: 8 # adjust based on GPU memory ppo_epochs: 4 # number of PPO update passes per batch clip_ratio: 0.2 # PPO clipping — standard value kl_penalty_coeff: 0.1 # KL penalty to prevent reward hacking entropy_coeff: 0.01 # small entropy bonus for exploration rollout: temperature: 0.7 top_p: 0.9 n: 1 # 1 response per prompt per rollout tensor_model_parallel_size: 1 ref: log_prob_micro_batch_size: 8 critic: model: path: ./reward_model/final # initialize critic from reward model optim: lr: 1e-5 ppo_micro_batch_size: 8 reward_model: path: ./reward_model/final micro_batch_size: 8 trainer: total_training_steps: 500 save_freq: 100 test_freq: 50 project_name: post_training_walkthrough logger: wandb A few important hyperparameters in the config: The learning rate is an [truncated for AI cost control]

展開要點與分析

文章情報

工程師進階

要點

  • AI 服務暫時不可用,系統已先保留來源內容與降級後設資料。
  • This is the final post in a four-part series about post-training. If you missed them, check out part 1, part 2, and part 3. Time to get your hands dirty! I’ll take you through imp…

要點與分析由自動化流程生成,可能有誤,請結合原始來源核實。