Fixes for VL MoE and v5 transformers

This commit is contained in:
Datta Nimmaturi 2026-01-18 12:19:26 +00:00
commit efd9d3ba5e
3 changed files with 16 additions and 5 deletions

View file

@ -2384,6 +2384,10 @@ def is_moe_model(model) -> bool:
"""
config = getattr(model, "config", model)
num_experts = getattr(config, "num_experts", None)
# Check text_config for VL models
if num_experts is None and hasattr(config, "text_config"):
num_experts = getattr(config.text_config, "num_experts", None)
return num_experts is not None and num_experts > 0
@ -2410,7 +2414,9 @@ def get_moe_target_parameters(model, target_modules=None) -> Optional[List[str]]
return None
config = getattr(model, "config", model)
num_experts = getattr(config, "num_experts", 0)
num_experts = getattr(config, "num_experts", None)
if num_experts is None and hasattr(config, "text_config"):
num_experts = getattr(config.text_config, "num_experts", 0)
# Determine which MoE parameters to include based on target_modules
moe_params = []
@ -2418,14 +2424,18 @@ def get_moe_target_parameters(model, target_modules=None) -> Optional[List[str]]
# 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"}
target_set = {"gate_proj", "up_proj", "down_proj", "gate_up_proj"}
elif isinstance(target_modules, str):
target_set = {target_modules}
# Heuristic for regex matching MLPs
if "proj" in target_modules and ("mlp" in target_modules or "ffn" in target_modules):
target_set.update({"gate_proj", "up_proj", "down_proj", "gate_up_proj"})
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:
# Also match "gate_up_proj" directly since users may specify the fused name
if "gate_proj" in target_set or "up_proj" in target_set or "gate_up_proj" in target_set:
moe_params.append("mlp.experts.gate_up_proj")
if "down_proj" in target_set:

View file

@ -676,6 +676,7 @@ class FastModel(FastBaseModel):
qat_scheme = None,
load_in_fp8 = False, # fp8 LoRA (True, False, 'block')
unsloth_tiled_mlp = False,
target_parameters = None, # For MoE expert parameters
*args,
**kwargs,
):

View file

@ -109,7 +109,7 @@ PRE_COMPILE_INFERENCE = [
"gpt_oss",
]
from transformers import GenerationConfig, CompileConfig, HybridCache, AutoConfig
from transformers import GenerationConfig, CompileConfig, AutoConfig
try:
from transformers import PreTrainedConfig
@ -120,7 +120,7 @@ except:
HAS_TORCH_DTYPE = "torch_dtype" in PretrainedConfig.__doc__
from transformers import GenerationConfig, CompileConfig, HybridCache
from transformers import GenerationConfig, CompileConfig
_compile_config = CompileConfig(
fullgraph = False,