From de057cdc3013ced385af57b126745e3899997c18 Mon Sep 17 00:00:00 2001 From: 51616 Date: Wed, 3 Sep 2025 13:14:17 +0900 Subject: [PATCH] emb rescale + alpha 1 --- src/ctx_to_lora/configs.py | 3 ++ src/ctx_to_lora/model_loading.py | 2 +- src/ctx_to_lora/modeling/hypernet.py | 48 ++++++++++++++++++---------- 3 files changed, 36 insertions(+), 17 deletions(-) diff --git a/src/ctx_to_lora/configs.py b/src/ctx_to_lora/configs.py index c806c56..933f78f 100644 --- a/src/ctx_to_lora/configs.py +++ b/src/ctx_to_lora/configs.py @@ -492,6 +492,9 @@ class HypernetArguments: default=False, metadata={"help": "Whether to use per-rank generation."}, ) + use_bias: bool = field( + default=True, metadata={"help": "Whether to include data-dependent LoRA"} + ) use_per_rank_bias: bool = field( default=False, metadata={"help": "Whether to use per-rank bias."} ) diff --git a/src/ctx_to_lora/model_loading.py b/src/ctx_to_lora/model_loading.py index dad159b..1f33110 100644 --- a/src/ctx_to_lora/model_loading.py +++ b/src/ctx_to_lora/model_loading.py @@ -174,7 +174,7 @@ def get_lora_config(model_dir, **kwargs): base_model_name_or_path=model_dir, task_type="CAUSAL_LM", lora_dropout=kwargs.get("lora_dropout", 0.0), - lora_alpha=2 / r**0.5, + lora_alpha=1, # 2 / r**0.5, ) peft_conf_kwargs.update(kwargs) diff --git a/src/ctx_to_lora/modeling/hypernet.py b/src/ctx_to_lora/modeling/hypernet.py index 20d1cfb..09fc2e2 100644 --- a/src/ctx_to_lora/modeling/hypernet.py +++ b/src/ctx_to_lora/modeling/hypernet.py @@ -69,6 +69,7 @@ class HypernetConfig: light_weight_latent_size: int per_rank_gen: bool use_per_rank_bias: bool + use_bias: bool per_layer_processing: bool use_token_mixing: bool num_pre_head_layers: int @@ -276,19 +277,29 @@ class HyperLoRA(nn.Module): self.d_lora = max(self.d_in[m] + self.d_out[m] for m in self.target_modules) - self.bias_a = nn.ParameterDict( - { - m: nn.Parameter( - torch.normal( - 0, - 0.1 / (self.d_in[m] * self.r) ** 0.5, - (self.n_layers, self.r, self.d_in[m]), + if self.config.use_bias: + self.bias_A = nn.ParameterDict( + { + m: nn.Parameter( + torch.normal( + 0, + 0.2 / (self.d_in[m] * self.r) ** 0.5, + (self.n_layers, self.r, self.d_in[m]), + ) ) - ) - for m in self.target_modules - } - ) - self.bias_b = nn.ParameterDict( + for m in self.target_modules + } + ) + else: + self.bias_A = nn.ParameterDict( + { + m: nn.Parameter( + torch.zeros((self.n_layers, self.r, self.d_in[m])) + ) + for m in self.target_modules + } + ) + self.bias_B = nn.ParameterDict( { m: nn.Parameter(torch.zeros((self.n_layers, self.r, self.d_out[m]))) for m in self.target_modules @@ -552,8 +563,8 @@ class HyperLoRA(nn.Module): def get_head_bias(self): bias_dict = dict() for module in self.target_modules: - bias_A = self.bias_a[module] - bias_B = self.bias_b[module] + bias_A = self.bias_A[module] + bias_B = self.bias_B[module] # transpose B # bias_B = rearrange(bias_B, "bs n_layers r d_out -> bs n_layers d_out r") @@ -639,7 +650,9 @@ class HyperLoRA(nn.Module): flat_loras = None if self.target_modules: lora_emb = self.layers(lora_emb) - norm_lora_emb = lora_emb / torch.norm(lora_emb, dim=-1, keepdim=True) + d = lora_emb.shape[-1] + norm = torch.norm(lora_emb, dim=-1, keepdim=True) + norm_lora_emb = lora_emb / norm * sqrt(d) # is this too big?? flat_loras = self.head(norm_lora_emb) flat_layernorms = None @@ -709,6 +722,8 @@ class ModulatedPretrainedModel(nn.Module): hypernet_config.num_pre_head_layers = 4 if getattr(hypernet_config, "use_per_rank_bias", None) is None: hypernet_config.use_per_rank_bias = False + if getattr(hypernet_config, "use_bias", None) is None: + hypernet_config.use_bias = True ctx_encoder_args = state_dict["ctx_encoder_args"] model = cls(base_model, hypernet_config, ctx_encoder_args, **kwargs) model.load_state_dict(state_dict) @@ -786,7 +801,8 @@ class ModulatedPretrainedModel(nn.Module): nn.init.normal_( self.hypernet.head.weight, mean=0, - std=2 / sqrt(self.hypernet.config.latent_size + self.hypernet.d_lora), + std=1 + / sqrt(self.hypernet.config.latent_size + self.hypernet.d_lora * r), # the head outputs per rank lora --> divide by r to scale down grad ) # nn.init.orthogonal_(self.hypernet.head.weight, gain=1.0)