diff --git a/src/ctx_to_lora/data/processing.py b/src/ctx_to_lora/data/processing.py index e42c328..b0bbd9e 100644 --- a/src/ctx_to_lora/data/processing.py +++ b/src/ctx_to_lora/data/processing.py @@ -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()}, ] ) diff --git a/src/ctx_to_lora/modeling/idefics2.py b/src/ctx_to_lora/modeling/idefics2.py index 0c7bde8..155c487 100644 --- a/src/ctx_to_lora/modeling/idefics2.py +++ b/src/ctx_to_lora/modeling/idefics2.py @@ -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 diff --git a/src/ctx_to_lora/trainer.py b/src/ctx_to_lora/trainer.py index 65fb425..ef1c382 100644 --- a/src/ctx_to_lora/trainer.py +++ b/src/ctx_to_lora/trainer.py @@ -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