mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
emb rescale + alpha 1
This commit is contained in:
parent
3505b006ad
commit
de057cdc30
3 changed files with 36 additions and 17 deletions
|
|
@ -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."}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue