[v1] refactor registry plugin structure and params (#10641)

This commit is contained in:
Jiaqi
2026-07-24 15:23:21 +08:00
committed by GitHub
parent 19e9fe3ced
commit 3f77101580
50 changed files with 843 additions and 1007 deletions

View File

@@ -28,6 +28,7 @@ from llamafactory.v1.trainers.dpo_trainer import DPOTrainer, compute_sigmoid_dpo
# Mock helpers
# ==============================================================================
def _make_mock_v1(
pref_beta: float = 0.1,
dpo_label_smoothing: float = 0.0,
@@ -67,6 +68,7 @@ R_REJECTED = torch.tensor([-3.2, -2.7, -4.2, -1.8])
# Test 1 — Core loss correctness (pure function ↔ v1 instance ↔ v0/TRL)
# ==============================================================================
def test_sigmoid_dpo_loss_correctness():
"""Comprehensive correctness check for compute_sigmoid_dpo_loss and its wrapper."""
# ---- 1a: pure function matches instance method ----
@@ -78,7 +80,12 @@ def test_sigmoid_dpo_loss_correctness():
# ---- 1b: v1 matches v0 (TRL) on fixed inputs ----
v0 = _make_mock_v0_dpo(beta=0.1)
v0_losses, _, _ = CustomDPOTrainer.dpo_loss(
v0, P_CHOSEN, P_REJECTED, R_CHOSEN, R_REJECTED, loss_type="sigmoid",
v0,
P_CHOSEN,
P_REJECTED,
R_CHOSEN,
R_REJECTED,
loss_type="sigmoid",
)
torch.testing.assert_close(actual, v0_losses, rtol=1e-6, atol=1e-6)
@@ -87,7 +94,12 @@ def test_sigmoid_dpo_loss_correctness():
v0b = _make_mock_v0_dpo(beta=beta)
v1b = _make_mock_v1(pref_beta=beta)
vl, _, _ = CustomDPOTrainer.dpo_loss(
v0b, P_CHOSEN, P_REJECTED, R_CHOSEN, R_REJECTED, loss_type="sigmoid",
v0b,
P_CHOSEN,
P_REJECTED,
R_CHOSEN,
R_REJECTED,
loss_type="sigmoid",
)
v1l = DPOTrainer._sigmoid_dpo_loss(v1b, P_CHOSEN, P_REJECTED, R_CHOSEN, R_REJECTED)
torch.testing.assert_close(v1l, vl, rtol=1e-6, atol=1e-6)
@@ -97,7 +109,12 @@ def test_sigmoid_dpo_loss_correctness():
v0s = _make_mock_v0_dpo(beta=0.1, label_smoothing=ls)
v1s = _make_mock_v1(pref_beta=0.1, dpo_label_smoothing=ls)
vl, _, _ = CustomDPOTrainer.dpo_loss(
v0s, P_CHOSEN, P_REJECTED, R_CHOSEN, R_REJECTED, loss_type="sigmoid",
v0s,
P_CHOSEN,
P_REJECTED,
R_CHOSEN,
R_REJECTED,
loss_type="sigmoid",
)
v1l = DPOTrainer._sigmoid_dpo_loss(v1s, P_CHOSEN, P_REJECTED, R_CHOSEN, R_REJECTED)
torch.testing.assert_close(v1l, vl, rtol=1e-6, atol=1e-6)
@@ -112,13 +129,17 @@ def test_sigmoid_dpo_loss_correctness():
v1c = _make_mock_v1(pref_beta=0.1)
loss_good = DPOTrainer._sigmoid_dpo_loss(
v1c,
torch.tensor([-1.0]), torch.tensor([-10.0]),
torch.tensor([-3.0]), torch.tensor([-3.0]),
torch.tensor([-1.0]),
torch.tensor([-10.0]),
torch.tensor([-3.0]),
torch.tensor([-3.0]),
)
loss_bad = DPOTrainer._sigmoid_dpo_loss(
v1c,
torch.tensor([-10.0]), torch.tensor([-1.0]),
torch.tensor([-3.0]), torch.tensor([-3.0]),
torch.tensor([-10.0]),
torch.tensor([-1.0]),
torch.tensor([-3.0]),
torch.tensor([-3.0]),
)
assert loss_good.item() < loss_bad.item()
@@ -147,6 +168,7 @@ def test_sigmoid_dpo_loss_correctness():
# Test 2 — Random cross-validation & reward equivalence
# ==============================================================================
def test_cross_validate_and_rewards():
"""Randomised v0↔v1 cross-validation (50 seeds) + reward-margin check."""
torch.manual_seed(42)
@@ -162,7 +184,12 @@ def test_cross_validate_and_rewards():
v1 = _make_mock_v1(pref_beta=beta, dpo_label_smoothing=ls)
v0_loss, _, _ = CustomDPOTrainer.dpo_loss(
v0, pc, pr, rc, rr, loss_type="sigmoid",
v0,
pc,
pr,
rc,
rr,
loss_type="sigmoid",
)
v1_loss = DPOTrainer._sigmoid_dpo_loss(v1, pc, pr, rc, rr)
torch.testing.assert_close(v1_loss, v0_loss, rtol=1e-5, atol=1e-5)
@@ -184,6 +211,7 @@ def test_cross_validate_and_rewards():
# Test 3 — End-to-end: log-prob extraction + synthetic batch + LD-DPO
# ==============================================================================
def _make_batch(num_pairs, seq_len, vocab_size, prompt_len=3, chosen_len=None, rejected_len=None):
if chosen_len is None or rejected_len is None:
rlen = (seq_len - prompt_len) // 2
@@ -198,8 +226,8 @@ def _make_batch(num_pairs, seq_len, vocab_size, prompt_len=3, chosen_len=None, r
labels[:, :prompt_len] = IGNORE_INDEX
token_type_ids = torch.zeros(num_pairs, actual, dtype=torch.long)
token_type_ids[:, prompt_len:prompt_len + chosen_len] = 1
token_type_ids[:, prompt_len + chosen_len:] = 2
token_type_ids[:, prompt_len : prompt_len + chosen_len] = 1
token_type_ids[:, prompt_len + chosen_len :] = 2
torch.manual_seed(99)
logits = torch.randn(num_pairs, actual, vocab_size)
@@ -226,7 +254,12 @@ def test_logp_extraction_and_e2e_loss():
# --- unequal-length (LD-DPO) batch ---
ids2, labels2, tt_ids2, logits2 = _make_batch(
1, 11, 64, prompt_len=2, chosen_len=6, rejected_len=3,
1,
11,
64,
prompt_len=2,
chosen_len=6,
rejected_len=3,
)
v1_ld = _make_mock_v1(pref_beta=0.1, ld_alpha=0.5)

View File

@@ -33,7 +33,6 @@ template: qwen3_nothink
kernel_config:
name: auto
include_kernels: auto
quant_config: null

View File

@@ -30,7 +30,6 @@ model_class: llm
kernel_config:
name: auto
include_kernels: auto
quant_config: null