2.0.0 initial

This commit is contained in:
siriuz42
2025-09-12 00:18:08 +00:00
parent d70708d42a
commit 7d8f3d971d
52 changed files with 1882 additions and 394 deletions
+18
View File
@@ -0,0 +1,18 @@
# Copyright 2024 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""adapter init file."""
from .dora_layers import DoraAttentionProjection, DoraCombinedQKVProjection, DoraLinear
from .lora_layers import LoraAttentionProjection, LoraCombinedQKVProjection, LoraLinear
+202
View File
@@ -0,0 +1,202 @@
# Copyright 2024 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from jax import numpy as jnp
from praxis import base_layer
from praxis.layers import attentions, linears
WeightInit = base_layer.WeightInit
WeightHParams = base_layer.WeightHParams
class DoraTheta(base_layer.Theta):
def __init__(self, module):
self.module = module
def _dora_initialized(self):
if (
self.module.has_variable("params", "lora_a")
and self.module.has_variable("params", "lora_b")
and self.module.has_variable("params", "dora_m")
and "lora_a" in self.module._weight_hparams
and "lora_b" in self.module._weight_hparams
and "dora_m" in self.module._weight_hparams
):
return True
else:
return False
def _dorafy_var(self, w):
lora_a = super().__getattr__("lora_a")
lora_b = super().__getattr__("lora_b")
dora_m = super().__getattr__("dora_m")
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
return w_prime
def __getattr__(self, k):
var = super().__getattr__(k)
if not self._dora_initialized():
return var
if k == "w":
return self._dorafy_var(var)
return var
def __getitem__(self, k):
var = super().__getattr__(k)
if not self._dora_initialized():
return var
if k == "w":
return self._dorafy_var(var)
return var
class DoraThetaDescriptor:
"""Dot syntax accession descriptor."""
def __get__(self, obj, objtype=None):
return DoraTheta(obj)
class DoraLinear(linears.Linear):
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
def setup(self) -> None:
lora_init = self.lora_init if self.lora_init else self.weight_init
super().setup()
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dims, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.output_dims, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[1, self.output_dims],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
class DoraAttentionProjection(attentions.AttentionProjection):
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[1, self.num_heads, self.dim_per_head],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
class DoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer):
rank: int = 0
lora_init: WeightInit | None = None
theta = DoraThetaDescriptor()
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable(
"lora_a",
WeightHParams(
shape=[3, self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[3, self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"dora_m",
WeightHParams(
shape=[3, 1, self.num_heads, self.dim_per_head],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None, None],
),
)
+166
View File
@@ -0,0 +1,166 @@
# Copyright 2024 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from jax import numpy as jnp
from praxis import base_layer
from praxis.layers import attentions, linears
WeightInit = base_layer.WeightInit
WeightHParams = base_layer.WeightHParams
class LoraTheta(base_layer.Theta):
def __init__(self, module):
self.module = module
def _lora_initialized(self):
if (
self.module.has_variable("params", "lora_a")
and self.module.has_variable("params", "lora_b")
and "lora_a" in self.module._weight_hparams
and "lora_b" in self.module._weight_hparams
):
return True
else:
return False
def _lorafy_var(self, w):
lora_a = super().__getattr__("lora_a")
lora_b = super().__getattr__("lora_b")
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
return w_prime
def __getattr__(self, k):
var = super().__getattr__(k)
if not self._lora_initialized():
return var
if k == "w":
return self._lorafy_var(var)
return var
def __getitem__(self, k):
var = super().__getattr__(k)
if not self._lora_initialized():
return var
if k == "w":
return self._lorafy_var(var)
return var
class LoraThetaDescriptor:
"""Dot syntax accession descriptor."""
def __get__(self, obj, objtype=None):
return LoraTheta(obj)
class LoraLinear(linears.Linear):
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
def setup(self) -> None:
lora_init = self.lora_init if self.lora_init else self.weight_init
super().setup()
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dims, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.output_dims, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None],
),
)
class LoraAttentionProjection(attentions.AttentionProjection):
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable(
"lora_a",
WeightHParams(
shape=[self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[
None,
None,
],
),
)
class LoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer):
rank: int = 0
lora_init: WeightInit | None = None
theta = LoraThetaDescriptor()
def setup(self) -> None:
super().setup()
w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable(
"lora_a",
WeightHParams(
shape=[3, self.input_dim, self.rank],
init=lora_init,
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
self.create_variable(
"lora_b",
WeightHParams(
shape=[3, self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None],
),
)
+487
View File
@@ -0,0 +1,487 @@
# Copyright 2024 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
This file provides functionality for loading and merging adapter weights
in timesfm model, specifically for LoRA and DoRA.
LoRA: https://arxiv.org/abs/2106.09685
DoRA: https://arxiv.org/abs/2402.09353v4
"""
import time
import jax
import jax.numpy as jnp
from paxml import checkpoints, tasks_lib
from paxml.train_states import TrainState
from praxis import pax_fiddle
from adapter.dora_layers import (
DoraAttentionProjection,
DoraCombinedQKVProjection,
DoraLinear,
)
from adapter.lora_layers import (
LoraAttentionProjection,
LoraCombinedQKVProjection,
LoraLinear,
)
from timesfm import TimesFm
def get_adapter_params(
params: dict, lora_target_modules: str, num_layers: int, use_dora: bool = False
) -> dict:
"""
Extracts adapter parameters from the given model parameters for saving the checkpoint.
Args:
params (dict): The full model parameters.
lora_target_modules (str): Target modules for LoRA/DoRA adaptation.
num_layers (int): Number of transformer layers.
use_dora (bool, optional): Whether DoRA was used or not. Defaults to False.
Returns:
dict: A dictionary containing the extracted adapter parameters.
"""
adapter_params = {}
for i in range(num_layers):
layer_key = f"x_layers_{i}"
adapter_params[layer_key] = {}
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = params["params"]["core_layer"]["stacked_transformer_layer"][
layer_key
]["ff_layer"][ff_layer_key]["linear"]
lora_a = linear["lora_a"]
lora_b = linear["lora_b"]
adapter_params[layer_key][ff_layer_key] = {
"lora_a": lora_a,
"lora_b": lora_b,
}
if use_dora:
adapter_params[layer_key][ff_layer_key]["dora_m"] = linear["dora_m"]
if lora_target_modules in ["all", "attention"]:
attention = params["params"]["core_layer"]["stacked_transformer_layer"][
layer_key
]["self_attention"]
for component in ["key", "query", "value", "post"]:
lora_a = attention[component]["lora_a"]
lora_b = attention[component]["lora_b"]
adapter_params[layer_key][component] = {
"lora_a": lora_a,
"lora_b": lora_b,
}
if use_dora:
adapter_params[layer_key][component]["dora_m"] = attention[
component
]["dora_m"]
return adapter_params
def load_adapter_checkpoint(
model: TimesFm,
adapter_checkpoint_path: str,
lora_rank: int,
lora_target_modules: str,
use_dora: bool,
) -> None:
"""
Loads an adapter checkpoint and merges it with the original model weights.
Args:
model (TimesFm): The model to update.
adapter_checkpoint_path (str): Path to the adapter checkpoint.
lora_rank (int): Rank of the LoRA adaptation.
lora_target_modules (str): Target modules for adaptation.
use_dora (bool): Whether DoRA was used or not.
Returns:
None
"""
"""
currently loading and initializing the model with adapter layers first and then merging the
adapter weights to original weights and replacing the adapter layers back to original layer.
# NOTE: refactor this. there should be a better way to load the LoRA checkpoint.
"""
model._logging(f"Restoring adapter checkpoint from {adapter_checkpoint_path}.")
start_time = time.time()
original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl = (
load_adapter_layer(
mdl_vars=model._train_state.mdl_vars,
model=model._model,
lora_rank=lora_rank,
lora_target_modules=lora_target_modules,
use_dora=use_dora,
)
)
var_weight_hparams = model._model.abstract_init_with_metadata(
model._get_sample_inputs(), do_eval=True
)
adapter_weight_hparams = _get_adapter_weight_params(
var_weight_hparams=var_weight_hparams,
lora_target_modules=lora_target_modules,
num_layers=model._model.stacked_transformer_params_tpl.num_layers,
use_dora=use_dora,
)
adapter_state_partition_specs = tasks_lib.create_state_partition_specs(
adapter_weight_hparams,
mesh_shape=model.mesh_shape,
mesh_axis_names=model.mesh_name,
discard_opt_states=True,
learners=None,
)
adapter_state_local_shapes = tasks_lib.create_state_unpadded_shapes(
adapter_weight_hparams,
discard_opt_states=True,
learners=None,
)
adapter_train_state = checkpoints.restore_checkpoint(
state_global_shapes=adapter_state_local_shapes,
checkpoint_dir=adapter_checkpoint_path,
checkpoint_type=checkpoints.CheckpointType.FLAX,
state_specs=adapter_state_partition_specs,
step=None,
)
# add adapter weights to the original weights
_merge_adapter_weights(
model=model,
adapter_train_state=adapter_train_state,
lora_target_modules=lora_target_modules,
num_layers=model._model.stacked_transformer_params_tpl.num_layers,
use_dora=use_dora,
)
# replace back with the original model layer
if lora_target_modules in ["all", "mlp"]:
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = (
original_linear_tpl
)
if lora_target_modules in ["all", "attention"]:
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = (
original_attn_tpl
)
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
original_combined_qkv_tpl
)
model._logging(
f"Restored adapter checkpoint in {time.time() - start_time:.2f} seconds."
)
# jit compile the model
model.jit_decode()
def _merge_adapter_weights(
model: TimesFm,
adapter_train_state: TrainState,
lora_target_modules: str,
num_layers: int,
use_dora: bool,
) -> None:
"""
Merges adapter weights with the original model weights.
Args:
model (TimesFm): The model to update.
adapter_train_state (TrainState): The adapter's train state.
lora_target_modules (str): Target modules for adaptation.
num_layers (int): Number of transformer layers.
use_dora (bool): Whether DoRA was used or not.
"""
for i in range(num_layers):
layer_key = f"x_layers_{i}"
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = model._train_state.mdl_vars["params"][
"stacked_transformer_layer"
][layer_key]["ff_layer"][ff_layer_key]["linear"]
params = adapter_train_state.mdl_vars[layer_key][ff_layer_key]
lora_a = params["lora_a"]
lora_b = params["lora_b"]
w = linear["w"]
lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
if use_dora:
dora_m = params["dora_m"]
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
linear["w"] = w_prime
del linear["dora_m"]
else:
linear["w"] = w_prime
del linear["lora_a"]
del linear["lora_b"]
if lora_target_modules in ["all", "attention"]:
attention = model._train_state.mdl_vars["params"][
"stacked_transformer_layer"
][layer_key]["self_attention"]
for component in ["key", "query", "value", "post"]:
params = adapter_train_state.mdl_vars[layer_key][component]
lora_a = params["lora_a"]
lora_b = params["lora_b"]
w = attention[component]["w"]
lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
if use_dora:
dora_m = params["dora_m"]
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
attention[component]["w"] = w_prime
del attention[component]["dora_m"]
else:
attention[component]["w"] = w_prime
del attention[component]["lora_a"]
del attention[component]["lora_b"]
def _get_adapter_weight_params(
var_weight_hparams: dict, lora_target_modules: str, num_layers: int, use_dora: bool
) -> dict:
"""
Extracts adapter weight parameters from the given variable weight hyperparameters.
Args:
var_weight_hparams (dict): Variable weight hyperparameters.
lora_target_modules (str): Target modules for adaptation.
num_layers (int): Number of transformer layers.
use_dora (bool): Whether DoRA was used or not.
Returns:
dict: A dictionary containing the extracted adapter weight parameters.
"""
adapter_params = {}
for i in range(num_layers):
layer = f"x_layers_{i}"
adapter_params[layer] = {}
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
adapter_weight_params = var_weight_hparams["params"][
"stacked_transformer_layer"
][layer]["ff_layer"][ff_layer_key]["linear"]
adapter_params[layer][ff_layer_key] = {
"lora_a": adapter_weight_params["lora_a"],
"lora_b": adapter_weight_params["lora_b"],
}
if use_dora:
adapter_params[layer][ff_layer_key]["dora_m"] = (
adapter_weight_params["dora_m"]
)
if lora_target_modules in ["all", "attention"]:
for component in ["key", "value", "query", "post"]:
adapter_weight_params = var_weight_hparams["params"][
"stacked_transformer_layer"
][layer]["self_attention"][component]
adapter_params[layer][component] = {
"lora_a": adapter_weight_params["lora_a"],
"lora_b": adapter_weight_params["lora_b"],
}
if use_dora:
adapter_params[layer][component]["dora_m"] = adapter_weight_params[
"dora_m"
]
return adapter_params
def load_adapter_layer(
mdl_vars: dict,
model: pax_fiddle.Config,
lora_rank: int,
lora_target_modules: str,
use_dora: bool = False,
) -> tuple[pax_fiddle.Config, pax_fiddle.Config]:
"""
Updates target modules with adapter layers.
Args:
mdl_vars (dict): Model variables.
model (pax_fiddle.Config): Model configuration.
lora_rank (int): Rank of the LoRA adaptation.
lora_target_modules (str): Target modules for adaptation.
use_dora (bool, optional): Whether DoRA was used or not.
Returns:
tuple[pax_fiddle.Config, pax_fiddle.Config]: Updated model configurations.
"""
original_linear_tpl = original_attn_tpl = original_combined_qkv_tpl = None
if lora_target_modules in ["all", "mlp"]:
original_linear_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl
)
adapter_linear_tpl = (
pax_fiddle.Config(
DoraLinear,
rank=lora_rank,
)
if use_dora
else pax_fiddle.Config(
LoraLinear,
rank=lora_rank,
)
)
adapter_linear_tpl.copy_fields_from(original_linear_tpl)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = (
adapter_linear_tpl
)
if lora_target_modules in ["all", "attention"]:
original_attn_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl
)
adapter_attn_tpl = (
pax_fiddle.Config(DoraAttentionProjection, rank=lora_rank)
if use_dora
else pax_fiddle.Config(LoraAttentionProjection, rank=lora_rank)
)
adapter_attn_tpl.copy_fields_from(original_attn_tpl)
original_combined_qkv_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl
)
adapter_combined_qkv_tpl = (
pax_fiddle.Config(DoraCombinedQKVProjection, rank=lora_rank)
if use_dora
else pax_fiddle.Config(LoraCombinedQKVProjection, rank=lora_rank)
)
adapter_combined_qkv_tpl.copy_fields_from(original_combined_qkv_tpl)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = (
adapter_attn_tpl
)
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
adapter_combined_qkv_tpl
)
# initialize and add adapter weights
_initialize_adapter_params(
mdl_vars=mdl_vars,
num_layers=model.stacked_transformer_params_tpl.num_layers,
lora_rank=lora_rank,
lora_target_modules=lora_target_modules,
use_dora=use_dora,
)
return original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl
def _initialize_adapter_params(
mdl_vars: dict,
num_layers,
lora_rank: int,
lora_target_modules: str,
use_dora: bool = False,
seed: int = 1234,
) -> dict:
"""
Initializes and adds adapter parameters to target modules.
Args:
mdl_vars (dict): Model variables.
num_layers (int): Number of transformer layers.
lora_rank (int): Rank of the LoRA adaptation.
lora_target_modules (str): Target modules for adaptation.
use_dora (bool, optional): Whether DoRA was used or not.
seed (int, optional): Random seed for initialization. Defaults to 1234.
Returns:
dict: Updated model variables with initialized adapter parameters.
"""
for i in range(num_layers):
layer_key = f"x_layers_{i}"
if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = mdl_vars["params"]["stacked_transformer_layer"][layer_key][
"ff_layer"
][ff_layer_key]["linear"]
original_w = linear["w"]
input_dim, output_dim = original_w.shape
std_dev = 1 / jnp.sqrt(lora_rank)
normal_initializer = jax.nn.initializers.normal(std_dev)
lora_a = normal_initializer(
jax.random.key(seed), (input_dim, lora_rank), jnp.float32
)
lora_b = jnp.zeros((output_dim, lora_rank))
linear["lora_a"] = lora_a
linear["lora_b"] = lora_b
if use_dora:
norm = jnp.linalg.norm(original_w, ord=2, axis=0, keepdims=True)
linear["dora_m"] = norm
if lora_target_modules in ["all", "attention"]:
attention = mdl_vars["params"]["stacked_transformer_layer"][layer_key][
"self_attention"
]
for component in ["key", "query", "value", "post"]:
original_w = attention[component]["w"]
w_dim = original_w.shape[0]
std_dev = 1 / jnp.sqrt(lora_rank)
normal_initializer = jax.nn.initializers.normal(std_dev)
lora_a = normal_initializer(
jax.random.key(seed), (w_dim, lora_rank), jnp.float32
)
lora_b = jnp.zeros((w_dim, lora_rank))
attention[component]["lora_a"] = lora_a
attention[component]["lora_b"] = lora_b
if use_dora:
norm = jnp.linalg.norm(
original_w, ord=2, axis=0, keepdims=True
).astype(jnp.float32)
attention[component]["dora_m"] = norm
return mdl_vars