Add troubleshooting.md file and revert yapf changes

This commit is contained in:
Funto-Adeyemi
2023-11-12 21:38:11 +00:00
parent 1f9cb2bf92
commit 63f0bf4162
14 changed files with 979 additions and 919 deletions
+14 -12
View File
@@ -34,8 +34,9 @@ def get_seasonality(freq: str) -> int:
return _get_seasonality(freq, seasonalities={"D": 7}) return _get_seasonality(freq, seasonalities={"D": 7})
def maybe_convert_col_to_datetime(df: pd.DataFrame, def maybe_convert_col_to_datetime(
col_name: str) -> pd.DataFrame: df: pd.DataFrame, col_name: str
) -> pd.DataFrame:
if not pd.api.types.is_datetime64_any_dtype(df[col_name]): if not pd.api.types.is_datetime64_any_dtype(df[col_name]):
df = df.copy() df = df.copy()
df[col_name] = pd.to_datetime(df[col_name]) 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, end=start_date,
periods=min_length - len(subset) + 1, periods=min_length - len(subset) + 1,
freq=freq, # 'MS' for month start freq=freq, # 'MS' for month start
)[:-1] # Exclude the start_date itself )[
:-1
] # Exclude the start_date itself
# 2c. Create padding data # 2c. Create padding data
padding_df = pd.DataFrame({ padding_df = pd.DataFrame(
"ds": padding_dates, {"ds": padding_dates, "unique_id": unique_id, "y": 0} # Zero padding
"unique_id": unique_id, )
"y": 0
} # Zero padding
)
# 2d. Combine original and padding data, and append to the list # 2d. Combine original and padding data, and append to the list
padded_data.append(pd.concat([padding_df, subset]).sort_values("ds")) 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)): for _, (cutoffs, train, valid) in tqdm(enumerate(splits)):
if len(valid.columns) > 3: if len(valid.columns) > 3:
raise NotImplementedError( raise NotImplementedError(
"Cross validation with exogenous variables is not yet supported.") "Cross validation with exogenous variables is not yet supported."
)
y_pred = self.forecast( y_pred = self.forecast(
df=train, df=train,
h=h, h=h,
@@ -137,7 +138,8 @@ class Forecaster:
raise ValueError( raise ValueError(
"Cross validation result produced less results than expected." "Cross validation result produced less results than expected."
" Please verify that the frequency parameter (freq) matches your" " 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) results.append(result)
out = vertical_concat(results) out = vertical_concat(results)
out = drop_index_if_pandas(out) out = drop_index_if_pandas(out)
@@ -201,7 +203,7 @@ class TimeGPT(Forecaster):
all_unique_ids = df["unique_id"].unique() all_unique_ids = df["unique_id"].unique()
all_fcst_df = [] all_fcst_df = []
for i in range(0, len(all_unique_ids), chunk_size): 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)] chunk_df = df[df["unique_id"].isin(chunk_ids)]
fct_chunk_df = client.forecast( fct_chunk_df = client.forecast(
df=chunk_df, df=chunk_df,
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
"""Evaluation script for timegpt.""" """Evaluation script for timegpt."""
import os import os
@@ -24,6 +25,7 @@ import pandas as pd
from ..baselines.timegpt_pipeline import run_timegpt from ..baselines.timegpt_pipeline import run_timegpt
from .utils import ExperimentHandler from .utils import ExperimentHandler
dataset_names = [ dataset_names = [
"m1_monthly", "m1_monthly",
"m1_quarterly", "m1_quarterly",
@@ -61,6 +63,7 @@ _MODEL_NAME = flags.DEFINE_string(
) )
_SAVE_DIR = flags.DEFINE_string("save_dir", "./results", "Save directory") _SAVE_DIR = flags.DEFINE_string("save_dir", "./results", "Save directory")
QUANTILES = list(np.arange(1, 10) / 10.0) QUANTILES = list(np.arange(1, 10) / 10.0)
@@ -87,9 +90,9 @@ def main():
) )
time_df = pd.DataFrame({"time": [total_time], "model": model_name}) time_df = pd.DataFrame({"time": [total_time], "model": model_name})
fcsts_df = exp.fcst_from_level_to_quantiles(fcsts_df, model_name) fcsts_df = exp.fcst_from_level_to_quantiles(fcsts_df, model_name)
results = exp.evaluate_from_predictions(models=[model_name], results = exp.evaluate_from_predictions(
fcsts_df=fcsts_df, models=[model_name], fcsts_df=fcsts_df, times_df=time_df
times_df=time_df) )
print(results, flush=True) print(results, flush=True)
results_list.append(results) results_list.append(results)
results_full = pd.concat(results_list) results_full = pd.concat(results_list)
@@ -54,6 +54,7 @@ dataset_names = [
"hospital", "hospital",
] ]
context_dict_v2 = {} context_dict_v2 = {}
context_dict_v1 = { context_dict_v1 = {
+36 -24
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
"""Forked from https://github.com/Nixtla/nixtla/blob/main/experiments/amazon-chronos/src/utils.py.""" """Forked from https://github.com/Nixtla/nixtla/blob/main/experiments/amazon-chronos/src/utils.py."""
from functools import partial from functools import partial
@@ -45,9 +46,11 @@ def quantile_loss(
target_col: str = "y", target_col: str = "y",
) -> pd.DataFrame: ) -> pd.DataFrame:
delta_y = df[models].sub(df[target_col], axis=0) delta_y = df[models].sub(df[target_col], axis=0)
res = (np.maximum(q * delta_y, res = (
(q - 1) * delta_y).groupby(df[id_col], np.maximum(q * delta_y, (q - 1) * delta_y)
observed=True).mean()) .groupby(df[id_col], observed=True)
.mean()
)
res.index.name = id_col res.index.name = id_col
res = res.reset_index() res = res.reset_index()
return res return res
@@ -63,8 +66,10 @@ class ExperimentHandler:
models_dir: str = "./models", models_dir: str = "./models",
): ):
if dataset not in gluonts_datasets: if dataset not in gluonts_datasets:
raise Exception(f"dataset {dataset} not found in gluonts " raise Exception(
f"available datasets: {', '.join(gluonts_datasets)}") f"dataset {dataset} not found in gluonts "
f"available datasets: {', '.join(gluonts_datasets)}"
)
self.dataset = dataset self.dataset = dataset
self.quantiles = quantiles self.quantiles = quantiles
self.level = self._transform_quantiles_to_levels(quantiles) self.level = self._transform_quantiles_to_levels(quantiles)
@@ -75,8 +80,10 @@ class ExperimentHandler:
gluonts_dataset = get_dataset(self.dataset) gluonts_dataset = get_dataset(self.dataset)
self.horizon = gluonts_dataset.metadata.prediction_length self.horizon = gluonts_dataset.metadata.prediction_length
if self.horizon is None: if self.horizon is None:
raise Exception(f"horizon not found for dataset {self.dataset} " raise Exception(
"experiment cannot be run") f"horizon not found for dataset {self.dataset} "
"experiment cannot be run"
)
self.freq = gluonts_dataset.metadata.freq self.freq = gluonts_dataset.metadata.freq
# get_seasonality() returns 1 for freq='D', override this to 7. This significantly improves the accuracy of # 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 # statistical models on datasets like m5/nn5_daily. The models like AutoARIMA/AutoETS can still set
@@ -115,8 +122,9 @@ class ExperimentHandler:
@staticmethod @staticmethod
def _transform_quantiles_to_levels(quantiles: List[float]) -> List[int]: def _transform_quantiles_to_levels(quantiles: List[float]) -> List[int]:
level = [int(100 - 200 * q) for q in quantiles if q < 0.5 level = [
] # in this case mean=mediain int(100 - 200 * q) for q in quantiles if q < 0.5
] # in this case mean=mediain
level = sorted(list(set(level))) level = sorted(list(set(level)))
return level return level
@@ -145,8 +153,9 @@ class ExperimentHandler:
last_n: int | None = None, last_n: int | None = None,
) -> pd.DataFrame: ) -> pd.DataFrame:
with multiprocessing.Pool(os.cpu_count()) as pool: # Create a process pool with multiprocessing.Pool(os.cpu_count()) as pool: # Create a process pool
results = pool.map(parallel_transform, zip(gluonts_dataset, results = pool.map(
repeat(last_n))) parallel_transform, zip(gluonts_dataset, repeat(last_n))
)
df = pd.concat(results) df = pd.concat(results)
df = df.reset_index(drop=True) df = df.reset_index(drop=True)
return df return df
@@ -168,8 +177,9 @@ class ExperimentHandler:
def save_dataframe(self, df: pd.DataFrame, file_name: str): def save_dataframe(self, df: pd.DataFrame, file_name: str):
df.to_csv(f"{self.results_dir}/{file_name}", index=False) df.to_csv(f"{self.results_dir}/{file_name}", index=False)
def save_results(self, fcst_df: pd.DataFrame, total_time: float, def save_results(
model_name: str): self, fcst_df: pd.DataFrame, total_time: float, model_name: str
):
self.save_dataframe( self.save_dataframe(
fcst_df, fcst_df,
f"{model_name}-{self.dataset}-fcst.csv", f"{model_name}-{self.dataset}-fcst.csv",
@@ -205,21 +215,23 @@ class ExperimentHandler:
times_df = [] times_df = []
for model in models: for model in models:
fcst_method_df = pd.read_csv( fcst_method_df = pd.read_csv(
f"{self.results_dir}/{model}-{self.dataset}-fcst.csv").set_index( f"{self.results_dir}/{model}-{self.dataset}-fcst.csv"
["unique_id", "ds"]) ).set_index(["unique_id", "ds"])
fcsts_df.append(fcst_method_df) fcsts_df.append(fcst_method_df)
time_method_df = pd.read_csv( 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) times_df.append(time_method_df)
fcsts_df = pd.concat(fcsts_df, axis=1).reset_index() fcsts_df = pd.concat(fcsts_df, axis=1).reset_index()
fcsts_df["ds"] = pd.to_datetime(fcsts_df["ds"]) fcsts_df["ds"] = pd.to_datetime(fcsts_df["ds"])
times_df = pd.concat(times_df) times_df = pd.concat(times_df)
return self.evaluate_from_predictions(models=models, return self.evaluate_from_predictions(
fcsts_df=fcsts_df, models=models, fcsts_df=fcsts_df, times_df=times_df
times_df=times_df) )
def evaluate_from_predictions(self, models: List[str], fcsts_df: pd.DataFrame, def evaluate_from_predictions(
times_df: pd.DataFrame) -> pd.DataFrame: self, models: List[str], fcsts_df: pd.DataFrame, times_df: pd.DataFrame
) -> pd.DataFrame:
test_df = self.test_df test_df = self.test_df
train_df = self.train_df train_df = self.train_df
test_df = test_df.merge(fcsts_df, how="left") test_df = test_df.merge(fcsts_df, how="left")
@@ -250,9 +262,9 @@ class ExperimentHandler:
eval_prob_df["metric"] = "scaled_crps" eval_prob_df["metric"] = "scaled_crps"
eval_df = pd.concat([eval_df, eval_prob_df]).reset_index(drop=True) 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.groupby("metric").mean(numeric_only=True).reset_index()
eval_df = eval_df.melt(id_vars="metric", eval_df = eval_df.melt(
value_name="value", id_vars="metric", value_name="value", var_name="model"
var_name="model") )
times_df.insert(0, "metric", "time") times_df.insert(0, "metric", "time")
times_df = times_df.rename(columns={"time": "value"}) times_df = times_df.rename(columns={"time": "value"})
eval_df = pd.concat([eval_df, times_df]) eval_df = pd.concat([eval_df, times_df])
+286 -283
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
""" """
Finetune pipeline. Finetune pipeline.
""" """
@@ -38,11 +39,13 @@ from timesfm import TimesFm, data_loader, patched_decoder
NestedMap = py_utils.NestedMap NestedMap = py_utils.NestedMap
warnings.filterwarnings("ignore") warnings.filterwarnings("ignore")
cmdstanpy_logger = logging.getLogger("cmdstanpy") cmdstanpy_logger = logging.getLogger("cmdstanpy")
absl_logger = logging.getLogger("absl") absl_logger = logging.getLogger("absl")
cmdstanpy_logger.disabled = True cmdstanpy_logger.disabled = True
absl_logger.disabled = True absl_logger.disabled = True
""" """
TimesFM model config. These are fixed since pre-training was done TimesFM model config. These are fixed since pre-training was done
with this configuration. with this configuration.
@@ -59,24 +62,20 @@ RANDOM_SEED = 1234
def finetune( def finetune(
*, *,
model_name: Annotated[str, model_name: Annotated[
typer.Option( str, typer.Option(help="Specify the name of the huggingface model.")
help="Specify the name of the huggingface model." ] = "google/timesfm-1.0-200m",
)] = "google/timesfm-1.0-200m",
checkpoint_path: Annotated[ checkpoint_path: Annotated[
str, str, typer.Option(help="The path to the local model checkpoint.")
typer.Option(help="The path to the local model checkpoint.")] = None, ] = None,
datetime_col: Annotated[str, datetime_col: Annotated[str, typer.Option(help="Column having datetime.")] = "ds",
typer.Option( ts_cols: Annotated[
help="Column having datetime.")] = "ds", list[str], typer.Option(help="Columns of time-series features.")
ts_cols: Annotated[list[str], ] = [],
typer.Option( normalize: Annotated[
help="Columns of time-series features.")] = [], bool, typer.Option(help="Normalize data for eval or not")
normalize: Annotated[bool, ] = True,
typer.Option( context_len: Annotated[int, typer.Option(help="Length of the context window")],
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.")], horizon_len: Annotated[int, typer.Option(help="Prediction length.")],
freq: Annotated[ freq: Annotated[
str, str,
@@ -88,312 +87,316 @@ def finetune(
data_path: Annotated[str, typer.Option(help="Path to dataset csv")], data_path: Annotated[str, typer.Option(help="Path to dataset csv")],
boundaries: Annotated[ boundaries: Annotated[
Tuple[int, int, int], 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), ] = (0, 0, 0),
backend: Annotated[str, backend: Annotated[str, typer.Option(help="Backend device: cpu, gpu, tpu")],
typer.Option(help="Backend device: cpu, gpu, tpu")],
batch_size: Annotated[ batch_size: Annotated[
int, int, typer.Option(help="Batch size for the randomly sampled batch")
typer.Option(help="Batch size for the randomly sampled batch")] = 16, ] = 16,
num_epochs: Annotated[int, typer.Option(help="Number of epochs")], num_epochs: Annotated[int, typer.Option(help="Number of epochs")],
learning_rate: Annotated[float, learning_rate: Annotated[float, typer.Option(help="adam optimizer learning rate")],
typer.Option(help="adam optimizer learning rate")], adam_epsilon: Annotated[float, typer.Option(help="adam optimizer epsilon")],
adam_epsilon: Annotated[float, adam_clip_threshold: Annotated[
typer.Option(help="adam optimizer epsilon")], float, typer.Option(help="adam optimizer clip threshold")
adam_clip_threshold: Annotated[float, ],
typer.Option( cos_initial_decay_value: Annotated[
help="adam optimizer clip threshold")], float, typer.Option(help="cosine initial decay value")
cos_initial_decay_value: Annotated[float, ],
typer.Option( cos_final_decay_value: Annotated[
help="cosine initial decay value")], float, typer.Option(help="cosine final decay value")
cos_final_decay_value: Annotated[float, ],
typer.Option( cos_decay_steps: Annotated[int, typer.Option(help="Number of cosine decay steps")],
help="cosine final decay value")], ema_decay: Annotated[float, typer.Option(help="Exponential moving average decay")],
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[ early_stop_patience: Annotated[
int, typer.Option(..., help="Early stopping patience")] = 5, int, typer.Option(..., help="Early stopping patience")
] = 5,
use_lora: Annotated[ use_lora: Annotated[
bool, bool,
typer. typer.Option(
Option(help="Train low rank adapters for stacked transformer block",), help="Train low rank adapters for stacked transformer block",
),
] = False, ] = False,
lora_rank: Annotated[ lora_rank: Annotated[
int, int,
typer.Option(help="LoRA Rank",), typer.Option(
help="LoRA Rank",
),
] = 8, ] = 8,
lora_target_modules: Annotated[ lora_target_modules: Annotated[
str, str,
typer.Option( typer.Option(
help= help="LoRA target modules of the transformer block. Allowed values: [all, attention, mlp]"
"LoRA target modules of the transformer block. Allowed values: [all, attention, mlp]"
), ),
] = "all", ] = "all",
use_dora: Annotated[ use_dora: Annotated[
bool, bool,
typer.Option(help="Apply DoRA strategy along with LoRA.",), typer.Option(
help="Apply DoRA strategy along with LoRA.",
),
] = False, ] = False,
use_linear_probing: Annotated[ use_linear_probing: Annotated[
bool, bool,
typer.Option( typer.Option(
help= help="Linear Probing. Train only input/output and embedding params. Freeze params in stack transformer block.",
"Linear Probing. Train only input/output and embedding params. Freeze params in stack transformer block.",
), ),
] = False, ] = False,
checkpoint_dir: Annotated[ checkpoint_dir: Annotated[
str, typer.Option(help="Checkpoint directory")] = "./checkpoints", str, typer.Option(help="Checkpoint directory")
wandb_project: Annotated[str, ] = "./checkpoints",
typer.Option(help="Weights & Biases project name" wandb_project: Annotated[
)] = "google_timesfm_finetune", str, typer.Option(help="Weights & Biases project name")
] = "google_timesfm_finetune",
) -> None: ) -> None:
key = jax.random.PRNGKey(seed=RANDOM_SEED) key = jax.random.PRNGKey(seed=RANDOM_SEED)
wandb.init(project=wandb_project, config=locals()) 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): if boundaries == (0, 0, 0):
# Default boundaries: train 60%, val 20%, test 20% # Default boundaries: train 60%, val 20%, test 20%
boundaries = [ boundaries = [
int(len(data_df) * 0.6), int(len(data_df) * 0.6),
int(len(data_df) * 0.8), int(len(data_df) * 0.8),
len(data_df) - 1, 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( dtl = data_loader.TimeSeriesdata(
data_path=data_path, data_path=data_path,
datetime_col=datetime_col, datetime_col=datetime_col,
num_cov_cols=None, num_cov_cols=None,
cat_cov_cols=None, cat_cov_cols=None,
ts_cols=np.array(ts_cols), ts_cols=np.array(ts_cols),
train_range=[0, boundaries[0]], train_range=[0, boundaries[0]],
val_range=[boundaries[0], boundaries[1]], val_range=[boundaries[0], boundaries[1]],
test_range=[boundaries[1], boundaries[2]], test_range=[boundaries[1], boundaries[2]],
hist_len=context_len, hist_len=context_len,
pred_len=horizon_len, pred_len=horizon_len,
batch_size=batch_size, batch_size=batch_size,
freq=freq, freq=freq,
normalize=normalize, normalize=normalize,
epoch_len=None, epoch_len=None,
holiday=False, holiday=False,
permute=False, permute=False,
)
train_batches = dtl.tf_dataset(mode="train", shift=1).batch(batch_size)
val_batches = dtl.tf_dataset(mode="val", shift=horizon_len)
for tbatch in tqdm(train_batches.as_numpy_iterator()):
pass
tfm = TimesFm(
context_len=context_len,
horizon_len=horizon_len,
input_patch_len=INPUT_PATCH_LEN,
output_patch_len=OUTPUT_PATCH_LEN,
num_layers=NUM_LAYERS,
model_dims=MODEL_DIMS,
backend=backend,
per_core_batch_size=batch_size,
quantiles=QUANTILES,
)
if checkpoint_path:
tfm.load_from_checkpoint(
checkpoint_path=checkpoint_path,
checkpoint_type=checkpoints.CheckpointType.FLAX,
)
else:
tfm.load_from_checkpoint(
repo_id=model_name,
checkpoint_type=checkpoints.CheckpointType.FLAX,
) )
model = pax_fiddle.Config( train_batches = dtl.tf_dataset(mode="train", shift=1).batch(batch_size)
patched_decoder.PatchedDecoderFinetuneModel, val_batches = dtl.tf_dataset(mode="val", shift=horizon_len)
name="patched_decoder_finetune",
core_layer_tpl=tfm.model_p,
)
if use_lora: for tbatch in tqdm(train_batches.as_numpy_iterator()):
load_adapter_layer( pass
mdl_vars=tfm._train_state.mdl_vars,
model=model.core_layer_tpl, tfm = TimesFm(
lora_rank=lora_rank, context_len=context_len,
lora_target_modules=lora_target_modules, horizon_len=horizon_len,
use_dora=use_dora, 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: if use_lora:
bprop_variable_inclusion.append(r"^.*lora.*$") load_adapter_layer(
if use_dora: mdl_vars=tfm._train_state.mdl_vars,
bprop_variable_inclusion.append(r"^.*dora.*$") model=model.core_layer_tpl,
elif use_linear_probing: lora_rank=lora_rank,
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, lora_target_modules=lora_target_modules,
num_layers=NUM_LAYERS,
use_dora=use_dora, use_dora=use_dora,
) )
jax_state_for_saving.mdl_vars["params"] = adapter_params
checkpoints.save_checkpoint(jax_state_for_saving, @pax_fiddle.auto_config
checkpoint_dir, def build_learner() -> learners.Learner:
overwrite=True) 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 return pax_fiddle.Config(
del jax_state_for_saving learners.Learner,
gc.collect() name="learner",
else: loss_name="avg_qloss",
patience += 1 optimizer=optimizers.Adam(
print(f"patience: {patience}") epsilon=adam_epsilon,
print("Fine-tuning completed.") 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__": if __name__ == "__main__":
typer.run(finetune) typer.run(finetune)
+1
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
"""adapter init file.""" """adapter init file."""
from .dora_layers import DoraAttentionProjection, DoraCombinedQKVProjection, DoraLinear from .dora_layers import DoraAttentionProjection, DoraCombinedQKVProjection, DoraLinear
+149 -148
View File
@@ -21,181 +21,182 @@ WeightHParams = base_layer.WeightHParams
class DoraTheta(base_layer.Theta): class DoraTheta(base_layer.Theta):
def __init__(self, module):
self.module = module
def __init__(self, module): def _dora_initialized(self):
self.module = module 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): def _dorafy_var(self, w):
if (self.module.has_variable("params", "lora_a") and lora_a = super().__getattr__("lora_a")
self.module.has_variable("params", "lora_b") and lora_b = super().__getattr__("lora_b")
self.module.has_variable("params", "dora_m") and dora_m = super().__getattr__("dora_m")
"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_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_a = super().__getattr__("lora_a") lora_delta = jnp.reshape(lora_delta, w.shape)
lora_b = super().__getattr__("lora_b")
dora_m = super().__getattr__("dora_m")
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) w_prime = w + lora_delta
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted
return w_prime
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) def __getattr__(self, k):
norm_adapted = w_prime / column_norm var = super().__getattr__(k)
w_prime = dora_m * norm_adapted if not self._dora_initialized():
return w_prime return var
def __getattr__(self, k): if k == "w":
var = super().__getattr__(k) return self._dorafy_var(var)
if not self._dora_initialized():
return var
if k == "w": return var
return self._dorafy_var(var)
return var def __getitem__(self, k):
var = super().__getattr__(k)
if not self._dora_initialized():
return var
def __getitem__(self, k): if k == "w":
var = super().__getattr__(k) return self._dorafy_var(var)
if not self._dora_initialized():
return var
if k == "w": return var
return self._dorafy_var(var)
return var
class DoraThetaDescriptor: class DoraThetaDescriptor:
"""Dot syntax accession descriptor.""" """Dot syntax accession descriptor."""
def __get__(self, obj, objtype=None): def __get__(self, obj, objtype=None):
return DoraTheta(obj) return DoraTheta(obj)
class DoraLinear(linears.Linear): class DoraLinear(linears.Linear):
rank: int = 0 rank: int = 0
lora_init: WeightInit | None = None lora_init: WeightInit | None = None
theta = DoraThetaDescriptor() theta = DoraThetaDescriptor()
def setup(self) -> None: def setup(self) -> None:
lora_init = self.lora_init if self.lora_init else self.weight_init lora_init = self.lora_init if self.lora_init else self.weight_init
super().setup() super().setup()
self.create_variable( self.create_variable(
"lora_a", "lora_a",
WeightHParams( WeightHParams(
shape=[self.input_dims, self.rank], shape=[self.input_dims, self.rank],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None], tensor_split_dims_mapping=[None, None],
), ),
) )
self.create_variable( self.create_variable(
"lora_b", "lora_b",
WeightHParams( WeightHParams(
shape=[self.output_dims, self.rank], shape=[self.output_dims, self.rank],
init=WeightInit.Constant(scale=0.0), init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None], tensor_split_dims_mapping=[None, None],
), ),
) )
self.create_variable( self.create_variable(
"dora_m", "dora_m",
WeightHParams( WeightHParams(
shape=[1, self.output_dims], shape=[1, self.output_dims],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None], tensor_split_dims_mapping=[None, None],
), ),
) )
class DoraAttentionProjection(attentions.AttentionProjection): class DoraAttentionProjection(attentions.AttentionProjection):
rank: int = 0 rank: int = 0
lora_init: WeightInit | None = None lora_init: WeightInit | None = None
theta = DoraThetaDescriptor() theta = DoraThetaDescriptor()
def setup(self) -> None: def setup(self) -> None:
super().setup() super().setup()
w_weight_params = self._weight_hparams["w"] w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable( self.create_variable(
"lora_a", "lora_a",
WeightHParams( WeightHParams(
shape=[self.input_dim, self.rank], shape=[self.input_dim, self.rank],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[ tensor_split_dims_mapping=[
None, None,
None, None,
], ],
), ),
) )
self.create_variable( self.create_variable(
"lora_b", "lora_b",
WeightHParams( WeightHParams(
shape=[self.dim_per_head * self.num_heads, self.rank], shape=[self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0), init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[ tensor_split_dims_mapping=[
None, None,
None, None,
], ],
), ),
) )
self.create_variable( self.create_variable(
"dora_m", "dora_m",
WeightHParams( WeightHParams(
shape=[1, self.num_heads, self.dim_per_head], shape=[1, self.num_heads, self.dim_per_head],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None], tensor_split_dims_mapping=[None, None, None],
), ),
) )
class DoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer): class DoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer):
rank: int = 0 rank: int = 0
lora_init: WeightInit | None = None lora_init: WeightInit | None = None
theta = DoraThetaDescriptor() theta = DoraThetaDescriptor()
def setup(self) -> None: def setup(self) -> None:
super().setup() super().setup()
w_weight_params = self._weight_hparams["w"] w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable( self.create_variable(
"lora_a", "lora_a",
WeightHParams( WeightHParams(
shape=[3, self.input_dim, self.rank], shape=[3, self.input_dim, self.rank],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None], tensor_split_dims_mapping=[None, None, None],
), ),
) )
self.create_variable( self.create_variable(
"lora_b", "lora_b",
WeightHParams( WeightHParams(
shape=[3, self.dim_per_head * self.num_heads, self.rank], shape=[3, self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0), init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None], tensor_split_dims_mapping=[None, None, None],
), ),
) )
self.create_variable( self.create_variable(
"dora_m", "dora_m",
WeightHParams( WeightHParams(
shape=[3, 1, self.num_heads, self.dim_per_head], shape=[3, 1, self.num_heads, self.dim_per_head],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None, None], tensor_split_dims_mapping=[None, None, None, None],
), ),
) )
+116 -115
View File
@@ -21,145 +21,146 @@ WeightHParams = base_layer.WeightHParams
class LoraTheta(base_layer.Theta): class LoraTheta(base_layer.Theta):
def __init__(self, module):
self.module = module
def __init__(self, module): def _lora_initialized(self):
self.module = module 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): def _lorafy_var(self, w):
if (self.module.has_variable("params", "lora_a") and lora_a = super().__getattr__("lora_a")
self.module.has_variable("params", "lora_b") and lora_b = super().__getattr__("lora_b")
"lora_a" in self.module._weight_hparams and lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b)
"lora_b" in self.module._weight_hparams): lora_delta = jnp.reshape(lora_delta, w.shape)
return True w_prime = w + lora_delta
else: return w_prime
return False
def _lorafy_var(self, w): def __getattr__(self, k):
lora_a = super().__getattr__("lora_a") var = super().__getattr__(k)
lora_b = super().__getattr__("lora_b") if not self._lora_initialized():
lora_delta = self.module.einsum("...dr,...nr->...dn", lora_a, lora_b) return var
lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta
return w_prime
def __getattr__(self, k): if k == "w":
var = super().__getattr__(k) return self._lorafy_var(var)
if not self._lora_initialized():
return var
if k == "w": return var
return self._lorafy_var(var)
return var def __getitem__(self, k):
var = super().__getattr__(k)
if not self._lora_initialized():
return var
def __getitem__(self, k): if k == "w":
var = super().__getattr__(k) return self._lorafy_var(var)
if not self._lora_initialized():
return var
if k == "w": return var
return self._lorafy_var(var)
return var
class LoraThetaDescriptor: class LoraThetaDescriptor:
"""Dot syntax accession descriptor.""" """Dot syntax accession descriptor."""
def __get__(self, obj, objtype=None): def __get__(self, obj, objtype=None):
return LoraTheta(obj) return LoraTheta(obj)
class LoraLinear(linears.Linear): class LoraLinear(linears.Linear):
rank: int = 0 rank: int = 0
lora_init: WeightInit | None = None lora_init: WeightInit | None = None
theta = LoraThetaDescriptor() theta = LoraThetaDescriptor()
def setup(self) -> None: def setup(self) -> None:
lora_init = self.lora_init if self.lora_init else self.weight_init lora_init = self.lora_init if self.lora_init else self.weight_init
super().setup() super().setup()
self.create_variable( self.create_variable(
"lora_a", "lora_a",
WeightHParams( WeightHParams(
shape=[self.input_dims, self.rank], shape=[self.input_dims, self.rank],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None], tensor_split_dims_mapping=[None, None],
), ),
) )
self.create_variable( self.create_variable(
"lora_b", "lora_b",
WeightHParams( WeightHParams(
shape=[self.output_dims, self.rank], shape=[self.output_dims, self.rank],
init=WeightInit.Constant(scale=0.0), init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None], tensor_split_dims_mapping=[None, None],
), ),
) )
class LoraAttentionProjection(attentions.AttentionProjection): class LoraAttentionProjection(attentions.AttentionProjection):
rank: int = 0 rank: int = 0
lora_init: WeightInit | None = None lora_init: WeightInit | None = None
theta = LoraThetaDescriptor() theta = LoraThetaDescriptor()
def setup(self) -> None: def setup(self) -> None:
super().setup() super().setup()
w_weight_params = self._weight_hparams["w"] w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable( self.create_variable(
"lora_a", "lora_a",
WeightHParams( WeightHParams(
shape=[self.input_dim, self.rank], shape=[self.input_dim, self.rank],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[ tensor_split_dims_mapping=[
None, None,
None, None,
], ],
), ),
) )
self.create_variable( self.create_variable(
"lora_b", "lora_b",
WeightHParams( WeightHParams(
shape=[self.dim_per_head * self.num_heads, self.rank], shape=[self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0), init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[ tensor_split_dims_mapping=[
None, None,
None, None,
], ],
), ),
) )
class LoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer): class LoraCombinedQKVProjection(attentions.CombinedQKVProjectionLayer):
rank: int = 0 rank: int = 0
lora_init: WeightInit | None = None lora_init: WeightInit | None = None
theta = LoraThetaDescriptor() theta = LoraThetaDescriptor()
def setup(self) -> None: def setup(self) -> None:
super().setup() super().setup()
w_weight_params = self._weight_hparams["w"] w_weight_params = self._weight_hparams["w"]
lora_init = self.lora_init if self.lora_init else w_weight_params.init lora_init = self.lora_init if self.lora_init else w_weight_params.init
self.create_variable( self.create_variable(
"lora_a", "lora_a",
WeightHParams( WeightHParams(
shape=[3, self.input_dim, self.rank], shape=[3, self.input_dim, self.rank],
init=lora_init, init=lora_init,
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None], tensor_split_dims_mapping=[None, None, None],
), ),
) )
self.create_variable( self.create_variable(
"lora_b", "lora_b",
WeightHParams( WeightHParams(
shape=[3, self.dim_per_head * self.num_heads, self.rank], shape=[3, self.dim_per_head * self.num_heads, self.rank],
init=WeightInit.Constant(scale=0.0), init=WeightInit.Constant(scale=0.0),
mesh_shape=self.mesh_shape, mesh_shape=self.mesh_shape,
tensor_split_dims_mapping=[None, None, None], tensor_split_dims_mapping=[None, None, None],
), ),
) )
+286 -256
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
""" """
This file provides functionality for loading and merging adapter weights This file provides functionality for loading and merging adapter weights
in timesfm model, specifically for LoRA and DoRA. in timesfm model, specifically for LoRA and DoRA.
@@ -39,11 +40,10 @@ from adapter.lora_layers import (
from timesfm import TimesFm from timesfm import TimesFm
def get_adapter_params(params: dict, def get_adapter_params(
lora_target_modules: str, params: dict, lora_target_modules: str, num_layers: int, use_dora: bool = False
num_layers: int, ) -> dict:
use_dora: bool = False) -> dict: """
"""
Extracts adapter parameters from the given model parameters for saving the checkpoint. Extracts adapter parameters from the given model parameters for saving the checkpoint.
Args: Args:
@@ -55,44 +55,47 @@ def get_adapter_params(params: dict,
Returns: Returns:
dict: A dictionary containing the extracted adapter parameters. dict: A dictionary containing the extracted adapter parameters.
""" """
adapter_params = {} adapter_params = {}
for i in range(num_layers): for i in range(num_layers):
layer_key = f"x_layers_{i}" layer_key = f"x_layers_{i}"
adapter_params[layer_key] = {} adapter_params[layer_key] = {}
if lora_target_modules in ["all", "mlp"]: if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]: for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = params["params"]["core_layer"]["stacked_transformer_layer"][ linear = params["params"]["core_layer"]["stacked_transformer_layer"][
layer_key]["ff_layer"][ff_layer_key]["linear"] layer_key
]["ff_layer"][ff_layer_key]["linear"]
lora_a = linear["lora_a"] lora_a = linear["lora_a"]
lora_b = linear["lora_b"] lora_b = linear["lora_b"]
adapter_params[layer_key][ff_layer_key] = { adapter_params[layer_key][ff_layer_key] = {
"lora_a": lora_a, "lora_a": lora_a,
"lora_b": lora_b, "lora_b": lora_b,
} }
if use_dora: if use_dora:
adapter_params[layer_key][ff_layer_key]["dora_m"] = linear["dora_m"] adapter_params[layer_key][ff_layer_key]["dora_m"] = linear["dora_m"]
if lora_target_modules in ["all", "attention"]: if lora_target_modules in ["all", "attention"]:
attention = params["params"]["core_layer"]["stacked_transformer_layer"][ attention = params["params"]["core_layer"]["stacked_transformer_layer"][
layer_key]["self_attention"] layer_key
]["self_attention"]
for component in ["key", "query", "value", "post"]: for component in ["key", "query", "value", "post"]:
lora_a = attention[component]["lora_a"] lora_a = attention[component]["lora_a"]
lora_b = attention[component]["lora_b"] lora_b = attention[component]["lora_b"]
adapter_params[layer_key][component] = { adapter_params[layer_key][component] = {
"lora_a": lora_a, "lora_a": lora_a,
"lora_b": lora_b, "lora_b": lora_b,
} }
if use_dora: if use_dora:
adapter_params[layer_key][component]["dora_m"] = attention[component][ adapter_params[layer_key][component]["dora_m"] = attention[
"dora_m"] component
return adapter_params ]["dora_m"]
return adapter_params
def load_adapter_checkpoint( def load_adapter_checkpoint(
@@ -102,7 +105,7 @@ def load_adapter_checkpoint(
lora_target_modules: str, lora_target_modules: str,
use_dora: bool, use_dora: bool,
) -> None: ) -> None:
""" """
Loads an adapter checkpoint and merges it with the original model weights. Loads an adapter checkpoint and merges it with the original model weights.
Args: Args:
@@ -115,77 +118,83 @@ def load_adapter_checkpoint(
Returns: Returns:
None None
""" """
"""
"""
currently loading and initializing the model with adapter layers first and then merging the 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. 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. # NOTE: refactor this. there should be a better way to load the LoRA checkpoint.
""" """
model._logging( model._logging(f"Restoring adapter checkpoint from {adapter_checkpoint_path}.")
f"Restoring adapter checkpoint from {adapter_checkpoint_path}.") start_time = time.time()
start_time = time.time() original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl = (
original_linear_tpl, original_attn_tpl, original_combined_qkv_tpl = ( load_adapter_layer(
load_adapter_layer( mdl_vars=model._train_state.mdl_vars,
mdl_vars=model._train_state.mdl_vars, model=model._model,
model=model._model, lora_rank=lora_rank,
lora_rank=lora_rank, lora_target_modules=lora_target_modules,
lora_target_modules=lora_target_modules, use_dora=use_dora,
use_dora=use_dora, )
)) )
var_weight_hparams = model._model.abstract_init_with_metadata( var_weight_hparams = model._model.abstract_init_with_metadata(
model._get_sample_inputs(), do_eval=True) model._get_sample_inputs(), do_eval=True
)
adapter_weight_hparams = _get_adapter_weight_params( adapter_weight_hparams = _get_adapter_weight_params(
var_weight_hparams=var_weight_hparams, var_weight_hparams=var_weight_hparams,
lora_target_modules=lora_target_modules, lora_target_modules=lora_target_modules,
num_layers=model._model.stacked_transformer_params_tpl.num_layers, num_layers=model._model.stacked_transformer_params_tpl.num_layers,
use_dora=use_dora, use_dora=use_dora,
) )
adapter_state_partition_specs = tasks_lib.create_state_partition_specs( adapter_state_partition_specs = tasks_lib.create_state_partition_specs(
adapter_weight_hparams, adapter_weight_hparams,
mesh_shape=model.mesh_shape, mesh_shape=model.mesh_shape,
mesh_axis_names=model.mesh_name, mesh_axis_names=model.mesh_name,
discard_opt_states=True, discard_opt_states=True,
learners=None, learners=None,
) )
adapter_state_local_shapes = tasks_lib.create_state_unpadded_shapes( adapter_state_local_shapes = tasks_lib.create_state_unpadded_shapes(
adapter_weight_hparams, adapter_weight_hparams,
discard_opt_states=True, discard_opt_states=True,
learners=None, learners=None,
) )
adapter_train_state = checkpoints.restore_checkpoint( adapter_train_state = checkpoints.restore_checkpoint(
state_global_shapes=adapter_state_local_shapes, state_global_shapes=adapter_state_local_shapes,
checkpoint_dir=adapter_checkpoint_path, checkpoint_dir=adapter_checkpoint_path,
checkpoint_type=checkpoints.CheckpointType.FLAX, checkpoint_type=checkpoints.CheckpointType.FLAX,
state_specs=adapter_state_partition_specs, state_specs=adapter_state_partition_specs,
step=None, step=None,
) )
# add adapter weights to the original weights # add adapter weights to the original weights
_merge_adapter_weights( _merge_adapter_weights(
model=model, model=model,
adapter_train_state=adapter_train_state, adapter_train_state=adapter_train_state,
lora_target_modules=lora_target_modules, lora_target_modules=lora_target_modules,
num_layers=model._model.stacked_transformer_params_tpl.num_layers, num_layers=model._model.stacked_transformer_params_tpl.num_layers,
use_dora=use_dora, use_dora=use_dora,
) )
# replace back with the original model layer # replace back with the original model layer
if lora_target_modules in ["all", "mlp"]: if lora_target_modules in ["all", "mlp"]:
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = ( model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = (
original_linear_tpl) original_linear_tpl
)
if lora_target_modules in ["all", "attention"]: if lora_target_modules in ["all", "attention"]:
model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = ( model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = (
original_attn_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._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
model._logging( original_combined_qkv_tpl
f"Restored adapter checkpoint in {time.time() - start_time:.2f} seconds.") )
model._logging(
f"Restored adapter checkpoint in {time.time() - start_time:.2f} seconds."
)
# jit compile the model # jit compile the model
model.jit_decode() model.jit_decode()
def _merge_adapter_weights( def _merge_adapter_weights(
@@ -195,7 +204,7 @@ def _merge_adapter_weights(
num_layers: int, num_layers: int,
use_dora: bool, use_dora: bool,
) -> None: ) -> None:
""" """
Merges adapter weights with the original model weights. Merges adapter weights with the original model weights.
Args: Args:
@@ -205,73 +214,74 @@ def _merge_adapter_weights(
num_layers (int): Number of transformer layers. num_layers (int): Number of transformer layers.
use_dora (bool): Whether DoRA was used or not. use_dora (bool): Whether DoRA was used or not.
""" """
for i in range(num_layers): for i in range(num_layers):
layer_key = f"x_layers_{i}" layer_key = f"x_layers_{i}"
if lora_target_modules in ["all", "mlp"]: if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]: for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = model._train_state.mdl_vars["params"][ linear = model._train_state.mdl_vars["params"][
"stacked_transformer_layer"][layer_key]["ff_layer"][ff_layer_key][ "stacked_transformer_layer"
"linear"] ][layer_key]["ff_layer"][ff_layer_key]["linear"]
params = adapter_train_state.mdl_vars[layer_key][ff_layer_key] params = adapter_train_state.mdl_vars[layer_key][ff_layer_key]
lora_a = params["lora_a"] lora_a = params["lora_a"]
lora_b = params["lora_b"] 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.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape) lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta w_prime = w + lora_delta
if use_dora: if use_dora:
dora_m = params["dora_m"] dora_m = params["dora_m"]
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted w_prime = dora_m * norm_adapted
linear["w"] = w_prime linear["w"] = w_prime
del linear["dora_m"] del linear["dora_m"]
else: else:
linear["w"] = w_prime linear["w"] = w_prime
del linear["lora_a"] del linear["lora_a"]
del linear["lora_b"] del linear["lora_b"]
if lora_target_modules in ["all", "attention"]: if lora_target_modules in ["all", "attention"]:
attention = model._train_state.mdl_vars["params"][ attention = model._train_state.mdl_vars["params"][
"stacked_transformer_layer"][layer_key]["self_attention"] "stacked_transformer_layer"
][layer_key]["self_attention"]
for component in ["key", "query", "value", "post"]: for component in ["key", "query", "value", "post"]:
params = adapter_train_state.mdl_vars[layer_key][component] params = adapter_train_state.mdl_vars[layer_key][component]
lora_a = params["lora_a"] lora_a = params["lora_a"]
lora_b = params["lora_b"] 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.einsum("...dr,...nr->...dn", lora_a, lora_b)
lora_delta = jnp.reshape(lora_delta, w.shape) lora_delta = jnp.reshape(lora_delta, w.shape)
w_prime = w + lora_delta w_prime = w + lora_delta
if use_dora: if use_dora:
dora_m = params["dora_m"] dora_m = params["dora_m"]
column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True) column_norm = jnp.linalg.norm(w_prime, ord=2, axis=0, keepdims=True)
norm_adapted = w_prime / column_norm norm_adapted = w_prime / column_norm
w_prime = dora_m * norm_adapted w_prime = dora_m * norm_adapted
attention[component]["w"] = w_prime attention[component]["w"] = w_prime
del attention[component]["dora_m"] del attention[component]["dora_m"]
else: else:
attention[component]["w"] = w_prime attention[component]["w"] = w_prime
del attention[component]["lora_a"] del attention[component]["lora_a"]
del attention[component]["lora_b"] del attention[component]["lora_b"]
def _get_adapter_weight_params(var_weight_hparams: dict, def _get_adapter_weight_params(
lora_target_modules: str, num_layers: int, var_weight_hparams: dict, lora_target_modules: str, num_layers: int, use_dora: bool
use_dora: bool) -> dict: ) -> dict:
""" """
Extracts adapter weight parameters from the given variable weight hyperparameters. Extracts adapter weight parameters from the given variable weight hyperparameters.
Args: Args:
@@ -283,39 +293,42 @@ def _get_adapter_weight_params(var_weight_hparams: dict,
Returns: Returns:
dict: A dictionary containing the extracted adapter weight parameters. dict: A dictionary containing the extracted adapter weight parameters.
""" """
adapter_params = {} adapter_params = {}
for i in range(num_layers): for i in range(num_layers):
layer = f"x_layers_{i}" layer = f"x_layers_{i}"
adapter_params[layer] = {} adapter_params[layer] = {}
if lora_target_modules in ["all", "mlp"]: if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]: for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
adapter_weight_params = var_weight_hparams["params"][ adapter_weight_params = var_weight_hparams["params"][
"stacked_transformer_layer"][layer]["ff_layer"][ff_layer_key][ "stacked_transformer_layer"
"linear"] ][layer]["ff_layer"][ff_layer_key]["linear"]
adapter_params[layer][ff_layer_key] = { adapter_params[layer][ff_layer_key] = {
"lora_a": adapter_weight_params["lora_a"], "lora_a": adapter_weight_params["lora_a"],
"lora_b": adapter_weight_params["lora_b"], "lora_b": adapter_weight_params["lora_b"],
} }
if use_dora: if use_dora:
adapter_params[layer][ff_layer_key]["dora_m"] = ( adapter_params[layer][ff_layer_key]["dora_m"] = (
adapter_weight_params["dora_m"]) adapter_weight_params["dora_m"]
)
if lora_target_modules in ["all", "attention"]: if lora_target_modules in ["all", "attention"]:
for component in ["key", "value", "query", "post"]: for component in ["key", "value", "query", "post"]:
adapter_weight_params = var_weight_hparams["params"][ adapter_weight_params = var_weight_hparams["params"][
"stacked_transformer_layer"][layer]["self_attention"][component] "stacked_transformer_layer"
adapter_params[layer][component] = { ][layer]["self_attention"][component]
"lora_a": adapter_weight_params["lora_a"], adapter_params[layer][component] = {
"lora_b": adapter_weight_params["lora_b"], "lora_a": adapter_weight_params["lora_a"],
} "lora_b": adapter_weight_params["lora_b"],
}
if use_dora: if use_dora:
adapter_params[layer][component]["dora_m"] = adapter_weight_params[ adapter_params[layer][component]["dora_m"] = adapter_weight_params[
"dora_m"] "dora_m"
]
return adapter_params return adapter_params
def load_adapter_layer( def load_adapter_layer(
@@ -325,7 +338,7 @@ def load_adapter_layer(
lora_target_modules: str, lora_target_modules: str,
use_dora: bool = False, use_dora: bool = False,
) -> tuple[pax_fiddle.Config, pax_fiddle.Config]: ) -> tuple[pax_fiddle.Config, pax_fiddle.Config]:
""" """
Updates target modules with adapter layers. Updates target modules with adapter layers.
Args: Args:
@@ -338,55 +351,67 @@ def load_adapter_layer(
Returns: Returns:
tuple[pax_fiddle.Config, pax_fiddle.Config]: Updated model configurations. tuple[pax_fiddle.Config, pax_fiddle.Config]: Updated model configurations.
""" """
original_linear_tpl = original_attn_tpl = original_combined_qkv_tpl = None original_linear_tpl = original_attn_tpl = original_combined_qkv_tpl = None
if lora_target_modules in ["all", "mlp"]: if lora_target_modules in ["all", "mlp"]:
original_linear_tpl = ( original_linear_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl. model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl
tr_fflayer_tpl.fflayer_tpl.linear_tpl) )
adapter_linear_tpl = (pax_fiddle.Config( adapter_linear_tpl = (
DoraLinear, pax_fiddle.Config(
rank=lora_rank, DoraLinear,
) if use_dora else pax_fiddle.Config( rank=lora_rank,
LoraLinear, )
rank=lora_rank, if use_dora
)) else pax_fiddle.Config(
adapter_linear_tpl.copy_fields_from(original_linear_tpl) LoraLinear,
model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = ( rank=lora_rank,
adapter_linear_tpl) )
)
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"]: if lora_target_modules in ["all", "attention"]:
original_attn_tpl = (model.stacked_transformer_params_tpl. original_attn_tpl = (
transformer_layer_params_tpl.tr_atten_tpl.proj_tpl) model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl
)
adapter_attn_tpl = ( adapter_attn_tpl = (
pax_fiddle.Config(DoraAttentionProjection, rank=lora_rank) if use_dora pax_fiddle.Config(DoraAttentionProjection, rank=lora_rank)
else pax_fiddle.Config(LoraAttentionProjection, rank=lora_rank)) if use_dora
adapter_attn_tpl.copy_fields_from(original_attn_tpl) else pax_fiddle.Config(LoraAttentionProjection, rank=lora_rank)
)
adapter_attn_tpl.copy_fields_from(original_attn_tpl)
original_combined_qkv_tpl = ( original_combined_qkv_tpl = (
model.stacked_transformer_params_tpl.transformer_layer_params_tpl. model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl
tr_atten_tpl.combined_qkv_proj_tpl) )
adapter_combined_qkv_tpl = ( adapter_combined_qkv_tpl = (
pax_fiddle.Config(DoraCombinedQKVProjection, rank=lora_rank) if use_dora pax_fiddle.Config(DoraCombinedQKVProjection, rank=lora_rank)
else pax_fiddle.Config(LoraCombinedQKVProjection, rank=lora_rank)) if use_dora
adapter_combined_qkv_tpl.copy_fields_from(original_combined_qkv_tpl) 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 = ( model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.proj_tpl = (
adapter_attn_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.combined_qkv_proj_tpl = (
adapter_combined_qkv_tpl
)
# initialize and add adapter weights # initialize and add adapter weights
_initialize_adapter_params( _initialize_adapter_params(
mdl_vars=mdl_vars, mdl_vars=mdl_vars,
num_layers=model.stacked_transformer_params_tpl.num_layers, num_layers=model.stacked_transformer_params_tpl.num_layers,
lora_rank=lora_rank, lora_rank=lora_rank,
lora_target_modules=lora_target_modules, lora_target_modules=lora_target_modules,
use_dora=use_dora, 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( def _initialize_adapter_params(
@@ -397,7 +422,7 @@ def _initialize_adapter_params(
use_dora: bool = False, use_dora: bool = False,
seed: int = 1234, seed: int = 1234,
) -> dict: ) -> dict:
""" """
Initializes and adds adapter parameters to target modules. Initializes and adds adapter parameters to target modules.
Args: Args:
@@ -411,47 +436,52 @@ def _initialize_adapter_params(
Returns: Returns:
dict: Updated model variables with initialized adapter parameters. dict: Updated model variables with initialized adapter parameters.
""" """
for i in range(num_layers): for i in range(num_layers):
layer_key = f"x_layers_{i}" layer_key = f"x_layers_{i}"
if lora_target_modules in ["all", "mlp"]: if lora_target_modules in ["all", "mlp"]:
for ff_layer_key in ["ffn_layer1", "ffn_layer2"]: for ff_layer_key in ["ffn_layer1", "ffn_layer2"]:
linear = mdl_vars["params"]["stacked_transformer_layer"][layer_key][ linear = mdl_vars["params"]["stacked_transformer_layer"][layer_key][
"ff_layer"][ff_layer_key]["linear"] "ff_layer"
original_w = linear["w"] ][ff_layer_key]["linear"]
input_dim, output_dim = original_w.shape original_w = linear["w"]
std_dev = 1 / jnp.sqrt(lora_rank) input_dim, output_dim = original_w.shape
std_dev = 1 / jnp.sqrt(lora_rank)
normal_initializer = jax.nn.initializers.normal(std_dev) normal_initializer = jax.nn.initializers.normal(std_dev)
lora_a = normal_initializer(jax.random.key(seed), lora_a = normal_initializer(
(input_dim, lora_rank), jnp.float32) jax.random.key(seed), (input_dim, lora_rank), jnp.float32
lora_b = jnp.zeros((output_dim, lora_rank)) )
lora_b = jnp.zeros((output_dim, lora_rank))
linear["lora_a"] = lora_a linear["lora_a"] = lora_a
linear["lora_b"] = lora_b linear["lora_b"] = lora_b
if use_dora: if use_dora:
norm = jnp.linalg.norm(original_w, ord=2, axis=0, keepdims=True) norm = jnp.linalg.norm(original_w, ord=2, axis=0, keepdims=True)
linear["dora_m"] = norm linear["dora_m"] = norm
if lora_target_modules in ["all", "attention"]: if lora_target_modules in ["all", "attention"]:
attention = mdl_vars["params"]["stacked_transformer_layer"][layer_key][ attention = mdl_vars["params"]["stacked_transformer_layer"][layer_key][
"self_attention"] "self_attention"
]
for component in ["key", "query", "value", "post"]: for component in ["key", "query", "value", "post"]:
original_w = attention[component]["w"] original_w = attention[component]["w"]
w_dim = original_w.shape[0] w_dim = original_w.shape[0]
std_dev = 1 / jnp.sqrt(lora_rank) std_dev = 1 / jnp.sqrt(lora_rank)
normal_initializer = jax.nn.initializers.normal(std_dev) normal_initializer = jax.nn.initializers.normal(std_dev)
lora_a = normal_initializer(jax.random.key(seed), (w_dim, lora_rank), lora_a = normal_initializer(
jnp.float32) jax.random.key(seed), (w_dim, lora_rank), jnp.float32
lora_b = jnp.zeros((w_dim, lora_rank)) )
lora_b = jnp.zeros((w_dim, lora_rank))
attention[component]["lora_a"] = lora_a attention[component]["lora_a"] = lora_a
attention[component]["lora_b"] = lora_b attention[component]["lora_b"] = lora_b
if use_dora: if use_dora:
norm = jnp.linalg.norm(original_w, ord=2, axis=0, norm = jnp.linalg.norm(
keepdims=True).astype(jnp.float32) original_w, ord=2, axis=0, keepdims=True
attention[component]["dora_m"] = norm ).astype(jnp.float32)
return mdl_vars attention[component]["dora_m"] = norm
return mdl_vars
+4 -5
View File
@@ -43,11 +43,11 @@ flags.DEFINE_list(
) )
flags.DEFINE_string( flags.DEFINE_string(
"local_model_path", None, "local_model_path",
None,
"Path to a local .safetensors model file. If provided, overrides Hugging Face download." "Path to a local .safetensors model file. If provided, overrides Hugging Face download."
) )
class TimeSeriesDataset(Dataset): class TimeSeriesDataset(Dataset):
"""Dataset for time series data compatible with TimesFM.""" """Dataset for time series data compatible with TimesFM."""
@@ -157,12 +157,11 @@ def get_model(load_weights: bool = False):
else: else:
repo_id = "google/timesfm-2.0-500m-pytorch" repo_id = "google/timesfm-2.0-500m-pytorch"
tfm = TimesFm(hparams=hparams, tfm = TimesFm(hparams=hparams,
checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id)) checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))
tfm_config = tfm._model_config tfm_config = tfm._model_config
model = PatchedTimeSeriesDecoder(tfm_config) model = PatchedTimeSeriesDecoder(tfm_config)
checkpoint_path = path.join(snapshot_download(repo_id), checkpoint_path = path.join(snapshot_download(repo_id), "torch_model.ckpt")
"torch_model.ckpt")
loaded_checkpoint = torch.load(checkpoint_path, weights_only=True) loaded_checkpoint = torch.load(checkpoint_path, weights_only=True)
model.load_state_dict(loaded_checkpoint) model.load_state_dict(loaded_checkpoint)
+5 -7
View File
@@ -25,13 +25,11 @@ from timesfm.timesfm_base import (
import sys import sys
try: try:
from timesfm.timesfm_jax import TimesFmJax as TimesFm from timesfm.timesfm_jax import TimesFmJax as TimesFm
from timesfm import data_loader 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 _: except Exception as _:
from timesfm.timesfm_torch import TimesFmTorch as TimesFm from timesfm.timesfm_torch import TimesFmTorch as TimesFm
print( print(f"Loaded PyTorch TimesFM, likely because python version is {sys.version}.")
f"Loaded PyTorch TimesFM, likely because python version is {sys.version}."
)
+14 -14
View File
@@ -11,6 +11,7 @@
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
"""Directory to extract time covariates. """Directory to extract time covariates.
Extract time covariates from datetime. Extract time covariates from datetime.
@@ -35,6 +36,7 @@ from pandas.tseries.offsets import Easter
from sklearn.preprocessing import StandardScaler from sklearn.preprocessing import StandardScaler
from tqdm import tqdm from tqdm import tqdm
# This is 183 to cover half a year (in both directions), also for leap years # 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 # + 17 as Eastern can be between March, 22 - April, 25
MAX_WINDOW = 183 + 17 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),
index + pd.Timedelta(days=MAX_WINDOW), index + pd.Timedelta(days=MAX_WINDOW),
) )
assert (len(holiday_date) != 0 # pylint: disable=g-explicit-length-test assert (
), f"No closest holiday for the date index {index} found." 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 # 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. # holiday. In this case, the smaller distance (182 days) is returned.
return (index - holiday_date[0]).days return (index - holiday_date[0]).days
@@ -57,19 +60,16 @@ def _distance_to_holiday(holiday):
return _distance_to_day return _distance_to_day
EasterSunday = Holiday("Easter Sunday", EasterSunday = Holiday(
month=1, "Easter Sunday", month=1, day=1, offset=[Easter(), Day(0)]
day=1, )
offset=[Easter(), Day(0)])
NewYearsDay = Holiday("New Years Day", month=1, day=1) NewYearsDay = Holiday("New Years Day", month=1, day=1)
SuperBowl = Holiday("Superbowl", SuperBowl = Holiday(
month=2, "Superbowl", month=2, day=1, offset=DateOffset(weekday=SU(1))
day=1, )
offset=DateOffset(weekday=SU(1))) MothersDay = Holiday(
MothersDay = Holiday("Mothers Day", "Mothers Day", month=5, day=1, offset=DateOffset(weekday=SU(2))
month=5, )
day=1,
offset=DateOffset(weekday=SU(2)))
IndependenceDay = Holiday("Independence Day", month=7, day=4) IndependenceDay = Holiday("Independence Day", month=7, day=4)
ChristmasEve = Holiday("Christmas", month=12, day=24) ChristmasEve = Holiday("Christmas", month=12, day=24)
ChristmasDay = Holiday("Christmas", month=12, day=25) ChristmasDay = Holiday("Christmas", month=12, day=25)
+16 -9
View File
@@ -25,12 +25,12 @@ import pandas as pd
from utilsforecast.processing import make_future_dataframe from utilsforecast.processing import make_future_dataframe
if TYPE_CHECKING: if TYPE_CHECKING:
from . import xreg_lib from . import xreg_lib
Category = xreg_lib.Category Category = xreg_lib.Category
XRegMode = xreg_lib.XRegMode XRegMode = xreg_lib.XRegMode
else: else:
Category = int | str Category = int | str
XRegMode = str XRegMode = str
_TOL = 1e-6 _TOL = 1e-6
DEFAULT_QUANTILES = (0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9) DEFAULT_QUANTILES = (0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9)
@@ -57,11 +57,18 @@ def freq_map(freq: str):
return 1 return 1
elif freq.endswith(("H", "T", "MIN", "D", "B", "U", "S")): elif freq.endswith(("H", "T", "MIN", "D", "B", "U", "S")):
return 0 return 0
elif (freq.endswith(("W", "M")) or freq.startswith("W-") or elif (
(freq.startswith("M") and len(freq) == 2)): freq.endswith(("W", "M"))
or freq.startswith("W-")
or (freq.startswith("M") and len(freq) == 2)
):
return 1 return 1
elif (freq.endswith(("Y", "Q", "A")) or freq.startswith("Y-") or elif (
freq.startswith("Q-") or freq.startswith("A-")): freq.endswith(("Y", "Q", "A"))
or freq.startswith("Y-")
or freq.startswith("Q-")
or freq.startswith("A-")
):
return 2 return 2
else: else:
raise ValueError(f"Invalid frequency: {freq}") raise ValueError(f"Invalid frequency: {freq}")
+43 -41
View File
@@ -12,6 +12,7 @@
# See the License for the specific language governing permissions and # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
from datetime import datetime, timedelta from datetime import datetime, timedelta
import numpy as np import numpy as np
@@ -21,10 +22,10 @@ import pytest
import timesfm import timesfm
def create_sample_dataframe(start_date: datetime, def create_sample_dataframe(
end_date: datetime, start_date: datetime, end_date: datetime, freq: str = "D"
freq: str = "D") -> pd.DataFrame: ) -> pd.DataFrame:
""" """
Create a sample DataFrame with time series data. Create a sample DataFrame with time series data.
Args: Args:
@@ -35,10 +36,10 @@ def create_sample_dataframe(start_date: datetime,
Returns: Returns:
pd.DataFrame: DataFrame with columns 'unique_id', 'ds', and 'ts'. pd.DataFrame: DataFrame with columns 'unique_id', 'ds', and 'ts'.
""" """
date_range = pd.date_range(start=start_date, end=end_date, freq=freq) date_range = pd.date_range(start=start_date, end=end_date, freq=freq)
ts_data = np.random.randn(len(date_range)) ts_data = np.random.randn(len(date_range))
df = pd.DataFrame({"unique_id": "ts-1", "ds": date_range, "ts": ts_data}) df = pd.DataFrame({"unique_id": "ts-1", "ds": date_range, "ts": ts_data})
return df return df
@pytest.mark.parametrize("context_length", [128, 256, 512]) @pytest.mark.parametrize("context_length", [128, 256, 512])
@@ -49,41 +50,42 @@ def test_timesfm_forecast_on_df(
prediction_length: int, prediction_length: int,
freq: str, freq: str,
) -> None: ) -> None:
model = timesfm.TimesFm( model = timesfm.TimesFm(
context_len=context_length, context_len=context_length,
horizon_len=prediction_length, horizon_len=prediction_length,
input_patch_len=32, input_patch_len=32,
output_patch_len=128, output_patch_len=128,
num_layers=20, num_layers=20,
model_dims=1280, model_dims=1280,
backend="cpu", backend="cpu",
) )
model.load_from_checkpoint(repo_id="google/timesfm-1.0-200m") model.load_from_checkpoint(repo_id="google/timesfm-1.0-200m")
end_date = datetime.now() end_date = datetime.now()
start_date = end_date - timedelta(days=context_length) start_date = end_date - timedelta(days=context_length)
input_df = create_sample_dataframe(start_date, end_date, freq) input_df = create_sample_dataframe(start_date, end_date, freq)
forecast_df = model.forecast_on_df( forecast_df = model.forecast_on_df(
inputs=input_df, inputs=input_df,
freq=freq, freq=freq,
value_name="ts", value_name="ts",
num_jobs=-1, num_jobs=-1,
) )
assert ( assert (
len(forecast_df) == prediction_length len(forecast_df) == prediction_length
), f"Expected forecast length of {prediction_length}, but got {len(forecast_df)}" ), f"Expected forecast length of {prediction_length}, but got {len(forecast_df)}"
assert ("timesfm" in forecast_df.columns assert (
), "Forecast DataFrame should contain 'timesfm' column" "timesfm" in forecast_df.columns
), "Forecast DataFrame should contain 'timesfm' column"
last_input_date = input_df["ds"].max() last_input_date = input_df["ds"].max()
first_forecast_date = forecast_df["ds"].min() first_forecast_date = forecast_df["ds"].min()
expected_first_forecast_date = last_input_date + pd.Timedelta(1, unit=freq) expected_first_forecast_date = last_input_date + pd.Timedelta(1, unit=freq)
assert ( assert (
first_forecast_date == expected_first_forecast_date first_forecast_date == expected_first_forecast_date
), f"Forecast should start from {expected_first_forecast_date}, but starts from {first_forecast_date}" ), f"Forecast should start from {expected_first_forecast_date}, but starts from {first_forecast_date}"
print( print(
f"Successful forecast with context_length={context_length}, prediction_length={prediction_length}, freq={freq}" f"Successful forecast with context_length={context_length}, prediction_length={prediction_length}, freq={freq}"
) )