Taimwe commited on
Commit
1e16000
·
verified ·
1 Parent(s): ba00a08

Fix MoE router LoRA target + attention-only fallback

Browse files
Files changed (1) hide show
  1. train_securecoder.py +24 -9
train_securecoder.py CHANGED
@@ -571,10 +571,14 @@ def print_stats(stats: list[dict], records: list[dict], tokenizer=None) -> None:
571
  # --------------------------------------------------------------------------
572
  # Model + training
573
  # --------------------------------------------------------------------------
574
- # Attention + router only by default: on a 128-expert MoE, adapting every expert
575
- # MLP means ~800M trainable parameters, which dominates VRAM and step time.
576
- # Pass --target-modules all-linear when you want the MLP/expert capacity too.
577
- ATTENTION_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj", "gate"]
 
 
 
 
578
 
579
 
580
  def load_model_and_tokenizer(args):
@@ -591,10 +595,8 @@ def load_model_and_tokenizer(args):
591
  if isinstance(targets, str) and targets != "all-linear":
592
  targets = [t.strip() for t in targets.split(",") if t.strip()]
593
 
594
- model = FastLanguageModel.get_peft_model(
595
- model,
596
  r=args.lora_r,
597
- target_modules=targets,
598
  lora_alpha=args.lora_alpha,
599
  lora_dropout=0.0,
600
  bias="none",
@@ -602,9 +604,22 @@ def load_model_and_tokenizer(args):
602
  random_state=args.seed,
603
  use_rslora=False,
604
  )
 
 
 
 
 
 
 
 
 
 
 
 
605
  trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
606
- log.info("trainable parameters: %s (%.2f%% of the model)",
607
- f"{trainable:,}", 100 * trainable / max(sum(p.numel() for p in model.parameters()), 1))
 
608
  return model, tokenizer
609
 
610
 
 
571
  # --------------------------------------------------------------------------
572
  # Model + training
573
  # --------------------------------------------------------------------------
574
+ # Attention projections only by default, for two reasons:
575
+ # * Qwen3's MoE router is a custom `Qwen3MoeTopKRouter` module, not nn.Linear, so
576
+ # listing it as a LoRA target dies with "Target module ... is not supported".
577
+ # * adapting all 128 experts x 32 layers is ~800M trainable parameters, which
578
+ # dominates VRAM and step time.
579
+ # Use --target-modules all-linear to include the expert MLPs (PEFT then skips the
580
+ # modules it cannot adapt instead of failing).
581
+ ATTENTION_TARGETS = ["q_proj", "k_proj", "v_proj", "o_proj"]
582
 
583
 
584
  def load_model_and_tokenizer(args):
 
595
  if isinstance(targets, str) and targets != "all-linear":
596
  targets = [t.strip() for t in targets.split(",") if t.strip()]
597
 
598
+ peft_kwargs: dict[str, Any] = dict(
 
599
  r=args.lora_r,
 
600
  lora_alpha=args.lora_alpha,
601
  lora_dropout=0.0,
602
  bias="none",
 
604
  random_state=args.seed,
605
  use_rslora=False,
606
  )
607
+
608
+ try:
609
+ model = FastLanguageModel.get_peft_model(model, target_modules=targets, **peft_kwargs)
610
+ except ValueError as exc:
611
+ if "is not supported" not in str(exc) or targets == ATTENTION_TARGETS:
612
+ raise
613
+ log.warning("LoRA targets rejected (%s); retrying with attention projections only",
614
+ str(exc).splitlines()[0][:180])
615
+ model = FastLanguageModel.get_peft_model(
616
+ model, target_modules=ATTENTION_TARGETS, **peft_kwargs
617
+ )
618
+
619
  trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)
620
+ total = sum(p.numel() for p in model.parameters())
621
+ log.info("trainable parameters: %s (%.2f%% of the model)", f"{trainable:,}",
622
+ 100 * trainable / max(total, 1))
623
  return model, tokenizer
624
 
625