From 6e77e0ca98f394694aec21193def48de5af41a48 Mon Sep 17 00:00:00 2001 From: Rajat Sen Date: Mon, 8 Jul 2024 18:31:35 +0000 Subject: [PATCH 1/2] Adding finetuning example + package restructuring --- .../long_horizon_benchmarks/run_eval.py | 2 +- notebooks/finetuning.ipynb | 612 ++++++++++++++++++ src/timesfm/__init__.py | 17 + src/timesfm/data_loader.py | 261 ++++++++ src/timesfm/patched_decoder.py | 521 +++++++++++++++ src/timesfm/time_features.py | 215 ++++++ src/timesfm/timesfm.py | 605 +++++++++++++++++ 7 files changed, 2232 insertions(+), 1 deletion(-) create mode 100644 notebooks/finetuning.ipynb create mode 100644 src/timesfm/__init__.py create mode 100644 src/timesfm/data_loader.py create mode 100644 src/timesfm/patched_decoder.py create mode 100644 src/timesfm/time_features.py create mode 100644 src/timesfm/timesfm.py diff --git a/experiments/long_horizon_benchmarks/run_eval.py b/experiments/long_horizon_benchmarks/run_eval.py index accfbd8..a079496 100644 --- a/experiments/long_horizon_benchmarks/run_eval.py +++ b/experiments/long_horizon_benchmarks/run_eval.py @@ -26,7 +26,7 @@ from paxml import checkpoints import timesfm import torch import tqdm -from . import data_loader +from timesfm import data_loader FLAGS = flags.FLAGS diff --git a/notebooks/finetuning.ipynb b/notebooks/finetuning.ipynb new file mode 100644 index 0000000..31da8f3 --- /dev/null +++ b/notebooks/finetuning.ipynb @@ -0,0 +1,612 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Importing relevant packages for finetuning" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import os\n", + "os.environ['XLA_PYTHON_CLIENT_PREALLOCATE'] = 'false'\n", + "os.environ['JAX_PMAP_USE_TENSORSTORE'] = 'false'" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import timesfm\n", + "import gc\n", + "import numpy as np\n", + "import pandas as pd\n", + "from timesfm import patched_decoder\n", + "from timesfm import data_loader" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "from tqdm import tqdm\n", + "import dataclasses\n", + "import IPython\n", + "import IPython.display\n", + "import matplotlib as mpl\n", + "import matplotlib.pyplot as plt\n", + "mpl.rcParams['figure.figsize'] = (8, 6)\n", + "mpl.rcParams['axes.grid'] = False" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Loading TimesFM pretrained checkpoint" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "tfm = timesfm.TimesFm(\n", + " context_len=512,\n", + " horizon_len=128,\n", + " input_patch_len=32,\n", + " output_patch_len=128,\n", + " num_layers=20,\n", + " model_dims=1280,\n", + " backend=\"gpu\",\n", + ")\n", + "tfm.load_from_checkpoint(repo_id=\"google/timesfm-1.0-200m\")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Evaluating pretrained checkpoint on ETT datasets" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "DATA_DICT = {\n", + " \"ettm2\": {\n", + " \"boundaries\": [34560, 46080, 57600],\n", + " \"data_path\": \"../datasets/ETT-small/ETTm2.csv\",\n", + " \"freq\": \"15min\",\n", + " },\n", + " \"ettm1\": {\n", + " \"boundaries\": [34560, 46080, 57600],\n", + " \"data_path\": \"../datasets/ETT-small/ETTm1.csv\",\n", + " \"freq\": \"15min\",\n", + " },\n", + " \"etth2\": {\n", + " \"boundaries\": [8640, 11520, 14400],\n", + " \"data_path\": \"../datasets/ETT-small/ETTh2.csv\",\n", + " \"freq\": \"H\",\n", + " },\n", + " \"etth1\": {\n", + " \"boundaries\": [8640, 11520, 14400],\n", + " \"data_path\": \"../datasets/ETT-small/ETTh1.csv\",\n", + " \"freq\": \"H\",\n", + " },\n", + " \"elec\": {\n", + " \"boundaries\": [18413, 21044, 26304],\n", + " \"data_path\": \"../datasets/electricity/electricity.csv\",\n", + " \"freq\": \"H\",\n", + " },\n", + " \"traffic\": {\n", + " \"boundaries\": [12280, 14036, 17544],\n", + " \"data_path\": \"../datasets/traffic/traffic.csv\",\n", + " \"freq\": \"H\",\n", + " },\n", + " \"weather\": {\n", + " \"boundaries\": [36887, 42157, 52696],\n", + " \"data_path\": \"../datasets/weather/weather.csv\",\n", + " \"freq\": \"10min\",\n", + " },\n", + "}" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "dataset = \"ettm1\"\n", + "data_path = DATA_DICT[dataset][\"data_path\"]\n", + "freq = DATA_DICT[dataset][\"freq\"]\n", + "int_freq = timesfm.freq_map(freq)\n", + "boundaries = DATA_DICT[dataset][\"boundaries\"]\n", + "\n", + "data_df = pd.read_csv(open(data_path, \"r\"))\n", + "\n", + "\n", + "ts_cols = [col for col in data_df.columns if col != \"date\"]\n", + "num_cov_cols = None\n", + "cat_cov_cols = None\n", + "\n", + "context_len = 512\n", + "pred_len = 96\n", + "\n", + "num_ts = len(ts_cols)\n", + "batch_size = 16\n", + "\n", + "dtl = data_loader.TimeSeriesdata(\n", + " data_path=data_path,\n", + " datetime_col=\"date\",\n", + " num_cov_cols=num_cov_cols,\n", + " cat_cov_cols=cat_cov_cols,\n", + " ts_cols=np.array(ts_cols),\n", + " train_range=[0, boundaries[0]],\n", + " val_range=[boundaries[0], boundaries[1]],\n", + " test_range=[boundaries[1], boundaries[2]],\n", + " hist_len=context_len,\n", + " pred_len=pred_len,\n", + " batch_size=num_ts,\n", + " freq=freq,\n", + " normalize=True,\n", + " epoch_len=None,\n", + " holiday=False,\n", + " permute=True,\n", + " )" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "train_batches = dtl.tf_dataset(mode=\"train\", shift=1).batch(batch_size)\n", + "val_batches = dtl.tf_dataset(mode=\"val\", shift=pred_len)\n", + "test_batches = dtl.tf_dataset(mode=\"test\", shift=pred_len)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "for tbatch in tqdm(train_batches.as_numpy_iterator()):\n", + " pass\n", + "print(tbatch[0].shape)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### MAE on the test split for the pretrained TimesFM model" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "mae_losses = []\n", + "for batch in tqdm(test_batches.as_numpy_iterator()):\n", + " past = batch[0]\n", + " actuals = batch[3]\n", + " _, forecasts = tfm.forecast(list(past), [0] * past.shape[0])\n", + " forecasts = forecasts[:, 0 : actuals.shape[1], 5]\n", + " mae_losses.append(np.abs(forecasts - actuals).mean())\n", + "\n", + "print(f\"MAE: {np.mean(mae_losses)}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Finetuning the model on the ETT dataset" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import jax\n", + "from jax import numpy as jnp\n", + "from praxis import pax_fiddle\n", + "from praxis import py_utils\n", + "from praxis import pytypes\n", + "from praxis import base_model\n", + "from praxis import optimizers\n", + "from praxis import schedules\n", + "from praxis import base_hyperparams\n", + "from praxis import base_layer\n", + "from paxml import tasks_lib\n", + "from paxml import trainer_lib\n", + "from paxml import checkpoints\n", + "from paxml import learners\n", + "from paxml import partitioning\n", + "from paxml import checkpoint_types" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "# PAX shortcuts\n", + "NestedMap = py_utils.NestedMap\n", + "WeightInit = base_layer.WeightInit\n", + "WeightHParams = base_layer.WeightHParams\n", + "InstantiableParams = py_utils.InstantiableParams\n", + "JTensor = pytypes.JTensor\n", + "NpTensor = pytypes.NpTensor\n", + "WeightedScalars = pytypes.WeightedScalars\n", + "instantiate = base_hyperparams.instantiate\n", + "LayerTpl = pax_fiddle.Config[base_layer.BaseLayer]\n", + "AuxLossStruct = base_layer.AuxLossStruct\n", + "\n", + "AUX_LOSS = base_layer.AUX_LOSS\n", + "template_field = base_layer.template_field\n", + "\n", + "# Standard prng key names\n", + "PARAMS = base_layer.PARAMS\n", + "RANDOM = base_layer.RANDOM\n", + "\n", + "key = jax.random.PRNGKey(seed=1234)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "model = pax_fiddle.Config(\n", + " patched_decoder.PatchedDecoderFinetuneModel,\n", + " name='patched_decoder_finetune',\n", + " core_layer_tpl=tfm.model_p,\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### We will hold the transformer layers fixed while finetuning, while training all other components." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "@pax_fiddle.auto_config\n", + "def build_learner() -> learners.Learner:\n", + " return pax_fiddle.Config(\n", + " learners.Learner,\n", + " name='learner',\n", + " loss_name='avg_qloss',\n", + " optimizer=optimizers.Adam(\n", + " epsilon=1e-7,\n", + " clip_threshold=1e2,\n", + " learning_rate=1e-2,\n", + " lr_schedule=pax_fiddle.Config(\n", + " schedules.Cosine,\n", + " initial_value=1e-3,\n", + " final_value=1e-4,\n", + " total_steps=40000,\n", + " ),\n", + " ema_decay=0.9999,\n", + " ),\n", + " # Linear probing i.e we hold the transformer layers fixed.\n", + " bprop_variable_exclusion=['.*/stacked_transformer_layer/.*'],\n", + " )" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "task_p = tasks_lib.SingleTask(\n", + " name='ts-learn',\n", + " model=model,\n", + " train=tasks_lib.SingleTask.Train(\n", + " learner=build_learner(),\n", + " ),\n", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "task_p.model.ici_mesh_shape = [1, 1, 1]\n", + "task_p.model.mesh_axis_names = ['replica', 'data', 'mdl']\n", + "\n", + "DEVICES = np.array(jax.devices()).reshape([1, 1, 1])\n", + "MESH = jax.sharding.Mesh(DEVICES, ['replica', 'data', 'mdl'])\n", + "\n", + "num_devices = jax.local_device_count()\n", + "print(f'num_devices: {num_devices}')\n", + "print(f'device kind: {jax.local_devices()[0].device_kind}')" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "jax_task = task_p\n", + "key, init_key = jax.random.split(key)\n", + "\n", + "# To correctly prepare a batch of data for model initialization (now that shape\n", + "# inference is merged), we take one devices*batch_size tensor tuple of data,\n", + "# slice out just one batch, then run the prepare_input_batch function over it.\n", + "\n", + "\n", + "def process_train_batch(batch):\n", + " past_ts = batch[0].reshape(batch_size * num_ts, -1)\n", + " actual_ts = batch[3].reshape(batch_size * num_ts, -1)\n", + " return NestedMap(input_ts=past_ts, actual_ts=actual_ts)\n", + "\n", + "\n", + "def process_eval_batch(batch):\n", + " past_ts = batch[0]\n", + " actual_ts = batch[3]\n", + " return NestedMap(input_ts=past_ts, actual_ts=actual_ts)\n", + "\n", + "\n", + "jax_model_states, _ = trainer_lib.initialize_model_state(\n", + " jax_task,\n", + " init_key,\n", + " process_train_batch(tbatch),\n", + " checkpoint_type=checkpoint_types.CheckpointType.GDA,\n", + ")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Setting the initial model weights to the pretrained TimesFM parameters." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "jax_model_states.mdl_vars['params']['core_layer'] = tfm._train_state.mdl_vars['params']\n", + "jax_vars = jax_model_states.mdl_vars\n", + "gc.collect()" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Training loop" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "jax_task = task_p\n", + "\n", + "\n", + "def train_step(states, prng_key, inputs):\n", + " return trainer_lib.train_step_single_learner(\n", + " jax_task, states, prng_key, inputs\n", + " )\n", + "\n", + "\n", + "def eval_step(states, prng_key, inputs):\n", + " states = states.to_eval_state()\n", + " return trainer_lib.eval_step_single_learner(\n", + " jax_task, states, prng_key, inputs\n", + " )\n", + "\n", + "key, train_key, eval_key = jax.random.split(key, 3)\n", + "train_prng_seed = jax.random.split(train_key, num=jax.local_device_count())\n", + "eval_prng_seed = jax.random.split(eval_key, num=jax.local_device_count())\n", + "\n", + "p_train_step = jax.pmap(train_step, axis_name='batch')\n", + "p_eval_step = jax.pmap(eval_step, axis_name='batch')" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "replicated_jax_states = trainer_lib.replicate_model_state(jax_model_states)\n", + "replicated_jax_vars = replicated_jax_states.mdl_vars" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "best_eval_loss = 1e7\n", + "step_count = 0\n", + "patience = 0\n", + "NUM_EPOCHS = 100\n", + "PATIENCE = 5\n", + "TRAIN_STEPS_PER_EVAL = 1000\n", + "CHECKPOINT_DIR='/home/senrajat_google_com/ettm1_finetune'" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "def reshape_batch_for_pmap(batch, num_devices):\n", + " def _reshape(input_tensor):\n", + " bsize = input_tensor.shape[0]\n", + " residual_shape = list(input_tensor.shape[1:])\n", + " nbsize = bsize // num_devices\n", + " return jnp.reshape(input_tensor, [num_devices, nbsize] + residual_shape)\n", + "\n", + " return jax.tree.map(_reshape, batch)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "for epoch in range(NUM_EPOCHS):\n", + " print(f\"__________________Epoch: {epoch}__________________\", flush=True)\n", + " train_its = train_batches.as_numpy_iterator()\n", + " if patience >= PATIENCE:\n", + " print(\"Early stopping.\", flush=True)\n", + " break\n", + " for batch in tqdm(train_its):\n", + " train_losses = []\n", + " if patience >= PATIENCE:\n", + " print(\"Early stopping.\", flush=True)\n", + " break\n", + " tbatch = process_train_batch(batch)\n", + " tbatch = reshape_batch_for_pmap(tbatch, num_devices)\n", + " replicated_jax_states, step_fun_out = p_train_step(\n", + " replicated_jax_states, train_prng_seed, tbatch\n", + " )\n", + " train_losses.append(step_fun_out.loss[0])\n", + " if step_count % TRAIN_STEPS_PER_EVAL == 0:\n", + " print(\n", + " f\"Train loss at step {step_count}: {np.mean(train_losses)}\",\n", + " flush=True,\n", + " )\n", + " train_losses = []\n", + " print(\"Starting eval.\", flush=True)\n", + " val_its = val_batches.as_numpy_iterator()\n", + " eval_losses = []\n", + " for ev_batch in tqdm(val_its):\n", + " ebatch = process_eval_batch(ev_batch)\n", + " ebatch = reshape_batch_for_pmap(ebatch, num_devices)\n", + " _, step_fun_out = p_eval_step(\n", + " replicated_jax_states, eval_prng_seed, ebatch\n", + " )\n", + " eval_losses.append(step_fun_out.loss[0])\n", + " mean_loss = np.mean(eval_losses)\n", + " print(f\"Eval loss at step {step_count}: {mean_loss}\", flush=True)\n", + " if mean_loss < best_eval_loss or np.isnan(mean_loss):\n", + " best_eval_loss = mean_loss\n", + " print(\"Saving checkpoint.\")\n", + " jax_state_for_saving = py_utils.maybe_unreplicate_for_fully_replicated(\n", + " replicated_jax_states\n", + " )\n", + " checkpoints.save_checkpoint(\n", + " jax_state_for_saving, CHECKPOINT_DIR, overwrite=True\n", + " )\n", + " patience = 0\n", + " del jax_state_for_saving\n", + " gc.collect()\n", + " else:\n", + " patience += 1\n", + " print(f\"patience: {patience}\")\n", + " step_count += 1" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Loading and evaluating the best (according to validation loss) finetuned checkpoint" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "train_state = checkpoints.restore_checkpoint(jax_model_states, CHECKPOINT_DIR)\n", + "print(train_state.step)\n", + "tfm._train_state.mdl_vars['params'] = train_state.mdl_vars['params']['core_layer']\n", + "tfm.jit_decode()\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "mae_losses = []\n", + "for batch in tqdm(test_batches.as_numpy_iterator()):\n", + " past = batch[0]\n", + " actuals = batch[3]\n", + " _, forecasts = tfm.forecast(list(past), [0] * past.shape[0])\n", + " forecasts = forecasts[:, 0 : actuals.shape[1], 5]\n", + " mae_losses.append(np.abs(forecasts - actuals).mean())\n", + "\n", + "print(f\"MAE: {np.mean(mae_losses)}\")" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## There is around a __9%__ reduction in MAE from finetuning." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "tfm_env_v3", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.10.14" + } + }, + "nbformat": 4, + "nbformat_minor": 2 +} diff --git a/src/timesfm/__init__.py b/src/timesfm/__init__.py new file mode 100644 index 0000000..866e848 --- /dev/null +++ b/src/timesfm/__init__.py @@ -0,0 +1,17 @@ +# Copyright 2024 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""TimesFM init file.""" + +from .timesfm import TimesFm, freq_map diff --git a/src/timesfm/data_loader.py b/src/timesfm/data_loader.py new file mode 100644 index 0000000..eeace7b --- /dev/null +++ b/src/timesfm/data_loader.py @@ -0,0 +1,261 @@ +# Copyright 2024 The Google Research Authors. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""TF dataloaders for general timeseries datasets. + +The expected input format is csv file with a datetime index. +""" + + +from absl import logging +import numpy as np +import pandas as pd +from sklearn.preprocessing import StandardScaler +import tensorflow as tf +from . import time_features + + +class TimeSeriesdata(object): + """Data loader class.""" + + def __init__( + self, + data_path, + datetime_col, + num_cov_cols, + cat_cov_cols, + ts_cols, + train_range, + val_range, + test_range, + hist_len, + pred_len, + batch_size, + freq='H', + normalize=True, + epoch_len=None, + holiday=False, + permute=True, + ): + """Initialize objects. + + Args: + data_path: path to csv file + datetime_col: column name for datetime col + num_cov_cols: list of numerical global covariates + cat_cov_cols: list of categorical global covariates + ts_cols: columns corresponding to ts + train_range: tuple of train ranges + val_range: tuple of validation ranges + test_range: tuple of test ranges + hist_len: historical context + pred_len: prediction length + batch_size: batch size (number of ts in a batch) + freq: freq of original data + normalize: std. normalize data or not + epoch_len: num iters in an epoch + holiday: use holiday features or not + permute: permute ts in train batches or not + + Returns: + None + """ + self.data_df = pd.read_csv(open(data_path, 'r')) + if not num_cov_cols: + self.data_df['ncol'] = np.zeros(self.data_df.shape[0]) + num_cov_cols = ['ncol'] + if not cat_cov_cols: + self.data_df['ccol'] = np.zeros(self.data_df.shape[0]) + cat_cov_cols = ['ccol'] + self.data_df.fillna(0, inplace=True) + self.data_df.set_index( + pd.DatetimeIndex(self.data_df[datetime_col]), inplace=True + ) + self.num_cov_cols = num_cov_cols + self.cat_cov_cols = cat_cov_cols + self.ts_cols = ts_cols + self.train_range = train_range + self.val_range = val_range + self.test_range = test_range + data_df_idx = self.data_df.index + date_index = data_df_idx.union( + pd.date_range( + data_df_idx[-1] + pd.Timedelta(1, freq=freq), + periods=pred_len + 1, + freq=freq, + ) + ) + self.time_df = time_features.TimeCovariates( + date_index, holiday=holiday + ).get_covariates() + self.hist_len = hist_len + self.pred_len = pred_len + self.batch_size = batch_size + self.freq = freq + self.normalize = normalize + self.data_mat = self.data_df[self.ts_cols].to_numpy().transpose() + self.data_mat = self.data_mat[:, 0 : self.test_range[1]] + self.time_mat = self.time_df.to_numpy().transpose() + self.num_feat_mat = self.data_df[num_cov_cols].to_numpy().transpose() + self.cat_feat_mat, self.cat_sizes = self._get_cat_cols(cat_cov_cols) + self.normalize = normalize + if normalize: + self._normalize_data() + logging.info( + 'Data Shapes: %s, %s, %s, %s', + self.data_mat.shape, + self.time_mat.shape, + self.num_feat_mat.shape, + self.cat_feat_mat.shape, + ) + self.epoch_len = epoch_len + self.permute = permute + + def _get_cat_cols(self, cat_cov_cols): + """Get categorical columns.""" + cat_vars = [] + cat_sizes = [] + for col in cat_cov_cols: + dct = {x: i for i, x in enumerate(self.data_df[col].unique())} + cat_sizes.append(len(dct)) + mapped = self.data_df[col].map(lambda x: dct[x]).to_numpy().transpose() # pylint: disable=cell-var-from-loop + cat_vars.append(mapped) + return np.vstack(cat_vars), cat_sizes + + def _normalize_data(self): + self.scaler = StandardScaler() + train_mat = self.data_mat[:, self.train_range[0] : self.train_range[1]] + self.scaler = self.scaler.fit(train_mat.transpose()) + self.data_mat = self.scaler.transform(self.data_mat.transpose()).transpose() + + def train_gen(self): + """Generator for training data.""" + num_ts = len(self.ts_cols) + perm = np.arange( + self.train_range[0] + self.hist_len, + self.train_range[1] - self.pred_len, + ) + perm = np.random.permutation(perm) + hist_len = self.hist_len + logging.info('Hist len: %s', hist_len) + if not self.epoch_len: + epoch_len = len(perm) + else: + epoch_len = self.epoch_len + for idx in perm[0:epoch_len]: + for _ in range(num_ts // self.batch_size + 1): + if self.permute: + tsidx = np.random.choice(num_ts, size=self.batch_size, replace=False) + else: + tsidx = np.arange(num_ts) + dtimes = np.arange(idx - hist_len, idx + self.pred_len) + ( + bts_train, + bts_pred, + bfeats_train, + bfeats_pred, + bcf_train, + bcf_pred, + ) = self._get_features_and_ts(dtimes, tsidx, hist_len) + + all_data = [ + bts_train, + bfeats_train, + bcf_train, + bts_pred, + bfeats_pred, + bcf_pred, + tsidx, + ] + yield tuple(all_data) + + def test_val_gen(self, mode='val', shift=1): + """Generator for validation/test data.""" + if mode == 'val': + start = self.val_range[0] + end = self.val_range[1] - self.pred_len + 1 + elif mode == 'test': + start = self.test_range[0] + end = self.test_range[1] - self.pred_len + 1 + else: + raise NotImplementedError('Eval mode not implemented') + num_ts = len(self.ts_cols) + hist_len = self.hist_len + logging.info('Hist len: %s', hist_len) + perm = np.arange(start, end) + if self.epoch_len: + epoch_len = self.epoch_len + else: + epoch_len = len(perm) + for i in range(0, epoch_len, shift): + idx = perm[i] + for batch_idx in range(0, num_ts, self.batch_size): + tsidx = np.arange(batch_idx, min(batch_idx + self.batch_size, num_ts)) + dtimes = np.arange(idx - hist_len, idx + self.pred_len) + ( + bts_train, + bts_pred, + bfeats_train, + bfeats_pred, + bcf_train, + bcf_pred, + ) = self._get_features_and_ts(dtimes, tsidx, hist_len) + all_data = [ + bts_train, + bfeats_train, + bcf_train, + bts_pred, + bfeats_pred, + bcf_pred, + tsidx, + ] + yield tuple(all_data) + + def _get_features_and_ts(self, dtimes, tsidx, hist_len=None): + """Get features and ts in specified windows.""" + if hist_len is None: + hist_len = self.hist_len + data_times = dtimes[dtimes < self.data_mat.shape[1]] + bdata = self.data_mat[:, data_times] + bts = bdata[tsidx, :] + bnf = self.num_feat_mat[:, data_times] + bcf = self.cat_feat_mat[:, data_times] + btf = self.time_mat[:, dtimes] + if bnf.shape[1] < btf.shape[1]: + rem_len = btf.shape[1] - bnf.shape[1] + rem_rep = np.repeat(bnf[:, [-1]], repeats=rem_len) + rem_rep_cat = np.repeat(bcf[:, [-1]], repeats=rem_len) + bnf = np.hstack([bnf, rem_rep.reshape(bnf.shape[0], -1)]) + bcf = np.hstack([bcf, rem_rep_cat.reshape(bcf.shape[0], -1)]) + bfeats = np.vstack([btf, bnf]) + bts_train = bts[:, 0:hist_len] + bts_pred = bts[:, hist_len:] + bfeats_train = bfeats[:, 0:hist_len] + bfeats_pred = bfeats[:, hist_len:] + bcf_train = bcf[:, 0:hist_len] + bcf_pred = bcf[:, hist_len:] + return bts_train, bts_pred, bfeats_train, bfeats_pred, bcf_train, bcf_pred + + def tf_dataset(self, mode='train', shift=1): + """Tensorflow Dataset.""" + if mode == 'train': + gen_fn = self.train_gen + else: + gen_fn = lambda: self.test_val_gen(mode, shift) + output_types = tuple( + [tf.float32] * 2 + [tf.int32] + [tf.float32] * 2 + [tf.int32] * 2 + ) + dataset = tf.data.Dataset.from_generator(gen_fn, output_types) + dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE) + return dataset diff --git a/src/timesfm/patched_decoder.py b/src/timesfm/patched_decoder.py new file mode 100644 index 0000000..5dbdcbc --- /dev/null +++ b/src/timesfm/patched_decoder.py @@ -0,0 +1,521 @@ +# Copyright 2024 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Pax ML model for patched time-series decoder. + +The file implements Residual MLPs, Patched Decoder layers and PAX ML models. +""" + +import dataclasses +from typing import Optional, Tuple + +import einshape as es +from jax import lax +import jax.numpy as jnp +from praxis import base_layer +from praxis import base_model +from praxis import layers +from praxis import pax_fiddle +from praxis import py_utils +from praxis import pytypes +from praxis.layers import activations +from praxis.layers import embedding_softmax +from praxis.layers import linears +from praxis.layers import normalizations +from praxis.layers import stochastics +from praxis.layers import transformers + + +# PAX shortcuts +NestedMap = py_utils.NestedMap +JTensor = pytypes.JTensor + +LayerTpl = pax_fiddle.Config[base_layer.BaseLayer] +template_field = base_layer.template_field + + +PAD_VAL = 1123581321.0 +DEFAULT_QUANTILES = [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9] + +# NestedMap keys +_INPUT_TS = "input_ts" +_TARGET_FUTURE = "actual_ts" +_INPUT_PADDING = "input_padding" +_OUTPUT_TS = "output_ts" +_FREQ = "freq" +_OUTPUT_TOKENS = "output_tokens" +_STATS = "stats" + + +# Small numerical value. +_TOLERANCE = 1e-7 + + +def _shift_padded_seq(mask: JTensor, seq: JTensor) -> JTensor: + """Shifts rows of seq based on the first 0 in each row of the mask.""" + num = seq.shape[1] + + # Find the index of the first 0 in each row of the mask + first_zero_idx = jnp.argmin(mask, axis=1) + + # Create a range array for indexing + idx_range = jnp.arange(num) + + def shift_row(carry, x): + seq_row, shift = x + shifted_idx = (idx_range - shift) % num + shifted_row = seq_row[shifted_idx] + return carry, shifted_row + + # Use lax.scan to shift each row of seq based on the corresponding + # first_zero_idx. + _, shifted_seq = lax.scan(shift_row, None, (seq, first_zero_idx)) + + return shifted_seq + + +class ResidualBlock(base_layer.BaseLayer): + """Simple feedforward block with residual connection. + + Attributes: + input_dims: input dimension. + hidden_dims: hidden dimension. + output_dims: output dimension. + dropout_prob: dropout probability. + layer_norm: whether to use layer norm or not. + dropout_tpl: config for dropout. + ln_tpl: config for layer norm. + act_tpl: config for activation in hidden layer. + """ + + input_dims: int = 0 + hidden_dims: int = 0 + output_dims: int = 0 + dropout_prob: float = 0.0 + layer_norm: bool = False + dropout_tpl: LayerTpl = template_field(stochastics.Dropout) + ln_tpl: LayerTpl = template_field(normalizations.LayerNorm) + act_tpl: LayerTpl = template_field(activations.Swish) + + def setup(self): + lnorm_tpl = self.ln_tpl.clone() + lnorm_tpl.dim = self.output_dims + self.create_child("ln_layer", lnorm_tpl) + + dropout_tpl = self.dropout_tpl.clone() + dropout_tpl.keep_prob = 1.0 - self.dropout_prob + self.create_child("dropout", dropout_tpl) + + self.create_child( + "hidden_layer", + pax_fiddle.Config( + linears.FeedForward, + input_dims=self.input_dims, + output_dims=self.hidden_dims, + activation_tpl=self.act_tpl.clone(), + ), + ) + + self.create_child( + "output_layer", + pax_fiddle.Config( + linears.FeedForward, + input_dims=self.hidden_dims, + output_dims=self.output_dims, + activation_tpl=pax_fiddle.Config(activations.Identity), + ), + ) + + self.create_child( + "residual_layer", + pax_fiddle.Config( + linears.FeedForward, + input_dims=self.input_dims, + output_dims=self.output_dims, + activation_tpl=pax_fiddle.Config(activations.Identity), + ), + ) + + def __call__(self, inputs: JTensor) -> JTensor: + hidden = self.hidden_layer(inputs) + output = self.output_layer(hidden) + output = self.dropout(output) + residual = self.residual_layer(inputs) + if self.layer_norm: + return self.ln_layer(output + residual) + else: + return output + residual + + +def _masked_mean_std(inputs: JTensor, padding: JTensor) -> Tuple[JTensor, JTensor]: + """Calculates mean and standard deviation of arr across axis 1. + + It should exclude values where pad is 1. + + Args: + inputs: A JAX array of shape [b, n, p]. + padding: A JAX array of shape [b, n, p] with values 0 or 1. + + Returns: + A tuple containing the mean and standard deviation of arr. We return the + statistics of the first patch with more than three non-padded values. + """ + # Selecting the first pad with more than 3 unpadded values. + pad_sum = jnp.sum(1 - padding, axis=2) + + def _get_patch_index(arr: JTensor): + indices = jnp.argmax(arr >= 3, axis=1) + row_sum = (arr >= 3).sum(axis=1) + return jnp.where(row_sum == 0, arr.shape[1] - 1, indices) + + patch_indices = _get_patch_index(pad_sum) + bidxs = jnp.arange(inputs.shape[0]) + + arr = inputs[bidxs, patch_indices, :] + pad = padding[bidxs, patch_indices, :] + + # Create a mask where P is 0 + mask = 1 - pad + + # Calculate the number of valid elements + num_valid_elements = jnp.sum(mask, axis=1) + + num_valid_elements = jnp.where(num_valid_elements == 0, 1, num_valid_elements) + + # Calculate the masked sum and squared sum of M + masked_sum = jnp.sum(arr * mask, axis=1) + masked_squared_sum = jnp.sum((arr * mask) ** 2, axis=1) + + # Calculate the masked mean and standard deviation + masked_mean = masked_sum / num_valid_elements + masked_var = masked_squared_sum / num_valid_elements - masked_mean**2 + masked_var = jnp.where(masked_var < 0.0, 0.0, masked_var) + masked_std = jnp.sqrt(masked_var) + + return masked_mean, masked_std + + +def _create_quantiles() -> list[float]: + """Returns the quantiles for forecasting.""" + return DEFAULT_QUANTILES + + +class PatchedTimeSeriesDecoder(base_layer.BaseLayer): + """Patch decoder layer for time-series foundation model. + + Attributes: + patch_len: length of input patches. + horizon_len: length of output patches. Referred to as `output_patch_len` + during inference. + model_dims: model dimension of stacked transformer layer. + hidden_dims: hidden dimensions in fully connected layers. + quantiles: list of quantiles for non prob model. + residual_block_tpl: config for residual block. + stacked_transformer_params_tpl: config for stacked transformer. + use_freq: whether to use frequency encoding. + + In all of what followed, except specified otherwise, B is batch size, T is + sequence length of time-series. N is the number of input patches that can be + obtained from T. P is the input patch length and H is the horizon length. Q is + number of output logits. D is model dimension. + """ + + patch_len: int = 0 + horizon_len: int = 0 + model_dims: int = 0 + hidden_dims: int = 0 + quantiles: list[float] = dataclasses.field(default_factory=_create_quantiles) + residual_block_tpl: LayerTpl = template_field(ResidualBlock) + stacked_transformer_params_tpl: LayerTpl = template_field( + transformers.StackedTransformer + ) + use_freq: bool = True + + def setup(self) -> None: + """Construct the model.""" + num_outputs = len(self.quantiles) + 1 + + stl = self.stacked_transformer_params_tpl.clone() + stl.model_dims = self.model_dims + stl.hidden_dims = self.hidden_dims + stl.mask_self_attention = True + + self.create_child("stacked_transformer_layer", stl) + + input_resl = self.residual_block_tpl.clone() + ff_in_dims = 2 * self.patch_len + input_resl.input_dims = ff_in_dims + input_resl.hidden_dims = self.hidden_dims + input_resl.output_dims = self.model_dims + self.create_child( + "input_ff_layer", + input_resl, + ) + + horizon_resl = self.residual_block_tpl.clone() + horizon_resl.input_dims = self.model_dims + horizon_resl.hidden_dims = self.hidden_dims + horizon_resl.output_dims = self.horizon_len * num_outputs + self.create_child( + "horizon_ff_layer", + horizon_resl, + ) + + self.create_child( + "position_emb", + pax_fiddle.Config( + layers.PositionalEmbedding, embedding_dims=self.model_dims + ), + ) + + if self.use_freq: + self.create_child( + "freq_emb", + pax_fiddle.Config( + embedding_softmax.Embedding, + num_classes=3, + input_dims=self.model_dims, + ), + ) + + def transform_decode_state( + self, transform_fn: base_layer.DecodeStateTransformFn + ) -> None: + """Transforms all decode state variables based on transform_fn.""" + self.stacked_transformer_layer.transform_decode_state(transform_fn) + + def _forward_transform( + self, inputs: JTensor, patched_pads: JTensor + ) -> Tuple[JTensor, Tuple[JTensor, JTensor]]: + """Input is of shape [B, N, P].""" + mu, sigma = _masked_mean_std(inputs, patched_pads) + sigma = jnp.where(sigma < _TOLERANCE, 1.0, sigma) + # Normalize each patch. + outputs = (inputs - mu[:, None, None]) / sigma[:, None, None] + outputs = jnp.where(jnp.abs(inputs - PAD_VAL) < _TOLERANCE, PAD_VAL, outputs) + return outputs, (mu, sigma) + + def _reverse_transform( + self, outputs: JTensor, stats: Tuple[JTensor, JTensor] + ) -> JTensor: + """Output is of shape [B, N, P, Q].""" + mu, sigma = stats + return outputs * sigma[:, None, None, None] + mu[:, None, None, None] + + def _preprocess_input( + self, + input_ts: JTensor, + input_padding: JTensor, + pos_emb: Optional[JTensor] = None, + ) -> Tuple[JTensor, JTensor, Optional[Tuple[JTensor, JTensor]], JTensor]: + """Preprocess input for stacked transformer.""" + # Reshape into patches. + patched_inputs = es.jax_einshape("b(np)->bnp", input_ts, p=self.patch_len) + input_padding = jnp.where( + jnp.abs(input_ts - PAD_VAL) < _TOLERANCE, 1, input_padding + ) + patched_pads = es.jax_einshape("b(np)->bnp", input_padding, p=self.patch_len) + patched_inputs, stats = self._forward_transform(patched_inputs, patched_pads) + # B x N x D + patched_inputs = patched_inputs * (1.0 - patched_pads) + concat_inputs = jnp.concatenate([patched_inputs, patched_pads], axis=-1) + model_input = self.input_ff_layer(concat_inputs) + # A patch should not be padded even if there is at least one zero. + patched_padding = jnp.min(patched_pads, axis=-1) + + if pos_emb is None: + position_emb = self.position_emb(seq_length=model_input.shape[1]) + else: + position_emb = pos_emb + if self.do_eval: + if position_emb.shape[0] != model_input.shape[0]: + position_emb = jnp.repeat(position_emb, model_input.shape[0], axis=0) + position_emb = _shift_padded_seq(patched_padding, position_emb) + model_input += position_emb + + return model_input, patched_padding, stats, patched_inputs + + def _postprocess_output( + self, + model_output: JTensor, + num_outputs: int, + stats: Tuple[JTensor, JTensor], + ) -> JTensor: + """Postprocess output of stacked transformer.""" + # B x N x (H.Q) + output_ts = self.horizon_ff_layer(model_output) + output_ts = es.jax_einshape( + "bn(hq)->bnhq", output_ts, q=num_outputs, h=self.horizon_len + ) + return self._reverse_transform(output_ts, stats) + + def __call__(self, inputs: NestedMap) -> NestedMap: + """PatchTST call. + + Args: + inputs: A NestedMap containing (1) input_ts: input sequence of shape [B, + T] where T must be multiple of patch_length; (2) input_padding: that + contains padding map. + + Returns: + A nested map with two keys: + (1) 'output_tokens' of shape [B, N, D]. + (2) 'output_ts' of shape [B, N, H, Q] + (3) 'stats' a Tuple of statistics for renormalization. + """ + input_ts, input_padding = inputs[_INPUT_TS], inputs[_INPUT_PADDING] + num_outputs = len(self.quantiles) + 1 + model_input, patched_padding, stats, _ = self._preprocess_input( + input_ts=input_ts, + input_padding=input_padding, + ) + if self.use_freq: + freq = inputs[_FREQ].astype(jnp.int32) + f_emb = self.freq_emb(freq) # B x 1 x D + f_emb = jnp.repeat(f_emb, model_input.shape[1], axis=1) + model_input += f_emb + model_output = self.stacked_transformer_layer(model_input, patched_padding) + + output_ts = self._postprocess_output(model_output, num_outputs, stats) + return NestedMap( + {_OUTPUT_TOKENS: model_output, _OUTPUT_TS: output_ts, _STATS: stats} + ) + + def decode( + self, + inputs: NestedMap, + horizon_len: int, + output_patch_len: Optional[int] = None, + max_len: int = 512, + ) -> tuple[JTensor, JTensor]: + """Auto-regressive decoding without caching. + + Args: + inputs: input time-series and paddings. Time-series shape B x C, padding + shape shape B x (C + H) where H is the prediction length. + horizon_len: prediction length. + output_patch_len: output length to be fetched from one step of + auto-regressive decoding. + max_len: maximum training context length. + + Returns: + Tuple of two forecasting results: + - Point (mean) output predictions as a tensor with shape B x H. + - Full predictions (mean and quantiles) as a tensor with shape + B x H x (1 + # quantiles). + """ + final_out = inputs[_INPUT_TS] + inp_time_len = final_out.shape[1] + paddings = inputs[_INPUT_PADDING] + if self.use_freq: + freq = inputs[_FREQ].astype(jnp.int32) + else: + freq = jnp.zeros([final_out.shape[0], 1], dtype=jnp.int32) + full_outputs = [] + if paddings.shape[1] != final_out.shape[1] + horizon_len: + raise ValueError( + "Length of paddings must match length of input + horizon_len:" + f" {paddings.shape[1]} != {final_out.shape[1]} + {horizon_len}" + ) + if output_patch_len is None: + output_patch_len = self.horizon_len + num_decode_patches = (horizon_len + output_patch_len - 1) // output_patch_len + for _ in range(num_decode_patches): + current_padding = paddings[:, 0 : final_out.shape[1]] + input_ts = final_out[:, -max_len:] + input_padding = current_padding[:, -max_len:] + model_input = NestedMap( + input_ts=input_ts, + input_padding=input_padding, + freq=freq, + ) + fprop_outputs = self(model_input)[_OUTPUT_TS] + # (full batch, last patch, output_patch_len, index of mean forecast = 0) + new_ts = fprop_outputs[:, -1, :output_patch_len, 0] + # (full batch, last patch, output_patch_len, all output indices) + full_outputs.append(fprop_outputs[:, -1, :output_patch_len, :]) + final_out = jnp.concatenate([final_out, new_ts], axis=-1) + + return ( + final_out[:, inp_time_len : inp_time_len + horizon_len], + jnp.concatenate(full_outputs, axis=1)[:, 0:horizon_len, :], + ) + + +class PatchedDecoderFinetuneModel(base_model.BaseModel): + """Model class for finetuning patched time-series decoder. + + Attributes: + core_layer_tpl: config for core layer. + freq: freq to finetune on. + """ + + core_layer_tpl: LayerTpl = template_field(PatchedTimeSeriesDecoder) + freq: int = 0 + + def setup(self) -> None: + self.create_child("core_layer", self.core_layer_tpl) + + def compute_predictions(self, input_batch: NestedMap) -> NestedMap: + input_ts = input_batch[_INPUT_TS] + input_padding = jnp.zeros_like(input_ts) + context_len = input_ts.shape[1] + input_patch_len = self.core_layer_tpl.patch_len + context_pad = ( + (context_len + input_patch_len - 1) // input_patch_len + ) * input_patch_len - context_len + + input_ts = jnp.pad(input_ts, [(0, 0), (context_pad, 0)]) + input_padding = jnp.pad( + input_padding, [(0, 0), (context_pad, 0)], constant_values=1 + ) + freq = jnp.ones([input_ts.shape[0], 1], dtype=jnp.int32) * self.freq + new_input_batch = NestedMap( + input_ts=input_ts, + input_padding=input_padding, + freq=freq, + ) + return self.core_layer(new_input_batch) + + def _quantile_loss( + self, pred: JTensor, actual: JTensor, quantile: float + ) -> JTensor: + """Calculates quantile loss. + + Args: + pred: B x T + actual: B x T + quantile: quantile at which loss is computed. + + Returns: + per coordinate loss. + """ + dev = actual - pred + loss_first = dev * quantile + loss_second = -dev * (1.0 - quantile) + return 2 * jnp.where(loss_first >= 0, loss_first, loss_second) + + def compute_loss( + self, prediction_output: NestedMap, input_batch: NestedMap + ) -> Tuple[NestedMap, NestedMap]: + output_ts = prediction_output[_OUTPUT_TS] + actual_ts = input_batch[_TARGET_FUTURE] + pred_ts = output_ts[:, -1, 0 : actual_ts.shape[1], :] + loss = jnp.square(pred_ts[:, :, 0] - actual_ts) + for i, quantile in enumerate(self.core_layer.quantiles): + loss += self._quantile_loss(pred_ts[:, :, i + 1], actual_ts, quantile) + loss = loss.mean() + loss_weight = jnp.array(1.0, dtype=jnp.float32) + per_example_out = NestedMap() + return {"avg_qloss": (loss, loss_weight)}, per_example_out diff --git a/src/timesfm/time_features.py b/src/timesfm/time_features.py new file mode 100644 index 0000000..0bd90a9 --- /dev/null +++ b/src/timesfm/time_features.py @@ -0,0 +1,215 @@ +# Copyright 2024 The Google Research Authors. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Directory to extract time covariates. + +Extract time covariates from datetime. +""" + +import numpy as np +import pandas as pd +from pandas.tseries.holiday import EasterMonday +from pandas.tseries.holiday import GoodFriday +from pandas.tseries.holiday import Holiday +from pandas.tseries.holiday import SU +from pandas.tseries.holiday import TH +from pandas.tseries.holiday import USColumbusDay +from pandas.tseries.holiday import USLaborDay +from pandas.tseries.holiday import USMartinLutherKingJr +from pandas.tseries.holiday import USMemorialDay +from pandas.tseries.holiday import USPresidentsDay +from pandas.tseries.holiday import USThanksgivingDay +from pandas.tseries.offsets import DateOffset +from pandas.tseries.offsets import Day +from pandas.tseries.offsets import Easter +from sklearn.preprocessing import StandardScaler +from tqdm import tqdm + + +# This is 183 to cover half a year (in both directions), also for leap years +# + 17 as Eastern can be between March, 22 - April, 25 +MAX_WINDOW = 183 + 17 + + +def _distance_to_holiday(holiday): + """Return distance to given holiday.""" + + def _distance_to_day(index): + holiday_date = holiday.dates( + index - pd.Timedelta(days=MAX_WINDOW), + index + pd.Timedelta(days=MAX_WINDOW), + ) + assert ( + len(holiday_date) != 0 # pylint: disable=g-explicit-length-test + ), f"No closest holiday for the date index {index} found." + # It sometimes returns two dates if it is exactly half a year after the + # holiday. In this case, the smaller distance (182 days) is returned. + return (index - holiday_date[0]).days + + return _distance_to_day + + +EasterSunday = Holiday( + "Easter Sunday", month=1, day=1, offset=[Easter(), Day(0)] +) +NewYearsDay = Holiday("New Years Day", month=1, day=1) +SuperBowl = Holiday( + "Superbowl", month=2, day=1, offset=DateOffset(weekday=SU(1)) +) +MothersDay = Holiday( + "Mothers Day", month=5, day=1, offset=DateOffset(weekday=SU(2)) +) +IndependenceDay = Holiday("Independence Day", month=7, day=4) +ChristmasEve = Holiday("Christmas", month=12, day=24) +ChristmasDay = Holiday("Christmas", month=12, day=25) +NewYearsEve = Holiday("New Years Eve", month=12, day=31) +BlackFriday = Holiday( + "Black Friday", + month=11, + day=1, + offset=[pd.DateOffset(weekday=TH(4)), Day(1)], +) +CyberMonday = Holiday( + "Cyber Monday", + month=11, + day=1, + offset=[pd.DateOffset(weekday=TH(4)), Day(4)], +) + +HOLIDAYS = [ + EasterMonday, + GoodFriday, + USColumbusDay, + USLaborDay, + USMartinLutherKingJr, + USMemorialDay, + USPresidentsDay, + USThanksgivingDay, + EasterSunday, + NewYearsDay, + SuperBowl, + MothersDay, + IndependenceDay, + ChristmasEve, + ChristmasDay, + NewYearsEve, + BlackFriday, + CyberMonday, +] + + +class TimeCovariates(object): + """Extract all time covariates except for holidays.""" + + def __init__( + self, + datetimes, + normalized=True, + holiday=False, + ): + """Init function. + + Args: + datetimes: pandas DatetimeIndex (lowest granularity supported is min) + normalized: whether to normalize features or not + holiday: fetch holiday features or not + + Returns: + None + """ + self.normalized = normalized + self.dti = datetimes + self.holiday = holiday + + def _minute_of_hour(self): + minutes = np.array(self.dti.minute, dtype=np.float32) + if self.normalized: + minutes = minutes / 59.0 - 0.5 + return minutes + + def _hour_of_day(self): + hours = np.array(self.dti.hour, dtype=np.float32) + if self.normalized: + hours = hours / 23.0 - 0.5 + return hours + + def _day_of_week(self): + day_week = np.array(self.dti.dayofweek, dtype=np.float32) + if self.normalized: + day_week = day_week / 6.0 - 0.5 + return day_week + + def _day_of_month(self): + day_month = np.array(self.dti.day, dtype=np.float32) + if self.normalized: + day_month = day_month / 30.0 - 0.5 + return day_month + + def _day_of_year(self): + day_year = np.array(self.dti.dayofyear, dtype=np.float32) + if self.normalized: + day_year = day_year / 364.0 - 0.5 + return day_year + + def _month_of_year(self): + month_year = np.array(self.dti.month, dtype=np.float32) + if self.normalized: + month_year = month_year / 11.0 - 0.5 + return month_year + + def _week_of_year(self): + week_year = np.array(self.dti.strftime("%U").astype(int), dtype=np.float32) + if self.normalized: + week_year = week_year / 51.0 - 0.5 + return week_year + + def _get_holidays(self): + dti_series = self.dti.to_series() + hol_variates = np.vstack([ + dti_series.apply(_distance_to_holiday(h)).values for h in tqdm(HOLIDAYS) + ]) + # hol_variates is (num_holiday, num_time_steps), the normalization should be + # performed in the num_time_steps dimension. + return StandardScaler().fit_transform(hol_variates.T).T + + def get_covariates(self): + """Get all time covariates.""" + moh = self._minute_of_hour().reshape(1, -1) + hod = self._hour_of_day().reshape(1, -1) + dom = self._day_of_month().reshape(1, -1) + dow = self._day_of_week().reshape(1, -1) + doy = self._day_of_year().reshape(1, -1) + moy = self._month_of_year().reshape(1, -1) + woy = self._week_of_year().reshape(1, -1) + + all_covs = [ + moh, + hod, + dom, + dow, + doy, + moy, + woy, + ] + columns = ["moh", "hod", "dom", "dow", "doy", "moy", "woy"] + if self.holiday: + hol_covs = self._get_holidays() + all_covs.append(hol_covs) + columns += [f"hol_{i}" for i in range(len(HOLIDAYS))] + + return pd.DataFrame( + data=np.vstack(all_covs).transpose(), + columns=columns, + index=self.dti, + ) diff --git a/src/timesfm/timesfm.py b/src/timesfm/timesfm.py new file mode 100644 index 0000000..5dc7611 --- /dev/null +++ b/src/timesfm/timesfm.py @@ -0,0 +1,605 @@ +# Copyright 2024 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""TimesFM forecast API for inference.""" + +import logging +import multiprocessing +from os import path +import time +from typing import Any, Literal, Optional, Sequence + +import einshape as es +import jax +import jax.numpy as jnp +import numpy as np +import pandas as pd +from huggingface_hub import snapshot_download +from paxml import checkpoints +from paxml import tasks_lib +from praxis import base_hyperparams +from praxis import base_layer +from praxis import pax_fiddle +from praxis import py_utils +from praxis import pytypes +from praxis.layers import normalizations +from praxis.layers import transformers +from . import patched_decoder +from utilsforecast.processing import make_future_dataframe + +instantiate = base_hyperparams.instantiate +NestedMap = py_utils.NestedMap +JTensor = pytypes.JTensor + + +def process_group(key, group, value_name, forecast_context_len): + group = group.tail(forecast_context_len) + return np.array(group[value_name], dtype=np.float32), key + + +def moving_average(arr, window_size): + """Calculates the moving average using NumPy's convolution function.""" + # Pad with zeros to handle initial window positions + arr_padded = np.pad(arr, (window_size - 1, 0), "constant") + smoothed_arr = ( + np.convolve(arr_padded, np.ones(window_size), "valid") / window_size + ) + return [smoothed_arr, arr - smoothed_arr] + + +def freq_map(freq: str): + """Returns the frequency map for the given frequency string.""" + freq = str.upper(freq) + if ( + freq.endswith("H") + or freq.endswith("T") + or freq.endswith("MIN") + or freq.endswith("D") + or freq.endswith("B") + or freq.endswith("U") + ): + return 0 + elif freq.endswith(("W", "M", "MS")): + return 1 + elif freq.endswith("Y") or freq.endswith("Q"): + return 2 + else: + raise ValueError(f"Invalid frequency: {freq}") + + +class TimesFm: + """TimesFM forecast API for inference. + + This class is the scaffolding for calling TimesFM forecast. To properly use: + 1. Create an instance with the correct hyperparameters of a TimesFM model. + 2. Call `load_from_checkpoint` to load a compatible checkpoint. + 3. Call `forecast` for inference. + + Given the model size, this API does not shard the model weights for SPMD. All + parallelism happens on the data dimension. + + Compilation happens during the first time `forecast` is called and uses the + `per_core_batch_size` to set and freeze the input signature. Subsequent calls + to `forecast` reflect the actual inference latency. + + Attributes: + per_core_batch_size: Batch size on each core for data parallelism. + backend: One of "cpu", "gpu" or "tpu". + num_devices: Number of cores provided the backend. + global_batch_size: per_core_batch_size * num_devices. Each batch of + inference task will be padded with respect to global_batch_size to + minimize latency. + context_len: Largest context length the model allows for each decode call. + This technically can be any large, but practically should set to the + context length the checkpoint was trained with. + horizon_len: Forecast horizon. + input_patch_len: Input patch len. + output_patch_len: Output patch len. How many timepoints is taken from a + single step of autoregressive decoding. Can be set as the training horizon + of the checkpoint. + mesh_shape: Shape of the data parallelism mesh. + mesh_name: Names of the data parallelism mesh. + model_p: Configuration of the TimesFM model deduced from the hparams. + """ + + def _logging(self, s): + if self._verbose: + print(s) + + def __init__( + self, + context_len: int, + horizon_len: int, + input_patch_len: int, + output_patch_len: int, + num_layers: int, + model_dims: int, + per_core_batch_size: int = 32, + backend: Literal["cpu", "gpu", "tpu"] = "cpu", + quantiles: Sequence[float] | None = None, + verbose: bool = True, + ) -> None: + """Initializes the TimesFM forecast API. + + Args: + context_len: Largest context length the model allows for each decode call. + This technically can be any large, but practically should set to the + context length the checkpoint was trained with. + horizon_len: Forecast horizon. + input_patch_len: Input patch len. + output_patch_len: Output patch len. How many timepoints is taken from a + single step of autoregressive decoding. Can be set as the training + horizon of the checkpoint. + num_layers: Number of transformer layers. + model_dims: Model dimension. + per_core_batch_size: Batch size on each core for data parallelism. + backend: One of "cpu", "gpu" or "tpu". + quantiles: list of output quantiles supported by the model. + verbose: Whether to print logging messages. + """ + self.per_core_batch_size = per_core_batch_size + self.backend = backend + self.num_devices = jax.local_device_count(self.backend) + self.global_batch_size = self.per_core_batch_size * self.num_devices + + self.context_len = context_len + self.horizon_len = horizon_len + self.input_patch_len = input_patch_len + self.output_patch_len = output_patch_len + + self.mesh_shape = [1, self.num_devices, 1] + self.mesh_name = ["replica", "data", "mdl"] + if quantiles is None: + quantiles = patched_decoder.DEFAULT_QUANTILES + + self.model_p = pax_fiddle.Config( + patched_decoder.PatchedTimeSeriesDecoder, + name="patched_decoder", + horizon_len=self.output_patch_len, + patch_len=input_patch_len, + model_dims=model_dims, + hidden_dims=model_dims, + residual_block_tpl=pax_fiddle.Config(patched_decoder.ResidualBlock), + quantiles=quantiles, + use_freq=True, + stacked_transformer_params_tpl=pax_fiddle.Config( + transformers.StackedTransformer, + num_heads=16, + num_layers=num_layers, + transformer_layer_params_tpl=pax_fiddle.Config( + transformers.Transformer, + ln_tpl=pax_fiddle.Config( + normalizations.RmsNorm, + ), + ), + ), + ) + + self._key1, self._key2 = jax.random.split(jax.random.PRNGKey(42)) + self._model = None + self._train_state = None + self._pmapped_decode = None + self._verbose = verbose + self._eval_context = base_layer.JaxContext.HParams(do_eval=True) + try: + multiprocessing.set_start_method("spawn") + except RuntimeError: + print("Multiprocessing context has already been set.") + + def _get_sample_inputs(self): + return { + "input_ts": jnp.zeros( + ( + self.per_core_batch_size, + self.context_len + self.output_patch_len, + ), + dtype=jnp.float32, + ), + "input_padding": jnp.zeros( + ( + self.per_core_batch_size, + self.context_len + self.output_patch_len, + ), + dtype=jnp.float32, + ), + "freq": jnp.zeros( + ( + self.per_core_batch_size, + 1, + ), + dtype=jnp.int32, + ), + } + + def load_from_checkpoint( + self, + checkpoint_path: Optional[str] = None, + repo_id: str = "google/timesfm-1.0-200m", + checkpoint_type: checkpoints.CheckpointType = checkpoints.CheckpointType.FLAX, + step: int | None = None, + ) -> None: + """Loads a checkpoint and compiles the decoder. + + Args: + checkpoint_path: Optional path to the checkpoint directory. + repo_id: Hugging Face Hub repo id. + checkpoint_type: type of PAX checkpoint + step: step of the checkpoint to load. If `None`, load latest checkpoint. + """ + # Download the checkpoint from Hugging Face Hub if not given + if checkpoint_path is None: + checkpoint_path = path.join(snapshot_download(repo_id), "checkpoints") + + # Initialize the model weights. + self._logging("Constructing model weights.") + start_time = time.time() + self._model = instantiate(self.model_p) + var_weight_hparams = self._model.abstract_init_with_metadata( + self._get_sample_inputs(), do_eval=True + ) + train_state_partition_specs = tasks_lib.create_state_partition_specs( + var_weight_hparams, + mesh_shape=self.mesh_shape, + mesh_axis_names=self.mesh_name, + discard_opt_states=True, + learners=None, + ) + train_state_local_shapes = tasks_lib.create_state_unpadded_shapes( + var_weight_hparams, + discard_opt_states=True, + learners=None, + ) + self._logging( + f"Constructed model weights in {time.time() - start_time:.2f} seconds." + ) + + # Load the model weights. + self._logging(f"Restoring checkpoint from {checkpoint_path}.") + start_time = time.time() + self._train_state = checkpoints.restore_checkpoint( + train_state_local_shapes, + checkpoint_dir=checkpoint_path, + checkpoint_type=checkpoint_type, + state_specs=train_state_partition_specs, + step=step, + ) + self._logging( + f"Restored checkpoint in {time.time() - start_time:.2f} seconds." + ) + self.jit_decode() + + def jit_decode(self): + """Jitting decoding function.""" + # Initialize and jit the decode fn. + def _decode(inputs): + assert self._model is not None + assert self._train_state is not None + return self._model.apply( + self._train_state.mdl_vars, + inputs, + horizon_len=self.horizon_len, + output_patch_len=self.output_patch_len, + max_len=self.context_len, + rngs={ + base_layer.PARAMS: self._key1, + base_layer.RANDOM: self._key2, + }, + method=self._model.decode, + ) + + self._logging("Jitting decoding.") + start_time = time.time() + self._pmapped_decode = jax.pmap( + _decode, + axis_name="batch", + devices=jax.devices(self.backend), + backend=self.backend, + axis_size=self.num_devices, + ) + with base_layer.JaxContext.new_context(hparams=self._eval_context): + _ = self._pmapped_decode( + NestedMap({ + "input_ts": jnp.zeros( + ( + self.num_devices, + self.per_core_batch_size, + self.context_len, + ), + dtype=jnp.float32, + ), + "input_padding": jnp.zeros( + ( + self.num_devices, + self.per_core_batch_size, + self.context_len + self.horizon_len, + ), + dtype=jnp.float32, + ), + "date_features": None, + "freq": jnp.zeros( + (self.num_devices, self.per_core_batch_size, 1), + dtype=jnp.int32, + ), + }) + ) + self._logging(f"Jitted decoding in {time.time() - start_time:.2f} seconds.") + + def _preprocess( + self, inputs: Sequence[np.array], freq: Sequence[int] + ) -> tuple[np.array, np.array, int]: + """Formats and pads raw inputs to feed into the model. + + This function both pads each time series to match the context length, and + pads the inputs to meet the SPMD shape requirement. + + Args: + inputs: A list of 1d JTensors. Each JTensor is the context time series of + a single forecast task. + freq: list of frequencies + + Returns: + A tuple of: + - the padded input time series to meet the model required context. + - the padding indicator. + - the number of padded examples for SPMD so that each core has the same + number (a multiple of `batch_size`) of examples. + """ + + input_ts, input_padding, inp_freq = [], [], [] + + pmap_pad = ( + (len(inputs) - 1) // self.global_batch_size + 1 + ) * self.global_batch_size - len(inputs) + + for i, ts in enumerate(inputs): + input_len = ts.shape[0] + padding = np.zeros(shape=(input_len + self.horizon_len,), dtype=float) + if input_len < self.context_len: + num_front_pad = self.context_len - input_len + ts = np.concatenate( + [np.zeros(shape=(num_front_pad,), dtype=float), ts], axis=0 + ) + padding = np.concatenate( + [np.ones(shape=(num_front_pad,), dtype=float), padding], axis=0 + ) + elif input_len > self.context_len: + ts = ts[-self.context_len :] + padding = padding[-(self.context_len + self.horizon_len) :] + + input_ts.append(ts) + input_padding.append(padding) + inp_freq.append(freq[i]) + + # Padding the remainder batch. + for _ in range(pmap_pad): + input_ts.append(input_ts[-1]) + input_padding.append(input_padding[-1]) + inp_freq.append(inp_freq[-1]) + + return ( + np.stack(input_ts, axis=0), + np.stack(input_padding, axis=0), + np.array(inp_freq).astype(np.int32).reshape(-1, 1), + pmap_pad, + ) + + def forecast( + self, + inputs: Sequence[Any], + freq: Sequence[int] | None = None, + window_size: int | None = None, + forecast_context_len: int | None = None, + ) -> tuple[JTensor, JTensor]: + """Forecasts on a list of time series. + + Args: + inputs: list of time series forecast contexts. Each context time series + should be in a format convertible to JTensor by `jnp.array`. + freq: frequency of each context time series. 0 for high frequency + (default), 1 for medium, and 2 for low. Notice this is different from + the `freq` required by `forecast_on_df`. + window_size: window size of trend + residual decomposition. If None then + we do not do decomposition. + forecast_context_len: optional max context length. + + Returns: + A tuple for JTensors: + - the mean forecast of size (# inputs, # forecast horizon), + - the full forecast (mean + quantiles) of size + (# inputs, # forecast horizon, 1 + # quantiles). + + Raises: + ValueError: If the checkpoint is not properly loaded. + """ + if not self._train_state or not self._model: + raise ValueError( + "Checkpoint not loaded. Call `load_from_checkpoint` before" + " `forecast`." + ) + if forecast_context_len is None: + forecast_context_len = self.context_len + inputs = [np.array(ts)[-forecast_context_len:] for ts in inputs] + inp_min = np.min([np.min(ts) for ts in inputs]) + + if window_size is not None: + new_inputs = [] + for ts in inputs: + new_inputs.extend(moving_average(ts, window_size)) + inputs = new_inputs + + if freq is None: + logging.info("No frequency provided via `freq`. Default to high (0).") + freq = [0] * len(inputs) + + input_ts, input_padding, inp_freq, pmap_pad = self._preprocess(inputs, freq) + with base_layer.JaxContext.new_context(hparams=self._eval_context): + mean_outputs = [] + full_outputs = [] + assert input_ts.shape[0] % self.global_batch_size == 0 + for i in range(input_ts.shape[0] // self.global_batch_size): + input_ts_in = jnp.array( + input_ts[ + i * self.global_batch_size : (i + 1) * self.global_batch_size + ] + ) + input_padding_in = jnp.array( + input_padding[ + i * self.global_batch_size : (i + 1) * self.global_batch_size + ], + ) + inp_freq_in = jnp.array( + inp_freq[ + i * self.global_batch_size : (i + 1) * self.global_batch_size, : + ], + dtype=jnp.int32, + ) + pmapped_inputs = NestedMap({ + "input_ts": es.jax_einshape( + "(db)...->db...", + input_ts_in, + d=self.num_devices, + ), + "input_padding": es.jax_einshape( + "(db)...->db...", + input_padding_in, + d=self.num_devices, + ), + "date_features": None, + "freq": es.jax_einshape( + "(db)...->db...", + inp_freq_in, + d=self.num_devices, + ), + }) + mean_output, full_output = self._pmapped_decode(pmapped_inputs) + mean_output = es.jax_einshape( + "db...->(db)...", mean_output, d=self.num_devices + ) + full_output = es.jax_einshape( + "db...->(db)...", full_output, d=self.num_devices + ) + mean_output = np.array(mean_output) + full_output = np.array(full_output) + mean_outputs.append(mean_output) + full_outputs.append(full_output) + + mean_outputs = np.concatenate(mean_outputs, axis=0) + full_outputs = np.concatenate(full_outputs, axis=0) + + if pmap_pad > 0: + mean_outputs = mean_outputs[:-pmap_pad, ...] + full_outputs = full_outputs[:-pmap_pad, ...] + + if window_size is not None: + mean_outputs = mean_outputs[0::2, ...] + mean_outputs[1::2, ...] + full_outputs = full_outputs[0::2, ...] + full_outputs[1::2, ...] + if inp_min >= 0: + mean_outputs = np.maximum(mean_outputs, 0.0) + full_outputs = np.maximum(full_outputs, 0.0) + return mean_outputs, full_outputs + + def forecast_on_df( + self, + inputs: pd.DataFrame, + freq: str, + forecast_context_len: int = 0, + value_name: str = "values", + model_name: str = "timesfm", + window_size: int | None = None, + num_jobs: int = 1, + ) -> pd.DataFrame: + """Forecasts on a list of time series. + + Args: + inputs: A pd.DataFrame of all time series. The dataframe should have a + `unique_id` column for identifying the time series, a `ds` column for + timestamps and a value column for the time series values. + freq: string valued `freq` of data. Notice this is different from the + `freq` required by `forecast`. See `freq_map` for allowed values. + forecast_context_len: If provided none zero, we take the last + `forecast_context_len` time-points from each series as the forecast + context instead of the `context_len` set by the model. + value_name: The name of the value column. + model_name: name of the model to be written into future df. + window_size: window size of trend + residual decomposition. If None then + we do not do decomposition. + num_jobs: number of parallel processes to use for dataframe processing. + + Returns: + Future forecasts dataframe. + """ + if not ( + "unique_id" in inputs.columns + and "ds" in inputs.columns + and value_name in inputs.columns + ): + raise ValueError( + f"DataFrame must have unique_id, ds and {value_name} columns." + ) + if not forecast_context_len: + forecast_context_len = self.context_len + logging.info("Preprocessing dataframe.") + df_sorted = inputs.sort_values(by=["unique_id", "ds"]) + new_inputs = [] + uids = [] + if num_jobs == 1: + print("Processing dataframe with single process.") + for key, group in df_sorted.groupby("unique_id"): + inp, uid = process_group( + key, + group, + value_name, + forecast_context_len, + ) + new_inputs.append(inp) + uids.append(uid) + else: + if num_jobs == -1: + num_jobs = multiprocessing.cpu_count() + print("Processing dataframe with multiple processes.") + with multiprocessing.Pool(processes=num_jobs) as pool: + results = pool.starmap( + process_group, + [ + (key, group, value_name, forecast_context_len) + for key, group in df_sorted.groupby("unique_id") + ], + ) + new_inputs, uids = zip(*results) + print("Finished preprocessing dataframe.") + freq_inps = [freq_map(freq)] * len(new_inputs) + _, full_forecast = self.forecast( + new_inputs, freq=freq_inps, window_size=window_size + ) + print("Finished forecasting.") + fcst_df = make_future_dataframe( + uids=uids, + last_times=df_sorted.groupby("unique_id")["ds"].tail(1), + h=self.horizon_len, + freq=freq, + ) + fcst_df[model_name] = full_forecast[:, 0 : self.horizon_len, 0].reshape( + -1, 1 + ) + + if self._model.quantiles is not None: + for i, q in enumerate(self._model.quantiles): + q_col = f"{model_name}-q-{q}" + fcst_df[q_col] = full_forecast[:, 0 : self.horizon_len, 1 + i].reshape( + -1, 1 + ) + if q == 0.5: + fcst_df[model_name] = fcst_df[q_col] + logging.info("Finished creating output dataframe.") + return fcst_df From 739d74e19911e5baafc9c5d23bf3bcd154ce9c24 Mon Sep 17 00:00:00 2001 From: Rajat Sen Date: Mon, 8 Jul 2024 18:37:00 +0000 Subject: [PATCH 2/2] deleting old files --- __init__.py | 21 - .../long_horizon_benchmarks/data_loader.py | 261 -------- .../long_horizon_benchmarks/time_features.py | 215 ------- src/patched_decoder.py | 461 -------------- src/timesfm.py | 602 ------------------ 5 files changed, 1560 deletions(-) delete mode 100644 __init__.py delete mode 100644 experiments/long_horizon_benchmarks/data_loader.py delete mode 100644 experiments/long_horizon_benchmarks/time_features.py delete mode 100644 src/patched_decoder.py delete mode 100644 src/timesfm.py diff --git a/__init__.py b/__init__.py deleted file mode 100644 index 8275b16..0000000 --- a/__init__.py +++ /dev/null @@ -1,21 +0,0 @@ -# Copyright 2024 Google LLC -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""TimesFM init file.""" - -from __future__ import absolute_import - -from .src.patched_decoder import PatchedTimeSeriesDecoder -from .src.timesfm import TimesFm -from .src.timesfm import freq_map diff --git a/experiments/long_horizon_benchmarks/data_loader.py b/experiments/long_horizon_benchmarks/data_loader.py deleted file mode 100644 index eeace7b..0000000 --- a/experiments/long_horizon_benchmarks/data_loader.py +++ /dev/null @@ -1,261 +0,0 @@ -# Copyright 2024 The Google Research Authors. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""TF dataloaders for general timeseries datasets. - -The expected input format is csv file with a datetime index. -""" - - -from absl import logging -import numpy as np -import pandas as pd -from sklearn.preprocessing import StandardScaler -import tensorflow as tf -from . import time_features - - -class TimeSeriesdata(object): - """Data loader class.""" - - def __init__( - self, - data_path, - datetime_col, - num_cov_cols, - cat_cov_cols, - ts_cols, - train_range, - val_range, - test_range, - hist_len, - pred_len, - batch_size, - freq='H', - normalize=True, - epoch_len=None, - holiday=False, - permute=True, - ): - """Initialize objects. - - Args: - data_path: path to csv file - datetime_col: column name for datetime col - num_cov_cols: list of numerical global covariates - cat_cov_cols: list of categorical global covariates - ts_cols: columns corresponding to ts - train_range: tuple of train ranges - val_range: tuple of validation ranges - test_range: tuple of test ranges - hist_len: historical context - pred_len: prediction length - batch_size: batch size (number of ts in a batch) - freq: freq of original data - normalize: std. normalize data or not - epoch_len: num iters in an epoch - holiday: use holiday features or not - permute: permute ts in train batches or not - - Returns: - None - """ - self.data_df = pd.read_csv(open(data_path, 'r')) - if not num_cov_cols: - self.data_df['ncol'] = np.zeros(self.data_df.shape[0]) - num_cov_cols = ['ncol'] - if not cat_cov_cols: - self.data_df['ccol'] = np.zeros(self.data_df.shape[0]) - cat_cov_cols = ['ccol'] - self.data_df.fillna(0, inplace=True) - self.data_df.set_index( - pd.DatetimeIndex(self.data_df[datetime_col]), inplace=True - ) - self.num_cov_cols = num_cov_cols - self.cat_cov_cols = cat_cov_cols - self.ts_cols = ts_cols - self.train_range = train_range - self.val_range = val_range - self.test_range = test_range - data_df_idx = self.data_df.index - date_index = data_df_idx.union( - pd.date_range( - data_df_idx[-1] + pd.Timedelta(1, freq=freq), - periods=pred_len + 1, - freq=freq, - ) - ) - self.time_df = time_features.TimeCovariates( - date_index, holiday=holiday - ).get_covariates() - self.hist_len = hist_len - self.pred_len = pred_len - self.batch_size = batch_size - self.freq = freq - self.normalize = normalize - self.data_mat = self.data_df[self.ts_cols].to_numpy().transpose() - self.data_mat = self.data_mat[:, 0 : self.test_range[1]] - self.time_mat = self.time_df.to_numpy().transpose() - self.num_feat_mat = self.data_df[num_cov_cols].to_numpy().transpose() - self.cat_feat_mat, self.cat_sizes = self._get_cat_cols(cat_cov_cols) - self.normalize = normalize - if normalize: - self._normalize_data() - logging.info( - 'Data Shapes: %s, %s, %s, %s', - self.data_mat.shape, - self.time_mat.shape, - self.num_feat_mat.shape, - self.cat_feat_mat.shape, - ) - self.epoch_len = epoch_len - self.permute = permute - - def _get_cat_cols(self, cat_cov_cols): - """Get categorical columns.""" - cat_vars = [] - cat_sizes = [] - for col in cat_cov_cols: - dct = {x: i for i, x in enumerate(self.data_df[col].unique())} - cat_sizes.append(len(dct)) - mapped = self.data_df[col].map(lambda x: dct[x]).to_numpy().transpose() # pylint: disable=cell-var-from-loop - cat_vars.append(mapped) - return np.vstack(cat_vars), cat_sizes - - def _normalize_data(self): - self.scaler = StandardScaler() - train_mat = self.data_mat[:, self.train_range[0] : self.train_range[1]] - self.scaler = self.scaler.fit(train_mat.transpose()) - self.data_mat = self.scaler.transform(self.data_mat.transpose()).transpose() - - def train_gen(self): - """Generator for training data.""" - num_ts = len(self.ts_cols) - perm = np.arange( - self.train_range[0] + self.hist_len, - self.train_range[1] - self.pred_len, - ) - perm = np.random.permutation(perm) - hist_len = self.hist_len - logging.info('Hist len: %s', hist_len) - if not self.epoch_len: - epoch_len = len(perm) - else: - epoch_len = self.epoch_len - for idx in perm[0:epoch_len]: - for _ in range(num_ts // self.batch_size + 1): - if self.permute: - tsidx = np.random.choice(num_ts, size=self.batch_size, replace=False) - else: - tsidx = np.arange(num_ts) - dtimes = np.arange(idx - hist_len, idx + self.pred_len) - ( - bts_train, - bts_pred, - bfeats_train, - bfeats_pred, - bcf_train, - bcf_pred, - ) = self._get_features_and_ts(dtimes, tsidx, hist_len) - - all_data = [ - bts_train, - bfeats_train, - bcf_train, - bts_pred, - bfeats_pred, - bcf_pred, - tsidx, - ] - yield tuple(all_data) - - def test_val_gen(self, mode='val', shift=1): - """Generator for validation/test data.""" - if mode == 'val': - start = self.val_range[0] - end = self.val_range[1] - self.pred_len + 1 - elif mode == 'test': - start = self.test_range[0] - end = self.test_range[1] - self.pred_len + 1 - else: - raise NotImplementedError('Eval mode not implemented') - num_ts = len(self.ts_cols) - hist_len = self.hist_len - logging.info('Hist len: %s', hist_len) - perm = np.arange(start, end) - if self.epoch_len: - epoch_len = self.epoch_len - else: - epoch_len = len(perm) - for i in range(0, epoch_len, shift): - idx = perm[i] - for batch_idx in range(0, num_ts, self.batch_size): - tsidx = np.arange(batch_idx, min(batch_idx + self.batch_size, num_ts)) - dtimes = np.arange(idx - hist_len, idx + self.pred_len) - ( - bts_train, - bts_pred, - bfeats_train, - bfeats_pred, - bcf_train, - bcf_pred, - ) = self._get_features_and_ts(dtimes, tsidx, hist_len) - all_data = [ - bts_train, - bfeats_train, - bcf_train, - bts_pred, - bfeats_pred, - bcf_pred, - tsidx, - ] - yield tuple(all_data) - - def _get_features_and_ts(self, dtimes, tsidx, hist_len=None): - """Get features and ts in specified windows.""" - if hist_len is None: - hist_len = self.hist_len - data_times = dtimes[dtimes < self.data_mat.shape[1]] - bdata = self.data_mat[:, data_times] - bts = bdata[tsidx, :] - bnf = self.num_feat_mat[:, data_times] - bcf = self.cat_feat_mat[:, data_times] - btf = self.time_mat[:, dtimes] - if bnf.shape[1] < btf.shape[1]: - rem_len = btf.shape[1] - bnf.shape[1] - rem_rep = np.repeat(bnf[:, [-1]], repeats=rem_len) - rem_rep_cat = np.repeat(bcf[:, [-1]], repeats=rem_len) - bnf = np.hstack([bnf, rem_rep.reshape(bnf.shape[0], -1)]) - bcf = np.hstack([bcf, rem_rep_cat.reshape(bcf.shape[0], -1)]) - bfeats = np.vstack([btf, bnf]) - bts_train = bts[:, 0:hist_len] - bts_pred = bts[:, hist_len:] - bfeats_train = bfeats[:, 0:hist_len] - bfeats_pred = bfeats[:, hist_len:] - bcf_train = bcf[:, 0:hist_len] - bcf_pred = bcf[:, hist_len:] - return bts_train, bts_pred, bfeats_train, bfeats_pred, bcf_train, bcf_pred - - def tf_dataset(self, mode='train', shift=1): - """Tensorflow Dataset.""" - if mode == 'train': - gen_fn = self.train_gen - else: - gen_fn = lambda: self.test_val_gen(mode, shift) - output_types = tuple( - [tf.float32] * 2 + [tf.int32] + [tf.float32] * 2 + [tf.int32] * 2 - ) - dataset = tf.data.Dataset.from_generator(gen_fn, output_types) - dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE) - return dataset diff --git a/experiments/long_horizon_benchmarks/time_features.py b/experiments/long_horizon_benchmarks/time_features.py deleted file mode 100644 index 0bd90a9..0000000 --- a/experiments/long_horizon_benchmarks/time_features.py +++ /dev/null @@ -1,215 +0,0 @@ -# Copyright 2024 The Google Research Authors. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Directory to extract time covariates. - -Extract time covariates from datetime. -""" - -import numpy as np -import pandas as pd -from pandas.tseries.holiday import EasterMonday -from pandas.tseries.holiday import GoodFriday -from pandas.tseries.holiday import Holiday -from pandas.tseries.holiday import SU -from pandas.tseries.holiday import TH -from pandas.tseries.holiday import USColumbusDay -from pandas.tseries.holiday import USLaborDay -from pandas.tseries.holiday import USMartinLutherKingJr -from pandas.tseries.holiday import USMemorialDay -from pandas.tseries.holiday import USPresidentsDay -from pandas.tseries.holiday import USThanksgivingDay -from pandas.tseries.offsets import DateOffset -from pandas.tseries.offsets import Day -from pandas.tseries.offsets import Easter -from sklearn.preprocessing import StandardScaler -from tqdm import tqdm - - -# This is 183 to cover half a year (in both directions), also for leap years -# + 17 as Eastern can be between March, 22 - April, 25 -MAX_WINDOW = 183 + 17 - - -def _distance_to_holiday(holiday): - """Return distance to given holiday.""" - - def _distance_to_day(index): - holiday_date = holiday.dates( - index - pd.Timedelta(days=MAX_WINDOW), - index + pd.Timedelta(days=MAX_WINDOW), - ) - assert ( - len(holiday_date) != 0 # pylint: disable=g-explicit-length-test - ), f"No closest holiday for the date index {index} found." - # It sometimes returns two dates if it is exactly half a year after the - # holiday. In this case, the smaller distance (182 days) is returned. - return (index - holiday_date[0]).days - - return _distance_to_day - - -EasterSunday = Holiday( - "Easter Sunday", month=1, day=1, offset=[Easter(), Day(0)] -) -NewYearsDay = Holiday("New Years Day", month=1, day=1) -SuperBowl = Holiday( - "Superbowl", month=2, day=1, offset=DateOffset(weekday=SU(1)) -) -MothersDay = Holiday( - "Mothers Day", month=5, day=1, offset=DateOffset(weekday=SU(2)) -) -IndependenceDay = Holiday("Independence Day", month=7, day=4) -ChristmasEve = Holiday("Christmas", month=12, day=24) -ChristmasDay = Holiday("Christmas", month=12, day=25) -NewYearsEve = Holiday("New Years Eve", month=12, day=31) -BlackFriday = Holiday( - "Black Friday", - month=11, - day=1, - offset=[pd.DateOffset(weekday=TH(4)), Day(1)], -) -CyberMonday = Holiday( - "Cyber Monday", - month=11, - day=1, - offset=[pd.DateOffset(weekday=TH(4)), Day(4)], -) - -HOLIDAYS = [ - EasterMonday, - GoodFriday, - USColumbusDay, - USLaborDay, - USMartinLutherKingJr, - USMemorialDay, - USPresidentsDay, - USThanksgivingDay, - EasterSunday, - NewYearsDay, - SuperBowl, - MothersDay, - IndependenceDay, - ChristmasEve, - ChristmasDay, - NewYearsEve, - BlackFriday, - CyberMonday, -] - - -class TimeCovariates(object): - """Extract all time covariates except for holidays.""" - - def __init__( - self, - datetimes, - normalized=True, - holiday=False, - ): - """Init function. - - Args: - datetimes: pandas DatetimeIndex (lowest granularity supported is min) - normalized: whether to normalize features or not - holiday: fetch holiday features or not - - Returns: - None - """ - self.normalized = normalized - self.dti = datetimes - self.holiday = holiday - - def _minute_of_hour(self): - minutes = np.array(self.dti.minute, dtype=np.float32) - if self.normalized: - minutes = minutes / 59.0 - 0.5 - return minutes - - def _hour_of_day(self): - hours = np.array(self.dti.hour, dtype=np.float32) - if self.normalized: - hours = hours / 23.0 - 0.5 - return hours - - def _day_of_week(self): - day_week = np.array(self.dti.dayofweek, dtype=np.float32) - if self.normalized: - day_week = day_week / 6.0 - 0.5 - return day_week - - def _day_of_month(self): - day_month = np.array(self.dti.day, dtype=np.float32) - if self.normalized: - day_month = day_month / 30.0 - 0.5 - return day_month - - def _day_of_year(self): - day_year = np.array(self.dti.dayofyear, dtype=np.float32) - if self.normalized: - day_year = day_year / 364.0 - 0.5 - return day_year - - def _month_of_year(self): - month_year = np.array(self.dti.month, dtype=np.float32) - if self.normalized: - month_year = month_year / 11.0 - 0.5 - return month_year - - def _week_of_year(self): - week_year = np.array(self.dti.strftime("%U").astype(int), dtype=np.float32) - if self.normalized: - week_year = week_year / 51.0 - 0.5 - return week_year - - def _get_holidays(self): - dti_series = self.dti.to_series() - hol_variates = np.vstack([ - dti_series.apply(_distance_to_holiday(h)).values for h in tqdm(HOLIDAYS) - ]) - # hol_variates is (num_holiday, num_time_steps), the normalization should be - # performed in the num_time_steps dimension. - return StandardScaler().fit_transform(hol_variates.T).T - - def get_covariates(self): - """Get all time covariates.""" - moh = self._minute_of_hour().reshape(1, -1) - hod = self._hour_of_day().reshape(1, -1) - dom = self._day_of_month().reshape(1, -1) - dow = self._day_of_week().reshape(1, -1) - doy = self._day_of_year().reshape(1, -1) - moy = self._month_of_year().reshape(1, -1) - woy = self._week_of_year().reshape(1, -1) - - all_covs = [ - moh, - hod, - dom, - dow, - doy, - moy, - woy, - ] - columns = ["moh", "hod", "dom", "dow", "doy", "moy", "woy"] - if self.holiday: - hol_covs = self._get_holidays() - all_covs.append(hol_covs) - columns += [f"hol_{i}" for i in range(len(HOLIDAYS))] - - return pd.DataFrame( - data=np.vstack(all_covs).transpose(), - columns=columns, - index=self.dti, - ) diff --git a/src/patched_decoder.py b/src/patched_decoder.py deleted file mode 100644 index b0decf8..0000000 --- a/src/patched_decoder.py +++ /dev/null @@ -1,461 +0,0 @@ -# Copyright 2024 Google LLC -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Pax ML model for patched time-series decoder. - -The file implements Residual MLPs, Patched Decoder layers and PAX ML models. -""" - -import dataclasses -from typing import Optional, Tuple - -import einshape as es -from jax import lax -import jax.numpy as jnp -from praxis import base_layer -from praxis import layers -from praxis import pax_fiddle -from praxis import py_utils -from praxis import pytypes -from praxis.layers import activations -from praxis.layers import embedding_softmax -from praxis.layers import linears -from praxis.layers import normalizations -from praxis.layers import stochastics -from praxis.layers import transformers - - -# PAX shortcuts -NestedMap = py_utils.NestedMap -JTensor = pytypes.JTensor - -LayerTpl = pax_fiddle.Config[base_layer.BaseLayer] -template_field = base_layer.template_field - - -PAD_VAL = 1123581321.0 -DEFAULT_QUANTILES = [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9] - -# NestedMap keys -_INPUT_TS = "input_ts" -_INPUT_PADDING = "input_padding" -_OUTPUT_TS = "output_ts" -_FREQ = "freq" -_OUTPUT_TOKENS = "output_tokens" -_STATS = "stats" - - -# Small numerical value. -_TOLERANCE = 1e-7 - - -def _shift_padded_seq(mask: JTensor, seq: JTensor) -> JTensor: - """Shifts rows of seq based on the first 0 in each row of the mask.""" - num = seq.shape[1] - - # Find the index of the first 0 in each row of the mask - first_zero_idx = jnp.argmin(mask, axis=1) - - # Create a range array for indexing - idx_range = jnp.arange(num) - - def shift_row(carry, x): - seq_row, shift = x - shifted_idx = (idx_range - shift) % num - shifted_row = seq_row[shifted_idx] - return carry, shifted_row - - # Use lax.scan to shift each row of seq based on the corresponding - # first_zero_idx. - _, shifted_seq = lax.scan(shift_row, None, (seq, first_zero_idx)) - - return shifted_seq - - -class ResidualBlock(base_layer.BaseLayer): - """Simple feedforward block with residual connection. - - Attributes: - input_dims: input dimension. - hidden_dims: hidden dimension. - output_dims: output dimension. - dropout_prob: dropout probability. - layer_norm: whether to use layer norm or not. - dropout_tpl: config for dropout. - ln_tpl: config for layer norm. - act_tpl: config for activation in hidden layer. - """ - - input_dims: int = 0 - hidden_dims: int = 0 - output_dims: int = 0 - dropout_prob: float = 0.0 - layer_norm: bool = False - dropout_tpl: LayerTpl = template_field(stochastics.Dropout) - ln_tpl: LayerTpl = template_field(normalizations.LayerNorm) - act_tpl: LayerTpl = template_field(activations.Swish) - - def setup(self): - lnorm_tpl = self.ln_tpl.clone() - lnorm_tpl.dim = self.output_dims - self.create_child("ln_layer", lnorm_tpl) - - dropout_tpl = self.dropout_tpl.clone() - dropout_tpl.keep_prob = 1.0 - self.dropout_prob - self.create_child("dropout", dropout_tpl) - - self.create_child( - "hidden_layer", - pax_fiddle.Config( - linears.FeedForward, - input_dims=self.input_dims, - output_dims=self.hidden_dims, - activation_tpl=self.act_tpl.clone(), - ), - ) - - self.create_child( - "output_layer", - pax_fiddle.Config( - linears.FeedForward, - input_dims=self.hidden_dims, - output_dims=self.output_dims, - activation_tpl=pax_fiddle.Config(activations.Identity), - ), - ) - - self.create_child( - "residual_layer", - pax_fiddle.Config( - linears.FeedForward, - input_dims=self.input_dims, - output_dims=self.output_dims, - activation_tpl=pax_fiddle.Config(activations.Identity), - ), - ) - - def __call__(self, inputs: JTensor) -> JTensor: - hidden = self.hidden_layer(inputs) - output = self.output_layer(hidden) - output = self.dropout(output) - residual = self.residual_layer(inputs) - if self.layer_norm: - return self.ln_layer(output + residual) - else: - return output + residual - - -def _masked_mean_std( - inputs: JTensor, padding: JTensor -) -> Tuple[JTensor, JTensor]: - """Calculates mean and standard deviation of arr across axis 1. - - It should exclude values where pad is 1. - - Args: - inputs: A JAX array of shape [b, n, p]. - padding: A JAX array of shape [b, n, p] with values 0 or 1. - - Returns: - A tuple containing the mean and standard deviation of arr. We return the - statistics of the first patch with more than three non-padded values. - """ - # Selecting the first pad with more than 3 unpadded values. - pad_sum = jnp.sum(1 - padding, axis=2) - - def _get_patch_index(arr: JTensor): - indices = jnp.argmax(arr >= 3, axis=1) - row_sum = (arr >= 3).sum(axis=1) - return jnp.where(row_sum == 0, arr.shape[1] - 1, indices) - - patch_indices = _get_patch_index(pad_sum) - bidxs = jnp.arange(inputs.shape[0]) - - arr = inputs[bidxs, patch_indices, :] - pad = padding[bidxs, patch_indices, :] - - # Create a mask where P is 0 - mask = 1 - pad - - # Calculate the number of valid elements - num_valid_elements = jnp.sum(mask, axis=1) - - num_valid_elements = jnp.where(num_valid_elements == 0, 1, num_valid_elements) - - # Calculate the masked sum and squared sum of M - masked_sum = jnp.sum(arr * mask, axis=1) - masked_squared_sum = jnp.sum((arr * mask) ** 2, axis=1) - - # Calculate the masked mean and standard deviation - masked_mean = masked_sum / num_valid_elements - masked_var = masked_squared_sum / num_valid_elements - masked_mean**2 - masked_var = jnp.where(masked_var < 0.0, 0.0, masked_var) - masked_std = jnp.sqrt(masked_var) - - return masked_mean, masked_std - - -def _create_quantiles() -> list[float]: - """Returns the quantiles for forecasting.""" - return DEFAULT_QUANTILES - - -class PatchedTimeSeriesDecoder(base_layer.BaseLayer): - """Patch decoder layer for time-series foundation model. - - Attributes: - patch_len: length of input patches. - horizon_len: length of output patches. Referred to as `output_patch_len` - during inference. - model_dims: model dimension of stacked transformer layer. - hidden_dims: hidden dimensions in fully connected layers. - quantiles: list of quantiles for non prob model. - residual_block_tpl: config for residual block. - stacked_transformer_params_tpl: config for stacked transformer. - use_freq: whether to use frequency encoding. - - In all of what followed, except specified otherwise, B is batch size, T is - sequence length of time-series. N is the number of input patches that can be - obtained from T. P is the input patch length and H is the horizon length. Q is - number of output logits. D is model dimension. - """ - - patch_len: int = 0 - horizon_len: int = 0 - model_dims: int = 0 - hidden_dims: int = 0 - quantiles: list[float] = dataclasses.field(default_factory=_create_quantiles) - residual_block_tpl: LayerTpl = template_field(ResidualBlock) - stacked_transformer_params_tpl: LayerTpl = template_field( - transformers.StackedTransformer - ) - use_freq: bool = True - - def setup(self) -> None: - """Construct the model.""" - num_outputs = len(self.quantiles) + 1 - - stl = self.stacked_transformer_params_tpl.clone() - stl.model_dims = self.model_dims - stl.hidden_dims = self.hidden_dims - stl.mask_self_attention = True - - self.create_child("stacked_transformer_layer", stl) - - input_resl = self.residual_block_tpl.clone() - ff_in_dims = 2 * self.patch_len - input_resl.input_dims = ff_in_dims - input_resl.hidden_dims = self.hidden_dims - input_resl.output_dims = self.model_dims - self.create_child( - "input_ff_layer", - input_resl, - ) - - horizon_resl = self.residual_block_tpl.clone() - horizon_resl.input_dims = self.model_dims - horizon_resl.hidden_dims = self.hidden_dims - horizon_resl.output_dims = self.horizon_len * num_outputs - self.create_child( - "horizon_ff_layer", - horizon_resl, - ) - - self.create_child( - "position_emb", - pax_fiddle.Config( - layers.PositionalEmbedding, embedding_dims=self.model_dims - ), - ) - - if self.use_freq: - self.create_child( - "freq_emb", - pax_fiddle.Config( - embedding_softmax.Embedding, - num_classes=3, - input_dims=self.model_dims, - ), - ) - - def transform_decode_state( - self, transform_fn: base_layer.DecodeStateTransformFn - ) -> None: - """Transforms all decode state variables based on transform_fn.""" - self.stacked_transformer_layer.transform_decode_state(transform_fn) - - def _forward_transform( - self, inputs: JTensor, patched_pads: JTensor - ) -> Tuple[JTensor, Tuple[JTensor, JTensor]]: - """Input is of shape [B, N, P].""" - mu, sigma = _masked_mean_std(inputs, patched_pads) - sigma = jnp.where(sigma < _TOLERANCE, 1.0, sigma) - # Normalize each patch. - outputs = (inputs - mu[:, None, None]) / sigma[:, None, None] - outputs = jnp.where( - jnp.abs(inputs - PAD_VAL) < _TOLERANCE, PAD_VAL, outputs - ) - return outputs, (mu, sigma) - - def _reverse_transform( - self, outputs: JTensor, stats: Tuple[JTensor, JTensor] - ) -> JTensor: - """Output is of shape [B, N, P, Q].""" - mu, sigma = stats - return outputs * sigma[:, None, None, None] + mu[:, None, None, None] - - def _preprocess_input( - self, - input_ts: JTensor, - input_padding: JTensor, - pos_emb: Optional[JTensor] = None, - ) -> Tuple[JTensor, JTensor, Optional[Tuple[JTensor, JTensor]], JTensor]: - """Preprocess input for stacked transformer.""" - # Reshape into patches. - patched_inputs = es.jax_einshape("b(np)->bnp", input_ts, p=self.patch_len) - input_padding = jnp.where( - jnp.abs(input_ts - PAD_VAL) < _TOLERANCE, 1, input_padding - ) - patched_pads = es.jax_einshape( - "b(np)->bnp", input_padding, p=self.patch_len - ) - patched_inputs, stats = self._forward_transform( - patched_inputs, patched_pads - ) - # B x N x D - patched_inputs = patched_inputs * (1.0 - patched_pads) - concat_inputs = jnp.concatenate([patched_inputs, patched_pads], axis=-1) - model_input = self.input_ff_layer(concat_inputs) - # A patch should not be padded even if there is at least one zero. - patched_padding = jnp.min(patched_pads, axis=-1) - - if pos_emb is None: - position_emb = self.position_emb(seq_length=model_input.shape[1]) - else: - position_emb = pos_emb - if self.do_eval: - if position_emb.shape[0] != model_input.shape[0]: - position_emb = jnp.repeat(position_emb, model_input.shape[0], axis=0) - position_emb = _shift_padded_seq(patched_padding, position_emb) - model_input += position_emb - - return model_input, patched_padding, stats, patched_inputs - - def _postprocess_output( - self, - model_output: JTensor, - num_outputs: int, - stats: Tuple[JTensor, JTensor], - ) -> JTensor: - """Postprocess output of stacked transformer.""" - # B x N x (H.Q) - output_ts = self.horizon_ff_layer(model_output) - output_ts = es.jax_einshape( - "bn(hq)->bnhq", output_ts, q=num_outputs, h=self.horizon_len - ) - return self._reverse_transform(output_ts, stats) - - def __call__(self, inputs: NestedMap) -> NestedMap: - """PatchTST call. - - Args: - inputs: A NestedMap containing (1) input_ts: input sequence of shape [B, - T] where T must be multiple of patch_length; (2) input_padding: that - contains padding map. - - Returns: - A nested map with two keys: - (1) 'output_tokens' of shape [B, N, D]. - (2) 'output_ts' of shape [B, N, H, Q] - (3) 'stats' a Tuple of statistics for renormalization. - """ - input_ts, input_padding = inputs[_INPUT_TS], inputs[_INPUT_PADDING] - num_outputs = len(self.quantiles) + 1 - model_input, patched_padding, stats, _ = self._preprocess_input( - input_ts=input_ts, - input_padding=input_padding, - ) - if self.use_freq: - freq = inputs[_FREQ].astype(jnp.int32) - f_emb = self.freq_emb(freq) # B x 1 x D - f_emb = jnp.repeat(f_emb, model_input.shape[1], axis=1) - model_input += f_emb - model_output = self.stacked_transformer_layer(model_input, patched_padding) - - output_ts = self._postprocess_output(model_output, num_outputs, stats) - return NestedMap( - {_OUTPUT_TOKENS: model_output, _OUTPUT_TS: output_ts, _STATS: stats} - ) - - def decode( - self, - inputs: NestedMap, - horizon_len: int, - output_patch_len: Optional[int] = None, - max_len: int = 512, - ) -> tuple[JTensor, JTensor]: - """Auto-regressive decoding without caching. - - Args: - inputs: input time-series and paddings. Time-series shape B x C, padding - shape shape B x (C + H) where H is the prediction length. - horizon_len: prediction length. - output_patch_len: output length to be fetched from one step of - auto-regressive decoding. - max_len: maximum training context length. - - Returns: - Tuple of two forecasting results: - - Point (mean) output predictions as a tensor with shape B x H. - - Full predictions (mean and quantiles) as a tensor with shape - B x H x (1 + # quantiles). - """ - final_out = inputs[_INPUT_TS] - inp_time_len = final_out.shape[1] - paddings = inputs[_INPUT_PADDING] - if self.use_freq: - freq = inputs[_FREQ].astype(jnp.int32) - else: - freq = jnp.zeros([final_out.shape[0], 1], dtype=jnp.int32) - full_outputs = [] - if paddings.shape[1] != final_out.shape[1] + horizon_len: - raise ValueError( - "Length of paddings must match length of input + horizon_len:" - f" {paddings.shape[1]} != {final_out.shape[1]} + {horizon_len}" - ) - if output_patch_len is None: - output_patch_len = self.horizon_len - num_decode_patches = ( - horizon_len + output_patch_len - 1 - ) // output_patch_len - for _ in range(num_decode_patches): - current_padding = paddings[:, 0 : final_out.shape[1]] - input_ts = final_out[:, -max_len:] - input_padding = current_padding[:, -max_len:] - model_input = NestedMap( - input_ts=input_ts, - input_padding=input_padding, - freq=freq, - ) - fprop_outputs = self(model_input)[_OUTPUT_TS] - # (full batch, last patch, output_patch_len, index of mean forecast = 0) - new_ts = fprop_outputs[:, -1, :output_patch_len, 0] - # (full batch, last patch, output_patch_len, all output indices) - full_outputs.append(fprop_outputs[:, -1, :output_patch_len, :]) - final_out = jnp.concatenate([final_out, new_ts], axis=-1) - - return ( - final_out[:, inp_time_len : inp_time_len + horizon_len], - jnp.concatenate(full_outputs, axis=1)[:, 0:horizon_len, :], - ) diff --git a/src/timesfm.py b/src/timesfm.py deleted file mode 100644 index 7dc9d0a..0000000 --- a/src/timesfm.py +++ /dev/null @@ -1,602 +0,0 @@ -# Copyright 2024 Google LLC -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""TimesFM forecast API for inference.""" - -import logging -import multiprocessing -from os import path -import time -from typing import Any, Literal, Optional, Sequence - -import einshape as es -import jax -import jax.numpy as jnp -import numpy as np -import pandas as pd -from huggingface_hub import snapshot_download -from paxml import checkpoints -from paxml import tasks_lib -from praxis import base_hyperparams -from praxis import base_layer -from praxis import pax_fiddle -from praxis import py_utils -from praxis import pytypes -from praxis.layers import normalizations -from praxis.layers import transformers -import patched_decoder -from utilsforecast.processing import make_future_dataframe - -instantiate = base_hyperparams.instantiate -NestedMap = py_utils.NestedMap -JTensor = pytypes.JTensor - - -def process_group(key, group, value_name, forecast_context_len): - group = group.tail(forecast_context_len) - return np.array(group[value_name], dtype=np.float32), key - - -def moving_average(arr, window_size): - """Calculates the moving average using NumPy's convolution function.""" - # Pad with zeros to handle initial window positions - arr_padded = np.pad(arr, (window_size - 1, 0), "constant") - smoothed_arr = ( - np.convolve(arr_padded, np.ones(window_size), "valid") / window_size - ) - return [smoothed_arr, arr - smoothed_arr] - - -def freq_map(freq: str): - """Returns the frequency map for the given frequency string.""" - freq = str.upper(freq) - if ( - freq.endswith("H") - or freq.endswith("T") - or freq.endswith("MIN") - or freq.endswith("D") - or freq.endswith("B") - or freq.endswith("U") - ): - return 0 - elif freq.endswith(("W", "M", "MS")): - return 1 - elif freq.endswith("Y") or freq.endswith("Q"): - return 2 - else: - raise ValueError(f"Invalid frequency: {freq}") - - -class TimesFm: - """TimesFM forecast API for inference. - - This class is the scaffolding for calling TimesFM forecast. To properly use: - 1. Create an instance with the correct hyperparameters of a TimesFM model. - 2. Call `load_from_checkpoint` to load a compatible checkpoint. - 3. Call `forecast` for inference. - - Given the model size, this API does not shard the model weights for SPMD. All - parallelism happens on the data dimension. - - Compilation happens during the first time `forecast` is called and uses the - `per_core_batch_size` to set and freeze the input signature. Subsequent calls - to `forecast` reflect the actual inference latency. - - Attributes: - per_core_batch_size: Batch size on each core for data parallelism. - backend: One of "cpu", "gpu" or "tpu". - num_devices: Number of cores provided the backend. - global_batch_size: per_core_batch_size * num_devices. Each batch of - inference task will be padded with respect to global_batch_size to - minimize latency. - context_len: Largest context length the model allows for each decode call. - This technically can be any large, but practically should set to the - context length the checkpoint was trained with. - horizon_len: Forecast horizon. - input_patch_len: Input patch len. - output_patch_len: Output patch len. How many timepoints is taken from a - single step of autoregressive decoding. Can be set as the training horizon - of the checkpoint. - mesh_shape: Shape of the data parallelism mesh. - mesh_name: Names of the data parallelism mesh. - model_p: Configuration of the TimesFM model deduced from the hparams. - """ - - def _logging(self, s): - if self._verbose: - print(s) - - def __init__( - self, - context_len: int, - horizon_len: int, - input_patch_len: int, - output_patch_len: int, - num_layers: int, - model_dims: int, - per_core_batch_size: int = 32, - backend: Literal["cpu", "gpu", "tpu"] = "cpu", - quantiles: Sequence[float] | None = None, - verbose: bool = True, - ) -> None: - """Initializes the TimesFM forecast API. - - Args: - context_len: Largest context length the model allows for each decode call. - This technically can be any large, but practically should set to the - context length the checkpoint was trained with. - horizon_len: Forecast horizon. - input_patch_len: Input patch len. - output_patch_len: Output patch len. How many timepoints is taken from a - single step of autoregressive decoding. Can be set as the training - horizon of the checkpoint. - num_layers: Number of transformer layers. - model_dims: Model dimension. - per_core_batch_size: Batch size on each core for data parallelism. - backend: One of "cpu", "gpu" or "tpu". - quantiles: list of output quantiles supported by the model. - verbose: Whether to print logging messages. - """ - self.per_core_batch_size = per_core_batch_size - self.backend = backend - self.num_devices = jax.local_device_count(self.backend) - self.global_batch_size = self.per_core_batch_size * self.num_devices - - self.context_len = context_len - self.horizon_len = horizon_len - self.input_patch_len = input_patch_len - self.output_patch_len = output_patch_len - - self.mesh_shape = [1, self.num_devices, 1] - self.mesh_name = ["replica", "data", "mdl"] - if quantiles is None: - quantiles = patched_decoder.DEFAULT_QUANTILES - - self.model_p = pax_fiddle.Config( - patched_decoder.PatchedTimeSeriesDecoder, - name="patched_decoder", - horizon_len=self.output_patch_len, - patch_len=input_patch_len, - model_dims=model_dims, - hidden_dims=model_dims, - residual_block_tpl=pax_fiddle.Config(patched_decoder.ResidualBlock), - quantiles=quantiles, - use_freq=True, - stacked_transformer_params_tpl=pax_fiddle.Config( - transformers.StackedTransformer, - num_heads=16, - num_layers=num_layers, - transformer_layer_params_tpl=pax_fiddle.Config( - transformers.Transformer, - ln_tpl=pax_fiddle.Config( - normalizations.RmsNorm, - ), - ), - ), - ) - - self._key1, self._key2 = jax.random.split(jax.random.PRNGKey(42)) - self._model = None - self._train_state = None - self._pmapped_decode = None - self._verbose = verbose - self._eval_context = base_layer.JaxContext.HParams(do_eval=True) - try: - multiprocessing.set_start_method("spawn") - except RuntimeError: - print("Multiprocessing context has already been set.") - - def _get_sample_inputs(self): - return { - "input_ts": jnp.zeros( - ( - self.per_core_batch_size, - self.context_len + self.output_patch_len, - ), - dtype=jnp.float32, - ), - "input_padding": jnp.zeros( - ( - self.per_core_batch_size, - self.context_len + self.output_patch_len, - ), - dtype=jnp.float32, - ), - "freq": jnp.zeros( - ( - self.per_core_batch_size, - 1, - ), - dtype=jnp.int32, - ), - } - - def load_from_checkpoint( - self, - checkpoint_path: Optional[str] = None, - repo_id: str = "google/timesfm-1.0-200m", - checkpoint_type: checkpoints.CheckpointType = checkpoints.CheckpointType.FLAX, - step: int | None = None, - ) -> None: - """Loads a checkpoint and compiles the decoder. - - Args: - checkpoint_path: Optional path to the checkpoint directory. - repo_id: Hugging Face Hub repo id. - checkpoint_type: type of PAX checkpoint - step: step of the checkpoint to load. If `None`, load latest checkpoint. - """ - # Download the checkpoint from Hugging Face Hub if not given - if checkpoint_path is None: - checkpoint_path = path.join(snapshot_download(repo_id), "checkpoints") - - # Initialize the model weights. - self._logging("Constructing model weights.") - start_time = time.time() - self._model = instantiate(self.model_p) - var_weight_hparams = self._model.abstract_init_with_metadata( - self._get_sample_inputs(), do_eval=True - ) - train_state_partition_specs = tasks_lib.create_state_partition_specs( - var_weight_hparams, - mesh_shape=self.mesh_shape, - mesh_axis_names=self.mesh_name, - discard_opt_states=True, - learners=None, - ) - train_state_local_shapes = tasks_lib.create_state_unpadded_shapes( - var_weight_hparams, - discard_opt_states=True, - learners=None, - ) - self._logging( - f"Constructed model weights in {time.time() - start_time:.2f} seconds." - ) - - # Load the model weights. - self._logging(f"Restoring checkpoint from {checkpoint_path}.") - start_time = time.time() - self._train_state = checkpoints.restore_checkpoint( - train_state_local_shapes, - checkpoint_dir=checkpoint_path, - checkpoint_type=checkpoint_type, - state_specs=train_state_partition_specs, - step=step, - ) - self._logging( - f"Restored checkpoint in {time.time() - start_time:.2f} seconds." - ) - - # Initialize and jit the decode fn. - def _decode(inputs): - assert self._model is not None - assert self._train_state is not None - return self._model.apply( - self._train_state.mdl_vars, - inputs, - horizon_len=self.horizon_len, - output_patch_len=self.output_patch_len, - max_len=self.context_len, - rngs={ - base_layer.PARAMS: self._key1, - base_layer.RANDOM: self._key2, - }, - method=self._model.decode, - ) - - self._logging("Jitting decoding.") - start_time = time.time() - self._pmapped_decode = jax.pmap( - _decode, - axis_name="batch", - devices=jax.devices(self.backend), - backend=self.backend, - axis_size=self.num_devices, - ) - with base_layer.JaxContext.new_context(hparams=self._eval_context): - _ = self._pmapped_decode( - NestedMap({ - "input_ts": jnp.zeros( - ( - self.num_devices, - self.per_core_batch_size, - self.context_len, - ), - dtype=jnp.float32, - ), - "input_padding": jnp.zeros( - ( - self.num_devices, - self.per_core_batch_size, - self.context_len + self.horizon_len, - ), - dtype=jnp.float32, - ), - "date_features": None, - "freq": jnp.zeros( - (self.num_devices, self.per_core_batch_size, 1), - dtype=jnp.int32, - ), - }) - ) - self._logging(f"Jitted decoding in {time.time() - start_time:.2f} seconds.") - - def _preprocess( - self, inputs: Sequence[np.array], freq: Sequence[int] - ) -> tuple[np.array, np.array, int]: - """Formats and pads raw inputs to feed into the model. - - This function both pads each time series to match the context length, and - pads the inputs to meet the SPMD shape requirement. - - Args: - inputs: A list of 1d JTensors. Each JTensor is the context time series of - a single forecast task. - freq: list of frequencies - - Returns: - A tuple of: - - the padded input time series to meet the model required context. - - the padding indicator. - - the number of padded examples for SPMD so that each core has the same - number (a multiple of `batch_size`) of examples. - """ - - input_ts, input_padding, inp_freq = [], [], [] - - pmap_pad = ( - (len(inputs) - 1) // self.global_batch_size + 1 - ) * self.global_batch_size - len(inputs) - - for i, ts in enumerate(inputs): - input_len = ts.shape[0] - padding = np.zeros(shape=(input_len + self.horizon_len,), dtype=float) - if input_len < self.context_len: - num_front_pad = self.context_len - input_len - ts = np.concatenate( - [np.zeros(shape=(num_front_pad,), dtype=float), ts], axis=0 - ) - padding = np.concatenate( - [np.ones(shape=(num_front_pad,), dtype=float), padding], axis=0 - ) - elif input_len > self.context_len: - ts = ts[-self.context_len :] - padding = padding[-(self.context_len + self.horizon_len) :] - - input_ts.append(ts) - input_padding.append(padding) - inp_freq.append(freq[i]) - - # Padding the remainder batch. - for _ in range(pmap_pad): - input_ts.append(input_ts[-1]) - input_padding.append(input_padding[-1]) - inp_freq.append(inp_freq[-1]) - - return ( - np.stack(input_ts, axis=0), - np.stack(input_padding, axis=0), - np.array(inp_freq).astype(np.int32).reshape(-1, 1), - pmap_pad, - ) - - def forecast( - self, - inputs: Sequence[Any], - freq: Sequence[int] | None = None, - window_size: int | None = None, - forecast_context_len: int | None = None, - ) -> tuple[JTensor, JTensor]: - """Forecasts on a list of time series. - - Args: - inputs: list of time series forecast contexts. Each context time series - should be in a format convertible to JTensor by `jnp.array`. - freq: frequency of each context time series. 0 for high frequency - (default), 1 for medium, and 2 for low. Notice this is different from - the `freq` required by `forecast_on_df`. - window_size: window size of trend + residual decomposition. If None then - we do not do decomposition. - forecast_context_len: optional max context length. - - Returns: - A tuple for JTensors: - - the mean forecast of size (# inputs, # forecast horizon), - - the full forecast (mean + quantiles) of size - (# inputs, # forecast horizon, 1 + # quantiles). - - Raises: - ValueError: If the checkpoint is not properly loaded. - """ - if not self._train_state or not self._model: - raise ValueError( - "Checkpoint not loaded. Call `load_from_checkpoint` before" - " `forecast`." - ) - if forecast_context_len is None: - forecast_context_len = self.context_len - inputs = [np.array(ts)[-forecast_context_len:] for ts in inputs] - inp_min = np.min([np.min(ts) for ts in inputs]) - - if window_size is not None: - new_inputs = [] - for ts in inputs: - new_inputs.extend(moving_average(ts, window_size)) - inputs = new_inputs - - if freq is None: - logging.info("No frequency provided via `freq`. Default to high (0).") - freq = [0] * len(inputs) - - input_ts, input_padding, inp_freq, pmap_pad = self._preprocess(inputs, freq) - with base_layer.JaxContext.new_context(hparams=self._eval_context): - mean_outputs = [] - full_outputs = [] - assert input_ts.shape[0] % self.global_batch_size == 0 - for i in range(input_ts.shape[0] // self.global_batch_size): - input_ts_in = jnp.array( - input_ts[ - i * self.global_batch_size : (i + 1) * self.global_batch_size - ] - ) - input_padding_in = jnp.array( - input_padding[ - i * self.global_batch_size : (i + 1) * self.global_batch_size - ], - ) - inp_freq_in = jnp.array( - inp_freq[ - i * self.global_batch_size : (i + 1) * self.global_batch_size, : - ], - dtype=jnp.int32, - ) - pmapped_inputs = NestedMap({ - "input_ts": es.jax_einshape( - "(db)...->db...", - input_ts_in, - d=self.num_devices, - ), - "input_padding": es.jax_einshape( - "(db)...->db...", - input_padding_in, - d=self.num_devices, - ), - "date_features": None, - "freq": es.jax_einshape( - "(db)...->db...", - inp_freq_in, - d=self.num_devices, - ), - }) - mean_output, full_output = self._pmapped_decode(pmapped_inputs) - mean_output = es.jax_einshape( - "db...->(db)...", mean_output, d=self.num_devices - ) - full_output = es.jax_einshape( - "db...->(db)...", full_output, d=self.num_devices - ) - mean_output = np.array(mean_output) - full_output = np.array(full_output) - mean_outputs.append(mean_output) - full_outputs.append(full_output) - - mean_outputs = np.concatenate(mean_outputs, axis=0) - full_outputs = np.concatenate(full_outputs, axis=0) - - if pmap_pad > 0: - mean_outputs = mean_outputs[:-pmap_pad, ...] - full_outputs = full_outputs[:-pmap_pad, ...] - - if window_size is not None: - mean_outputs = mean_outputs[0::2, ...] + mean_outputs[1::2, ...] - full_outputs = full_outputs[0::2, ...] + full_outputs[1::2, ...] - if inp_min >= 0: - mean_outputs = np.maximum(mean_outputs, 0.0) - full_outputs = np.maximum(full_outputs, 0.0) - return mean_outputs, full_outputs - - def forecast_on_df( - self, - inputs: pd.DataFrame, - freq: str, - forecast_context_len: int = 0, - value_name: str = "values", - model_name: str = "timesfm", - window_size: int | None = None, - num_jobs: int = 1, - ) -> pd.DataFrame: - """Forecasts on a list of time series. - - Args: - inputs: A pd.DataFrame of all time series. The dataframe should have a - `unique_id` column for identifying the time series, a `ds` column for - timestamps and a value column for the time series values. - freq: string valued `freq` of data. Notice this is different from the - `freq` required by `forecast`. See `freq_map` for allowed values. - forecast_context_len: If provided none zero, we take the last - `forecast_context_len` time-points from each series as the forecast - context instead of the `context_len` set by the model. - value_name: The name of the value column. - model_name: name of the model to be written into future df. - window_size: window size of trend + residual decomposition. If None then - we do not do decomposition. - num_jobs: number of parallel processes to use for dataframe processing. - - Returns: - Future forecasts dataframe. - """ - if not ( - "unique_id" in inputs.columns - and "ds" in inputs.columns - and value_name in inputs.columns - ): - raise ValueError( - f"DataFrame must have unique_id, ds and {value_name} columns." - ) - if not forecast_context_len: - forecast_context_len = self.context_len - logging.info("Preprocessing dataframe.") - df_sorted = inputs.sort_values(by=["unique_id", "ds"]) - new_inputs = [] - uids = [] - if num_jobs == 1: - print("Processing dataframe with single process.") - for key, group in df_sorted.groupby("unique_id"): - inp, uid = process_group( - key, - group, - value_name, - forecast_context_len, - ) - new_inputs.append(inp) - uids.append(uid) - else: - if num_jobs == -1: - num_jobs = multiprocessing.cpu_count() - print("Processing dataframe with multiple processes.") - with multiprocessing.Pool(processes=num_jobs) as pool: - results = pool.starmap( - process_group, - [ - (key, group, value_name, forecast_context_len) - for key, group in df_sorted.groupby("unique_id") - ], - ) - new_inputs, uids = zip(*results) - print("Finished preprocessing dataframe.") - freq_inps = [freq_map(freq)] * len(new_inputs) - _, full_forecast = self.forecast( - new_inputs, freq=freq_inps, window_size=window_size - ) - print("Finished forecasting.") - fcst_df = make_future_dataframe( - uids=uids, - last_times=df_sorted.groupby("unique_id")["ds"].tail(1), - h=self.horizon_len, - freq=freq, - ) - fcst_df[model_name] = full_forecast[:, 0 : self.horizon_len, 0].reshape( - -1, 1 - ) - - if self._model.quantiles is not None: - for i, q in enumerate(self._model.quantiles): - q_col = f"{model_name}-q-{q}" - fcst_df[q_col] = full_forecast[:, 0 : self.horizon_len, 1 + i].reshape( - -1, 1 - ) - if q == 0.5: - fcst_df[model_name] = fcst_df[q_col] - logging.info("Finished creating output dataframe.") - return fcst_df