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/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/__init__.py b/src/timesfm/__init__.py similarity index 78% rename from __init__.py rename to src/timesfm/__init__.py index 8275b16..866e848 100644 --- a/__init__.py +++ b/src/timesfm/__init__.py @@ -14,8 +14,4 @@ """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 +from .timesfm import TimesFm, freq_map diff --git a/experiments/long_horizon_benchmarks/data_loader.py b/src/timesfm/data_loader.py similarity index 100% rename from experiments/long_horizon_benchmarks/data_loader.py rename to src/timesfm/data_loader.py 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/experiments/long_horizon_benchmarks/time_features.py b/src/timesfm/time_features.py similarity index 100% rename from experiments/long_horizon_benchmarks/time_features.py rename to src/timesfm/time_features.py diff --git a/src/timesfm.py b/src/timesfm/timesfm.py similarity index 99% rename from src/timesfm.py rename to src/timesfm/timesfm.py index 7dc9d0a..5dc7611 100644 --- a/src/timesfm.py +++ b/src/timesfm/timesfm.py @@ -35,7 +35,7 @@ from praxis import py_utils from praxis import pytypes from praxis.layers import normalizations from praxis.layers import transformers -import patched_decoder +from . import patched_decoder from utilsforecast.processing import make_future_dataframe instantiate = base_hyperparams.instantiate @@ -277,7 +277,10 @@ class TimesFm: 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