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
+13 -11
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,14 +64,13 @@ 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
@@ -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 = {
+35 -23
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,7 +122,8 @@ 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 = [
int(100 - 200 * q) for q in quantiles if q < 0.5
] # in this case mean=mediain ] # 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])
+72 -69
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,66 +87,67 @@ 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())
@@ -261,7 +261,9 @@ def finetune(
task_p = tasks_lib.SingleTask( task_p = tasks_lib.SingleTask(
name="ts-learn", name="ts-learn",
model=model, model=model,
train=tasks_lib.SingleTask.Train(learner=build_learner(),), train=tasks_lib.SingleTask.Train(
learner=build_learner(),
),
) )
task_p.model.ici_mesh_shape = [1, 1, 1] task_p.model.ici_mesh_shape = [1, 1, 1]
@@ -294,19 +296,18 @@ def finetune(
checkpoint_type=checkpoint_types.CheckpointType.GDA, checkpoint_type=checkpoint_types.CheckpointType.GDA,
) )
jax_model_states.mdl_vars["params"]["core_layer"] = tfm._train_state.mdl_vars[ jax_model_states.mdl_vars["params"]["core_layer"] = tfm._train_state.mdl_vars[
"params"] "params"
]
gc.collect() gc.collect()
jax_task = task_p jax_task = task_p
def train_step(states, prng_key, inputs): def train_step(states, prng_key, inputs):
return trainer_lib.train_step_single_learner(jax_task, states, prng_key, return trainer_lib.train_step_single_learner(jax_task, states, prng_key, inputs)
inputs)
def eval_step(states, prng_key, inputs): def eval_step(states, prng_key, inputs):
states = states.to_eval_state() states = states.to_eval_state()
return trainer_lib.eval_step_single_learner(jax_task, states, prng_key, return trainer_lib.eval_step_single_learner(jax_task, states, prng_key, inputs)
inputs)
key, train_key, eval_key = jax.random.split(key, 3) key, train_key, eval_key = jax.random.split(key, 3)
train_prng_seed = jax.random.split(train_key, num=jax.local_device_count()) train_prng_seed = jax.random.split(train_key, num=jax.local_device_count())
@@ -318,7 +319,6 @@ def finetune(
replicated_jax_states = trainer_lib.replicate_model_state(jax_model_states) replicated_jax_states = trainer_lib.replicate_model_state(jax_model_states)
def reshape_batch_for_pmap(batch, num_devices): def reshape_batch_for_pmap(batch, num_devices):
def _reshape(input_tensor): def _reshape(input_tensor):
bsize = input_tensor.shape[0] bsize = input_tensor.shape[0]
residual_shape = list(input_tensor.shape[1:]) residual_shape = list(input_tensor.shape[1:])
@@ -341,7 +341,8 @@ def finetune(
tbatch = process_train_batch(batch) tbatch = process_train_batch(batch)
tbatch = reshape_batch_for_pmap(tbatch, num_devices) tbatch = reshape_batch_for_pmap(tbatch, num_devices)
replicated_jax_states, step_fun_out = p_train_step( replicated_jax_states, step_fun_out = p_train_step(
replicated_jax_states, train_prng_seed, tbatch) replicated_jax_states, train_prng_seed, tbatch
)
train_losses.append(step_fun_out.loss[0]) train_losses.append(step_fun_out.loss[0])
wandb.log({"train_step_loss": step_fun_out.loss[0]}) wandb.log({"train_step_loss": step_fun_out.loss[0]})
@@ -353,8 +354,7 @@ def finetune(
for ev_batch in tqdm(val_its): for ev_batch in tqdm(val_its):
ebatch = process_eval_batch(ev_batch) ebatch = process_eval_batch(ev_batch)
ebatch = reshape_batch_for_pmap(ebatch, num_devices) ebatch = reshape_batch_for_pmap(ebatch, num_devices)
_, step_fun_out = p_eval_step(replicated_jax_states, eval_prng_seed, _, step_fun_out = p_eval_step(replicated_jax_states, eval_prng_seed, ebatch)
ebatch)
eval_losses.append(step_fun_out.loss[0]) eval_losses.append(step_fun_out.loss[0])
wandb.log({"eval_step_loss": step_fun_out.loss[0]}) wandb.log({"eval_step_loss": step_fun_out.loss[0]})
@@ -362,17 +362,20 @@ def finetune(
print(f"Train Loss: {avg_train_loss}, Val Loss: {avg_eval_loss}") print(f"Train Loss: {avg_train_loss}, Val Loss: {avg_eval_loss}")
wandb.log({ wandb.log(
{
"epoch": epoch + 1, "epoch": epoch + 1,
"avg_train_loss": avg_train_loss, "avg_train_loss": avg_train_loss,
"avg_val_loss": avg_eval_loss, "avg_val_loss": avg_eval_loss,
}) }
)
if avg_eval_loss < best_eval_loss or np.isnan(avg_eval_loss): if avg_eval_loss < best_eval_loss or np.isnan(avg_eval_loss):
best_eval_loss = avg_eval_loss best_eval_loss = avg_eval_loss
print("Saving checkpoint.") print("Saving checkpoint.")
jax_state_for_saving = py_utils.maybe_unreplicate_for_fully_replicated( jax_state_for_saving = py_utils.maybe_unreplicate_for_fully_replicated(
replicated_jax_states) replicated_jax_states
)
if use_lora: if use_lora:
adapter_params = get_adapter_params( adapter_params = get_adapter_params(
params=jax_state_for_saving.mdl_vars, params=jax_state_for_saving.mdl_vars,
@@ -382,9 +385,9 @@ def finetune(
) )
jax_state_for_saving.mdl_vars["params"] = adapter_params jax_state_for_saving.mdl_vars["params"] = adapter_params
checkpoints.save_checkpoint(jax_state_for_saving, checkpoints.save_checkpoint(
checkpoint_dir, jax_state_for_saving, checkpoint_dir, overwrite=True
overwrite=True) )
patience = 0 patience = 0
del jax_state_for_saving del jax_state_for_saving
+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
+8 -7
View File
@@ -21,17 +21,18 @@ WeightHParams = base_layer.WeightHParams
class DoraTheta(base_layer.Theta): class DoraTheta(base_layer.Theta):
def __init__(self, module): def __init__(self, module):
self.module = module self.module = module
def _dora_initialized(self): def _dora_initialized(self):
if (self.module.has_variable("params", "lora_a") and if (
self.module.has_variable("params", "lora_b") and self.module.has_variable("params", "lora_a")
self.module.has_variable("params", "dora_m") and and self.module.has_variable("params", "lora_b")
"lora_a" in self.module._weight_hparams and and self.module.has_variable("params", "dora_m")
"lora_b" in self.module._weight_hparams and and "lora_a" in self.module._weight_hparams
"dora_m" in self.module._weight_hparams): and "lora_b" in self.module._weight_hparams
and "dora_m" in self.module._weight_hparams
):
return True return True
else: else:
return False return False
+6 -5
View File
@@ -21,15 +21,16 @@ WeightHParams = base_layer.WeightHParams
class LoraTheta(base_layer.Theta): class LoraTheta(base_layer.Theta):
def __init__(self, module): def __init__(self, module):
self.module = module self.module = module
def _lora_initialized(self): def _lora_initialized(self):
if (self.module.has_variable("params", "lora_a") and if (
self.module.has_variable("params", "lora_b") and self.module.has_variable("params", "lora_a")
"lora_a" in self.module._weight_hparams and and self.module.has_variable("params", "lora_b")
"lora_b" in self.module._weight_hparams): and "lora_a" in self.module._weight_hparams
and "lora_b" in self.module._weight_hparams
):
return True return True
else: else:
return False return False
+81 -51
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,10 +40,9 @@ 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.
@@ -63,7 +63,8 @@ def get_adapter_params(params: dict,
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"]
@@ -78,7 +79,8 @@ def get_adapter_params(params: dict,
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"]
@@ -90,8 +92,9 @@ def get_adapter_params(params: dict,
} }
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
]["dora_m"]
return adapter_params return adapter_params
@@ -115,13 +118,13 @@ 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(
@@ -130,10 +133,12 @@ def load_adapter_checkpoint(
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,
@@ -174,15 +179,19 @@ def load_adapter_checkpoint(
# 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 = ( model._model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
original_combined_qkv_tpl) original_combined_qkv_tpl
)
model._logging( model._logging(
f"Restored adapter checkpoint in {time.time() - start_time:.2f} seconds.") 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()
@@ -211,8 +220,8 @@ def _merge_adapter_weights(
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"]
@@ -240,7 +249,8 @@ def _merge_adapter_weights(
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]
@@ -268,9 +278,9 @@ def _merge_adapter_weights(
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.
@@ -291,8 +301,8 @@ def _get_adapter_weight_params(var_weight_hparams: dict,
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"],
@@ -300,12 +310,14 @@ def _get_adapter_weight_params(var_weight_hparams: dict,
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"
][layer]["self_attention"][component]
adapter_params[layer][component] = { adapter_params[layer][component] = {
"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"],
@@ -313,7 +325,8 @@ def _get_adapter_weight_params(var_weight_hparams: dict,
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
@@ -341,41 +354,53 @@ def load_adapter_layer(
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 = (
pax_fiddle.Config(
DoraLinear, DoraLinear,
rank=lora_rank, rank=lora_rank,
) if use_dora else pax_fiddle.Config( )
if use_dora
else pax_fiddle.Config(
LoraLinear, LoraLinear,
rank=lora_rank, rank=lora_rank,
)) )
)
adapter_linear_tpl.copy_fields_from(original_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 = ( model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_fflayer_tpl.fflayer_tpl.linear_tpl = (
adapter_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
else pax_fiddle.Config(LoraAttentionProjection, rank=lora_rank)
)
adapter_attn_tpl.copy_fields_from(original_attn_tpl) 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
else pax_fiddle.Config(LoraCombinedQKVProjection, rank=lora_rank)
)
adapter_combined_qkv_tpl.copy_fields_from(original_combined_qkv_tpl) 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 = ( model.stacked_transformer_params_tpl.transformer_layer_params_tpl.tr_atten_tpl.combined_qkv_proj_tpl = (
adapter_combined_qkv_tpl) adapter_combined_qkv_tpl
)
# initialize and add adapter weights # initialize and add adapter weights
_initialize_adapter_params( _initialize_adapter_params(
@@ -416,14 +441,16 @@ def _initialize_adapter_params(
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"
][ff_layer_key]["linear"]
original_w = linear["w"] original_w = linear["w"]
input_dim, output_dim = original_w.shape input_dim, output_dim = original_w.shape
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), 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
@@ -435,7 +462,8 @@ def _initialize_adapter_params(
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"]
@@ -443,15 +471,17 @@ def _initialize_adapter_params(
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
).astype(jnp.float32)
attention[component]["dora_m"] = norm attention[component]["dora_m"] = norm
return mdl_vars return mdl_vars
+3 -4
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."""
@@ -161,8 +161,7 @@ def get_model(load_weights: bool = False):
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)
+1 -3
View File
@@ -32,6 +32,4 @@ try:
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}."
)
+13 -13
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,7 +50,8 @@ 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 (
len(holiday_date) != 0 # pylint: disable=g-explicit-length-test
), f"No closest holiday for the date index {index} found." ), 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.
@@ -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)
+11 -4
View File
@@ -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}")
+6 -4
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,9 +22,9 @@ 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.
@@ -74,7 +75,8 @@ def test_timesfm_forecast_on_df(
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 (
"timesfm" in forecast_df.columns
), "Forecast DataFrame should contain 'timesfm' column" ), "Forecast DataFrame should contain 'timesfm' column"
last_input_date = input_df["ds"].max() last_input_date = input_df["ds"].max()