add extra_modules (for layernorm)

This commit is contained in:
51616 2025-01-27 10:13:11 +00:00
parent 411fb91a90
commit 1665d89eee
3 changed files with 243 additions and 61 deletions

View file

@ -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

View file

@ -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
)

View file

@ -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,