From 63f0bf41624c3a0700da7874d6a4b7411f46ed37 Mon Sep 17 00:00:00 2001 From: Funto-Adeyemi Date: Sun, 12 Nov 2023 21:38:11 +0000 Subject: [PATCH] Add troubleshooting.md file and revert yapf changes --- 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 +-- 14 files changed, 979 insertions(+), 919 deletions(-) diff --git a/experiments/baselines/timegpt_pipeline.py b/experiments/baselines/timegpt_pipeline.py index d66213f..8a2bfbd 100644 --- a/experiments/baselines/timegpt_pipeline.py +++ b/experiments/baselines/timegpt_pipeline.py @@ -34,8 +34,9 @@ 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]) @@ -63,15 +64,14 @@ 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,7 +121,8 @@ 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, @@ -137,7 +138,8 @@ 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) @@ -201,7 +203,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 0df159f..38d964d 100644 --- a/experiments/extended_benchmarks/run_timegpt.py +++ b/experiments/extended_benchmarks/run_timegpt.py @@ -11,6 +11,7 @@ # 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 @@ -24,6 +25,7 @@ import pandas as pd from ..baselines.timegpt_pipeline import run_timegpt from .utils import ExperimentHandler + dataset_names = [ "m1_monthly", "m1_quarterly", @@ -61,6 +63,7 @@ _MODEL_NAME = flags.DEFINE_string( ) _SAVE_DIR = flags.DEFINE_string("save_dir", "./results", "Save directory") + QUANTILES = list(np.arange(1, 10) / 10.0) @@ -87,9 +90,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 a0fd2c4..e8878c8 100644 --- a/experiments/extended_benchmarks/run_timesfm.py +++ b/experiments/extended_benchmarks/run_timesfm.py @@ -54,6 +54,7 @@ dataset_names = [ "hospital", ] + context_dict_v2 = {} context_dict_v1 = { diff --git a/experiments/extended_benchmarks/utils.py b/experiments/extended_benchmarks/utils.py index c9307f4..de0368b 100644 --- a/experiments/extended_benchmarks/utils.py +++ b/experiments/extended_benchmarks/utils.py @@ -11,6 +11,7 @@ # 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 @@ -45,9 +46,11 @@ 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 @@ -63,8 +66,10 @@ 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) @@ -75,8 +80,10 @@ 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 @@ -115,8 +122,9 @@ 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 @@ -145,8 +153,9 @@ 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 @@ -168,8 +177,9 @@ 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", @@ -205,21 +215,23 @@ 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") @@ -250,9 +262,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 3eee4f1..84c59f9 100644 --- a/peft/finetune.py +++ b/peft/finetune.py @@ -11,6 +11,7 @@ # 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. """ @@ -38,11 +39,13 @@ 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. @@ -59,24 +62,20 @@ 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, @@ -88,312 +87,316 @@ 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, - ) - - 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, + 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, ) - model = pax_fiddle.Config( - patched_decoder.PatchedDecoderFinetuneModel, - name="patched_decoder_finetune", - core_layer_tpl=tfm.model_p, - ) + train_batches = dtl.tf_dataset(mode="train", shift=1).batch(batch_size) + val_batches = dtl.tf_dataset(mode="val", shift=horizon_len) - 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, + 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, ) - @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/.*"] - - 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, + 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, - 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) + @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/.*"] - patience = 0 - del jax_state_for_saving - gc.collect() - else: - patience += 1 - print(f"patience: {patience}") - print("Fine-tuning completed.") + 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.") if __name__ == "__main__": - typer.run(finetune) + typer.run(finetune) diff --git a/src/adapter/__init__.py b/src/adapter/__init__.py index 5bd05be..6870b01 100644 --- a/src/adapter/__init__.py +++ b/src/adapter/__init__.py @@ -11,6 +11,7 @@ # 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 e49dee8..9a28911 100644 --- a/src/adapter/dora_layers.py +++ b/src/adapter/dora_layers.py @@ -21,181 +21,182 @@ WeightHParams = base_layer.WeightHParams class DoraTheta(base_layer.Theta): + def __init__(self, module): + self.module = module - 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 _dora_initialized(self): - if (self.module.has_variable("params", "lora_a") and - self.module.has_variable("params", "lora_b") and - self.module.has_variable("params", "dora_m") and - "lora_a" in self.module._weight_hparams and - "lora_b" in self.module._weight_hparams and - "dora_m" in self.module._weight_hparams): - return True - else: - return False + def _dorafy_var(self, w): + lora_a = super().__getattr__("lora_a") + lora_b = super().__getattr__("lora_b") + dora_m = super().__getattr__("dora_m") - def _dorafy_var(self, w): - lora_a = super().__getattr__("lora_a") - lora_b = super().__getattr__("lora_b") - dora_m = super().__getattr__("dora_m") + lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) + lora_delta = jnp.reshape(lora_delta, w.shape) - 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 - w_prime = w + lora_delta + column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) + norm_adapted = w_prime / column_norm + w_prime = dora_m * norm_adapted + return w_prime - column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) - norm_adapted = w_prime / column_norm - w_prime = dora_m * norm_adapted - return w_prime + def __getattr__(self, k): + var = super().__getattr__(k) + if not self._dora_initialized(): + return var - def __getattr__(self, k): - var = super().__getattr__(k) - if not self._dora_initialized(): - return var + if k == "w": + return self._dorafy_var(var) - if k == "w": - return self._dorafy_var(var) + return var - return var + def __getitem__(self, k): + var = super().__getattr__(k) + if not self._dora_initialized(): + 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) - if k == "w": - return self._dorafy_var(var) - - return 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 8eaeb2a..15df5a5 100644 --- a/src/adapter/lora_layers.py +++ b/src/adapter/lora_layers.py @@ -21,145 +21,146 @@ WeightHParams = base_layer.WeightHParams class LoraTheta(base_layer.Theta): + def __init__(self, module): + self.module = module - 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 _lora_initialized(self): - if (self.module.has_variable("params", "lora_a") and - self.module.has_variable("params", "lora_b") and - "lora_a" in self.module._weight_hparams and - "lora_b" in self.module._weight_hparams): - return True - else: - return False + def _lorafy_var(self, w): + lora_a = super().__getattr__("lora_a") + lora_b = super().__getattr__("lora_b") + lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) + lora_delta = jnp.reshape(lora_delta, w.shape) + w_prime = w + lora_delta + return w_prime - def _lorafy_var(self, w): - lora_a = super().__getattr__("lora_a") - lora_b = super().__getattr__("lora_b") - lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) - lora_delta = jnp.reshape(lora_delta, w.shape) - w_prime = w + lora_delta - return w_prime + def __getattr__(self, k): + var = super().__getattr__(k) + if not self._lora_initialized(): + return var - def __getattr__(self, k): - var = super().__getattr__(k) - if not self._lora_initialized(): - return var + if k == "w": + return self._lorafy_var(var) - if k == "w": - return self._lorafy_var(var) + return var - return var + def __getitem__(self, k): + var = super().__getattr__(k) + if not self._lora_initialized(): + 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) - if k == "w": - return self._lorafy_var(var) - - return 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 a26f381..4c3fc5b 100644 --- a/src/adapter/utils.py +++ b/src/adapter/utils.py @@ -11,6 +11,7 @@ # 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. @@ -39,11 +40,10 @@ 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,44 +55,47 @@ def get_adapter_params(params: dict, 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( @@ -102,7 +105,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: @@ -115,77 +118,83 @@ 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( @@ -195,7 +204,7 @@ def _merge_adapter_weights( num_layers: int, use_dora: bool, ) -> None: - """ + """ Merges adapter weights with the original model weights. Args: @@ -205,73 +214,74 @@ 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: @@ -283,39 +293,42 @@ def _get_adapter_weight_params(var_weight_hparams: dict, 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( @@ -325,7 +338,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: @@ -338,55 +351,67 @@ 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( @@ -397,7 +422,7 @@ def _initialize_adapter_params( use_dora: bool = False, seed: int = 1234, ) -> dict: - """ + """ Initializes and adds adapter parameters to target modules. Args: @@ -411,47 +436,52 @@ 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 747d6e9..f1d76af 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,12 +157,11 @@ 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 f929ada..739807a 100644 --- a/src/timesfm/__init__.py +++ b/src/timesfm/__init__.py @@ -25,13 +25,11 @@ 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 0cc0861..0bd90a9 100644 --- a/src/timesfm/time_features.py +++ b/src/timesfm/time_features.py @@ -11,6 +11,7 @@ # 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. @@ -35,6 +36,7 @@ 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 @@ -48,8 +50,9 @@ 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 @@ -57,19 +60,16 @@ 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 615e798..73b1495 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,11 +57,18 @@ 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 713a970..3277a9a 100644 --- a/tests/test_timesfm.py +++ b/tests/test_timesfm.py @@ -12,6 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. + from datetime import datetime, timedelta import numpy as np @@ -21,10 +22,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: @@ -35,10 +36,10 @@ def create_sample_dataframe(start_date: datetime, 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]) @@ -49,41 +50,42 @@ 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}" + )