weight decay param names (now ignore all layernorms + embeddings + bias)

This commit is contained in:
51616 2025-07-09 15:57:53 +00:00
parent 64f7b3d401
commit 1494ad2839
3 changed files with 30 additions and 9 deletions

View file

@ -994,7 +994,7 @@ def convert_ctx_prompt_response_to_messages(
[
{"role": "system", "content": system_msg.strip()},
{"role": "user", "content": user_msg.strip()},
{"role": "assistant", "content": response},
{"role": "assistant", "content": response.strip()},
]
)

View file

@ -498,10 +498,10 @@ class Idefics2PerceiverLayer(nn.Module):
self.rms_norm_eps = config.rms_norm_eps
self.is_cross_attn = is_cross_attn
self.input_latents_norm = Idefics2RMSNorm(
self.input_latents_layernorm = Idefics2RMSNorm(
self.hidden_size, eps=self.rms_norm_eps
)
self.input_context_norm = (
self.input_context_layernorm = (
Idefics2RMSNorm(self.hidden_size, eps=self.rms_norm_eps)
if self.is_cross_attn
else torch.nn.Identity()
@ -550,8 +550,8 @@ class Idefics2PerceiverLayer(nn.Module):
"""
residual = latents
latents = self.input_latents_norm(latents)
context = self.input_context_norm(context)
latents = self.input_latents_layernorm(latents)
context = self.input_context_layernorm(context)
latents, self_attn_weights, present_key_value = self.self_attn(
latents=latents,
@ -617,7 +617,7 @@ class Idefics2PerceiverResampler(Idefics2PreTrainedModel):
self.rms_norm_eps = config.rms_norm_eps
# Create Latents for Perceiver
self.latents = nn.Parameter(torch.randn(self.n_latents, self.hidden_size))
self.latents_q = nn.Parameter(torch.randn(self.n_latents, self.hidden_size))
# First block
assert config.num_blocks > 0
@ -662,7 +662,7 @@ class Idefics2PerceiverResampler(Idefics2PreTrainedModel):
# ]
# self.layers = nn.ModuleList(self.layers)
self.norm = Idefics2RMSNorm(self.hidden_size, eps=self.rms_norm_eps)
self.layernorm = Idefics2RMSNorm(self.hidden_size, eps=self.rms_norm_eps)
self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2"
assert self._use_flash_attention_2
@ -680,7 +680,7 @@ class Idefics2PerceiverResampler(Idefics2PreTrainedModel):
# flattened packed sequence
bsz = torch.where(position_ids == 0, 1, 0).sum()
latents = self.latents.unsqueeze(0).expand((bsz, *self.latents.size()))
latents = self.latents_q.unsqueeze(0).expand((bsz, *self.latents_q.size()))
# latent_attention_mask = torch.ones(
# (attention_mask.size(0), latents.size(1)),
@ -826,7 +826,7 @@ class Idefics2PerceiverResampler(Idefics2PreTrainedModel):
# compressed_context = layer_outputs[0]
compressed_context = self.norm(compressed_context)
compressed_context = self.layernorm(compressed_context)
return compressed_context

View file

@ -1,8 +1,10 @@
import logging
from torch import nn
from transformers import Trainer
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
from transformers.trainer import _is_peft_model
from transformers.trainer_pt_utils import get_parameter_names
from transformers.trainer_utils import IntervalStrategy
from ctx_to_lora.modeling.hypernet import ModulatedPretrainedModel
@ -90,6 +92,23 @@ class ModulatedModelTrainer(Trainer):
return (loss, outputs) if return_outputs else loss
def get_decay_parameter_names(model) -> list[str]:
"""
Get all parameter names that weight decay will be applied to.
This function filters out parameters in two ways:
1. By layer type (nn.Embedding)
2. By parameter name patterns (containing 'bias', 'layernorm', 'rmsnorm'
or 'latents_q' [perceiver's latent queries]).
"""
decay_parameters = get_parameter_names(
model,
[nn.Embedding, nn.LayerNorm],
["bias", "layernorm", "rmsnorm", "latents_q"],
)
return decay_parameters
def train_model(
model,
training_args,
@ -123,6 +142,8 @@ def train_model(
training_args.per_device_train_batch_size = 128
trainer = trainer_cls(**trainer_kwargs)
# MONKEY PATCH: remove embedding layers from weight decay
trainer.get_decay_parameter_names = get_decay_parameter_names
# Trainer loads the best model after training
# is done when load_best_model_at_end=True