mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
weight decay param names (now ignore all layernorms + embeddings + bias)
This commit is contained in:
parent
64f7b3d401
commit
1494ad2839
3 changed files with 30 additions and 9 deletions
|
|
@ -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()},
|
||||
]
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue