Quest 17 - DPO Loss Implementation
Quest 17: DPO Loss Implementation
hard 30 minutes🎯 Learning Objectives
- Understand DPO (Direct Preference Optimization) as a simplified alternative to RLHF
- Implement the DPO loss function using log-space arithmetic
- Apply numerical stability techniques to prevent overflow/underflow
- Recognize when policy and reference models produce identical outputs
📖 Concept: DPO Loss
DPO (Direct Preference Optimization) เป็นวิธี alignment ที่ปฏิวัติวงการ — มันกำจัด reward model ที่ซับซ้อนออก แล้ว optimize โดยตรงบน preference data
แทนที่จะ train reward model แล้วใช้ RL (RLHF), DPO คำนวณ loss โดยตรงจาก chosen vs rejected responses:
L = -log(σ(β × (log π(y_w|x)/π_ref(y_w|x) - log π(y_l|x)/π_ref(y_l|x))))ที่ไหน:
π= policy model (model ที่กำลัง train)π_ref= reference model (model เดิม)β= temperature parameter (ควบคุมความแรงของการ optimize)y_w= chosen response,y_l= rejected response
คิดเหมือน martial arts: DPO เปรียบเหมือนการตัดสินใจทันทีโดยไม่ต้องคิดผ่าน reward model — เลือก technique ที่ดีกว่าโดยตรง
⚙️ How It Works
The DPO Loss Formula
1. Compute log-probability ratios: ratio_chosen = log π(y_w|x) - log π_ref(y_w|x) ratio_rejected = log π(y_l|x) - log π_ref(y_l|x)
2. Scale by beta: scaled_diff = β × (ratio_chosen - ratio_rejected)
3. Apply sigmoid and log: L = -log(σ(scaled_diff))Why Log-Space Arithmetic?
Direct computation of probability ratios can overflow:
// DANGEROUS: may overflowconst prob = Math.exp(logProb); // exp(-1000) = 0const ratio = prob_chosen / prob_rejected; // 0 / 0 = NaN
// SAFE: log-spaceconst logRatio = logProb_chosen - logProb_rejected; // always finiteThe Sigmoid Function
function sigmoid(x) { return 1 / (1 + Math.exp(-x));}When policy = reference, the loss equals log(2) ≈ 0.693 — this is the “no improvement” baseline.
💡 Example: Implementing DPO Loss
Step 1: Implement the core formula
function dpoLoss(policyLogps, refLogps, beta) { // Log-space ratios const ratioChosen = policyLogps.chosen - refLogps.chosen; const ratioRejected = policyLogps.rejected - refLogps.rejected;
// Scaled difference const scaledDiff = beta * (ratioChosen - ratioRejected);
// DPO loss: -log(σ(scaledDiff)) return -Math.log(1 / (1 + Math.exp(-scaledDiff)));}Step 2: Verify with known values
// When policy = reference → loss = log(2)const loss = dpoLoss( { chosen: -1, rejected: -2 }, { chosen: -1, rejected: -2 }, 1.0);// loss ≈ 0.693 (log 2)Step 3: Run the tests
node test.js# Test: returns positive number# Test: policy=reference gives ~log(2)# Test: higher beta changes loss# Test: handles large log-probs without overflow⚠️ Common Mistakes
Mistake 1: Using exp() instead of log-space
Math.exp(logProb)for large negative logProbs gives 0 → NaN → Always work in log-space: subtract log-probs instead of dividing probs.
Mistake 2: Wrong formula sign
“L = log(σ(…))” instead of “L = -log(σ(…))” → The loss is NEGATIVE log-sigmoid. Forgetting the minus sign gives wrong gradients.
Mistake 3: Ignoring the beta parameter
“I’ll just hardcode beta = 1.0” → Beta controls how strongly you push away from reference. Different tasks need different beta values.
Mistake 4: Not returning a scalar
Returning an object instead of a single number → DPO loss must be a scalar number for backpropagation.
📝 Knowledge Check
📝 Knowledge Check
Q1:What is the DPO loss formula?
Q2:Why should you use log-space arithmetic instead of computing probabilities directly?
Q3:When the policy model matches the reference model exactly, what is the DPO loss?
🏋️ Quest: DPO Loss Implementation
Now it’s time to implement the DPO loss function!
-
Download ไฟล์เริ่มต้นของ quest:
Terminal window npx bluebeltdojo download quest-17-dpo-losscd quest-17-dpo-loss -
เปิด
problem.jsใน editor ของคุณพร้อมความช่วยเหลือของ AI -
Implement
dpoLoss(policyLogps, refLogps, beta)— คำนวณ DPO loss -
สำคัญ: ใช้ log-space arithmetic เพื่อป้องกัน numerical overflow
-
ตรวจสอบ solution ของคุณ:
Terminal window node test.js -
When all tests pass, submit your solution:
Terminal window npx bluebeltdojo submit
💡 Tip: DPO คือ algorithm ที่ทำให้ alignment ง่ายขึ้น — เข้าใจ loss function แล้วจะเข้าใจว่า model เรียนรู้จาก preferences อย่างไร
คำใบ้
- อ่าน instructions ใน
problem.jsอย่างละเอียด - ใช้ log-space arithmetic เสมอ — อย่าใช้ exp() กับ log-probs
- ตรวจสอบ: policy=reference ต้องได้ loss ≈ log(2) ≈ 0.693
- ถ้าติดขัด ลองอ่าน “Common Mistakes” อีกครั้ง — อย่าดู solution โดยตรง