mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-26 17:11:02 +02:00
add extra_modules (for layernorm)
This commit is contained in:
parent
411fb91a90
commit
1665d89eee
3 changed files with 243 additions and 61 deletions
|
|
@ -154,6 +154,10 @@ class TrainingArguments(TrainingArguments):
|
|||
default=0.95,
|
||||
metadata={"help": "Adam beta 2."},
|
||||
)
|
||||
adam_epsilon: float = field(
|
||||
default=1e-6,
|
||||
metadata={"help": "Adam epsilon."},
|
||||
)
|
||||
lr_scheduler_type: str = field(
|
||||
default="cosine_with_min_lr",
|
||||
metadata={"help": "Learning rate scheduler type."},
|
||||
|
|
@ -245,7 +249,7 @@ class LoRAArguments:
|
|||
metadata={"help": ("LoRA R value.")},
|
||||
)
|
||||
lora_dropout: Optional[float] = field(
|
||||
default=0,
|
||||
default=0.02,
|
||||
metadata={"help": ("LoRA dropout.")},
|
||||
)
|
||||
target_modules: Optional[list[str]] = field(
|
||||
|
|
@ -276,6 +280,10 @@ class CtxTrainingArguments:
|
|||
default=2**13,
|
||||
metadata={"help": "Maximum context length for training."},
|
||||
)
|
||||
use_multipack_sampler: bool = field(
|
||||
default=False,
|
||||
metadata={"help": "Whether to use multipack sampler."},
|
||||
)
|
||||
max_new_tokens: Optional[int] = field(
|
||||
default=2**10,
|
||||
metadata={"help": "Maximum new tokens for generation-based evaluation."},
|
||||
|
|
@ -344,10 +352,10 @@ class HypernetArguments:
|
|||
default=0.0,
|
||||
metadata={"help": "Dropout rate for HyperLoRA."},
|
||||
)
|
||||
# trainable_base_modules: Optional[list[str]] = field(
|
||||
# default=None,
|
||||
# metadata={"help": ("Modules to train of the base model.")},
|
||||
# )
|
||||
extra_modules: Optional[list[str]] = field(
|
||||
default=None,
|
||||
metadata={"help": "Extra modules to train."},
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
|
|
|||
|
|
@ -185,3 +185,50 @@ def add_generated_lora_hook(
|
|||
return newoutput
|
||||
|
||||
return apply_hook_to_layers(model, [module_name], [layer_index], post_hook=lora_hook)
|
||||
|
||||
|
||||
def add_generated_layernorm_hook(
|
||||
model: torch.nn.Module,
|
||||
module_name: str,
|
||||
layer_index: int,
|
||||
W: Float[Tensor, "bs hidden_size"],
|
||||
training: bool,
|
||||
) -> list[RemovableHandle]:
|
||||
"""
|
||||
Adds layer normalization hooks to specified modules and layers.
|
||||
|
||||
Args:
|
||||
model (torch.nn.Module): Model to hook
|
||||
module_name (str): Name of layernorm module (e.g., "input_layernorm")
|
||||
layer_index (int): Index of layer to modify
|
||||
W (Tensor): Learned weight tensor of shape [batch_size, hidden_size]
|
||||
training (bool): Whether model is in training mode
|
||||
|
||||
Returns:
|
||||
list[RemovableHandle]: Hook handles for removal
|
||||
"""
|
||||
|
||||
def layernorm_hook(
|
||||
module: torch.nn.Module,
|
||||
args: tuple,
|
||||
output: Float[Tensor, "bs seq_len hidden_size"],
|
||||
) -> Float[Tensor, "bs seq_len hidden_size"]:
|
||||
# For models that return tuples from layernorm (e.g., some attention implementations)
|
||||
if isinstance(output, tuple):
|
||||
main_output = output[0]
|
||||
rest = output[1:]
|
||||
else:
|
||||
main_output = output
|
||||
rest = None
|
||||
|
||||
x = args[0].to(W.dtype)
|
||||
# Apply learned weights to layernorm output
|
||||
# Unsqueeze to add seq_len dimension for broadcasting
|
||||
scaled_output = x * W.unsqueeze(1)
|
||||
new_output = main_output + scaled_output.to(output.dtype)
|
||||
|
||||
return (new_output, *rest) if rest else new_output
|
||||
|
||||
return apply_hook_to_layers(
|
||||
model, [module_name], [layer_index], post_hook=layernorm_hook
|
||||
)
|
||||
|
|
|
|||
|
|
@ -44,7 +44,11 @@ from ctx_to_lora.configs import (
|
|||
HypernetArguments,
|
||||
CtxEncoderArguments,
|
||||
)
|
||||
from ctx_to_lora.hooks import add_generated_lora_hook, remove_hook_handles
|
||||
from ctx_to_lora.hooks import (
|
||||
add_generated_layernorm_hook,
|
||||
add_generated_lora_hook,
|
||||
remove_hook_handles,
|
||||
)
|
||||
from ctx_to_lora.model_loading import get_lora_config, get_model, get_model_and_tokenizer
|
||||
from ctx_to_lora.pooling import POOL_FN, get_pooling_fn
|
||||
from ctx_to_lora.utils import (
|
||||
|
|
@ -86,6 +90,7 @@ def get_aggregator_config(
|
|||
model: PreTrainedModel,
|
||||
ctx_encoder_model_config: PretrainedConfig,
|
||||
output_size: int,
|
||||
num_modules: int,
|
||||
aggregator_args: AggregatorArguments,
|
||||
):
|
||||
lora_config = model.peft_config["default"]
|
||||
|
|
@ -93,7 +98,7 @@ def get_aggregator_config(
|
|||
feature_size=ctx_encoder_model_config.hidden_size,
|
||||
output_size=output_size,
|
||||
num_layers=get_num_layers(model),
|
||||
num_modules=len(lora_config.target_modules),
|
||||
num_modules=num_modules,
|
||||
**vars(aggregator_args),
|
||||
)
|
||||
|
||||
|
|
@ -107,7 +112,9 @@ class HypernetConfig:
|
|||
|
||||
lora_config: LoraConfig
|
||||
module_names: dict[str, list[str]]
|
||||
# trainable_base_modules: Optional[list[str]]
|
||||
extra_modules: Optional[list[str]]
|
||||
base_hidden_size: int
|
||||
|
||||
layer_indices: Iterable[int]
|
||||
feature_sizes: tuple[dict[str, int], dict[str, int]]
|
||||
aggregator_config: AggregatorConfig
|
||||
|
|
@ -120,9 +127,13 @@ def get_hypernet_config(
|
|||
aggregator_args: AggregatorArguments,
|
||||
):
|
||||
lora_config = model.peft_config["default"]
|
||||
num_modules = len(lora_config.target_modules) + len(
|
||||
hypernet_args.extra_modules or []
|
||||
)
|
||||
indices = torch.arange(get_num_layers(model), device=model.device)
|
||||
return HypernetConfig(
|
||||
**vars(hypernet_args),
|
||||
base_hidden_size=model.config.hidden_size,
|
||||
lora_config=lora_config,
|
||||
module_names=get_lora_module_names(model, lora_config.target_modules, indices),
|
||||
layer_indices=indices,
|
||||
|
|
@ -131,6 +142,7 @@ def get_hypernet_config(
|
|||
model,
|
||||
ctx_encoder_model_config,
|
||||
hypernet_args.latent_size,
|
||||
num_modules,
|
||||
aggregator_args,
|
||||
),
|
||||
)
|
||||
|
|
@ -358,7 +370,7 @@ class EarlyExit(nn.Module):
|
|||
def config(self):
|
||||
return self.base_model.config
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
def forward(self, **kwargs):
|
||||
# if len(kwargs["input_ids"].shape) == 1:
|
||||
# kwargs["input_ids"] = kwargs["input_ids"].unsqueeze(0)
|
||||
|
|
@ -422,6 +434,7 @@ class HyperLoRA(nn.Module):
|
|||
# by mixing the pooled features with layer embs and module embs (for pooling)
|
||||
# or via a perceiver w/ bottleneck size = n_modules * n_layers
|
||||
self.config = config
|
||||
logger.debug(f"HyperLoRA config: {self.config}")
|
||||
self._init_model()
|
||||
|
||||
def _init_model(self):
|
||||
|
|
@ -430,7 +443,14 @@ class HyperLoRA(nn.Module):
|
|||
|
||||
self.lora_config = self.config.lora_config
|
||||
|
||||
self.target_modules = self.lora_config.target_modules
|
||||
self.target_modules = (
|
||||
self.lora_config.target_modules if self.lora_config else None
|
||||
)
|
||||
self.num_modules = len(self.target_modules) if self.target_modules else 0
|
||||
self.extra_modules = (
|
||||
self.config.extra_modules if self.config.extra_modules else None
|
||||
)
|
||||
self.num_extra_modules = len(self.extra_modules) if self.extra_modules else 0
|
||||
self.layer_indices = self.config.layer_indices
|
||||
|
||||
self.d_in, self.d_out = self.config.feature_sizes
|
||||
|
|
@ -439,9 +459,19 @@ class HyperLoRA(nn.Module):
|
|||
input_size=self.config.latent_size,
|
||||
hidden_size=self.config.latent_size * 4,
|
||||
output_size=self.config.latent_size,
|
||||
dropout_rate=self.config.dropout_rate,
|
||||
dropout_rate=getattr(self.config, "dropout_rate", 0),
|
||||
)
|
||||
|
||||
self.pre_layers_norm = nn.LayerNorm(self.config.latent_size)
|
||||
self.norm = nn.LayerNorm(self.config.latent_size)
|
||||
if self.extra_modules:
|
||||
self.extra_layers = MLPResidualBlock(
|
||||
input_size=self.config.latent_size,
|
||||
hidden_size=self.config.latent_size * 4,
|
||||
output_size=self.config.latent_size,
|
||||
dropout_rate=getattr(self.config, "dropout_rate", 0),
|
||||
)
|
||||
self.pre_extra_layers_norm = nn.LayerNorm(self.config.latent_size)
|
||||
self.extra_norm = nn.LayerNorm(self.config.latent_size)
|
||||
if self.config.use_light_weight_lora:
|
||||
# light-weight lora projection (per module)
|
||||
|
||||
|
|
@ -480,32 +510,50 @@ class HyperLoRA(nn.Module):
|
|||
# has something to do with the bias shape (n_modules r d_lora)
|
||||
# when n_modules == 1, adamw_torch_fused complains about device/layout
|
||||
# but when n_modules > 1, it works fine
|
||||
if n_modules == 1:
|
||||
self.head = Mix(
|
||||
"bs n_layers n_modules d_latent -> bs n_layers n_modules r d_lora",
|
||||
weight_shape="d_latent r d_lora",
|
||||
# bias_shape=None, # no bias
|
||||
bias_shape="r d_lora",
|
||||
d_latent=self.config.latent_size,
|
||||
r=self.config.lora_config.r,
|
||||
d_lora=d_lora,
|
||||
)
|
||||
else:
|
||||
# each module processes d -> r d_out independently
|
||||
self.head = Mix(
|
||||
"bs n_layers n_modules d_latent -> bs n_layers n_modules r d_lora",
|
||||
weight_shape="n_modules d_latent r d_lora",
|
||||
# bias_shape=None, # no bias
|
||||
bias_shape="n_modules r d_lora",
|
||||
n_modules=len(self.target_modules),
|
||||
d_latent=self.config.latent_size,
|
||||
r=self.config.lora_config.r,
|
||||
d_lora=d_lora,
|
||||
)
|
||||
# print(self.head)
|
||||
# print(self.head.weight.shape)
|
||||
# print(self.head.bias.shape)
|
||||
# breakpoint()
|
||||
if n_modules > 0:
|
||||
if n_modules == 1:
|
||||
self.head = Mix(
|
||||
"bs n_layers n_modules d_latent -> bs n_layers n_modules r d_lora",
|
||||
weight_shape="d_latent r d_lora",
|
||||
# bias_shape=None, # no bias
|
||||
bias_shape="r d_lora",
|
||||
d_latent=self.config.latent_size,
|
||||
r=self.config.lora_config.r,
|
||||
d_lora=d_lora,
|
||||
)
|
||||
else:
|
||||
# each module processes d -> r d_out independently
|
||||
self.head = Mix(
|
||||
"bs n_layers n_modules d_latent -> bs n_layers n_modules r d_lora",
|
||||
weight_shape="n_modules d_latent r d_lora",
|
||||
# bias_shape=None, # no bias
|
||||
bias_shape="n_modules r d_lora",
|
||||
n_modules=len(self.target_modules),
|
||||
d_latent=self.config.latent_size,
|
||||
r=self.config.lora_config.r,
|
||||
d_lora=d_lora,
|
||||
)
|
||||
|
||||
# separate head for extra_modules, e.g., layernorm
|
||||
n_extra_modules = len(self.extra_modules) if self.extra_modules else 0
|
||||
if n_extra_modules > 0:
|
||||
if n_extra_modules == 1:
|
||||
self.extra_head = Mix(
|
||||
"bs n_layers n_modules d_latent -> bs n_layers n_modules hidden_size",
|
||||
weight_shape="d_latent hidden_size",
|
||||
bias_shape="hidden_size",
|
||||
d_latent=self.config.latent_size,
|
||||
hidden_size=self.config.base_hidden_size,
|
||||
)
|
||||
else:
|
||||
self.extra_head = Mix(
|
||||
"bs n_layers n_modules d_latent -> bs n_layers n_modules hidden_size",
|
||||
weight_shape="n_modules d_latent hidden_size",
|
||||
bias_shape="n_modules hidden_size",
|
||||
n_modules=n_extra_modules,
|
||||
d_latent=self.config.latent_size,
|
||||
hidden_size=self.config.base_hidden_size,
|
||||
)
|
||||
|
||||
def _to_lora_dict(
|
||||
self, flat_loras: Float[Tensor, "bs n_layers n_modules r max_io_dim"]
|
||||
|
|
@ -546,27 +594,55 @@ class HyperLoRA(nn.Module):
|
|||
|
||||
return lora_dict
|
||||
|
||||
def _to_layernorm_dict(
|
||||
self, flat_layernorms: Float[Tensor, "bs n_layers n_modules hidden_size"]
|
||||
) -> dict[str, Float[Tensor, "bs n_layers hidden_size"]]:
|
||||
if self.extra_modules is None:
|
||||
return None
|
||||
layernorms = unpack(
|
||||
flat_layernorms,
|
||||
[[] for _ in range(len(self.extra_modules))],
|
||||
"bs n_layers * hidden_size",
|
||||
)
|
||||
return {k: v for k, v in zip(self.extra_modules, layernorms)}
|
||||
|
||||
def forward(
|
||||
self,
|
||||
features: Float[Tensor, "bs seq_len feature_dim"],
|
||||
attn_mask: Optional[Integer[Tensor, "bs seq_len"]] = None,
|
||||
):
|
||||
|
||||
# [bs, n_layers, n_modules, feature_dim]
|
||||
# [bs, n_layers, n_total_modules, feature_dim]
|
||||
emb = self.aggregator(features.to(torch.float32), attn_mask)
|
||||
lora_emb, extra_emb = unpack(
|
||||
emb,
|
||||
[[self.num_modules], [self.num_extra_modules]],
|
||||
"bs n_layers * feature_dim",
|
||||
)
|
||||
|
||||
# [bs, n_layers, n_modules, r, max_in_d_outim]
|
||||
flat_loras = self.head(self.layers(emb))
|
||||
lora_emb = self.norm(self.layers(self.pre_layers_norm(lora_emb)))
|
||||
flat_loras = self.head(lora_emb)
|
||||
|
||||
return flat_loras
|
||||
flat_layernorms = None
|
||||
if self.num_extra_modules:
|
||||
# [bs, n_layers, n_extra_modules, base_hidden_size]
|
||||
emb = self.extra_norm(
|
||||
self.extra_layers(
|
||||
self.pre_extra_layers_norm(extra_emb[:, :, self.num_modules :])
|
||||
)
|
||||
)
|
||||
flat_layernorms = self.extra_head(extra_emb)
|
||||
|
||||
def generate_loras(
|
||||
return flat_loras, flat_layernorms
|
||||
|
||||
def generate_weights(
|
||||
self,
|
||||
features: Float[Tensor, "bs seq_len feature_dim"],
|
||||
attn_mask: Optional[Integer[Tensor, "bs seq_len"]] = None,
|
||||
):
|
||||
flat_loras = self.forward(features, attn_mask)
|
||||
return self._to_lora_dict(flat_loras)
|
||||
flat_loras, flat_layernorms = self.forward(features, attn_mask)
|
||||
return self._to_lora_dict(flat_loras), self._to_layernorm_dict(flat_layernorms)
|
||||
|
||||
|
||||
class ModulatedPretrainedModel(nn.Module):
|
||||
|
|
@ -676,6 +752,9 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
logger.debug(f"peft_weights: {peft_weights}")
|
||||
self.hypernet.head.weight.data[:] = 0
|
||||
self.hypernet.head.bias.data[:] = 0
|
||||
if self.hypernet.config.extra_modules:
|
||||
self.hypernet.extra_head.weight.data[:] = 0
|
||||
self.hypernet.extra_head.bias.data[:] = 0
|
||||
# d_in = self.hypernet.d_in
|
||||
# d_out = self.hypernet.d_out
|
||||
# m = max(self.hypernet.target_modules, key=lambda m: d_in[m] + d_out[m])
|
||||
|
|
@ -783,18 +862,18 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
# out["ctx_features"] = features
|
||||
# return out
|
||||
|
||||
def generate_loras(
|
||||
def generate_weights(
|
||||
self,
|
||||
ctx_ids: Integer[Tensor, "bs ctx_len"],
|
||||
ctx_attn_mask: Integer[Tensor, "bs ctx_len"],
|
||||
*args: Any,
|
||||
**kwargs: Any,
|
||||
):
|
||||
with torch.no_grad():
|
||||
with torch.inference_mode():
|
||||
ctx_features = self.ctx_encoder(
|
||||
input_ids=ctx_ids, attention_mask=ctx_attn_mask, *args, **kwargs
|
||||
)
|
||||
return self.hypernet.generate_loras(ctx_features, ctx_attn_mask)
|
||||
return self.hypernet.generate_weights(ctx_features, ctx_attn_mask)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
|
|
@ -834,11 +913,13 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
if "attention_mask" in model_inputs_kwargs
|
||||
else None
|
||||
)
|
||||
with torch.no_grad():
|
||||
with torch.inference_mode():
|
||||
ctx_features = self.ctx_encoder(
|
||||
input_ids=ctx_ids, attention_mask=ctx_attn_mask
|
||||
)
|
||||
generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask)
|
||||
generated_loras, generated_layernorms = self.hypernet.generate_weights(
|
||||
ctx_features, ctx_attn_mask
|
||||
)
|
||||
|
||||
# compute kl loss
|
||||
# - compute logits from the base model from tokenized chat [bs, chat_len, vocab_size]
|
||||
|
|
@ -887,11 +968,19 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
else:
|
||||
# input_ids in model_inputs_kwargs contains only
|
||||
# prompt + response (for hypernet training)
|
||||
with apply_generated_loras(
|
||||
self.base_model,
|
||||
generated_loras,
|
||||
self.hypernet.layer_indices,
|
||||
self.training,
|
||||
with (
|
||||
apply_generated_loras(
|
||||
self.base_model,
|
||||
generated_loras,
|
||||
self.hypernet.layer_indices,
|
||||
self.training,
|
||||
),
|
||||
apply_generated_layernorm(
|
||||
self.base_model,
|
||||
generated_layernorms,
|
||||
self.hypernet.layer_indices,
|
||||
self.training,
|
||||
),
|
||||
):
|
||||
model_outputs = self.base_model(
|
||||
*model_inputs_args, **model_inputs_kwargs
|
||||
|
|
@ -899,7 +988,7 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
|
||||
return model_outputs
|
||||
|
||||
@torch.no_grad()
|
||||
@torch.inference_mode()
|
||||
def generate(
|
||||
self,
|
||||
ctx_ids: Optional[Integer[Tensor, "bs ctx_length"]] = None,
|
||||
|
|
@ -934,15 +1023,25 @@ class ModulatedPretrainedModel(nn.Module):
|
|||
ctx_features = self.ctx_encoder(
|
||||
input_ids=ctx_ids, attention_mask=ctx_attn_mask
|
||||
)
|
||||
generated_loras = self.hypernet.generate_loras(ctx_features, ctx_attn_mask)
|
||||
generated_loras, generated_layernorms = self.hypernet.generate_weights(
|
||||
ctx_features, ctx_attn_mask
|
||||
)
|
||||
|
||||
# apply lora hook to the base model
|
||||
# self.apply_generated_loras(generated_loras)
|
||||
with apply_generated_loras(
|
||||
self.base_model,
|
||||
generated_loras,
|
||||
self.hypernet.layer_indices,
|
||||
self.training,
|
||||
with (
|
||||
apply_generated_loras(
|
||||
self.base_model,
|
||||
generated_loras,
|
||||
self.hypernet.layer_indices,
|
||||
self.training,
|
||||
),
|
||||
apply_generated_layernorm(
|
||||
self.base_model,
|
||||
generated_layernorms,
|
||||
self.hypernet.layer_indices,
|
||||
self.training,
|
||||
),
|
||||
):
|
||||
model_outputs = self.base_model.generate(
|
||||
*model_inputs_args, **model_inputs_kwargs
|
||||
|
|
@ -1037,6 +1136,33 @@ class ModulatedModelWithSharedInput(nn.Module):
|
|||
return self.modulated_model.generate(ctx_ids, ctx_attn_mask, *args, **kwargs)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def apply_generated_layernorm(
|
||||
base_model: nn.Module,
|
||||
generated_layernorms: Optional[dict[str, Float[Tensor, "bs n_layers h"]]] = None,
|
||||
layer_indices: Optional[Iterable[int]] = None,
|
||||
training: bool = False,
|
||||
):
|
||||
if generated_layernorms is None:
|
||||
yield base_model
|
||||
return
|
||||
|
||||
try:
|
||||
hooks = []
|
||||
for module_name in generated_layernorms:
|
||||
for layer_idx in layer_indices:
|
||||
hooks += add_generated_layernorm_hook(
|
||||
base_model,
|
||||
module_name,
|
||||
layer_idx,
|
||||
W=generated_layernorms[module_name][:, layer_idx],
|
||||
training=training,
|
||||
)
|
||||
yield base_model
|
||||
finally:
|
||||
remove_hook_handles(hooks)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def apply_generated_loras(
|
||||
base_model: nn.Module,
|
||||
|
|
@ -1052,6 +1178,7 @@ def apply_generated_loras(
|
|||
hooks = []
|
||||
for module_name in generated_loras:
|
||||
for layer_idx in layer_indices:
|
||||
# TODO: handle sequence packing???
|
||||
hooks += add_generated_lora_hook(
|
||||
base_model,
|
||||
module_name,
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue