From 04a5a99b4356d6418c53dd8f15beb0c64951275b Mon Sep 17 00:00:00 2001 From: Funto-Adeyemi Date: Sat, 11 Nov 2023 09:39:14 +0000 Subject: [PATCH] Add troubleshooting section to readme file --- .style.yapf | 4 + experiments/baselines/timegpt_pipeline.py | 26 +- .../extended_benchmarks/run_timegpt.py | 9 +- .../extended_benchmarks/run_timesfm.py | 1 - experiments/extended_benchmarks/utils.py | 60 +- peft/finetune.py | 569 +++++++++--------- src/adapter/__init__.py | 1 - src/adapter/dora_layers.py | 297 +++++---- src/adapter/lora_layers.py | 231 ++++--- src/adapter/utils.py | 542 ++++++++--------- src/finetuning/finetuning_example.py | 11 +- src/timesfm/__init__.py | 12 +- src/timesfm/time_features.py | 28 +- src/timesfm/timesfm_base.py | 27 +- tests/test_timesfm.py | 84 ++- 15 files changed, 923 insertions(+), 979 deletions(-) create mode 100644 .style.yapf diff --git a/.style.yapf b/.style.yapf new file mode 100644 index 0000000..725b3cc --- /dev/null +++ b/.style.yapf @@ -0,0 +1,4 @@ +[style] +based_on_style = google +indent_width = 2 +spaces_before_comment = 2 \ No newline at end of file diff --git a/experiments/baselines/timegpt_pipeline.py b/experiments/baselines/timegpt_pipeline.py index 8a2bfbd..d66213f 100644 --- a/experiments/baselines/timegpt_pipeline.py +++ b/experiments/baselines/timegpt_pipeline.py @@ -34,9 +34,8 @@ def get_seasonality(freq: str) -> int: return _get_seasonality(freq, seasonalities={"D": 7}) -def maybe_convert_col_to_datetime( - df: pd.DataFrame, col_name: str -) -> pd.DataFrame: +def maybe_convert_col_to_datetime(df: pd.DataFrame, + col_name: str) -> pd.DataFrame: if not pd.api.types.is_datetime64_any_dtype(df[col_name]): df = df.copy() df[col_name] = pd.to_datetime(df[col_name]) @@ -64,14 +63,15 @@ def zero_pad_time_series(df, freq, min_length=36): end=start_date, periods=min_length - len(subset) + 1, freq=freq, # 'MS' for month start - )[ - :-1 - ] # Exclude the start_date itself + )[:-1] # Exclude the start_date itself # 2c. Create padding data - padding_df = pd.DataFrame( - {"ds": padding_dates, "unique_id": unique_id, "y": 0} # Zero padding - ) + padding_df = pd.DataFrame({ + "ds": padding_dates, + "unique_id": unique_id, + "y": 0 + } # Zero padding + ) # 2d. Combine original and padding data, and append to the list padded_data.append(pd.concat([padding_df, subset]).sort_values("ds")) @@ -121,8 +121,7 @@ class Forecaster: for _, (cutoffs, train, valid) in tqdm(enumerate(splits)): if len(valid.columns) > 3: raise NotImplementedError( - "Cross validation with exogenous variables is not yet supported." - ) + "Cross validation with exogenous variables is not yet supported.") y_pred = self.forecast( df=train, h=h, @@ -138,8 +137,7 @@ class Forecaster: raise ValueError( "Cross validation result produced less results than expected." " Please verify that the frequency parameter (freq) matches your" - " series' and that there aren't any missing periods." - ) + " series' and that there aren't any missing periods.") results.append(result) out = vertical_concat(results) out = drop_index_if_pandas(out) @@ -203,7 +201,7 @@ class TimeGPT(Forecaster): all_unique_ids = df["unique_id"].unique() all_fcst_df = [] for i in range(0, len(all_unique_ids), chunk_size): - chunk_ids = all_unique_ids[i : i + chunk_size] + chunk_ids = all_unique_ids[i:i + chunk_size] chunk_df = df[df["unique_id"].isin(chunk_ids)] fct_chunk_df = client.forecast( df=chunk_df, diff --git a/experiments/extended_benchmarks/run_timegpt.py b/experiments/extended_benchmarks/run_timegpt.py index 38d964d..0df159f 100644 --- a/experiments/extended_benchmarks/run_timegpt.py +++ b/experiments/extended_benchmarks/run_timegpt.py @@ -11,7 +11,6 @@ # 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. - """Evaluation script for timegpt.""" import os @@ -25,7 +24,6 @@ import pandas as pd from ..baselines.timegpt_pipeline import run_timegpt from .utils import ExperimentHandler - dataset_names = [ "m1_monthly", "m1_quarterly", @@ -63,7 +61,6 @@ _MODEL_NAME = flags.DEFINE_string( ) _SAVE_DIR = flags.DEFINE_string("save_dir", "./results", "Save directory") - QUANTILES = list(np.arange(1, 10) / 10.0) @@ -90,9 +87,9 @@ def main(): ) time_df = pd.DataFrame({"time": [total_time], "model": model_name}) fcsts_df = exp.fcst_from_level_to_quantiles(fcsts_df, model_name) - results = exp.evaluate_from_predictions( - models=[model_name], fcsts_df=fcsts_df, times_df=time_df - ) + results = exp.evaluate_from_predictions(models=[model_name], + fcsts_df=fcsts_df, + times_df=time_df) print(results, flush=True) results_list.append(results) results_full = pd.concat(results_list) diff --git a/experiments/extended_benchmarks/run_timesfm.py b/experiments/extended_benchmarks/run_timesfm.py index e8878c8..a0fd2c4 100644 --- a/experiments/extended_benchmarks/run_timesfm.py +++ b/experiments/extended_benchmarks/run_timesfm.py @@ -54,7 +54,6 @@ dataset_names = [ "hospital", ] - context_dict_v2 = {} context_dict_v1 = { diff --git a/experiments/extended_benchmarks/utils.py b/experiments/extended_benchmarks/utils.py index de0368b..c9307f4 100644 --- a/experiments/extended_benchmarks/utils.py +++ b/experiments/extended_benchmarks/utils.py @@ -11,7 +11,6 @@ # 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. - """Forked from https://github.com/Nixtla/nixtla/blob/main/experiments/amazon-chronos/src/utils.py.""" from functools import partial @@ -46,11 +45,9 @@ def quantile_loss( target_col: str = "y", ) -> pd.DataFrame: delta_y = df[models].sub(df[target_col], axis=0) - res = ( - np.maximum(q * delta_y, (q - 1) * delta_y) - .groupby(df[id_col], observed=True) - .mean() - ) + res = (np.maximum(q * delta_y, + (q - 1) * delta_y).groupby(df[id_col], + observed=True).mean()) res.index.name = id_col res = res.reset_index() return res @@ -66,10 +63,8 @@ class ExperimentHandler: models_dir: str = "./models", ): if dataset not in gluonts_datasets: - raise Exception( - f"dataset {dataset} not found in gluonts " - f"available datasets: {', '.join(gluonts_datasets)}" - ) + raise Exception(f"dataset {dataset} not found in gluonts " + f"available datasets: {', '.join(gluonts_datasets)}") self.dataset = dataset self.quantiles = quantiles self.level = self._transform_quantiles_to_levels(quantiles) @@ -80,10 +75,8 @@ class ExperimentHandler: gluonts_dataset = get_dataset(self.dataset) self.horizon = gluonts_dataset.metadata.prediction_length if self.horizon is None: - raise Exception( - f"horizon not found for dataset {self.dataset} " - "experiment cannot be run" - ) + raise Exception(f"horizon not found for dataset {self.dataset} " + "experiment cannot be run") self.freq = gluonts_dataset.metadata.freq # get_seasonality() returns 1 for freq='D', override this to 7. This significantly improves the accuracy of # statistical models on datasets like m5/nn5_daily. The models like AutoARIMA/AutoETS can still set @@ -122,9 +115,8 @@ class ExperimentHandler: @staticmethod def _transform_quantiles_to_levels(quantiles: List[float]) -> List[int]: - level = [ - int(100 - 200 * q) for q in quantiles if q < 0.5 - ] # in this case mean=mediain + level = [int(100 - 200 * q) for q in quantiles if q < 0.5 + ] # in this case mean=mediain level = sorted(list(set(level))) return level @@ -153,9 +145,8 @@ class ExperimentHandler: last_n: int | None = None, ) -> pd.DataFrame: with multiprocessing.Pool(os.cpu_count()) as pool: # Create a process pool - results = pool.map( - parallel_transform, zip(gluonts_dataset, repeat(last_n)) - ) + results = pool.map(parallel_transform, zip(gluonts_dataset, + repeat(last_n))) df = pd.concat(results) df = df.reset_index(drop=True) return df @@ -177,9 +168,8 @@ class ExperimentHandler: def save_dataframe(self, df: pd.DataFrame, file_name: str): df.to_csv(f"{self.results_dir}/{file_name}", index=False) - def save_results( - self, fcst_df: pd.DataFrame, total_time: float, model_name: str - ): + def save_results(self, fcst_df: pd.DataFrame, total_time: float, + model_name: str): self.save_dataframe( fcst_df, f"{model_name}-{self.dataset}-fcst.csv", @@ -215,23 +205,21 @@ class ExperimentHandler: times_df = [] for model in models: fcst_method_df = pd.read_csv( - f"{self.results_dir}/{model}-{self.dataset}-fcst.csv" - ).set_index(["unique_id", "ds"]) + f"{self.results_dir}/{model}-{self.dataset}-fcst.csv").set_index( + ["unique_id", "ds"]) fcsts_df.append(fcst_method_df) time_method_df = pd.read_csv( - f"{self.results_dir}/{model}-{self.dataset}-time.csv" - ) + f"{self.results_dir}/{model}-{self.dataset}-time.csv") times_df.append(time_method_df) fcsts_df = pd.concat(fcsts_df, axis=1).reset_index() fcsts_df["ds"] = pd.to_datetime(fcsts_df["ds"]) times_df = pd.concat(times_df) - return self.evaluate_from_predictions( - models=models, fcsts_df=fcsts_df, times_df=times_df - ) + return self.evaluate_from_predictions(models=models, + fcsts_df=fcsts_df, + times_df=times_df) - def evaluate_from_predictions( - self, models: List[str], fcsts_df: pd.DataFrame, times_df: pd.DataFrame - ) -> pd.DataFrame: + def evaluate_from_predictions(self, models: List[str], fcsts_df: pd.DataFrame, + times_df: pd.DataFrame) -> pd.DataFrame: test_df = self.test_df train_df = self.train_df test_df = test_df.merge(fcsts_df, how="left") @@ -262,9 +250,9 @@ class ExperimentHandler: eval_prob_df["metric"] = "scaled_crps" eval_df = pd.concat([eval_df, eval_prob_df]).reset_index(drop=True) eval_df = eval_df.groupby("metric").mean(numeric_only=True).reset_index() - eval_df = eval_df.melt( - id_vars="metric", value_name="value", var_name="model" - ) + eval_df = eval_df.melt(id_vars="metric", + value_name="value", + var_name="model") times_df.insert(0, "metric", "time") times_df = times_df.rename(columns={"time": "value"}) eval_df = pd.concat([eval_df, times_df]) diff --git a/peft/finetune.py b/peft/finetune.py index 84c59f9..3eee4f1 100644 --- a/peft/finetune.py +++ b/peft/finetune.py @@ -11,7 +11,6 @@ # 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. """ @@ -39,13 +38,11 @@ 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. @@ -62,20 +59,24 @@ RANDOM_SEED = 1234 def finetune( *, - model_name: Annotated[ - str, typer.Option(help="Specify the name of the huggingface model.") - ] = "google/timesfm-1.0-200m", + model_name: Annotated[str, + typer.Option( + help="Specify the name of the huggingface model." + )] = "google/timesfm-1.0-200m", checkpoint_path: Annotated[ - str, typer.Option(help="The path to the local model checkpoint.") - ] = None, - 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")], + str, + typer.Option(help="The path to the local model checkpoint.")] = None, + 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, @@ -87,316 +88,312 @@ def finetune( 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", - ), + typer.Option(help="boundaries of dataset to train, val, test",), ] = (0, 0, 0), - backend: Annotated[str, typer.Option(help="Backend device: cpu, gpu, tpu")], + 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, + 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")], + 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, + int, typer.Option(..., help="Early stopping patience")] = 5, use_lora: Annotated[ bool, - typer.Option( - help="Train low rank adapters for stacked transformer block", - ), + typer. + Option(help="Train low rank adapters for stacked transformer block",), ] = False, lora_rank: Annotated[ int, - typer.Option( - help="LoRA Rank", - ), + typer.Option(help="LoRA Rank",), ] = 8, lora_target_modules: Annotated[ str, typer.Option( - help="LoRA target modules of the transformer block. Allowed values: [all, attention, mlp]" + help= + "LoRA target modules of the transformer block. Allowed values: [all, attention, mlp]" ), ] = "all", use_dora: Annotated[ bool, - typer.Option( - help="Apply DoRA strategy along with LoRA.", - ), + 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 stack transformer block.", + help= + "Linear Probing. Train only input/output and embedding params. Freeze params in stack transformer block.", ), ] = 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", + 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()) + key = jax.random.PRNGKey(seed=RANDOM_SEED) + wandb.init(project=wandb_project, config=locals()) - data_df = pd.read_csv(open(data_path, "r")) + 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, - ] + 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] + 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, + 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, ) - train_batches = dtl.tf_dataset(mode="train", shift=1).batch(batch_size) - val_batches = dtl.tf_dataset(mode="val", shift=horizon_len) + model = pax_fiddle.Config( + patched_decoder.PatchedDecoderFinetuneModel, + name="patched_decoder_finetune", + core_layer_tpl=tfm.model_p, + ) - 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 = [] + bprop_variable_exclusion = [] if use_lora: - load_adapter_layer( - mdl_vars=tfm._train_state.mdl_vars, - model=model.core_layer_tpl, - lora_rank=lora_rank, + bprop_variable_inclusion.append(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_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{wandb.run.id}" + for epoch in range(num_epochs): + if patience >= early_stop_patience: + print("Early stopping.") + break + print(f"Epoch: {epoch + 1}") + train_its = train_batches.as_numpy_iterator() + train_losses = [] + for batch in tqdm(train_its): + 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 - @pax_fiddle.auto_config - def build_learner() -> learners.Learner: - bprop_variable_inclusion = [] - bprop_variable_exclusion = [] - if use_lora: - bprop_variable_inclusion.append(r"^.*lora.*$") - if use_dora: - bprop_variable_inclusion.append(r"^.*dora.*$") - elif use_linear_probing: - bprop_variable_exclusion = [".*/stacked_transformer_layer/.*"] + checkpoints.save_checkpoint(jax_state_for_saving, + checkpoint_dir, + overwrite=True) - 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_{datetime.now().strftime('%Y%m%d_%H%M%S')}_{wandb.run.id}" - for epoch in range(num_epochs): - if patience >= early_stop_patience: - print("Early stopping.") - break - print(f"Epoch: {epoch + 1}") - train_its = train_batches.as_numpy_iterator() - train_losses = [] - for batch in tqdm(train_its): - 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.") + 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) + typer.run(finetune) diff --git a/src/adapter/__init__.py b/src/adapter/__init__.py index 6870b01..5bd05be 100644 --- a/src/adapter/__init__.py +++ b/src/adapter/__init__.py @@ -11,7 +11,6 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. - """adapter init file.""" from .dora_layers import DoraAttentionProjection, DoraCombinedQKVProjection, DoraLinear diff --git a/src/adapter/dora_layers.py b/src/adapter/dora_layers.py index 9a28911..e49dee8 100644 --- a/src/adapter/dora_layers.py +++ b/src/adapter/dora_layers.py @@ -21,182 +21,181 @@ WeightHParams = base_layer.WeightHParams class DoraTheta(base_layer.Theta): - def __init__(self, module): - self.module = module - def _dora_initialized(self): - if ( - self.module.has_variable("params", "lora_a") - and self.module.has_variable("params", "lora_b") - and self.module.has_variable("params", "dora_m") - and "lora_a" in self.module._weight_hparams - and "lora_b" in self.module._weight_hparams - and "dora_m" in self.module._weight_hparams - ): - return True - else: - return False + def __init__(self, module): + self.module = module - def _dorafy_var(self, w): - lora_a = super().__getattr__("lora_a") - lora_b = super().__getattr__("lora_b") - dora_m = super().__getattr__("dora_m") + 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 - lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) - lora_delta = jnp.reshape(lora_delta, w.shape) + def _dorafy_var(self, w): + lora_a = super().__getattr__("lora_a") + lora_b = super().__getattr__("lora_b") + dora_m = super().__getattr__("dora_m") - w_prime = w + lora_delta + lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) + lora_delta = jnp.reshape(lora_delta, w.shape) - column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) - norm_adapted = w_prime / column_norm - w_prime = dora_m * norm_adapted - return w_prime + w_prime = w + lora_delta - def __getattr__(self, k): - var = super().__getattr__(k) - if not self._dora_initialized(): - return var + column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) + norm_adapted = w_prime / column_norm + w_prime = dora_m * norm_adapted + return w_prime - if k == "w": - return self._dorafy_var(var) + def __getattr__(self, k): + var = super().__getattr__(k) + if not self._dora_initialized(): + return var - return var + if k == "w": + return self._dorafy_var(var) - def __getitem__(self, k): - var = super().__getattr__(k) - if not self._dora_initialized(): - return var + return var - if k == "w": - return self._dorafy_var(var) + def __getitem__(self, k): + var = super().__getattr__(k) + if not self._dora_initialized(): + return var - return var + if k == "w": + return self._dorafy_var(var) + + return var class DoraThetaDescriptor: - """Dot syntax accession descriptor.""" + """Dot syntax accession descriptor.""" - def __get__(self, obj, objtype=None): - return DoraTheta(obj) + def __get__(self, obj, objtype=None): + return DoraTheta(obj) class DoraLinear(linears.Linear): - rank: int = 0 - lora_init: WeightInit | None = None - theta = DoraThetaDescriptor() + 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 + 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], - ), - ) + super().setup() + self.create_variable( + "lora_a", + WeightHParams( + shape=[self.input_dims, self.rank], + init=lora_init, + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[None, None], + ), + ) + self.create_variable( + "lora_b", + WeightHParams( + shape=[self.output_dims, self.rank], + init=WeightInit.Constant(scale=0.0), + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[None, None], + ), + ) + self.create_variable( + "dora_m", + WeightHParams( + shape=[1, self.output_dims], + init=lora_init, + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[None, None], + ), + ) class DoraAttentionProjection(attentions.AttentionProjection): - rank: int = 0 - lora_init: WeightInit | None = None - theta = DoraThetaDescriptor() + 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 + 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], - ), - ) + self.create_variable( + "lora_a", + WeightHParams( + shape=[self.input_dim, self.rank], + init=lora_init, + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[ + None, + None, + ], + ), + ) + self.create_variable( + "lora_b", + WeightHParams( + shape=[self.dim_per_head * self.num_heads, self.rank], + init=WeightInit.Constant(scale=0.0), + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[ + None, + None, + ], + ), + ) + self.create_variable( + "dora_m", + WeightHParams( + shape=[1, self.num_heads, self.dim_per_head], + init=lora_init, + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[None, None, None], + ), + ) class DoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer): - rank: int = 0 - lora_init: WeightInit | None = None - theta = DoraThetaDescriptor() + 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 + 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], - ), - ) + 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 index 15df5a5..8eaeb2a 100644 --- a/src/adapter/lora_layers.py +++ b/src/adapter/lora_layers.py @@ -21,146 +21,145 @@ WeightHParams = base_layer.WeightHParams class LoraTheta(base_layer.Theta): - def __init__(self, module): - self.module = module - def _lora_initialized(self): - if ( - self.module.has_variable("params", "lora_a") - and self.module.has_variable("params", "lora_b") - and "lora_a" in self.module._weight_hparams - and "lora_b" in self.module._weight_hparams - ): - return True - else: - return False + def __init__(self, module): + self.module = module - def _lorafy_var(self, w): - lora_a = super().__getattr__("lora_a") - lora_b = super().__getattr__("lora_b") - lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) - lora_delta = jnp.reshape(lora_delta, w.shape) - w_prime = w + lora_delta - return w_prime + def _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 __getattr__(self, k): - var = super().__getattr__(k) - if not self._lora_initialized(): - return var + def _lorafy_var(self, w): + lora_a = super().__getattr__("lora_a") + lora_b = super().__getattr__("lora_b") + lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) + lora_delta = jnp.reshape(lora_delta, w.shape) + w_prime = w + lora_delta + return w_prime - if k == "w": - return self._lorafy_var(var) + def __getattr__(self, k): + var = super().__getattr__(k) + if not self._lora_initialized(): + return var - return var + if k == "w": + return self._lorafy_var(var) - def __getitem__(self, k): - var = super().__getattr__(k) - if not self._lora_initialized(): - return var + return var - if k == "w": - return self._lorafy_var(var) + def __getitem__(self, k): + var = super().__getattr__(k) + if not self._lora_initialized(): + return var - return var + if k == "w": + return self._lorafy_var(var) + + return var class LoraThetaDescriptor: - """Dot syntax accession descriptor.""" + """Dot syntax accession descriptor.""" - def __get__(self, obj, objtype=None): - return LoraTheta(obj) + def __get__(self, obj, objtype=None): + return LoraTheta(obj) class LoraLinear(linears.Linear): - rank: int = 0 - lora_init: WeightInit | None = None - theta = LoraThetaDescriptor() + 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 + 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], - ), - ) + super().setup() + self.create_variable( + "lora_a", + WeightHParams( + shape=[self.input_dims, self.rank], + init=lora_init, + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[None, None], + ), + ) + self.create_variable( + "lora_b", + WeightHParams( + shape=[self.output_dims, self.rank], + init=WeightInit.Constant(scale=0.0), + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[None, None], + ), + ) class LoraAttentionProjection(attentions.AttentionProjection): - rank: int = 0 - lora_init: WeightInit | None = None - theta = LoraThetaDescriptor() + 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 + 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( + "lora_a", + WeightHParams( + shape=[self.input_dim, self.rank], + init=lora_init, + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[ + None, + None, + ], + ), + ) + self.create_variable( + "lora_b", + WeightHParams( + shape=[self.dim_per_head * self.num_heads, self.rank], + init=WeightInit.Constant(scale=0.0), + mesh_shape=self.mesh_shape, + tensor_split_dims_mapping=[ + None, + None, + ], + ), + ) class LoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer): - rank: int = 0 - lora_init: WeightInit | None = None - theta = LoraThetaDescriptor() + 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 + 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( + "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 index 4c3fc5b..a26f381 100644 --- a/src/adapter/utils.py +++ b/src/adapter/utils.py @@ -11,7 +11,6 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. - """ This file provides functionality for loading and merging adapter weights in timesfm model, specifically for LoRA and DoRA. @@ -40,10 +39,11 @@ from adapter.lora_layers import ( from timesfm import TimesFm -def get_adapter_params( - params: dict, lora_target_modules: str, num_layers: int, use_dora: bool = False -) -> dict: - """ +def get_adapter_params(params: dict, + lora_target_modules: str, + num_layers: int, + use_dora: bool = False) -> dict: + """ Extracts adapter parameters from the given model parameters for saving the checkpoint. Args: @@ -55,47 +55,44 @@ def get_adapter_params( Returns: dict: A dictionary containing the extracted adapter parameters. """ - adapter_params = {} - for i in range(num_layers): - layer_key = f"x_layers_{i}" - adapter_params[layer_key] = {} + 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"] + 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"] + 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, - } + 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 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"] + 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"] + 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, - } + 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 + if use_dora: + adapter_params[layer_key][component]["dora_m"] = attention[component][ + "dora_m"] + return adapter_params def load_adapter_checkpoint( @@ -105,7 +102,7 @@ def load_adapter_checkpoint( lora_target_modules: str, use_dora: bool, ) -> None: - """ + """ Loads an adapter checkpoint and merges it with the original model weights. Args: @@ -118,83 +115,77 @@ def load_adapter_checkpoint( Returns: None """ - - """ + """ currently loading and initializing the model with adapter layers first and then merging the adapter weights to original weights and replacing the adapter layers back to original layer. # NOTE: refactor this. there should be a better way to load the LoRA checkpoint. """ - model._logging(f"Restoring adapter checkpoint from {adapter_checkpoint_path}.") - start_time = time.time() - original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl = ( - load_adapter_layer( - mdl_vars=model._train_state.mdl_vars, - model=model._model, - lora_rank=lora_rank, - lora_target_modules=lora_target_modules, - use_dora=use_dora, - ) - ) + 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 - ) + 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_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, - ) + 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, - ) + # 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 - ) + # 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." - ) + 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() + # jit compile the model + model.jit_decode() def _merge_adapter_weights( @@ -204,7 +195,7 @@ def _merge_adapter_weights( num_layers: int, use_dora: bool, ) -> None: - """ + """ Merges adapter weights with the original model weights. Args: @@ -214,74 +205,73 @@ def _merge_adapter_weights( num_layers (int): Number of transformer layers. use_dora (bool): Whether DoRA was used or not. """ - for i in range(num_layers): - layer_key = f"x_layers_{i}" + 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"] + 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"] + params = adapter_train_state.mdl_vars[layer_key][ff_layer_key] + lora_a = params["lora_a"] + lora_b = params["lora_b"] - w = linear["w"] + w = linear["w"] - lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b) - lora_delta = jnp.reshape(lora_delta, w.shape) - w_prime = w + lora_delta + lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b) + lora_delta = jnp.reshape(lora_delta, w.shape) + w_prime = w + lora_delta - if use_dora: - dora_m = params["dora_m"] - column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) - norm_adapted = w_prime / column_norm - w_prime = dora_m * norm_adapted - linear["w"] = w_prime - del linear["dora_m"] + if use_dora: + dora_m = params["dora_m"] + column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) + norm_adapted = w_prime / column_norm + w_prime = dora_m * norm_adapted + linear["w"] = w_prime + del linear["dora_m"] - else: - linear["w"] = w_prime + else: + linear["w"] = w_prime - del linear["lora_a"] - del linear["lora_b"] + 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"] + 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"] + for component in ["key", "query", "value", "post"]: + params = adapter_train_state.mdl_vars[layer_key][component] + lora_a = params["lora_a"] + lora_b = params["lora_b"] - w = attention[component]["w"] + w = attention[component]["w"] - lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b) - lora_delta = jnp.reshape(lora_delta, w.shape) - w_prime = w + lora_delta + lora_delta = jnp.einsum("...dr,...nr->...dn", lora_a, lora_b) + lora_delta = jnp.reshape(lora_delta, w.shape) + w_prime = w + lora_delta - if use_dora: - dora_m = params["dora_m"] - column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) - norm_adapted = w_prime / column_norm - w_prime = dora_m * norm_adapted - attention[component]["w"] = w_prime - del attention[component]["dora_m"] + if use_dora: + dora_m = params["dora_m"] + column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) + norm_adapted = w_prime / column_norm + w_prime = dora_m * norm_adapted + attention[component]["w"] = w_prime + del attention[component]["dora_m"] - else: - attention[component]["w"] = w_prime + else: + attention[component]["w"] = w_prime - del attention[component]["lora_a"] - del attention[component]["lora_b"] + 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: - """ +def _get_adapter_weight_params(var_weight_hparams: dict, + lora_target_modules: str, num_layers: int, + use_dora: bool) -> dict: + """ Extracts adapter weight parameters from the given variable weight hyperparameters. Args: @@ -293,42 +283,39 @@ def _get_adapter_weight_params( Returns: dict: A dictionary containing the extracted adapter weight parameters. """ - adapter_params = {} - for i in range(num_layers): - layer = f"x_layers_{i}" - adapter_params[layer] = {} + 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 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 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 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" - ] + if use_dora: + adapter_params[layer][component]["dora_m"] = adapter_weight_params[ + "dora_m"] - return adapter_params + return adapter_params def load_adapter_layer( @@ -338,7 +325,7 @@ def load_adapter_layer( lora_target_modules: str, use_dora: bool = False, ) -> tuple[pax_fiddle.Config, pax_fiddle.Config]: - """ + """ Updates target modules with adapter layers. Args: @@ -351,67 +338,55 @@ def load_adapter_layer( Returns: tuple[pax_fiddle.Config, pax_fiddle.Config]: Updated model configurations. """ - original_linear_tpl = original_attn_tpl = original_combined_qkv_tpl = None - if lora_target_modules in ["all", "mlp"]: - original_linear_tpl = ( - model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl - ) - adapter_linear_tpl = ( - pax_fiddle.Config( - DoraLinear, - rank=lora_rank, - ) - if use_dora - else pax_fiddle.Config( - LoraLinear, - rank=lora_rank, - ) - ) - adapter_linear_tpl.copy_fields_from(original_linear_tpl) - model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = ( - adapter_linear_tpl - ) + 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 - ) + 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) + 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 - ) + 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) + 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 - ) + 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, - ) + # 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 + return original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl def _initialize_adapter_params( @@ -422,7 +397,7 @@ def _initialize_adapter_params( use_dora: bool = False, seed: int = 1234, ) -> dict: - """ + """ Initializes and adds adapter parameters to target modules. Args: @@ -436,52 +411,47 @@ def _initialize_adapter_params( Returns: dict: Updated model variables with initialized adapter parameters. """ - for i in range(num_layers): - layer_key = f"x_layers_{i}" - if lora_target_modules in ["all", "mlp"]: - for ff_layer_key in ["ffn_layer1", "ffn_layer2"]: - linear = mdl_vars["params"]["stacked_transformer_layer"][layer_key][ - "ff_layer" - ][ff_layer_key]["linear"] - original_w = linear["w"] - input_dim, output_dim = original_w.shape - std_dev = 1 / jnp.sqrt(lora_rank) + 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)) + 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 + 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 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" - ] + 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) + 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)) + 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 + 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 + 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 diff --git a/src/finetuning/finetuning_example.py b/src/finetuning/finetuning_example.py index f1d76af..747d6e9 100644 --- a/src/finetuning/finetuning_example.py +++ b/src/finetuning/finetuning_example.py @@ -43,11 +43,11 @@ flags.DEFINE_list( ) flags.DEFINE_string( - "local_model_path", - None, + "local_model_path", None, "Path to a local .safetensors model file. If provided, overrides Hugging Face download." ) + class TimeSeriesDataset(Dataset): """Dataset for time series data compatible with TimesFM.""" @@ -148,7 +148,7 @@ def get_model(load_weights: bool = False): use_positional_embedding=False, context_len=192, ) - + if load_weights: if FLAGS.local_model_path: tfm_config = TimesFMConfig() @@ -157,11 +157,12 @@ def get_model(load_weights: bool = False): else: repo_id = "google/timesfm-2.0-500m-pytorch" tfm = TimesFm(hparams=hparams, - checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id)) + checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id)) tfm_config = tfm._model_config model = PatchedTimeSeriesDecoder(tfm_config) - checkpoint_path = path.join(snapshot_download(repo_id), "torch_model.ckpt") + checkpoint_path = path.join(snapshot_download(repo_id), + "torch_model.ckpt") loaded_checkpoint = torch.load(checkpoint_path, weights_only=True) model.load_state_dict(loaded_checkpoint) diff --git a/src/timesfm/__init__.py b/src/timesfm/__init__.py index 739807a..f929ada 100644 --- a/src/timesfm/__init__.py +++ b/src/timesfm/__init__.py @@ -25,11 +25,13 @@ from timesfm.timesfm_base import ( import sys try: - from timesfm.timesfm_jax import TimesFmJax as TimesFm - from timesfm import data_loader + from timesfm.timesfm_jax import TimesFmJax as TimesFm + from timesfm import data_loader - print(f"Loaded Jax TimesFM, likely because python version is {sys.version}.") + print(f"Loaded Jax TimesFM, likely because python version is {sys.version}.") except Exception as _: - from timesfm.timesfm_torch import TimesFmTorch as TimesFm + from timesfm.timesfm_torch import TimesFmTorch as TimesFm - print(f"Loaded PyTorch TimesFM, likely because python version is {sys.version}.") + print( + f"Loaded PyTorch TimesFM, likely because python version is {sys.version}." + ) diff --git a/src/timesfm/time_features.py b/src/timesfm/time_features.py index 0bd90a9..0cc0861 100644 --- a/src/timesfm/time_features.py +++ b/src/timesfm/time_features.py @@ -11,7 +11,6 @@ # 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. - """Directory to extract time covariates. Extract time covariates from datetime. @@ -36,7 +35,6 @@ from pandas.tseries.offsets import Easter from sklearn.preprocessing import StandardScaler from tqdm import tqdm - # This is 183 to cover half a year (in both directions), also for leap years # + 17 as Eastern can be between March, 22 - April, 25 MAX_WINDOW = 183 + 17 @@ -50,9 +48,8 @@ def _distance_to_holiday(holiday): index - pd.Timedelta(days=MAX_WINDOW), index + pd.Timedelta(days=MAX_WINDOW), ) - assert ( - len(holiday_date) != 0 # pylint: disable=g-explicit-length-test - ), f"No closest holiday for the date index {index} found." + assert (len(holiday_date) != 0 # pylint: disable=g-explicit-length-test + ), f"No closest holiday for the date index {index} found." # It sometimes returns two dates if it is exactly half a year after the # holiday. In this case, the smaller distance (182 days) is returned. return (index - holiday_date[0]).days @@ -60,16 +57,19 @@ def _distance_to_holiday(holiday): return _distance_to_day -EasterSunday = Holiday( - "Easter Sunday", month=1, day=1, offset=[Easter(), Day(0)] -) +EasterSunday = Holiday("Easter Sunday", + month=1, + day=1, + offset=[Easter(), Day(0)]) NewYearsDay = Holiday("New Years Day", month=1, day=1) -SuperBowl = Holiday( - "Superbowl", month=2, day=1, offset=DateOffset(weekday=SU(1)) -) -MothersDay = Holiday( - "Mothers Day", month=5, day=1, offset=DateOffset(weekday=SU(2)) -) +SuperBowl = Holiday("Superbowl", + month=2, + day=1, + offset=DateOffset(weekday=SU(1))) +MothersDay = Holiday("Mothers Day", + month=5, + day=1, + offset=DateOffset(weekday=SU(2))) IndependenceDay = Holiday("Independence Day", month=7, day=4) ChristmasEve = Holiday("Christmas", month=12, day=24) ChristmasDay = Holiday("Christmas", month=12, day=25) diff --git a/src/timesfm/timesfm_base.py b/src/timesfm/timesfm_base.py index 73b1495..615e798 100644 --- a/src/timesfm/timesfm_base.py +++ b/src/timesfm/timesfm_base.py @@ -25,12 +25,12 @@ import pandas as pd from utilsforecast.processing import make_future_dataframe if TYPE_CHECKING: - from . import xreg_lib - Category = xreg_lib.Category - XRegMode = xreg_lib.XRegMode + from . import xreg_lib + Category = xreg_lib.Category + XRegMode = xreg_lib.XRegMode else: - Category = int | str - XRegMode = str + Category = int | str + XRegMode = str _TOL = 1e-6 DEFAULT_QUANTILES = (0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9) @@ -45,7 +45,7 @@ def moving_average(arr, window_size): """Calculates the moving average using NumPy's convolution function.""" # Pad with zeros to handle initial window positions arr_padded = np.pad(arr, (window_size - 1, 0), "constant") - smoothed_arr = (np.convolve(arr_padded, np.ones(window_size), "valid") / + smoothed_arr = (np.convolve(arr_padded, np.ones(window_size), "valid") / window_size) return [smoothed_arr, arr - smoothed_arr] @@ -57,18 +57,11 @@ def freq_map(freq: str): return 1 elif freq.endswith(("H", "T", "MIN", "D", "B", "U", "S")): return 0 - elif ( - freq.endswith(("W", "M")) - or freq.startswith("W-") - or (freq.startswith("M") and len(freq) == 2) - ): + elif (freq.endswith(("W", "M")) or freq.startswith("W-") or + (freq.startswith("M") and len(freq) == 2)): return 1 - elif ( - freq.endswith(("Y", "Q", "A")) - or freq.startswith("Y-") - or freq.startswith("Q-") - or freq.startswith("A-") - ): + elif (freq.endswith(("Y", "Q", "A")) or freq.startswith("Y-") or + freq.startswith("Q-") or freq.startswith("A-")): return 2 else: raise ValueError(f"Invalid frequency: {freq}") diff --git a/tests/test_timesfm.py b/tests/test_timesfm.py index 3277a9a..713a970 100644 --- a/tests/test_timesfm.py +++ b/tests/test_timesfm.py @@ -12,7 +12,6 @@ # See the License for the specific language governing permissions and # limitations under the License. - from datetime import datetime, timedelta import numpy as np @@ -22,10 +21,10 @@ import pytest import timesfm -def create_sample_dataframe( - start_date: datetime, end_date: datetime, freq: str = "D" -) -> pd.DataFrame: - """ +def create_sample_dataframe(start_date: datetime, + end_date: datetime, + freq: str = "D") -> pd.DataFrame: + """ Create a sample DataFrame with time series data. Args: @@ -36,10 +35,10 @@ def create_sample_dataframe( Returns: pd.DataFrame: DataFrame with columns 'unique_id', 'ds', and 'ts'. """ - date_range = pd.date_range(start=start_date, end=end_date, freq=freq) - ts_data = np.random.randn(len(date_range)) - df = pd.DataFrame({"unique_id": "ts-1", "ds": date_range, "ts": ts_data}) - return df + date_range = pd.date_range(start=start_date, end=end_date, freq=freq) + ts_data = np.random.randn(len(date_range)) + df = pd.DataFrame({"unique_id": "ts-1", "ds": date_range, "ts": ts_data}) + return df @pytest.mark.parametrize("context_length", [128, 256, 512]) @@ -50,42 +49,41 @@ def test_timesfm_forecast_on_df( prediction_length: int, freq: str, ) -> None: - model = timesfm.TimesFm( - context_len=context_length, - horizon_len=prediction_length, - input_patch_len=32, - output_patch_len=128, - num_layers=20, - model_dims=1280, - backend="cpu", - ) - model.load_from_checkpoint(repo_id="google/timesfm-1.0-200m") + model = timesfm.TimesFm( + context_len=context_length, + horizon_len=prediction_length, + input_patch_len=32, + output_patch_len=128, + num_layers=20, + model_dims=1280, + backend="cpu", + ) + model.load_from_checkpoint(repo_id="google/timesfm-1.0-200m") - end_date = datetime.now() - start_date = end_date - timedelta(days=context_length) - input_df = create_sample_dataframe(start_date, end_date, freq) + end_date = datetime.now() + start_date = end_date - timedelta(days=context_length) + input_df = create_sample_dataframe(start_date, end_date, freq) - forecast_df = model.forecast_on_df( - inputs=input_df, - freq=freq, - value_name="ts", - num_jobs=-1, - ) + forecast_df = model.forecast_on_df( + inputs=input_df, + freq=freq, + value_name="ts", + num_jobs=-1, + ) - assert ( - len(forecast_df) == prediction_length - ), f"Expected forecast length of {prediction_length}, but got {len(forecast_df)}" - assert ( - "timesfm" in forecast_df.columns - ), "Forecast DataFrame should contain 'timesfm' column" + assert ( + len(forecast_df) == prediction_length + ), f"Expected forecast length of {prediction_length}, but got {len(forecast_df)}" + assert ("timesfm" in forecast_df.columns + ), "Forecast DataFrame should contain 'timesfm' column" - last_input_date = input_df["ds"].max() - first_forecast_date = forecast_df["ds"].min() - expected_first_forecast_date = last_input_date + pd.Timedelta(1, unit=freq) - assert ( - first_forecast_date == expected_first_forecast_date - ), f"Forecast should start from {expected_first_forecast_date}, but starts from {first_forecast_date}" + last_input_date = input_df["ds"].max() + first_forecast_date = forecast_df["ds"].min() + expected_first_forecast_date = last_input_date + pd.Timedelta(1, unit=freq) + assert ( + first_forecast_date == expected_first_forecast_date + ), f"Forecast should start from {expected_first_forecast_date}, but starts from {first_forecast_date}" - print( - f"Successful forecast with context_length={context_length}, prediction_length={prediction_length}, freq={freq}" - ) + print( + f"Successful forecast with context_length={context_length}, prediction_length={prediction_length}, freq={freq}" + )