[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

@@ -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

View File

@@ -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"):

View File

@@ -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__

View File

@@ -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 = {

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