diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index abc8380562..bcf92cd1f8 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -73,6 +73,8 @@ __all__ = [ "verify_fp8_support_if_applicable", "_get_inference_mode_context_manager", "hf_login", + "is_moe_model", + "get_moe_target_parameters", ] import torch @@ -2365,3 +2367,73 @@ def hf_login(token: Optional[str] = None) -> Optional[str]: except Exception as e: logger.info(f"Failed to login to huggingface using token with error: {e}") return token + + +# ============================================= +# MoE (Mixture of Experts) Detection and LoRA Utilities + +def is_moe_model(model) -> bool: + """ + Detect if a model is a Mixture of Experts (MoE) model. + + Args: + model: The model to check (can be HF model or config) + + Returns: + True if the model is an MoE model, False otherwise + """ + config = getattr(model, "config", model) + num_experts = getattr(config, "num_experts", None) + return num_experts is not None and num_experts > 0 + + +def get_moe_target_parameters(model, target_modules=None) -> Optional[List[str]]: + """ + Get the target_parameters for MoE expert layers if applicable. + + For MoE models, returns the parameter paths for expert weights + (gate_up_proj, down_proj) that should be targeted by PEFT's + target_parameters for LoRA on nn.Parameter. + + Only includes MoE parameters that match what's in target_modules: + - If "down_proj" is in target_modules -> includes "mlp.experts.down_proj" + - If "gate_proj" or "up_proj" is in target_modules -> includes "mlp.experts.gate_up_proj" + + Args: + model: The model to get target parameters for + target_modules: List/tuple of target module names to match against + + Returns: + List of parameter paths for MoE experts, or None if not an MoE model + """ + if not is_moe_model(model): + return None + + config = getattr(model, "config", model) + num_experts = getattr(config, "num_experts", 0) + + # Determine which MoE parameters to include based on target_modules + moe_params = [] + + # Normalize target_modules to a set for efficient lookup + if target_modules is None: + # If no target_modules specified, include all MoE params + target_set = {"gate_proj", "up_proj", "down_proj"} + elif isinstance(target_modules, str): + target_set = {target_modules} + else: + target_set = set(target_modules) if target_modules else set() + + # gate_up_proj combines both gate_proj and up_proj in MoE + if "gate_proj" in target_set or "up_proj" in target_set: + moe_params.append("mlp.experts.gate_up_proj") + + if "down_proj" in target_set: + moe_params.append("mlp.experts.down_proj") + + if moe_params: + print(f"Unsloth: Detected MoE model with {num_experts} experts - enabling LoRA on: {moe_params}") + return moe_params + + return None +# ============================================= diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 1d7695b9aa..deb31ec0dc 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -2600,6 +2600,7 @@ class FastLlamaModel: loftq_config = {}, temporary_location = "_unsloth_temporary_saved_buffers", qat_scheme = None, + target_parameters = None, # For MoE expert layers (nn.Parameter) **kwargs, ): if os.environ.get("UNSLOTH_USE_NEW_MODEL", "0") == "1": @@ -2629,6 +2630,7 @@ class FastLlamaModel: init_lora_weights = init_lora_weights, loftq_config = loftq_config, temporary_location = temporary_location, + target_parameters = target_parameters, **kwargs, ) if os.environ.get("UNSLOTH_ENABLE_FULL_FINETUNING", "0") == "1": @@ -2940,6 +2942,10 @@ class FastLlamaModel: # Does not get lora yet, so get name from model, not base model is_classification = "Classification" in str(type(model)) + # Auto-detect MoE models and populate target_parameters for expert layers + if target_parameters is None: + target_parameters = get_moe_target_parameters(model, target_modules) + arguments = dict( r = r, lora_alpha = lora_alpha, @@ -2952,6 +2958,7 @@ class FastLlamaModel: loftq_config = loftq_config, use_rslora = use_rslora, modules_to_save = modules_to_save, + target_parameters = target_parameters, **kwargs, ) if not SUPPORTS_LOFTQ: diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index f39acc20e4..795a7f22b5 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -937,6 +937,7 @@ class FastBaseModel: task_type = TaskType.CAUSAL_LM, temporary_location = "_unsloth_temporary_saved_buffers", qat_scheme = None, + target_parameters = None, # For MoE expert layers (nn.Parameter) **kwargs, ): if os.environ.get("UNSLOTH_ENABLE_FULL_FINETUNING", "0") == "1": @@ -1017,6 +1018,10 @@ class FastBaseModel: loftq_config, lora_dropout, bias, init_lora_weights, model ) + # Auto-detect MoE models and populate target_parameters for expert layers + if target_parameters is None: + target_parameters = get_moe_target_parameters(model, target_modules) + # Get only allowed parameters for LoraConfig local_variables = { **locals(),