From 6e77e0ca98f394694aec21193def48de5af41a48 Mon Sep 17 00:00:00 2001 From: Rajat Sen Date: Mon, 8 Jul 2024 18:31:35 +0000 Subject: [PATCH] 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