From 71d980267db352cfc20997d72f457f462710b94d Mon Sep 17 00:00:00 2001 From: tanmayshishodia Date: Tue, 16 Jul 2024 01:26:36 +0000 Subject: [PATCH] add parameter efficient finetuning pipeline Why? this commit adds a generic finetuning pipeline with LoRA and DoRA support --- .gitignore | 5 + environment.yml | 4 +- environment_cpu.yml | 2 + peft/finetune.py | 410 ++++++++++++++++++++++++++++++++++++ peft/usage.ipynb | 351 +++++++++++++++++++++++++++++++ src/adapter/__init__.py | 18 ++ src/adapter/dora_layers.py | 205 ++++++++++++++++++ src/adapter/lora_layers.py | 170 +++++++++++++++ src/adapter/utils.py | 411 +++++++++++++++++++++++++++++++++++++ 9 files changed, 1575 insertions(+), 1 deletion(-) create mode 100644 .gitignore create mode 100644 peft/finetune.py create mode 100644 peft/usage.ipynb create mode 100644 src/adapter/__init__.py create mode 100644 src/adapter/dora_layers.py create mode 100644 src/adapter/lora_layers.py create mode 100644 src/adapter/utils.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..04c1260 --- /dev/null +++ b/.gitignore @@ -0,0 +1,5 @@ +__pycache__/ +checkpoints/ +wandb/ +datasets/ +results/ \ No newline at end of file diff --git a/environment.yml b/environment.yml index 7a1d6fe..94cbfca 100644 --- a/environment.yml +++ b/environment.yml @@ -1,4 +1,4 @@ -name: tfm_env +name: ok_tfm_env channels: - conda-forge @@ -16,3 +16,5 @@ dependencies: - jax[cuda12]==0.4.26 - einshape - scikit-learn + - typer + - wandb diff --git a/environment_cpu.yml b/environment_cpu.yml index d808176..65b5883 100644 --- a/environment_cpu.yml +++ b/environment_cpu.yml @@ -16,3 +16,5 @@ dependencies: - jax[cpu]==0.4.26 - einshape - scikit-learn + - typer + - wandb diff --git a/peft/finetune.py b/peft/finetune.py new file mode 100644 index 0000000..747dac5 --- /dev/null +++ b/peft/finetune.py @@ -0,0 +1,410 @@ +# 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. + +""" +Finetune pipeline. +""" +import gc +import logging +import warnings +from datetime import datetime +from typing import Tuple + +import jax +import jax.numpy as jnp +import numpy as np +import pandas as pd +import typer +import wandb +from jax import numpy as jnp +from paxml import checkpoint_types, checkpoints, learners, tasks_lib, trainer_lib +from praxis import optimizers, pax_fiddle, py_utils, schedules +from rich import print +from tqdm import tqdm +from typing_extensions import Annotated + +from adapter.utils import get_adapter_params, load_adapter_layer +from timesfm import TimesFm, data_loader, patched_decoder + +NestedMap = py_utils.NestedMap + + +warnings.filterwarnings("ignore") +cmdstanpy_logger = logging.getLogger("cmdstanpy") +absl_logger = logging.getLogger("absl") +cmdstanpy_logger.disabled = True +absl_logger.disabled = True + +""" +TimesFM model config. These are fixed since pre-training was done +with this configuration. +""" +INPUT_PATCH_LEN = 32 +OUTPUT_PATCH_LEN = 128 +NUM_LAYERS = 20 +MODEL_DIMS = 1280 + +QUANTILES = list(np.arange(1, 10) / 10.0) +EPS = 1e-7 +RANDOM_SEED = 1234 + + +def get_forecasts(model, past: np.ndarray, freq: int) -> np.ndarray: + """Get forecasts.""" + lfreq = [freq] * past.shape[0] + _, out = model.forecast(list(past), lfreq) + out = out[:, :, 5] + return out + + +def finetune( + *, + checkpoint_path: Annotated[ + str, typer.Option(help="The path to the model checkpoint.") + ] = None, + model_name: Annotated[ + str, typer.Option(help="Specify the name of the huggingface model.") + ] = "google/timesfm-1.0-200m", + datetime_col: Annotated[str, typer.Option(help="Column having datetime.")] = "ds", + ts_cols: Annotated[ + list[str], typer.Option(help="Columns of time-series features.") + ] = [], + normalize: Annotated[ + bool, typer.Option(help="Normalize data for eval or not") + ] = True, + context_len: Annotated[int, typer.Option(help="Length of the context window")], + horizon_len: Annotated[int, typer.Option(help="Prediction length.")], + freq: Annotated[ + str, + typer.Option( + ..., + help="Frequency Map Str", + ), + ], + data_path: Annotated[str, typer.Option(help="Path to dataset csv")], + boundaries: Annotated[ + Tuple[int, int, int], + typer.Option( + help="boundaries of dataset to train, val, test", + ), + ] = (0, 0, 0), + backend: Annotated[str, typer.Option(help="Backend device: cpu, gpu, tpu")], + batch_size: Annotated[ + int, typer.Option(help="Batch size for the randomly sampled batch") + ] = 16, + num_epochs: Annotated[int, typer.Option(help="Number of epochs")], + learning_rate: Annotated[float, typer.Option(help="adam optimizer learning rate")], + adam_epsilon: Annotated[float, typer.Option(help="adam optimizer epsilon")], + adam_clip_threshold: Annotated[ + float, typer.Option(help="adam optimizer clip threshold") + ], + cos_initial_decay_value: Annotated[ + float, typer.Option(help="cosine initial decay value") + ], + cos_final_decay_value: Annotated[ + float, typer.Option(help="cosine final decay value") + ], + cos_decay_steps: Annotated[int, typer.Option(help="Number of cosine decay steps")], + ema_decay: Annotated[float, typer.Option(help="Exponential moving average decay")], + early_stop_patience: Annotated[ + int, typer.Option(..., help="Early stopping patience") + ] = 5, + use_lora: Annotated[ + bool, + typer.Option( + help="Train low rank adapters. Freeze all other params in model", + ), + ] = False, + lora_rank: Annotated[ + int, + typer.Option( + help="LoRA Rank", + ), + ] = 8, + lora_target_modules: Annotated[ + str, + typer.Option(help="LoRA target modules. Allowed values: [all, attention, mlp]"), + ] = "all", + use_dora: Annotated[ + bool, + typer.Option( + help="Apply DoRA strategy along with LoRA.", + ), + ] = False, + use_linear_probing: Annotated[ + bool, + typer.Option( + help="Linear Probing. Train only input/output and embedding params. Freeze params in self attention modules.", + ), + ] = False, + checkpoint_dir: Annotated[ + str, typer.Option(help="Checkpoint directory") + ] = "./checkpoints", + wandb_project: Annotated[ + str, typer.Option(help="Weights & Biases project name") + ] = "google_timesfm_finetune", +) -> None: + key = jax.random.PRNGKey(seed=RANDOM_SEED) + wandb.init(project=wandb_project, config=locals()) + + data_df = pd.read_csv(open(data_path, "r")) + + if boundaries == (0, 0, 0): + # Default boundaries: train 60%, val 20%, test 20% + boundaries = [ + int(len(data_df) * 0.6), + int(len(data_df) * 0.8), + len(data_df) - 1, + ] + + ts_cols = [col for col in data_df.columns if col != datetime_col] + + dtl = data_loader.TimeSeriesdata( + data_path=data_path, + datetime_col=datetime_col, + num_cov_cols=None, + cat_cov_cols=None, + ts_cols=np.array(ts_cols), + train_range=[0, boundaries[0]], + val_range=[boundaries[0], boundaries[1]], + test_range=[boundaries[1], boundaries[2]], + hist_len=context_len, + pred_len=horizon_len, + batch_size=batch_size, + freq=freq, + normalize=normalize, + epoch_len=None, + holiday=False, + permute=False, + ) + + train_batches = dtl.tf_dataset(mode="train", shift=1).batch(batch_size) + val_batches = dtl.tf_dataset(mode="val", shift=horizon_len) + + for tbatch in tqdm(train_batches.as_numpy_iterator()): + pass + + tfm = TimesFm( + context_len=context_len, + horizon_len=horizon_len, + input_patch_len=INPUT_PATCH_LEN, + output_patch_len=OUTPUT_PATCH_LEN, + num_layers=NUM_LAYERS, + model_dims=MODEL_DIMS, + backend=backend, + per_core_batch_size=batch_size, + quantiles=QUANTILES, + ) + + if checkpoint_path: + tfm.load_from_checkpoint( + checkpoint_path=checkpoint_path, + checkpoint_type=checkpoints.CheckpointType.FLAX, + ) + else: + tfm.load_from_checkpoint( + repo_id=model_name, + checkpoint_type=checkpoints.CheckpointType.FLAX, + ) + + model = pax_fiddle.Config( + patched_decoder.PatchedDecoderFinetuneModel, + name="patched_decoder_finetune", + core_layer_tpl=tfm.model_p, + ) + + if use_lora: + load_adapter_layer( + mdl_vars=tfm._train_state.mdl_vars, + model=model.core_layer_tpl, + lora_rank=lora_rank, + lora_target_modules=lora_target_modules, + use_dora=use_dora, + ) + + @pax_fiddle.auto_config + def build_learner() -> learners.Learner: + bprop_variable_inclusion = None + bprop_variable_exclusion = None + if use_lora: + bprop_variable_inclusion = [r"^.*lora.*$"] + if use_dora: + bprop_variable_inclusion.append(r"^.*dora.*$") + elif use_linear_probing: + bprop_variable_exclusion = [".*/stacked_transformer_layer/.*"] + + return pax_fiddle.Config( + learners.Learner, + name="learner", + loss_name="avg_qloss", + optimizer=optimizers.Adam( + epsilon=adam_epsilon, + clip_threshold=adam_clip_threshold, + learning_rate=learning_rate, + lr_schedule=pax_fiddle.Config( + schedules.Cosine, + initial_value=cos_initial_decay_value, + final_value=cos_final_decay_value, + total_steps=cos_decay_steps, + ), + ema_decay=ema_decay, + ), + bprop_variable_exclusion=bprop_variable_exclusion, + bprop_variable_inclusion=bprop_variable_inclusion, + ) + + task_p = tasks_lib.SingleTask( + name="ts-learn", + model=model, + train=tasks_lib.SingleTask.Train( + learner=build_learner(), + ), + ) + + task_p.model.ici_mesh_shape = [1, 1, 1] + task_p.model.mesh_axis_names = ["replica", "data", "mdl"] + + DEVICES = np.array(jax.devices()).reshape([1, 1, 1]) + jax.sharding.Mesh(DEVICES, ["replica", "data", "mdl"]) + + num_devices = jax.local_device_count() + print(f"num_devices: {num_devices}") + print(f"device kind: {jax.local_devices()[0].device_kind}") + + jax_task = task_p + key, init_key = jax.random.split(key) + + def process_train_batch(batch): + past_ts = batch[0].reshape(batch_size * len(ts_cols), -1) + actual_ts = batch[3].reshape(batch_size * len(ts_cols), -1) + return NestedMap(input_ts=past_ts, actual_ts=actual_ts) + + def process_eval_batch(batch): + past_ts = batch[0] + actual_ts = batch[3] + return NestedMap(input_ts=past_ts, actual_ts=actual_ts) + + jax_model_states, _ = trainer_lib.initialize_model_state( + jax_task, + init_key, + process_train_batch(tbatch), + checkpoint_type=checkpoint_types.CheckpointType.GDA, + ) + jax_model_states.mdl_vars["params"]["core_layer"] = tfm._train_state.mdl_vars[ + "params" + ] + gc.collect() + + jax_task = task_p + + def train_step(states, prng_key, inputs): + return trainer_lib.train_step_single_learner(jax_task, states, prng_key, inputs) + + def eval_step(states, prng_key, inputs): + states = states.to_eval_state() + return trainer_lib.eval_step_single_learner(jax_task, states, prng_key, inputs) + + key, train_key, eval_key = jax.random.split(key, 3) + train_prng_seed = jax.random.split(train_key, num=jax.local_device_count()) + eval_prng_seed = jax.random.split(eval_key, num=jax.local_device_count()) + + p_train_step = jax.pmap(train_step, axis_name="batch") + p_eval_step = jax.pmap(eval_step, axis_name="batch") + + replicated_jax_states = trainer_lib.replicate_model_state(jax_model_states) + + def reshape_batch_for_pmap(batch, num_devices): + def _reshape(input_tensor): + bsize = input_tensor.shape[0] + residual_shape = list(input_tensor.shape[1:]) + nbsize = bsize // num_devices + return jnp.reshape(input_tensor, [num_devices, nbsize] + residual_shape) + + return jax.tree.map(_reshape, batch) + + patience = 0 + best_eval_loss = 1e7 + checkpoint_dir = ( + f"{checkpoint_dir}/run_ignore_{datetime.now().strftime('%Y%m%d_%H%M%S')}" + ) + for epoch in range(num_epochs): + print(f"Epoch: {epoch + 1}") + train_its = train_batches.as_numpy_iterator() + train_losses = [] + for batch in tqdm(train_its): + if patience >= early_stop_patience: + print("Early stopping.") + break + tbatch = process_train_batch(batch) + tbatch = reshape_batch_for_pmap(tbatch, num_devices) + replicated_jax_states, step_fun_out = p_train_step( + replicated_jax_states, train_prng_seed, tbatch + ) + train_losses.append(step_fun_out.loss[0]) + wandb.log({"train_step_loss": step_fun_out.loss[0]}) + + avg_train_loss = np.mean(train_losses) + + print("Starting eval.") + val_its = val_batches.as_numpy_iterator() + eval_losses = [] + for ev_batch in tqdm(val_its): + ebatch = process_eval_batch(ev_batch) + ebatch = reshape_batch_for_pmap(ebatch, num_devices) + _, step_fun_out = p_eval_step(replicated_jax_states, eval_prng_seed, ebatch) + eval_losses.append(step_fun_out.loss[0]) + wandb.log({"eval_step_loss": step_fun_out.loss[0]}) + + avg_eval_loss = np.mean(eval_losses) + + print(f"Train Loss: {avg_train_loss}, Val Loss: {avg_eval_loss}") + + wandb.log( + { + "epoch": epoch + 1, + "avg_train_loss": avg_train_loss, + "avg_val_loss": avg_eval_loss, + } + ) + + if avg_eval_loss < best_eval_loss or np.isnan(avg_eval_loss): + best_eval_loss = avg_eval_loss + print("Saving checkpoint.") + jax_state_for_saving = py_utils.maybe_unreplicate_for_fully_replicated( + replicated_jax_states + ) + if use_lora: + adapter_params = get_adapter_params( + params=jax_state_for_saving.mdl_vars, + lora_target_modules=lora_target_modules, + num_layers=NUM_LAYERS, + use_dora=use_dora, + ) + jax_state_for_saving.mdl_vars["params"] = adapter_params + + checkpoints.save_checkpoint( + jax_state_for_saving, checkpoint_dir, overwrite=True + ) + + patience = 0 + del jax_state_for_saving + gc.collect() + else: + patience += 1 + print(f"patience: {patience}") + print("Fine-tuning completed.") + + +if __name__ == "__main__": + typer.run(finetune) diff --git a/peft/usage.ipynb b/peft/usage.ipynb new file mode 100644 index 0000000..58ae1c0 --- /dev/null +++ b/peft/usage.ipynb @@ -0,0 +1,351 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Load Base Model" + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "metadata": {}, + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "2024-07-16 00:36:17.861915: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcudart.so.11.0'; dlerror: libcudart.so.11.0: cannot open shared object file: No such file or directory\n", + "2024-07-16 00:36:25.276917: W external/xla/xla/service/gpu/nvptx_compiler.cc:718] The NVIDIA driver's CUDA version is 12.2 which is older than the ptxas CUDA version (12.5.40). Because the driver is older than the ptxas version, XLA is disabling parallel compilation, which may slow down compilation. You should update your NVIDIA driver or use the NVIDIA-provided CUDA forward compatibility packages.\n" + ] + }, + { + "data": { + "application/vnd.jupyter.widget-view+json": { + "model_id": "935a82ff713149179c103451e4837c27", + "version_major": 2, + "version_minor": 0 + }, + "text/plain": [ + "Fetching 5 files: 0%| | 0/5 [00:00\n", + "WARNING:absl:Configured `CheckpointManager` using deprecated legacy API. Please follow the instructions at https://orbax.readthedocs.io/en/latest/api_refactor.html to migrate by May 1st, 2024.\n", + "WARNING:absl:train_state_unpadded_shape_dtype_struct is not provided. We assume `train_state` is unpadded.\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Constructed model weights in 2.84 seconds.\n", + "Restoring checkpoint from /home/ubuntu/.cache/huggingface/hub/models--google--timesfm-1.0-200m/snapshots/8775f7531211ac864b739fe776b0b255c277e2be/checkpoints.\n" + ] + }, + { + "name": "stderr", + "output_type": "stream", + "text": [ + "ERROR:absl:For checkpoint version > 1.0, we require users to provide\n", + " `train_state_unpadded_shape_dtype_struct` during checkpoint\n", + " saving/restoring, to avoid potential silent bugs when loading\n", + " checkpoints to incompatible unpadded shapes of TrainState.\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Restored checkpoint in 0.86 seconds.\n", + "Jitting decoding.\n", + "Jitted decoding in 16.44 seconds.\n" + ] + } + ], + "source": [ + "from timesfm import TimesFm, freq_map, data_loader\n", + "from adapter.utils import load_adapter_checkpoint\n", + "from tqdm import tqdm\n", + "import numpy as np\n", + "import pandas as pd\n", + "\n", + "\n", + "tfm = TimesFm(\n", + " context_len=512,\n", + " horizon_len=128,\n", + " input_patch_len=32,\n", + " output_patch_len=128,\n", + " num_layers=20,\n", + " model_dims=1280,\n", + " backend=\"cpu\",\n", + ")\n", + "tfm.load_from_checkpoint(repo_id=\"google/timesfm-1.0-200m\")" + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "metadata": {}, + "outputs": [], + "source": [ + "DATA_DICT = {\n", + " \"ettm2\": {\n", + " \"boundaries\": [34560, 46080, 57600],\n", + " \"data_path\": \"../datasets/ETT-small/ETTm2.csv\",\n", + " \"freq\": \"15min\",\n", + " },\n", + " \"ettm1\": {\n", + " \"boundaries\": [34560, 46080, 57600],\n", + " \"data_path\": \"../datasets/ETT-small/ETTm1.csv\",\n", + " \"freq\": \"15min\",\n", + " },\n", + " \"etth2\": {\n", + " \"boundaries\": [8640, 11520, 14400],\n", + " \"data_path\": \"../datasets/ETT-small/ETTh2.csv\",\n", + " \"freq\": \"H\",\n", + " },\n", + " \"etth1\": {\n", + " \"boundaries\": [8640, 11520, 14400],\n", + " \"data_path\": \"../datasets/ETT-small/ETTh1.csv\",\n", + " \"freq\": \"H\",\n", + " },\n", + " \"elec\": {\n", + " \"boundaries\": [18413, 21044, 26304],\n", + " \"data_path\": \"../datasets/electricity/electricity.csv\",\n", + " \"freq\": \"H\",\n", + " },\n", + " \"traffic\": {\n", + " \"boundaries\": [12280, 14036, 17544],\n", + " \"data_path\": \"../datasets/traffic/traffic.csv\",\n", + " \"freq\": \"H\",\n", + " },\n", + " \"weather\": {\n", + " \"boundaries\": [36887, 42157, 52696],\n", + " \"data_path\": \"../datasets/weather/weather.csv\",\n", + " \"freq\": \"10min\",\n", + " },\n", + "}" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "print(tfm._train_state.mdl_vars)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Load Adapter Checkpoint\n", + "\n", + "Specify the adapter checkpoint path, rank and the modules used to train the adapters and whether dora was employed or not." + ] + }, + { + "cell_type": "code", + "execution_count": 3, + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Restoring adapter checkpoint from /home/ubuntu/tanmay/timesfm/adapter_finetuning/checkpoints/run_ignore_20240715_234116.\n" + ] + }, + { + "name": "stderr", + "output_type": "stream", + "text": [ + "WARNING:absl:No registered CheckpointArgs found for handler type: \n", + "WARNING:absl:Configured `CheckpointManager` using deprecated legacy API. Please follow the instructions at https://orbax.readthedocs.io/en/latest/api_refactor.html to migrate by May 1st, 2024.\n", + "WARNING:absl:train_state_unpadded_shape_dtype_struct is not provided. We assume `train_state` is unpadded.\n", + "WARNING:absl:A possible mismatch (could be spurious) between the saved checkpoint structure and the current one has been detected (PyTreeDef(CustomNode(TrainState[()], [*, {'params': {'x_layers_0': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_1': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_10': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_11': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_12': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_13': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_14': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_15': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_16': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_17': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_18': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_19': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_2': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_3': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_4': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_5': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_6': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_7': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_8': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_9': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}}}, [(CustomNode(NestedMap[('count',)], [*]), CustomNode(NestedMap[('count',)], [*]), CustomNode(NestedMap[('count', 'm', 'v')], [*, {'params': {'core_layer': {'freq_emb': {'emb_var': *}, 'horizon_ff_layer': {'hidden_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'output_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'residual_layer': {'bias': {'b': *}, 'linear': {'w': *}}}, 'input_ff_layer': {'hidden_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'output_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'residual_layer': {'bias': {'b': *}, 'linear': {'w': *}}}, 'stacked_transformer_layer': {'x_layers_0': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_1': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_10': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_11': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_12': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_13': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_14': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_15': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_16': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_17': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_18': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_19': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_2': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_3': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_4': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_5': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_6': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_7': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_8': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_9': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}}}}}, {'params': {'core_layer': {'freq_emb': {'emb_var': *}, 'horizon_ff_layer': {'hidden_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'output_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'residual_layer': {'bias': {'b': *}, 'linear': {'w': *}}}, 'input_ff_layer': {'hidden_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'output_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'residual_layer': {'bias': {'b': *}, 'linear': {'w': *}}}, 'stacked_transformer_layer': {'x_layers_0': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_1': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_10': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_11': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_12': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_13': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_14': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_15': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_16': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_17': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_18': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_19': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_2': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_3': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_4': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_5': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_6': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_7': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_8': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_9': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}}}}}]), CustomNode(NestedMap[('count',)], [*]), CustomNode(NestedMap[('count', 'ema')], [*, {'params': {'core_layer': {'freq_emb': {'emb_var': *}, 'horizon_ff_layer': {'hidden_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'output_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'residual_layer': {'bias': {'b': *}, 'linear': {'w': *}}}, 'input_ff_layer': {'hidden_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'output_layer': {'bias': {'b': *}, 'linear': {'w': *}}, 'residual_layer': {'bias': {'b': *}, 'linear': {'w': *}}}, 'stacked_transformer_layer': {'x_layers_0': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_1': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_10': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_11': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_12': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_13': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_14': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_15': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_16': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_17': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_18': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_19': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_2': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_3': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_4': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_5': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_6': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_7': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_8': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}, 'x_layers_9': {'ff_layer': {'ffn_layer1': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'ffn_layer2': {'bias': {'b': *}, 'linear': {'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}, 'layer_norm': {'bias': *, 'scale': *}}, 'layer_norm': {'scale': *}, 'self_attention': {'key': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'per_dim_scale': {'per_dim_scale': *}, 'post': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'query': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}, 'value': {'b': *, 'dora_m': *, 'lora_a': *, 'lora_b': *, 'w': *}}}}}}}]))], ()])) vs PyTreeDef(CustomNode(TrainState[()], [*, {'x_layers_0': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_1': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_10': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_11': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_12': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_13': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_14': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_15': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_16': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_17': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_18': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_19': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_2': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_3': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_4': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_5': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_6': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_7': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_8': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}, 'x_layers_9': {'ffn_layer1': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'ffn_layer2': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'key': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'post': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'query': {'dora_m': *, 'lora_a': *, 'lora_b': *}, 'value': {'dora_m': *, 'lora_a': *, 'lora_b': *}}}, [], ()]))).\n", + "ERROR:absl:For checkpoint version > 1.0, we require users to provide\n", + " `train_state_unpadded_shape_dtype_struct` during checkpoint\n", + " saving/restoring, to avoid potential silent bugs when loading\n", + " checkpoints to incompatible unpadded shapes of TrainState.\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Restored adapter checkpoint in 8.77 seconds.\n", + "Jitting decoding.\n", + "Jitted decoding in 14.77 seconds.\n" + ] + } + ], + "source": [ + "load_adapter_checkpoint(\n", + " model=tfm,\n", + " adapter_checkpoint_path=\"/home/ubuntu/tanmay/timesfm/adapter_finetuning/checkpoints/run_ignore_20240715_234116\",\n", + " lora_rank=1,\n", + " lora_target_modules=\"all\",\n", + " use_dora=True,\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Test Performance" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "metadata": {}, + "outputs": [], + "source": [ + "dataset = \"ettm1\"\n", + "data_path = DATA_DICT[dataset][\"data_path\"]\n", + "freq = DATA_DICT[dataset][\"freq\"]\n", + "int_freq = freq_map(freq)\n", + "boundaries = DATA_DICT[dataset][\"boundaries\"]\n", + "\n", + "data_df = pd.read_csv(open(data_path, \"r\"))\n", + "\n", + "ts_cols = [col for col in data_df.columns if col != \"date\"]\n", + "num_cov_cols = None\n", + "cat_cov_cols = None\n", + "\n", + "context_len = 512\n", + "pred_len = 96\n", + "\n", + "num_ts = len(ts_cols)\n", + "batch_size = 16\n", + "\n", + "dtl = data_loader.TimeSeriesdata(\n", + " data_path=data_path,\n", + " datetime_col=\"date\",\n", + " num_cov_cols=num_cov_cols,\n", + " cat_cov_cols=cat_cov_cols,\n", + " ts_cols=np.array(ts_cols),\n", + " train_range=[0, boundaries[0]],\n", + " val_range=[boundaries[0], boundaries[1]],\n", + " test_range=[boundaries[1], boundaries[2]],\n", + " hist_len=context_len,\n", + " pred_len=pred_len,\n", + " batch_size=num_ts,\n", + " freq=\"15min\",\n", + " normalize=True,\n", + " epoch_len=None,\n", + " holiday=False,\n", + " permute=True,\n", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "metadata": {}, + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "2024-07-16 00:37:10.346547: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcudart.so.11.0'; dlerror: libcudart.so.11.0: cannot open shared object file: No such file or directory\n", + "2024-07-16 00:37:10.346626: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcublas.so.11'; dlerror: libcublas.so.11: cannot open shared object file: No such file or directory\n", + "2024-07-16 00:37:10.346677: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcublasLt.so.11'; dlerror: libcublasLt.so.11: cannot open shared object file: No such file or directory\n", + "2024-07-16 00:37:10.346724: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcufft.so.10'; dlerror: libcufft.so.10: cannot open shared object file: No such file or directory\n", + "2024-07-16 00:37:10.350324: W tensorflow/stream_executor/platform/default/dso_loader.cc:64] Could not load dynamic library 'libcusparse.so.11'; dlerror: libcusparse.so.11: cannot open shared object file: No such file or directory\n", + "2024-07-16 00:37:10.350358: W tensorflow/core/common_runtime/gpu/gpu_device.cc:1850] Cannot dlopen some GPU libraries. Please make sure the missing libraries mentioned above are installed properly if you would like to use GPU. Follow the guide at https://www.tensorflow.org/install/gpu for how to download and setup the required libraries for your platform.\n", + "Skipping registering GPU devices...\n" + ] + } + ], + "source": [ + "test_batches = dtl.tf_dataset(mode=\"test\", shift=pred_len)" + ] + }, + { + "cell_type": "code", + "execution_count": 6, + "metadata": {}, + "outputs": [ + { + "name": "stderr", + "output_type": "stream", + "text": [ + "120it [01:20, 1.50it/s]\n" + ] + }, + { + "data": { + "text/html": [ + "
MAE: 0.32633697986602783\n",
+       "
\n" + ], + "text/plain": [ + "MAE: \u001b[1;36m0.32633697986602783\u001b[0m\n" + ] + }, + "metadata": {}, + "output_type": "display_data" + } + ], + "source": [ + "mae_losses = []\n", + "for batch in tqdm(test_batches.as_numpy_iterator()):\n", + " past = batch[0]\n", + " actuals = batch[3]\n", + " _, forecasts = tfm.forecast(list(past), [0] * past.shape[0])\n", + " forecasts = forecasts[:, 0 : actuals.shape[1], 5]\n", + " mae_losses.append(np.abs(forecasts - actuals).mean())\n", + "\n", + "print(f\"MAE: {np.mean(mae_losses)}\")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [] + } + ], + "metadata": { + "kernelspec": { + "display_name": "tanmay_tfm_env", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.10.14" + } + }, + "nbformat": 4, + "nbformat_minor": 2 +} diff --git a/src/adapter/__init__.py b/src/adapter/__init__.py new file mode 100644 index 0000000..4e672f0 --- /dev/null +++ b/src/adapter/__init__.py @@ -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. + +"""TimesFM init file.""" + +from .dora_layers import DoraAttentionProjection, DoraCombinedQKVProjection, DoraLinear +from .lora_layers import LoraAttentionProjection, LoraCombinedQKVProjection, LoraLinear diff --git a/src/adapter/dora_layers.py b/src/adapter/dora_layers.py new file mode 100644 index 0000000..ce35b71 --- /dev/null +++ b/src/adapter/dora_layers.py @@ -0,0 +1,205 @@ +# 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, pytypes +from praxis.layers.attentions import AttentionProjection, CombinedQKVProjectionLayer +from praxis.layers.linears import Linear + +WeightInit = base_layer.WeightInit +template_field = base_layer.template_field +WeightHParams = base_layer.WeightHParams +JTensor = pytypes.JTensor + + +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, var): + lora_a = super().__getattr__("lora_a") + lora_b = super().__getattr__("lora_b") + dora_m = super().__getattr__("dora_m") + + new_var = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) + new_var = jnp.reshape(new_var, var.shape) + + new_var += var + + column_norm = jnp.linalg.norm(new_var, ord=2, axis=0, keepdims=True) + norm_adapted = new_var / column_norm + w = dora_m * norm_adapted + return w + + 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(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(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(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], + ), + ) diff --git a/src/adapter/lora_layers.py b/src/adapter/lora_layers.py new file mode 100644 index 0000000..7669d06 --- /dev/null +++ b/src/adapter/lora_layers.py @@ -0,0 +1,170 @@ +# 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, pax_fiddle, pytypes +from praxis.layers.attentions import AttentionProjection, CombinedQKVProjectionLayer +from praxis.layers.linears import Linear + +WeightInit = base_layer.WeightInit +LayerTpl = pax_fiddle.Config[base_layer.BaseLayer] +template_field = base_layer.template_field +WeightHParams = base_layer.WeightHParams +JTensor = pytypes.JTensor + + +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, var): + lora_a = super().__getattr__("lora_a") + lora_b = super().__getattr__("lora_b") + new_var = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) + new_var = jnp.reshape(new_var, var.shape) + new_var += var + return new_var + + 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(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(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(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], + ), + ) diff --git a/src/adapter/utils.py b/src/adapter/utils.py new file mode 100644 index 0000000..ed27dd1 --- /dev/null +++ b/src/adapter/utils.py @@ -0,0 +1,411 @@ +# 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. + +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: + 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: + """ + 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: + 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"] + + var = linear["w"] + + new_var = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b) + new_var = jnp.reshape(new_var, var.shape) + new_var += var + + if use_dora: + dora_m = params["dora_m"] + column_norm = jnp.linalg.norm(new_var, ord=2, axis=0, keepdims=True) + norm_adapted = new_var / column_norm + calc_weights = dora_m * norm_adapted + linear["w"] = calc_weights + del linear["dora_m"] + + else: + linear["w"] = new_var + + 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"] + + var = attention[component]["w"] + + new_var = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b) + new_var = jnp.reshape(new_var, var.shape) + new_var += var + + if use_dora: + m = params["dora_m"] + column_norm = jnp.linalg.norm(new_var, ord=2, axis=0, keepdims=True) + norm_adapted = new_var / column_norm + calc_weights = m * norm_adapted + attention[component]["w"] = calc_weights + del attention[component]["dora_m"] + + else: + attention[component]["w"] = new_var + + 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: + 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]: + """ + update self attention modules with LoRA/DoRA layers + """ + 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: + """ + initialize and add LoRA params in self attention + """ + 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