[v1] support LoRA with FSDPTurbo expert parallelism (#10791)

Co-authored-by: Rose_is_Rosie <2362225813@qq.com>
This commit is contained in:
xvxuopop
2026-08-27 18:50:56 +08:00
committed by GitHub
parent 273a988ebe
commit 7fcf5b3b13
3 changed files with 40 additions and 12 deletions

View File

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

View File

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

View File

@@ -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
if spec is None: # the outer PeftModel for LoRA training and checkpointing.
raise ValueError(f"No FSDPTurbo EP spec is registered for model_type={_get_model_type(model)}.") 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_modules = spec.ep_modules ep_target_model = get_base_model()
model = spec.prepare(model) logger.info_rank0("Resolving FSDPTurbo EP/FSDP plans against the PEFT base model.")
ep_modules = []
if self.ep_size > 1: if self.ep_size > 1:
spec = FSDPTurboEPModelSpec.get(ep_target_model)
if spec is None:
raise ValueError(
f"No FSDPTurbo EP spec is registered for model_type={_get_model_type(ep_target_model)}."
)
ep_modules = spec.ep_modules
ep_target_model = spec.prepare(ep_target_model)
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)))