diff --git a/src/ctx_to_lora/configs.py b/src/ctx_to_lora/configs.py index ed7e260..cab89b6 100644 --- a/src/ctx_to_lora/configs.py +++ b/src/ctx_to_lora/configs.py @@ -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 diff --git a/src/ctx_to_lora/hooks.py b/src/ctx_to_lora/hooks.py index 5ee312b..a86a361 100644 --- a/src/ctx_to_lora/hooks.py +++ b/src/ctx_to_lora/hooks.py @@ -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 + ) diff --git a/src/ctx_to_lora/modeling_utils.py b/src/ctx_to_lora/modeling_utils.py index 8af2cb0..9eebd67 100644 --- a/src/ctx_to_lora/modeling_utils.py +++ b/src/ctx_to_lora/modeling_utils.py @@ -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,