Fix MoE router LoRA target + attention-only fallback
Browse files- 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
|
| 575 |
-
#
|
| 576 |
-
#
|
| 577 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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 |
-
|
| 607 |
-
|
|
|
|
| 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 |
|