Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
209 changes: 209 additions & 0 deletions tests/engine/test_dense_lora_train_engine.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,209 @@
import os
import shutil
import tempfile
import time

import parametrize
import torch
import torch.distributed as dist
from torch.optim.lr_scheduler import LambdaLR

from transformers import AutoTokenizer
from xtuner._testing import DeterministicDDPTestCase
from xtuner.v1.config import AdamWConfig, FSDPConfig, LRConfig
from xtuner.v1.engine.train_engine import TrainEngine
from xtuner.v1.loss.ce_loss import CELossConfig
from xtuner.v1.model.adapter.lora import LoraConfig
from xtuner.v1.model.base import ModelItem
from xtuner.v1.model.dense.qwen3 import Qwen3Dense8BConfig
from xtuner.v1.model.moe.moe import SequenceContext
from xtuner.v1.utils import pad_to_max_length
from xtuner.v1.utils.device import get_device
from xtuner.v1.utils.test_utils import init_data_mesh


# Qwen3 8B
QWEN3_PATH = os.environ["QWEN3_PATH"]
DEVICE = get_device()


class TestDenseEngine(DeterministicDDPTestCase):
@parametrize.parametrize(
"device,tp_size,sp_size",
[
("cuda", 1, 1),
("cuda", 1, 2),
],
)
def test_dense_engine_train(self, device, tp_size, sp_size):
pg = self.create_pg(device)

dense_cfg = Qwen3Dense8BConfig()
optim_cfg: AdamWConfig = AdamWConfig()
lr_cfg: LRConfig = LRConfig()
fsdp_cfg: FSDPConfig = FSDPConfig(
torch_compile=True,
cpu_offload=False,
tp_size=tp_size,
# hsdp_sharding_size=hsdp_sharding_size,
)

adapter_cfg = LoraConfig(
r=8,
lora_alpha=8,
lora_dropout=0,
target_modules=[
"q_proj",
"k_proj",
"v_proj",
"o_proj",
"gate_proj",
"up_proj",
"down_proj",
],
bias="none",
modules_to_save=["lm_head"],
)
engine = TrainEngine(
model_cfg=dense_cfg,
optim_cfg=optim_cfg,
fsdp_cfg=fsdp_cfg,
adapter_cfg=adapter_cfg,
)
engine.from_hf(hf_path=QWEN3_PATH)

loss_cfg = CELossConfig()

total_steps = 1000
warmup_steps = total_steps * lr_cfg.warmup_ratio

def warmup_fn(x):
return x / warmup_steps if x < warmup_steps else 1

lr_scheduler = LambdaLR(engine.optimizer, warmup_fn)

tok = AutoTokenizer.from_pretrained(QWEN3_PATH)
txt = "根据国际地球自转和参考系服务机构的数据,今年夏天是自2020年以来第六次地球自转加速。7月9日将成为有史以来最短的一天,比平时短1.3到1.6毫秒。 "
input_ids = tok.encode(txt, return_tensors="pt").view(1, -1)
labels = input_ids.clone()
input_ids = input_ids[:, :-1]
labels = labels[:, 1:]
pack_len = 8192 - input_ids.shape[1]
input_ids = pad_to_max_length(input_ids, 0, max_length=8192)
labels = pad_to_max_length(labels, -100, max_length=8192)
losses = []

sp_mesh = None
if sp_size > 1:
data_mesh = init_data_mesh(str(DEVICE), sp_size)
sp_mesh = data_mesh["sp"]

for _ in range(10):
seq_ctx = SequenceContext.from_input_ids((input_ids,), device=DEVICE)
labels = labels.to(DEVICE)
seq_ctx.num_padding = pack_len
if sp_mesh is not None:
seq_ctx = seq_ctx.split(sequence_parallel_mesh=sp_mesh)
loss_ctx = loss_cfg.build(data={"shifted_labels": labels}, sp_mesh=sp_mesh)
loss_ctx = loss_cfg.loss_ctx_cls.build_batches([loss_ctx])[0]
engine_input = [ModelItem(seq_ctx=seq_ctx, loss_ctx={"lm": loss_ctx})]
loss_log = engine.train_step(engine_input)["logs_info"]
grad_norm = engine.clip_grad_norm()
engine.step_optimizer(grad_norm)
lr_scheduler.step()
losses.append(loss_log["reduced_llm_loss"])
losses_ref = [2.57, 2.57, 2.57, 2.57, 2.57, 2.57, 2.56, 2.56, 2.54, 2.53]
for loss, loss_ref in zip(losses, losses_ref):
self.assertTrue(
abs(loss - loss_ref) < 0.02,
f"loss={loss}, loss_ref={loss_ref}, diff={abs(loss - loss_ref)}",
)

torch.cuda.empty_cache()
try:
dist.destroy_process_group(pg)
except Exception:
pass

@parametrize.parametrize(
"device,tp_size,hsdp_sharding_size",
[
("cuda", 1, 8), # todo: test ep8 and hsdp, OOM in 8 gpus
],
)
def test_save_and_load(self, device, tp_size, hsdp_sharding_size):
pg = self.create_pg(device)

temp_dir = tempfile.mkdtemp()
if dist.get_rank() == 0:
temp_dir = [temp_dir]
else:
temp_dir = [None]
dist.broadcast_object_list(temp_dir, src=0)
temp_dir = temp_dir[0]
moe_cfg = Qwen3Dense8BConfig()
optim_cfg: AdamWConfig = AdamWConfig()
fsdp_cfg: FSDPConfig = FSDPConfig(
torch_compile=True,
cpu_offload=False,
tp_size=tp_size,
hsdp_sharding_size=hsdp_sharding_size,
)
adapter_cfg = LoraConfig(
r=4,
lora_alpha=16,
lora_dropout=0,
target_modules=["q_proj", "v_proj"],
bias="none",
modules_to_save=["lm_head"],
)
engine = TrainEngine(
model_cfg=moe_cfg,
optim_cfg=optim_cfg,
fsdp_cfg=fsdp_cfg,
adapter_cfg=adapter_cfg,
)

engine.from_hf(hf_path=QWEN3_PATH)
engine.save_hf(
hf_dir=temp_dir,
save_dtype=torch.bfloat16,
)

dist.barrier()
time.sleep(1)

engine2 = TrainEngine(
model_cfg=moe_cfg,
optim_cfg=optim_cfg,
fsdp_cfg=fsdp_cfg,
adapter_cfg=adapter_cfg,
)
engine2.from_hf(hf_path=temp_dir, strict=True)

state_dict = engine.model.state_dict()
state_dict2 = engine2.model.state_dict()
for key, val in state_dict.items():
if "lora_A." not in key and "lora_B." not in key and "lm_head" not in key:
continue
val2 = state_dict2[key]
val = val.full_tensor().bfloat16()
val2 = val2.full_tensor().bfloat16()
self.assertTrue(torch.equal(val, val2), f"Mismatch in {key} after adapter round trip")

if dist.get_rank() == 0:
shutil.rmtree(temp_dir)

torch.cuda.empty_cache()
try:
dist.destroy_process_group(pg)
except Exception:
pass

@property
def world_size(self) -> int:
return int(os.getenv("XTUNER_TEST_WORLD_SIZE", "8"))

@property
def destroy_pg_upon_exit(self) -> bool:
return False
137 changes: 137 additions & 0 deletions tests/module/test_lora_linear.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
import torch
import torch.nn as nn

from xtuner.v1.model.adapter.lora import LoraConfig, LoraModel
from xtuner.v1.module.grouped_linear.moe_group_linear import GroupedLinear
from xtuner.v1.module.lora_linear.lora_grouped_linear import LoraGroupedLinear
from xtuner.v1.module.lora_linear.lora_linear import LoraLinear
from xtuner.v1.utils.load_spec import LoadEnum, LoadSpec


def test_lora_linear_merge_round_trip():
base_layer = nn.Linear(4, 3, bias=False)
lora_layer = LoraLinear(base_layer, rank=2, alpha=4)
x = torch.randn(5, 4)

torch.testing.assert_close(lora_layer(x), base_layer(x))

nn.init.normal_(lora_layer.lora_A.weight)
nn.init.normal_(lora_layer.lora_B.weight)
expected = lora_layer(x)
base_weight = base_layer.weight.detach().clone()

lora_layer.merge_lora()
torch.testing.assert_close(lora_layer(x), expected)
lora_layer.unmerge_lora()
torch.testing.assert_close(base_layer.weight, base_weight)


def test_lora_linear_can_initialize_after_meta_materialization():
with torch.device("meta"):
lora_layer = LoraLinear(nn.Linear(4, 3, bias=False), rank=2, alpha=4)

lora_layer.to_empty(device="cpu")
lora_layer.reset_parameters()
assert torch.count_nonzero(lora_layer.lora_A.weight) > 0
assert torch.count_nonzero(lora_layer.lora_B.weight) == 0


def test_lora_grouped_linear_merge_round_trip():
base_layer = GroupedLinear(in_features=4, out_features=3, num_routed_experts=2)
lora_layer = LoraGroupedLinear(base_layer, rank=2, alpha=4)

nn.init.normal_(base_layer.weight)
nn.init.normal_(lora_layer.lora_A.weight)
nn.init.normal_(lora_layer.lora_B.weight)
base_weight = base_layer.weight.detach().clone()
a = lora_layer.lora_A.weight.view(2, 2, 4)
b = lora_layer.lora_B.weight.view(2, 3, 2)
expected = base_weight.view(2, 3, 4) + torch.bmm(b, a) * lora_layer.scale

lora_layer.merge_lora()
torch.testing.assert_close(base_layer.weight.view(2, 3, 4), expected)
lora_layer.unmerge_lora()
torch.testing.assert_close(base_layer.weight, base_weight)


class _ToyModel(nn.Module):
fsdp_mesh = None

def __init__(self):
super().__init__()
self.q_proj = nn.Linear(4, 4)
self.lm_head = nn.Linear(4, 4)
self.forward_marker = object()

def to_hf_key_list(self, key: str) -> list[str]:
return [key]

def _init_load_spec(self):
self.load_spec_mapping = {}
for name, param in self.state_dict().items():
hf_key = self.to_hf_key_list(name)[0]
self.load_spec_mapping[name] = LoadSpec(
name=name,
hf_keys=[hf_key],
shape=tuple(param.shape),
load_enum=LoadEnum.SAME,
)

@staticmethod
def _clean_param_name(name: str) -> str:
return name

@staticmethod
def _load_same_hf_param(param, load_spec, checkpoint_loader):
tensor = checkpoint_loader.load(load_spec.hf_keys[0])
if tensor is None:
return load_spec.hf_keys.copy()
with torch.no_grad():
param.copy_(tensor)
return []

@staticmethod
def _load_fused_hf_param(*_):
raise AssertionError("Toy model has no fused parameters")

@staticmethod
def _load_shard_hf_param(*_):
raise AssertionError("Toy model has no sharded parameters")

def set_hf(self, hf_path):
self.hf_path = hf_path

def forwarded_method(self):
return self.forward_marker


def test_lora_model_forwards_base_model_api_and_preserves_modules_to_save():
base_model = _ToyModel()
model = LoraModel(
base_model,
LoraConfig(target_modules=["q_proj", "lm_head"], modules_to_save=["lm_head"]),
)

assert isinstance(base_model.q_proj, LoraLinear)
assert isinstance(base_model.lm_head, nn.Linear)
assert all(param.requires_grad for param in base_model.lm_head.parameters())
assert model.forwarded_method() is base_model.forward_marker


def test_adapter_save_load_round_trip(tmp_path):
config = LoraConfig(target_modules=["q_proj"], modules_to_save=["lm_head"], base_model_name_or_path="base")
model = LoraModel(_ToyModel(), config)
nn.init.normal_(model.base_model.q_proj.lora_A.weight)
nn.init.normal_(model.base_model.q_proj.lora_B.weight)
nn.init.normal_(model.base_model.lm_head.weight)
nn.init.normal_(model.base_model.lm_head.bias)
model._save_hf(tmp_path)

restored = LoraModel(_ToyModel(), config.model_copy(deep=True))
restored._load_adapter(tmp_path, strict=True)

expected = {name: param for name, param in model.base_model.named_parameters() if param.requires_grad}
actual = {name: param for name, param in restored.base_model.named_parameters() if param.requires_grad}
assert expected.keys() == actual.keys()
for name in expected:
torch.testing.assert_close(actual[name].bfloat16(), expected[name].bfloat16(), rtol=0, atol=0)
5 changes: 5 additions & 0 deletions xtuner/v1/engine/train_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
from xtuner.v1.config import FSDPConfig, OptimConfig
from xtuner.v1.data_proto.sequence_context import SequenceContext
from xtuner.v1.loss import LogProbContext
from xtuner.v1.model.adapter.lora import LoraConfig
from xtuner.v1.model.base import (
BaseModel,
BatchForwardInfo,
Expand Down Expand Up @@ -147,10 +148,12 @@ def __init__(
optim_cfg: OptimConfig,
fsdp_cfg: FSDPConfig,
intra_layer_micro_batch: int = 1,
adapter_cfg: LoraConfig | None = None,
) -> None:
self.model_cfg = model_cfg
self.optim_cfg = optim_cfg
self.fsdp_cfg = fsdp_cfg
self.adapter_cfg = adapter_cfg
self.model = self.build_model()
self.optimizer = self.build_optimizer(optim_cfg)
self.intra_layer_micro_batch = intra_layer_micro_batch
Expand All @@ -170,6 +173,8 @@ def __has_freeze_params(self) -> bool:
def build_model(self) -> BaseModel:
with torch.device("meta"):
model = self.model_cfg.build()
if self.adapter_cfg:
model = self.adapter_cfg.build(model)

model = model.fully_shard(self.fsdp_cfg)

Expand Down
4 changes: 4 additions & 0 deletions xtuner/v1/model/adapter/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
from .lora import LoraConfig, LoraModel


__all__ = ["LoraConfig", "LoraModel"]
Loading
Loading