mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-07-28 11:46:09 +08:00
[v1] refactor registry plugin structure and params (#10641)
This commit is contained in:
@@ -27,10 +27,9 @@ def test_get_args_from_yaml(tmp_path: Path):
|
||||
model_class: llm
|
||||
kernel_config:
|
||||
name: auto
|
||||
include_kernels: auto # choice: null/true/false/auto/kernel_id1,kernel_id2,kernel_id3, default is null
|
||||
peft_config:
|
||||
name: lora
|
||||
lora_rank: 0.8
|
||||
r: 8
|
||||
quant_config: null
|
||||
|
||||
### data
|
||||
@@ -60,9 +59,8 @@ def test_get_args_from_yaml(tmp_path: Path):
|
||||
assert data_args.train_dataset == "llamafactory/v1-sft-demo"
|
||||
assert model_args.model == "llamafactory/tiny-random-qwen3"
|
||||
assert model_args.kernel_config.name == "auto"
|
||||
assert model_args.kernel_config.get("include_kernels") == "auto"
|
||||
assert model_args.peft_config.name == "lora"
|
||||
assert model_args.peft_config.get("lora_rank") == 0.8
|
||||
assert model_args.peft_config.get("r") == 8
|
||||
assert training_args.output_dir == "outputs/test_run"
|
||||
assert training_args.micro_batch_size == 1
|
||||
assert training_args.global_batch_size == 1
|
||||
|
||||
@@ -30,9 +30,7 @@ def test_tiny_qwen():
|
||||
def test_tiny_qwen_with_kernel_plugin():
|
||||
from llamafactory.v1.plugins.model_plugins.kernels.ops.rms_norm.npu_rms_norm import npu_rms_norm_forward
|
||||
|
||||
model_args = ModelArguments(
|
||||
model="llamafactory/tiny-random-qwen3", kernel_config={"name": "auto", "include_kernels": "auto"}
|
||||
)
|
||||
model_args = ModelArguments(model="llamafactory/tiny-random-qwen3", kernel_config={"name": "auto"})
|
||||
model_engine = ModelEngine(model_args)
|
||||
# test enable apply kernel plugin
|
||||
if hasattr(torch, "npu"):
|
||||
|
||||
@@ -30,13 +30,13 @@ def _apply_kernel(rank) -> None:
|
||||
if k.startswith("llamafactory.v1.plugins.model_plugins.kernels"):
|
||||
del sys.modules[k]
|
||||
|
||||
from llamafactory.v1.plugins.model_plugins.kernels.interface import apply_default_kernels
|
||||
from llamafactory.v1.plugins.model_plugins.kernels.interface import apply_kernels
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained("llamafactory/tiny-random-qwen3")
|
||||
original_rmsnorm_forward = model.model.layers[0].input_layernorm.forward
|
||||
original_swiglu_forward = model.model.layers[0].mlp.forward
|
||||
|
||||
model = apply_default_kernels(model=model, include_kernels="npu_fused_rmsnorm")
|
||||
model = apply_kernels(model=model, config={"name": "npu_fused_rmsnorm"})
|
||||
|
||||
assert model.model.layers[0].input_layernorm.forward.__func__ is not original_rmsnorm_forward.__func__
|
||||
assert model.model.layers[0].mlp.forward.__func__ is original_swiglu_forward.__func__
|
||||
@@ -53,13 +53,13 @@ def _apply_all_kernels(rank) -> None:
|
||||
if k.startswith("llamafactory.v1.plugins.model_plugins.kernels"):
|
||||
del sys.modules[k]
|
||||
|
||||
from llamafactory.v1.plugins.model_plugins.kernels.interface import apply_default_kernels
|
||||
from llamafactory.v1.plugins.model_plugins.kernels.interface import apply_kernels
|
||||
|
||||
model = AutoModelForCausalLM.from_pretrained("llamafactory/tiny-random-qwen3")
|
||||
original_rmsnorm_forward = model.model.layers[0].input_layernorm.forward
|
||||
original_swiglu_forward = model.model.layers[0].mlp.forward
|
||||
|
||||
model = apply_default_kernels(model=model, include_kernels=True)
|
||||
model = apply_kernels(model=model, config={"name": "auto"})
|
||||
|
||||
assert model.model.layers[0].input_layernorm.forward.__func__ is not original_rmsnorm_forward.__func__
|
||||
assert model.model.layers[0].mlp.forward.__func__ is not original_swiglu_forward.__func__
|
||||
|
||||
@@ -18,6 +18,7 @@ import torch.multiprocessing as mp
|
||||
|
||||
from llamafactory.v1.accelerator.interface import DistributedInterface
|
||||
from llamafactory.v1.config.model_args import ModelArguments
|
||||
from llamafactory.v1.config.training_args import TrainingArguments
|
||||
from llamafactory.v1.core.model_engine import ModelEngine
|
||||
from llamafactory.v1.plugins.model_plugins.parallelization.sequence_parallel import (
|
||||
SequenceParallelModelPlugin,
|
||||
@@ -33,15 +34,14 @@ def _test_sequence_parallel_loss(
|
||||
with dist_env(local_rank, world_size, master_port):
|
||||
model_args = ModelArguments(model="llamafactory/tiny-random-qwen3")
|
||||
|
||||
# Initialize distributed interface with config
|
||||
dist_config = {"cp_mode": "ulysses", "cp_size": cp_size, "dp_size": dp_size}
|
||||
DistributedInterface(dist_config)
|
||||
training_args = TrainingArguments(cp_mode="ulysses", cp_size=cp_size, dp_size=dp_size)
|
||||
DistributedInterface(training_args)
|
||||
|
||||
# Now create model engine
|
||||
model_engine = ModelEngine(model_args=model_args)
|
||||
|
||||
# Apply sequence parallel plugin
|
||||
SequenceParallelModelPlugin(dist_config.get("cp_mode", "ulysses"))(model_engine.model, dist_config)
|
||||
SequenceParallelModelPlugin(training_args.cp_mode)(model_engine.model, training_args.cp_size)
|
||||
|
||||
input_ids = torch.arange(1, batch_size * 5 + 1, dtype=torch.long).view(batch_size, 5)
|
||||
model_inputs = {
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -33,7 +33,6 @@ template: qwen3_nothink
|
||||
|
||||
kernel_config:
|
||||
name: auto
|
||||
include_kernels: auto
|
||||
|
||||
quant_config: null
|
||||
|
||||
|
||||
@@ -30,7 +30,6 @@ model_class: llm
|
||||
|
||||
kernel_config:
|
||||
name: auto
|
||||
include_kernels: auto
|
||||
|
||||
quant_config: null
|
||||
|
||||
|
||||
Reference in New Issue
Block a user