mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-09-06 23:35:46 +08:00
[v1] support LoRA with FSDPTurbo expert parallelism (#10791)
Co-authored-by: Rose_is_Rosie <2362225813@qq.com>
This commit is contained in:
@@ -307,6 +307,11 @@ class BaseTrainer:
|
|||||||
self.model.parameters(), self.args.max_grad_norm, total_norm
|
self.model.parameters(), self.args.max_grad_norm, total_norm
|
||||||
)
|
)
|
||||||
grad_norm = total_norm.item()
|
grad_norm = total_norm.item()
|
||||||
|
# Do not retain a full generation of gradient tensors across optimizer
|
||||||
|
# steps. ``zero_grad(set_to_none=True)`` clears ``param.grad``, but this
|
||||||
|
# local list would otherwise keep every old gradient alive until the next
|
||||||
|
# assignment, doubling gradient memory during the following backward.
|
||||||
|
del grads
|
||||||
|
|
||||||
if not torch.isfinite(torch.tensor(grad_norm)): # type: ignore # pyright: ignore [reportUnknownReturnType]
|
if not torch.isfinite(torch.tensor(grad_norm)): # type: ignore # pyright: ignore [reportUnknownReturnType]
|
||||||
logger.warning_rank0(f"Gradient norm is not finite: {grad_norm}")
|
logger.warning_rank0(f"Gradient norm is not finite: {grad_norm}")
|
||||||
|
|||||||
@@ -94,6 +94,11 @@ def _make_norms_dtype_safe(model: HFModel) -> int:
|
|||||||
return n
|
return n
|
||||||
|
|
||||||
|
|
||||||
|
def is_lora_model(model: HFModel) -> bool:
|
||||||
|
"""Return whether PEFT LoRA layers have already been injected into the model."""
|
||||||
|
return any(isinstance(module, LoraLayer) for module in model.modules())
|
||||||
|
|
||||||
|
|
||||||
def get_transformer_layer_cls(model: HFModel) -> set[type[nn.Module]]:
|
def get_transformer_layer_cls(model: HFModel) -> set[type[nn.Module]]:
|
||||||
classes: set[type[nn.Module]] = set()
|
classes: set[type[nn.Module]] = set()
|
||||||
for module in model.modules():
|
for module in model.modules():
|
||||||
@@ -123,7 +128,8 @@ def save_model(model: HFModel, output_dir: str, processor: Processor) -> None:
|
|||||||
if DistributedInterface().get_rank() == 0:
|
if DistributedInterface().get_rank() == 0:
|
||||||
logger.info("Gathering state dict for saving...")
|
logger.info("Gathering state dict for saving...")
|
||||||
|
|
||||||
options = StateDictOptions(full_state_dict=True, cpu_offload=True)
|
lora_model = is_lora_model(model)
|
||||||
|
options = StateDictOptions(full_state_dict=True, cpu_offload=True, ignore_frozen_params=lora_model)
|
||||||
state_dict = get_model_state_dict(model, options=options)
|
state_dict = get_model_state_dict(model, options=options)
|
||||||
|
|
||||||
if DistributedInterface().get_rank() == 0:
|
if DistributedInterface().get_rank() == 0:
|
||||||
@@ -151,7 +157,8 @@ def save_checkpoint(model: HFModel, optimizer: torch.optim.Optimizer, ckpt_dir:
|
|||||||
if DistributedInterface().get_rank() == 0:
|
if DistributedInterface().get_rank() == 0:
|
||||||
logger.info("Gathering state dict for saving additional HF format checkpoint...")
|
logger.info("Gathering state dict for saving additional HF format checkpoint...")
|
||||||
|
|
||||||
hf_options = StateDictOptions(full_state_dict=True, cpu_offload=True)
|
lora_model = is_lora_model(model)
|
||||||
|
hf_options = StateDictOptions(full_state_dict=True, cpu_offload=True, ignore_frozen_params=lora_model)
|
||||||
hf_state_dict = get_model_state_dict(model, options=hf_options)
|
hf_state_dict = get_model_state_dict(model, options=hf_options)
|
||||||
|
|
||||||
if DistributedInterface().get_rank() == 0:
|
if DistributedInterface().get_rank() == 0:
|
||||||
@@ -218,7 +225,7 @@ class FSDP2Engine:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def is_lora_module_wrap(self, model) -> bool:
|
def is_lora_module_wrap(self, model) -> bool:
|
||||||
return any(isinstance(module, LoraLayer) for module in model.modules())
|
return is_lora_model(model)
|
||||||
|
|
||||||
def prepare_model(self, model: HFModel, ignored_params: set[nn.Parameter] | None = None) -> HFModel:
|
def prepare_model(self, model: HFModel, ignored_params: set[nn.Parameter] | None = None) -> HFModel:
|
||||||
if self.fsdp_mesh is None:
|
if self.fsdp_mesh is None:
|
||||||
|
|||||||
@@ -271,7 +271,7 @@ class FSDPTurboFSDP2Engine(FSDP2Engine):
|
|||||||
|
|
||||||
Design:
|
Design:
|
||||||
- FSDPTurbo owns EP / EFSDP only.
|
- FSDPTurbo owns EP / EFSDP only.
|
||||||
- LlamaFactory owns FSDP / CP / init-load lifecycle.
|
- LlamaFactory owns PEFT / FSDP / CP / init-load / checkpoint lifecycle.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, dist_config: dict, bf16: bool = False):
|
def __init__(self, dist_config: dict, bf16: bool = False):
|
||||||
@@ -361,14 +361,27 @@ class FSDPTurboFSDP2Engine(FSDP2Engine):
|
|||||||
from fsdp_turbo.fsdp_turbo_config import EPPlanConfig, FSDPPlanConfig
|
from fsdp_turbo.fsdp_turbo_config import EPPlanConfig, FSDPPlanConfig
|
||||||
from fsdp_turbo.utils.str_match import module_name_match
|
from fsdp_turbo.utils.str_match import module_name_match
|
||||||
|
|
||||||
spec = FSDPTurboEPModelSpec.get(model)
|
# Resolve FSDPTurbo plans on the PEFT base model while preserving
|
||||||
|
# the outer PeftModel for LoRA training and checkpointing.
|
||||||
|
ep_target_model = model
|
||||||
|
if self.is_lora_module_wrap(model):
|
||||||
|
get_base_model = getattr(model, "get_base_model", None)
|
||||||
|
if get_base_model is None:
|
||||||
|
raise RuntimeError("FSDPTurbo could not access the base model from the LoRA-wrapped model.")
|
||||||
|
|
||||||
|
ep_target_model = get_base_model()
|
||||||
|
logger.info_rank0("Resolving FSDPTurbo EP/FSDP plans against the PEFT base model.")
|
||||||
|
|
||||||
|
ep_modules = []
|
||||||
|
if self.ep_size > 1:
|
||||||
|
spec = FSDPTurboEPModelSpec.get(ep_target_model)
|
||||||
if spec is None:
|
if spec is None:
|
||||||
raise ValueError(f"No FSDPTurbo EP spec is registered for model_type={_get_model_type(model)}.")
|
raise ValueError(
|
||||||
|
f"No FSDPTurbo EP spec is registered for model_type={_get_model_type(ep_target_model)}."
|
||||||
|
)
|
||||||
|
|
||||||
ep_modules = spec.ep_modules
|
ep_modules = spec.ep_modules
|
||||||
model = spec.prepare(model)
|
ep_target_model = spec.prepare(ep_target_model)
|
||||||
|
|
||||||
if self.ep_size > 1:
|
|
||||||
ep_plan = EPPlanConfig(
|
ep_plan = EPPlanConfig(
|
||||||
apply_modules=ep_modules,
|
apply_modules=ep_modules,
|
||||||
dispatcher=self.dist_config.get("ep_dispatcher", "eager"),
|
dispatcher=self.dist_config.get("ep_dispatcher", "eager"),
|
||||||
@@ -394,13 +407,13 @@ class FSDPTurboFSDP2Engine(FSDP2Engine):
|
|||||||
logger.info(f"FSDPTurbo EP device mesh: {ep_mesh}")
|
logger.info(f"FSDPTurbo EP device mesh: {ep_mesh}")
|
||||||
logger.info(f"FSDPTurbo EP gradient divide factor: {ep_plan.gradient_divide_factor}")
|
logger.info(f"FSDPTurbo EP gradient divide factor: {ep_plan.gradient_divide_factor}")
|
||||||
|
|
||||||
model = expert_parallelize_modules(model, ep_mesh, ep_plan)
|
ep_target_model = expert_parallelize_modules(ep_target_model, ep_mesh, ep_plan)
|
||||||
|
|
||||||
if self.ep_fsdp_size > 1:
|
if self.ep_fsdp_size > 1:
|
||||||
if self.rank == 0:
|
if self.rank == 0:
|
||||||
logger.info(f"FSDPTurbo EFSDP apply patterns: {ep_plan.apply_efsdp_modules}")
|
logger.info(f"FSDPTurbo EFSDP apply patterns: {ep_plan.apply_efsdp_modules}")
|
||||||
logger.info(f"FSDPTurbo EFSDP device mesh: {efsdp_mesh}")
|
logger.info(f"FSDPTurbo EFSDP device mesh: {efsdp_mesh}")
|
||||||
model = expert_fully_shard_modules(model, efsdp_mesh, ep_plan, fsdp_plan)
|
ep_target_model = expert_fully_shard_modules(ep_target_model, efsdp_mesh, ep_plan, fsdp_plan)
|
||||||
|
|
||||||
# Collect ignored params for the outer FSDP wrap
|
# Collect ignored params for the outer FSDP wrap
|
||||||
fsdp_ignored_modules = list(self.dist_config.get("fsdp_ignored_modules", []))
|
fsdp_ignored_modules = list(self.dist_config.get("fsdp_ignored_modules", []))
|
||||||
@@ -409,7 +422,10 @@ class FSDPTurboFSDP2Engine(FSDP2Engine):
|
|||||||
|
|
||||||
ignored_params = set()
|
ignored_params = set()
|
||||||
if fsdp_ignored_modules:
|
if fsdp_ignored_modules:
|
||||||
for name, module in model.named_modules():
|
# Resolve patterns against the same unwrapped model used by the EP
|
||||||
|
# plan. The collected Parameter objects are shared with the outer
|
||||||
|
# PeftModel, so they can be passed directly to its FSDP2 wrapper.
|
||||||
|
for name, module in ep_target_model.named_modules():
|
||||||
for pattern in fsdp_ignored_modules:
|
for pattern in fsdp_ignored_modules:
|
||||||
if module_name_match(pattern, name):
|
if module_name_match(pattern, name):
|
||||||
ignored_params.update(list(module.parameters(recurse=True)))
|
ignored_params.update(list(module.parameters(recurse=True)))
|
||||||
|
|||||||
Reference in New Issue
Block a user