From 739d74e19911e5baafc9c5d23bf3bcd154ce9c24 Mon Sep 17 00:00:00 2001 From: Rajat Sen Date: Mon, 8 Jul 2024 18:37:00 +0000 Subject: [PATCH] deleting old files --- __init__.py | 21 - .../long_horizon_benchmarks/data_loader.py | 261 -------- .../long_horizon_benchmarks/time_features.py | 215 ------- src/patched_decoder.py | 461 -------------- src/timesfm.py | 602 ------------------ 5 files changed, 1560 deletions(-) delete mode 100644 __init__.py delete mode 100644 experiments/long_horizon_benchmarks/data_loader.py delete mode 100644 experiments/long_horizon_benchmarks/time_features.py delete mode 100644 src/patched_decoder.py delete mode 100644 src/timesfm.py diff --git a/__init__.py b/__init__.py deleted file mode 100644 index 8275b16..0000000 --- a/__init__.py +++ /dev/null @@ -1,21 +0,0 @@ -# Copyright 2024 Google LLC -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""TimesFM init file.""" - -from __future__ import absolute_import - -from .src.patched_decoder import PatchedTimeSeriesDecoder -from .src.timesfm import TimesFm -from .src.timesfm import freq_map diff --git a/experiments/long_horizon_benchmarks/data_loader.py b/experiments/long_horizon_benchmarks/data_loader.py deleted file mode 100644 index eeace7b..0000000 --- a/experiments/long_horizon_benchmarks/data_loader.py +++ /dev/null @@ -1,261 +0,0 @@ -# Copyright 2024 The Google Research Authors. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""TF dataloaders for general timeseries datasets. - -The expected input format is csv file with a datetime index. -""" - - -from absl import logging -import numpy as np -import pandas as pd -from sklearn.preprocessing import StandardScaler -import tensorflow as tf -from . import time_features - - -class TimeSeriesdata(object): - """Data loader class.""" - - def __init__( - self, - data_path, - datetime_col, - num_cov_cols, - cat_cov_cols, - ts_cols, - train_range, - val_range, - test_range, - hist_len, - pred_len, - batch_size, - freq='H', - normalize=True, - epoch_len=None, - holiday=False, - permute=True, - ): - """Initialize objects. - - Args: - data_path: path to csv file - datetime_col: column name for datetime col - num_cov_cols: list of numerical global covariates - cat_cov_cols: list of categorical global covariates - ts_cols: columns corresponding to ts - train_range: tuple of train ranges - val_range: tuple of validation ranges - test_range: tuple of test ranges - hist_len: historical context - pred_len: prediction length - batch_size: batch size (number of ts in a batch) - freq: freq of original data - normalize: std. normalize data or not - epoch_len: num iters in an epoch - holiday: use holiday features or not - permute: permute ts in train batches or not - - Returns: - None - """ - self.data_df = pd.read_csv(open(data_path, 'r')) - if not num_cov_cols: - self.data_df['ncol'] = np.zeros(self.data_df.shape[0]) - num_cov_cols = ['ncol'] - if not cat_cov_cols: - self.data_df['ccol'] = np.zeros(self.data_df.shape[0]) - cat_cov_cols = ['ccol'] - self.data_df.fillna(0, inplace=True) - self.data_df.set_index( - pd.DatetimeIndex(self.data_df[datetime_col]), inplace=True - ) - self.num_cov_cols = num_cov_cols - self.cat_cov_cols = cat_cov_cols - self.ts_cols = ts_cols - self.train_range = train_range - self.val_range = val_range - self.test_range = test_range - data_df_idx = self.data_df.index - date_index = data_df_idx.union( - pd.date_range( - data_df_idx[-1] + pd.Timedelta(1, freq=freq), - periods=pred_len + 1, - freq=freq, - ) - ) - self.time_df = time_features.TimeCovariates( - date_index, holiday=holiday - ).get_covariates() - self.hist_len = hist_len - self.pred_len = pred_len - self.batch_size = batch_size - self.freq = freq - self.normalize = normalize - self.data_mat = self.data_df[self.ts_cols].to_numpy().transpose() - self.data_mat = self.data_mat[:, 0 : self.test_range[1]] - self.time_mat = self.time_df.to_numpy().transpose() - self.num_feat_mat = self.data_df[num_cov_cols].to_numpy().transpose() - self.cat_feat_mat, self.cat_sizes = self._get_cat_cols(cat_cov_cols) - self.normalize = normalize - if normalize: - self._normalize_data() - logging.info( - 'Data Shapes: %s, %s, %s, %s', - self.data_mat.shape, - self.time_mat.shape, - self.num_feat_mat.shape, - self.cat_feat_mat.shape, - ) - self.epoch_len = epoch_len - self.permute = permute - - def _get_cat_cols(self, cat_cov_cols): - """Get categorical columns.""" - cat_vars = [] - cat_sizes = [] - for col in cat_cov_cols: - dct = {x: i for i, x in enumerate(self.data_df[col].unique())} - cat_sizes.append(len(dct)) - mapped = self.data_df[col].map(lambda x: dct[x]).to_numpy().transpose() # pylint: disable=cell-var-from-loop - cat_vars.append(mapped) - return np.vstack(cat_vars), cat_sizes - - def _normalize_data(self): - self.scaler = StandardScaler() - train_mat = self.data_mat[:, self.train_range[0] : self.train_range[1]] - self.scaler = self.scaler.fit(train_mat.transpose()) - self.data_mat = self.scaler.transform(self.data_mat.transpose()).transpose() - - def train_gen(self): - """Generator for training data.""" - num_ts = len(self.ts_cols) - perm = np.arange( - self.train_range[0] + self.hist_len, - self.train_range[1] - self.pred_len, - ) - perm = np.random.permutation(perm) - hist_len = self.hist_len - logging.info('Hist len: %s', hist_len) - if not self.epoch_len: - epoch_len = len(perm) - else: - epoch_len = self.epoch_len - for idx in perm[0:epoch_len]: - for _ in range(num_ts // self.batch_size + 1): - if self.permute: - tsidx = np.random.choice(num_ts, size=self.batch_size, replace=False) - else: - tsidx = np.arange(num_ts) - dtimes = np.arange(idx - hist_len, idx + self.pred_len) - ( - bts_train, - bts_pred, - bfeats_train, - bfeats_pred, - bcf_train, - bcf_pred, - ) = self._get_features_and_ts(dtimes, tsidx, hist_len) - - all_data = [ - bts_train, - bfeats_train, - bcf_train, - bts_pred, - bfeats_pred, - bcf_pred, - tsidx, - ] - yield tuple(all_data) - - def test_val_gen(self, mode='val', shift=1): - """Generator for validation/test data.""" - if mode == 'val': - start = self.val_range[0] - end = self.val_range[1] - self.pred_len + 1 - elif mode == 'test': - start = self.test_range[0] - end = self.test_range[1] - self.pred_len + 1 - else: - raise NotImplementedError('Eval mode not implemented') - num_ts = len(self.ts_cols) - hist_len = self.hist_len - logging.info('Hist len: %s', hist_len) - perm = np.arange(start, end) - if self.epoch_len: - epoch_len = self.epoch_len - else: - epoch_len = len(perm) - for i in range(0, epoch_len, shift): - idx = perm[i] - for batch_idx in range(0, num_ts, self.batch_size): - tsidx = np.arange(batch_idx, min(batch_idx + self.batch_size, num_ts)) - dtimes = np.arange(idx - hist_len, idx + self.pred_len) - ( - bts_train, - bts_pred, - bfeats_train, - bfeats_pred, - bcf_train, - bcf_pred, - ) = self._get_features_and_ts(dtimes, tsidx, hist_len) - all_data = [ - bts_train, - bfeats_train, - bcf_train, - bts_pred, - bfeats_pred, - bcf_pred, - tsidx, - ] - yield tuple(all_data) - - def _get_features_and_ts(self, dtimes, tsidx, hist_len=None): - """Get features and ts in specified windows.""" - if hist_len is None: - hist_len = self.hist_len - data_times = dtimes[dtimes < self.data_mat.shape[1]] - bdata = self.data_mat[:, data_times] - bts = bdata[tsidx, :] - bnf = self.num_feat_mat[:, data_times] - bcf = self.cat_feat_mat[:, data_times] - btf = self.time_mat[:, dtimes] - if bnf.shape[1] < btf.shape[1]: - rem_len = btf.shape[1] - bnf.shape[1] - rem_rep = np.repeat(bnf[:, [-1]], repeats=rem_len) - rem_rep_cat = np.repeat(bcf[:, [-1]], repeats=rem_len) - bnf = np.hstack([bnf, rem_rep.reshape(bnf.shape[0], -1)]) - bcf = np.hstack([bcf, rem_rep_cat.reshape(bcf.shape[0], -1)]) - bfeats = np.vstack([btf, bnf]) - bts_train = bts[:, 0:hist_len] - bts_pred = bts[:, hist_len:] - bfeats_train = bfeats[:, 0:hist_len] - bfeats_pred = bfeats[:, hist_len:] - bcf_train = bcf[:, 0:hist_len] - bcf_pred = bcf[:, hist_len:] - return bts_train, bts_pred, bfeats_train, bfeats_pred, bcf_train, bcf_pred - - def tf_dataset(self, mode='train', shift=1): - """Tensorflow Dataset.""" - if mode == 'train': - gen_fn = self.train_gen - else: - gen_fn = lambda: self.test_val_gen(mode, shift) - output_types = tuple( - [tf.float32] * 2 + [tf.int32] + [tf.float32] * 2 + [tf.int32] * 2 - ) - dataset = tf.data.Dataset.from_generator(gen_fn, output_types) - dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE) - return dataset diff --git a/experiments/long_horizon_benchmarks/time_features.py b/experiments/long_horizon_benchmarks/time_features.py deleted file mode 100644 index 0bd90a9..0000000 --- a/experiments/long_horizon_benchmarks/time_features.py +++ /dev/null @@ -1,215 +0,0 @@ -# Copyright 2024 The Google Research Authors. -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Directory to extract time covariates. - -Extract time covariates from datetime. -""" - -import numpy as np -import pandas as pd -from pandas.tseries.holiday import EasterMonday -from pandas.tseries.holiday import GoodFriday -from pandas.tseries.holiday import Holiday -from pandas.tseries.holiday import SU -from pandas.tseries.holiday import TH -from pandas.tseries.holiday import USColumbusDay -from pandas.tseries.holiday import USLaborDay -from pandas.tseries.holiday import USMartinLutherKingJr -from pandas.tseries.holiday import USMemorialDay -from pandas.tseries.holiday import USPresidentsDay -from pandas.tseries.holiday import USThanksgivingDay -from pandas.tseries.offsets import DateOffset -from pandas.tseries.offsets import Day -from pandas.tseries.offsets import Easter -from sklearn.preprocessing import StandardScaler -from tqdm import tqdm - - -# This is 183 to cover half a year (in both directions), also for leap years -# + 17 as Eastern can be between March, 22 - April, 25 -MAX_WINDOW = 183 + 17 - - -def _distance_to_holiday(holiday): - """Return distance to given holiday.""" - - def _distance_to_day(index): - holiday_date = holiday.dates( - index - pd.Timedelta(days=MAX_WINDOW), - index + pd.Timedelta(days=MAX_WINDOW), - ) - assert ( - len(holiday_date) != 0 # pylint: disable=g-explicit-length-test - ), f"No closest holiday for the date index {index} found." - # It sometimes returns two dates if it is exactly half a year after the - # holiday. In this case, the smaller distance (182 days) is returned. - return (index - holiday_date[0]).days - - return _distance_to_day - - -EasterSunday = Holiday( - "Easter Sunday", month=1, day=1, offset=[Easter(), Day(0)] -) -NewYearsDay = Holiday("New Years Day", month=1, day=1) -SuperBowl = Holiday( - "Superbowl", month=2, day=1, offset=DateOffset(weekday=SU(1)) -) -MothersDay = Holiday( - "Mothers Day", month=5, day=1, offset=DateOffset(weekday=SU(2)) -) -IndependenceDay = Holiday("Independence Day", month=7, day=4) -ChristmasEve = Holiday("Christmas", month=12, day=24) -ChristmasDay = Holiday("Christmas", month=12, day=25) -NewYearsEve = Holiday("New Years Eve", month=12, day=31) -BlackFriday = Holiday( - "Black Friday", - month=11, - day=1, - offset=[pd.DateOffset(weekday=TH(4)), Day(1)], -) -CyberMonday = Holiday( - "Cyber Monday", - month=11, - day=1, - offset=[pd.DateOffset(weekday=TH(4)), Day(4)], -) - -HOLIDAYS = [ - EasterMonday, - GoodFriday, - USColumbusDay, - USLaborDay, - USMartinLutherKingJr, - USMemorialDay, - USPresidentsDay, - USThanksgivingDay, - EasterSunday, - NewYearsDay, - SuperBowl, - MothersDay, - IndependenceDay, - ChristmasEve, - ChristmasDay, - NewYearsEve, - BlackFriday, - CyberMonday, -] - - -class TimeCovariates(object): - """Extract all time covariates except for holidays.""" - - def __init__( - self, - datetimes, - normalized=True, - holiday=False, - ): - """Init function. - - Args: - datetimes: pandas DatetimeIndex (lowest granularity supported is min) - normalized: whether to normalize features or not - holiday: fetch holiday features or not - - Returns: - None - """ - self.normalized = normalized - self.dti = datetimes - self.holiday = holiday - - def _minute_of_hour(self): - minutes = np.array(self.dti.minute, dtype=np.float32) - if self.normalized: - minutes = minutes / 59.0 - 0.5 - return minutes - - def _hour_of_day(self): - hours = np.array(self.dti.hour, dtype=np.float32) - if self.normalized: - hours = hours / 23.0 - 0.5 - return hours - - def _day_of_week(self): - day_week = np.array(self.dti.dayofweek, dtype=np.float32) - if self.normalized: - day_week = day_week / 6.0 - 0.5 - return day_week - - def _day_of_month(self): - day_month = np.array(self.dti.day, dtype=np.float32) - if self.normalized: - day_month = day_month / 30.0 - 0.5 - return day_month - - def _day_of_year(self): - day_year = np.array(self.dti.dayofyear, dtype=np.float32) - if self.normalized: - day_year = day_year / 364.0 - 0.5 - return day_year - - def _month_of_year(self): - month_year = np.array(self.dti.month, dtype=np.float32) - if self.normalized: - month_year = month_year / 11.0 - 0.5 - return month_year - - def _week_of_year(self): - week_year = np.array(self.dti.strftime("%U").astype(int), dtype=np.float32) - if self.normalized: - week_year = week_year / 51.0 - 0.5 - return week_year - - def _get_holidays(self): - dti_series = self.dti.to_series() - hol_variates = np.vstack([ - dti_series.apply(_distance_to_holiday(h)).values for h in tqdm(HOLIDAYS) - ]) - # hol_variates is (num_holiday, num_time_steps), the normalization should be - # performed in the num_time_steps dimension. - return StandardScaler().fit_transform(hol_variates.T).T - - def get_covariates(self): - """Get all time covariates.""" - moh = self._minute_of_hour().reshape(1, -1) - hod = self._hour_of_day().reshape(1, -1) - dom = self._day_of_month().reshape(1, -1) - dow = self._day_of_week().reshape(1, -1) - doy = self._day_of_year().reshape(1, -1) - moy = self._month_of_year().reshape(1, -1) - woy = self._week_of_year().reshape(1, -1) - - all_covs = [ - moh, - hod, - dom, - dow, - doy, - moy, - woy, - ] - columns = ["moh", "hod", "dom", "dow", "doy", "moy", "woy"] - if self.holiday: - hol_covs = self._get_holidays() - all_covs.append(hol_covs) - columns += [f"hol_{i}" for i in range(len(HOLIDAYS))] - - return pd.DataFrame( - data=np.vstack(all_covs).transpose(), - columns=columns, - index=self.dti, - ) diff --git a/src/patched_decoder.py b/src/patched_decoder.py deleted file mode 100644 index b0decf8..0000000 --- a/src/patched_decoder.py +++ /dev/null @@ -1,461 +0,0 @@ -# Copyright 2024 Google LLC -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Pax ML model for patched time-series decoder. - -The file implements Residual MLPs, Patched Decoder layers and PAX ML models. -""" - -import dataclasses -from typing import Optional, Tuple - -import einshape as es -from jax import lax -import jax.numpy as jnp -from praxis import base_layer -from praxis import layers -from praxis import pax_fiddle -from praxis import py_utils -from praxis import pytypes -from praxis.layers import activations -from praxis.layers import embedding_softmax -from praxis.layers import linears -from praxis.layers import normalizations -from praxis.layers import stochastics -from praxis.layers import transformers - - -# PAX shortcuts -NestedMap = py_utils.NestedMap -JTensor = pytypes.JTensor - -LayerTpl = pax_fiddle.Config[base_layer.BaseLayer] -template_field = base_layer.template_field - - -PAD_VAL = 1123581321.0 -DEFAULT_QUANTILES = [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9] - -# NestedMap keys -_INPUT_TS = "input_ts" -_INPUT_PADDING = "input_padding" -_OUTPUT_TS = "output_ts" -_FREQ = "freq" -_OUTPUT_TOKENS = "output_tokens" -_STATS = "stats" - - -# Small numerical value. -_TOLERANCE = 1e-7 - - -def _shift_padded_seq(mask: JTensor, seq: JTensor) -> JTensor: - """Shifts rows of seq based on the first 0 in each row of the mask.""" - num = seq.shape[1] - - # Find the index of the first 0 in each row of the mask - first_zero_idx = jnp.argmin(mask, axis=1) - - # Create a range array for indexing - idx_range = jnp.arange(num) - - def shift_row(carry, x): - seq_row, shift = x - shifted_idx = (idx_range - shift) % num - shifted_row = seq_row[shifted_idx] - return carry, shifted_row - - # Use lax.scan to shift each row of seq based on the corresponding - # first_zero_idx. - _, shifted_seq = lax.scan(shift_row, None, (seq, first_zero_idx)) - - return shifted_seq - - -class ResidualBlock(base_layer.BaseLayer): - """Simple feedforward block with residual connection. - - Attributes: - input_dims: input dimension. - hidden_dims: hidden dimension. - output_dims: output dimension. - dropout_prob: dropout probability. - layer_norm: whether to use layer norm or not. - dropout_tpl: config for dropout. - ln_tpl: config for layer norm. - act_tpl: config for activation in hidden layer. - """ - - input_dims: int = 0 - hidden_dims: int = 0 - output_dims: int = 0 - dropout_prob: float = 0.0 - layer_norm: bool = False - dropout_tpl: LayerTpl = template_field(stochastics.Dropout) - ln_tpl: LayerTpl = template_field(normalizations.LayerNorm) - act_tpl: LayerTpl = template_field(activations.Swish) - - def setup(self): - lnorm_tpl = self.ln_tpl.clone() - lnorm_tpl.dim = self.output_dims - self.create_child("ln_layer", lnorm_tpl) - - dropout_tpl = self.dropout_tpl.clone() - dropout_tpl.keep_prob = 1.0 - self.dropout_prob - self.create_child("dropout", dropout_tpl) - - self.create_child( - "hidden_layer", - pax_fiddle.Config( - linears.FeedForward, - input_dims=self.input_dims, - output_dims=self.hidden_dims, - activation_tpl=self.act_tpl.clone(), - ), - ) - - self.create_child( - "output_layer", - pax_fiddle.Config( - linears.FeedForward, - input_dims=self.hidden_dims, - output_dims=self.output_dims, - activation_tpl=pax_fiddle.Config(activations.Identity), - ), - ) - - self.create_child( - "residual_layer", - pax_fiddle.Config( - linears.FeedForward, - input_dims=self.input_dims, - output_dims=self.output_dims, - activation_tpl=pax_fiddle.Config(activations.Identity), - ), - ) - - def __call__(self, inputs: JTensor) -> JTensor: - hidden = self.hidden_layer(inputs) - output = self.output_layer(hidden) - output = self.dropout(output) - residual = self.residual_layer(inputs) - if self.layer_norm: - return self.ln_layer(output + residual) - else: - return output + residual - - -def _masked_mean_std( - inputs: JTensor, padding: JTensor -) -> Tuple[JTensor, JTensor]: - """Calculates mean and standard deviation of arr across axis 1. - - It should exclude values where pad is 1. - - Args: - inputs: A JAX array of shape [b, n, p]. - padding: A JAX array of shape [b, n, p] with values 0 or 1. - - Returns: - A tuple containing the mean and standard deviation of arr. We return the - statistics of the first patch with more than three non-padded values. - """ - # Selecting the first pad with more than 3 unpadded values. - pad_sum = jnp.sum(1 - padding, axis=2) - - def _get_patch_index(arr: JTensor): - indices = jnp.argmax(arr >= 3, axis=1) - row_sum = (arr >= 3).sum(axis=1) - return jnp.where(row_sum == 0, arr.shape[1] - 1, indices) - - patch_indices = _get_patch_index(pad_sum) - bidxs = jnp.arange(inputs.shape[0]) - - arr = inputs[bidxs, patch_indices, :] - pad = padding[bidxs, patch_indices, :] - - # Create a mask where P is 0 - mask = 1 - pad - - # Calculate the number of valid elements - num_valid_elements = jnp.sum(mask, axis=1) - - num_valid_elements = jnp.where(num_valid_elements == 0, 1, num_valid_elements) - - # Calculate the masked sum and squared sum of M - masked_sum = jnp.sum(arr * mask, axis=1) - masked_squared_sum = jnp.sum((arr * mask) ** 2, axis=1) - - # Calculate the masked mean and standard deviation - masked_mean = masked_sum / num_valid_elements - masked_var = masked_squared_sum / num_valid_elements - masked_mean**2 - masked_var = jnp.where(masked_var < 0.0, 0.0, masked_var) - masked_std = jnp.sqrt(masked_var) - - return masked_mean, masked_std - - -def _create_quantiles() -> list[float]: - """Returns the quantiles for forecasting.""" - return DEFAULT_QUANTILES - - -class PatchedTimeSeriesDecoder(base_layer.BaseLayer): - """Patch decoder layer for time-series foundation model. - - Attributes: - patch_len: length of input patches. - horizon_len: length of output patches. Referred to as `output_patch_len` - during inference. - model_dims: model dimension of stacked transformer layer. - hidden_dims: hidden dimensions in fully connected layers. - quantiles: list of quantiles for non prob model. - residual_block_tpl: config for residual block. - stacked_transformer_params_tpl: config for stacked transformer. - use_freq: whether to use frequency encoding. - - In all of what followed, except specified otherwise, B is batch size, T is - sequence length of time-series. N is the number of input patches that can be - obtained from T. P is the input patch length and H is the horizon length. Q is - number of output logits. D is model dimension. - """ - - patch_len: int = 0 - horizon_len: int = 0 - model_dims: int = 0 - hidden_dims: int = 0 - quantiles: list[float] = dataclasses.field(default_factory=_create_quantiles) - residual_block_tpl: LayerTpl = template_field(ResidualBlock) - stacked_transformer_params_tpl: LayerTpl = template_field( - transformers.StackedTransformer - ) - use_freq: bool = True - - def setup(self) -> None: - """Construct the model.""" - num_outputs = len(self.quantiles) + 1 - - stl = self.stacked_transformer_params_tpl.clone() - stl.model_dims = self.model_dims - stl.hidden_dims = self.hidden_dims - stl.mask_self_attention = True - - self.create_child("stacked_transformer_layer", stl) - - input_resl = self.residual_block_tpl.clone() - ff_in_dims = 2 * self.patch_len - input_resl.input_dims = ff_in_dims - input_resl.hidden_dims = self.hidden_dims - input_resl.output_dims = self.model_dims - self.create_child( - "input_ff_layer", - input_resl, - ) - - horizon_resl = self.residual_block_tpl.clone() - horizon_resl.input_dims = self.model_dims - horizon_resl.hidden_dims = self.hidden_dims - horizon_resl.output_dims = self.horizon_len * num_outputs - self.create_child( - "horizon_ff_layer", - horizon_resl, - ) - - self.create_child( - "position_emb", - pax_fiddle.Config( - layers.PositionalEmbedding, embedding_dims=self.model_dims - ), - ) - - if self.use_freq: - self.create_child( - "freq_emb", - pax_fiddle.Config( - embedding_softmax.Embedding, - num_classes=3, - input_dims=self.model_dims, - ), - ) - - def transform_decode_state( - self, transform_fn: base_layer.DecodeStateTransformFn - ) -> None: - """Transforms all decode state variables based on transform_fn.""" - self.stacked_transformer_layer.transform_decode_state(transform_fn) - - def _forward_transform( - self, inputs: JTensor, patched_pads: JTensor - ) -> Tuple[JTensor, Tuple[JTensor, JTensor]]: - """Input is of shape [B, N, P].""" - mu, sigma = _masked_mean_std(inputs, patched_pads) - sigma = jnp.where(sigma < _TOLERANCE, 1.0, sigma) - # Normalize each patch. - outputs = (inputs - mu[:, None, None]) / sigma[:, None, None] - outputs = jnp.where( - jnp.abs(inputs - PAD_VAL) < _TOLERANCE, PAD_VAL, outputs - ) - return outputs, (mu, sigma) - - def _reverse_transform( - self, outputs: JTensor, stats: Tuple[JTensor, JTensor] - ) -> JTensor: - """Output is of shape [B, N, P, Q].""" - mu, sigma = stats - return outputs * sigma[:, None, None, None] + mu[:, None, None, None] - - def _preprocess_input( - self, - input_ts: JTensor, - input_padding: JTensor, - pos_emb: Optional[JTensor] = None, - ) -> Tuple[JTensor, JTensor, Optional[Tuple[JTensor, JTensor]], JTensor]: - """Preprocess input for stacked transformer.""" - # Reshape into patches. - patched_inputs = es.jax_einshape("b(np)->bnp", input_ts, p=self.patch_len) - input_padding = jnp.where( - jnp.abs(input_ts - PAD_VAL) < _TOLERANCE, 1, input_padding - ) - patched_pads = es.jax_einshape( - "b(np)->bnp", input_padding, p=self.patch_len - ) - patched_inputs, stats = self._forward_transform( - patched_inputs, patched_pads - ) - # B x N x D - patched_inputs = patched_inputs * (1.0 - patched_pads) - concat_inputs = jnp.concatenate([patched_inputs, patched_pads], axis=-1) - model_input = self.input_ff_layer(concat_inputs) - # A patch should not be padded even if there is at least one zero. - patched_padding = jnp.min(patched_pads, axis=-1) - - if pos_emb is None: - position_emb = self.position_emb(seq_length=model_input.shape[1]) - else: - position_emb = pos_emb - if self.do_eval: - if position_emb.shape[0] != model_input.shape[0]: - position_emb = jnp.repeat(position_emb, model_input.shape[0], axis=0) - position_emb = _shift_padded_seq(patched_padding, position_emb) - model_input += position_emb - - return model_input, patched_padding, stats, patched_inputs - - def _postprocess_output( - self, - model_output: JTensor, - num_outputs: int, - stats: Tuple[JTensor, JTensor], - ) -> JTensor: - """Postprocess output of stacked transformer.""" - # B x N x (H.Q) - output_ts = self.horizon_ff_layer(model_output) - output_ts = es.jax_einshape( - "bn(hq)->bnhq", output_ts, q=num_outputs, h=self.horizon_len - ) - return self._reverse_transform(output_ts, stats) - - def __call__(self, inputs: NestedMap) -> NestedMap: - """PatchTST call. - - Args: - inputs: A NestedMap containing (1) input_ts: input sequence of shape [B, - T] where T must be multiple of patch_length; (2) input_padding: that - contains padding map. - - Returns: - A nested map with two keys: - (1) 'output_tokens' of shape [B, N, D]. - (2) 'output_ts' of shape [B, N, H, Q] - (3) 'stats' a Tuple of statistics for renormalization. - """ - input_ts, input_padding = inputs[_INPUT_TS], inputs[_INPUT_PADDING] - num_outputs = len(self.quantiles) + 1 - model_input, patched_padding, stats, _ = self._preprocess_input( - input_ts=input_ts, - input_padding=input_padding, - ) - if self.use_freq: - freq = inputs[_FREQ].astype(jnp.int32) - f_emb = self.freq_emb(freq) # B x 1 x D - f_emb = jnp.repeat(f_emb, model_input.shape[1], axis=1) - model_input += f_emb - model_output = self.stacked_transformer_layer(model_input, patched_padding) - - output_ts = self._postprocess_output(model_output, num_outputs, stats) - return NestedMap( - {_OUTPUT_TOKENS: model_output, _OUTPUT_TS: output_ts, _STATS: stats} - ) - - def decode( - self, - inputs: NestedMap, - horizon_len: int, - output_patch_len: Optional[int] = None, - max_len: int = 512, - ) -> tuple[JTensor, JTensor]: - """Auto-regressive decoding without caching. - - Args: - inputs: input time-series and paddings. Time-series shape B x C, padding - shape shape B x (C + H) where H is the prediction length. - horizon_len: prediction length. - output_patch_len: output length to be fetched from one step of - auto-regressive decoding. - max_len: maximum training context length. - - Returns: - Tuple of two forecasting results: - - Point (mean) output predictions as a tensor with shape B x H. - - Full predictions (mean and quantiles) as a tensor with shape - B x H x (1 + # quantiles). - """ - final_out = inputs[_INPUT_TS] - inp_time_len = final_out.shape[1] - paddings = inputs[_INPUT_PADDING] - if self.use_freq: - freq = inputs[_FREQ].astype(jnp.int32) - else: - freq = jnp.zeros([final_out.shape[0], 1], dtype=jnp.int32) - full_outputs = [] - if paddings.shape[1] != final_out.shape[1] + horizon_len: - raise ValueError( - "Length of paddings must match length of input + horizon_len:" - f" {paddings.shape[1]} != {final_out.shape[1]} + {horizon_len}" - ) - if output_patch_len is None: - output_patch_len = self.horizon_len - num_decode_patches = ( - horizon_len + output_patch_len - 1 - ) // output_patch_len - for _ in range(num_decode_patches): - current_padding = paddings[:, 0 : final_out.shape[1]] - input_ts = final_out[:, -max_len:] - input_padding = current_padding[:, -max_len:] - model_input = NestedMap( - input_ts=input_ts, - input_padding=input_padding, - freq=freq, - ) - fprop_outputs = self(model_input)[_OUTPUT_TS] - # (full batch, last patch, output_patch_len, index of mean forecast = 0) - new_ts = fprop_outputs[:, -1, :output_patch_len, 0] - # (full batch, last patch, output_patch_len, all output indices) - full_outputs.append(fprop_outputs[:, -1, :output_patch_len, :]) - final_out = jnp.concatenate([final_out, new_ts], axis=-1) - - return ( - final_out[:, inp_time_len : inp_time_len + horizon_len], - jnp.concatenate(full_outputs, axis=1)[:, 0:horizon_len, :], - ) diff --git a/src/timesfm.py b/src/timesfm.py deleted file mode 100644 index 7dc9d0a..0000000 --- a/src/timesfm.py +++ /dev/null @@ -1,602 +0,0 @@ -# Copyright 2024 Google LLC -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""TimesFM forecast API for inference.""" - -import logging -import multiprocessing -from os import path -import time -from typing import Any, Literal, Optional, Sequence - -import einshape as es -import jax -import jax.numpy as jnp -import numpy as np -import pandas as pd -from huggingface_hub import snapshot_download -from paxml import checkpoints -from paxml import tasks_lib -from praxis import base_hyperparams -from praxis import base_layer -from praxis import pax_fiddle -from praxis import py_utils -from praxis import pytypes -from praxis.layers import normalizations -from praxis.layers import transformers -import patched_decoder -from utilsforecast.processing import make_future_dataframe - -instantiate = base_hyperparams.instantiate -NestedMap = py_utils.NestedMap -JTensor = pytypes.JTensor - - -def process_group(key, group, value_name, forecast_context_len): - group = group.tail(forecast_context_len) - return np.array(group[value_name], dtype=np.float32), key - - -def moving_average(arr, window_size): - """Calculates the moving average using NumPy's convolution function.""" - # Pad with zeros to handle initial window positions - arr_padded = np.pad(arr, (window_size - 1, 0), "constant") - smoothed_arr = ( - np.convolve(arr_padded, np.ones(window_size), "valid") / window_size - ) - return [smoothed_arr, arr - smoothed_arr] - - -def freq_map(freq: str): - """Returns the frequency map for the given frequency string.""" - freq = str.upper(freq) - if ( - freq.endswith("H") - or freq.endswith("T") - or freq.endswith("MIN") - or freq.endswith("D") - or freq.endswith("B") - or freq.endswith("U") - ): - return 0 - elif freq.endswith(("W", "M", "MS")): - return 1 - elif freq.endswith("Y") or freq.endswith("Q"): - return 2 - else: - raise ValueError(f"Invalid frequency: {freq}") - - -class TimesFm: - """TimesFM forecast API for inference. - - This class is the scaffolding for calling TimesFM forecast. To properly use: - 1. Create an instance with the correct hyperparameters of a TimesFM model. - 2. Call `load_from_checkpoint` to load a compatible checkpoint. - 3. Call `forecast` for inference. - - Given the model size, this API does not shard the model weights for SPMD. All - parallelism happens on the data dimension. - - Compilation happens during the first time `forecast` is called and uses the - `per_core_batch_size` to set and freeze the input signature. Subsequent calls - to `forecast` reflect the actual inference latency. - - Attributes: - per_core_batch_size: Batch size on each core for data parallelism. - backend: One of "cpu", "gpu" or "tpu". - num_devices: Number of cores provided the backend. - global_batch_size: per_core_batch_size * num_devices. Each batch of - inference task will be padded with respect to global_batch_size to - minimize latency. - context_len: Largest context length the model allows for each decode call. - This technically can be any large, but practically should set to the - context length the checkpoint was trained with. - horizon_len: Forecast horizon. - input_patch_len: Input patch len. - output_patch_len: Output patch len. How many timepoints is taken from a - single step of autoregressive decoding. Can be set as the training horizon - of the checkpoint. - mesh_shape: Shape of the data parallelism mesh. - mesh_name: Names of the data parallelism mesh. - model_p: Configuration of the TimesFM model deduced from the hparams. - """ - - def _logging(self, s): - if self._verbose: - print(s) - - def __init__( - self, - context_len: int, - horizon_len: int, - input_patch_len: int, - output_patch_len: int, - num_layers: int, - model_dims: int, - per_core_batch_size: int = 32, - backend: Literal["cpu", "gpu", "tpu"] = "cpu", - quantiles: Sequence[float] | None = None, - verbose: bool = True, - ) -> None: - """Initializes the TimesFM forecast API. - - Args: - context_len: Largest context length the model allows for each decode call. - This technically can be any large, but practically should set to the - context length the checkpoint was trained with. - horizon_len: Forecast horizon. - input_patch_len: Input patch len. - output_patch_len: Output patch len. How many timepoints is taken from a - single step of autoregressive decoding. Can be set as the training - horizon of the checkpoint. - num_layers: Number of transformer layers. - model_dims: Model dimension. - per_core_batch_size: Batch size on each core for data parallelism. - backend: One of "cpu", "gpu" or "tpu". - quantiles: list of output quantiles supported by the model. - verbose: Whether to print logging messages. - """ - self.per_core_batch_size = per_core_batch_size - self.backend = backend - self.num_devices = jax.local_device_count(self.backend) - self.global_batch_size = self.per_core_batch_size * self.num_devices - - self.context_len = context_len - self.horizon_len = horizon_len - self.input_patch_len = input_patch_len - self.output_patch_len = output_patch_len - - self.mesh_shape = [1, self.num_devices, 1] - self.mesh_name = ["replica", "data", "mdl"] - if quantiles is None: - quantiles = patched_decoder.DEFAULT_QUANTILES - - self.model_p = pax_fiddle.Config( - patched_decoder.PatchedTimeSeriesDecoder, - name="patched_decoder", - horizon_len=self.output_patch_len, - patch_len=input_patch_len, - model_dims=model_dims, - hidden_dims=model_dims, - residual_block_tpl=pax_fiddle.Config(patched_decoder.ResidualBlock), - quantiles=quantiles, - use_freq=True, - stacked_transformer_params_tpl=pax_fiddle.Config( - transformers.StackedTransformer, - num_heads=16, - num_layers=num_layers, - transformer_layer_params_tpl=pax_fiddle.Config( - transformers.Transformer, - ln_tpl=pax_fiddle.Config( - normalizations.RmsNorm, - ), - ), - ), - ) - - self._key1, self._key2 = jax.random.split(jax.random.PRNGKey(42)) - self._model = None - self._train_state = None - self._pmapped_decode = None - self._verbose = verbose - self._eval_context = base_layer.JaxContext.HParams(do_eval=True) - try: - multiprocessing.set_start_method("spawn") - except RuntimeError: - print("Multiprocessing context has already been set.") - - def _get_sample_inputs(self): - return { - "input_ts": jnp.zeros( - ( - self.per_core_batch_size, - self.context_len + self.output_patch_len, - ), - dtype=jnp.float32, - ), - "input_padding": jnp.zeros( - ( - self.per_core_batch_size, - self.context_len + self.output_patch_len, - ), - dtype=jnp.float32, - ), - "freq": jnp.zeros( - ( - self.per_core_batch_size, - 1, - ), - dtype=jnp.int32, - ), - } - - def load_from_checkpoint( - self, - checkpoint_path: Optional[str] = None, - repo_id: str = "google/timesfm-1.0-200m", - checkpoint_type: checkpoints.CheckpointType = checkpoints.CheckpointType.FLAX, - step: int | None = None, - ) -> None: - """Loads a checkpoint and compiles the decoder. - - Args: - checkpoint_path: Optional path to the checkpoint directory. - repo_id: Hugging Face Hub repo id. - checkpoint_type: type of PAX checkpoint - step: step of the checkpoint to load. If `None`, load latest checkpoint. - """ - # Download the checkpoint from Hugging Face Hub if not given - if checkpoint_path is None: - checkpoint_path = path.join(snapshot_download(repo_id), "checkpoints") - - # Initialize the model weights. - self._logging("Constructing model weights.") - start_time = time.time() - self._model = instantiate(self.model_p) - var_weight_hparams = self._model.abstract_init_with_metadata( - self._get_sample_inputs(), do_eval=True - ) - train_state_partition_specs = tasks_lib.create_state_partition_specs( - var_weight_hparams, - mesh_shape=self.mesh_shape, - mesh_axis_names=self.mesh_name, - discard_opt_states=True, - learners=None, - ) - train_state_local_shapes = tasks_lib.create_state_unpadded_shapes( - var_weight_hparams, - discard_opt_states=True, - learners=None, - ) - self._logging( - f"Constructed model weights in {time.time() - start_time:.2f} seconds." - ) - - # Load the model weights. - self._logging(f"Restoring checkpoint from {checkpoint_path}.") - start_time = time.time() - self._train_state = checkpoints.restore_checkpoint( - train_state_local_shapes, - checkpoint_dir=checkpoint_path, - checkpoint_type=checkpoint_type, - state_specs=train_state_partition_specs, - step=step, - ) - self._logging( - f"Restored checkpoint in {time.time() - start_time:.2f} seconds." - ) - - # Initialize and jit the decode fn. - def _decode(inputs): - assert self._model is not None - assert self._train_state is not None - return self._model.apply( - self._train_state.mdl_vars, - inputs, - horizon_len=self.horizon_len, - output_patch_len=self.output_patch_len, - max_len=self.context_len, - rngs={ - base_layer.PARAMS: self._key1, - base_layer.RANDOM: self._key2, - }, - method=self._model.decode, - ) - - self._logging("Jitting decoding.") - start_time = time.time() - self._pmapped_decode = jax.pmap( - _decode, - axis_name="batch", - devices=jax.devices(self.backend), - backend=self.backend, - axis_size=self.num_devices, - ) - with base_layer.JaxContext.new_context(hparams=self._eval_context): - _ = self._pmapped_decode( - NestedMap({ - "input_ts": jnp.zeros( - ( - self.num_devices, - self.per_core_batch_size, - self.context_len, - ), - dtype=jnp.float32, - ), - "input_padding": jnp.zeros( - ( - self.num_devices, - self.per_core_batch_size, - self.context_len + self.horizon_len, - ), - dtype=jnp.float32, - ), - "date_features": None, - "freq": jnp.zeros( - (self.num_devices, self.per_core_batch_size, 1), - dtype=jnp.int32, - ), - }) - ) - self._logging(f"Jitted decoding in {time.time() - start_time:.2f} seconds.") - - def _preprocess( - self, inputs: Sequence[np.array], freq: Sequence[int] - ) -> tuple[np.array, np.array, int]: - """Formats and pads raw inputs to feed into the model. - - This function both pads each time series to match the context length, and - pads the inputs to meet the SPMD shape requirement. - - Args: - inputs: A list of 1d JTensors. Each JTensor is the context time series of - a single forecast task. - freq: list of frequencies - - Returns: - A tuple of: - - the padded input time series to meet the model required context. - - the padding indicator. - - the number of padded examples for SPMD so that each core has the same - number (a multiple of `batch_size`) of examples. - """ - - input_ts, input_padding, inp_freq = [], [], [] - - pmap_pad = ( - (len(inputs) - 1) // self.global_batch_size + 1 - ) * self.global_batch_size - len(inputs) - - for i, ts in enumerate(inputs): - input_len = ts.shape[0] - padding = np.zeros(shape=(input_len + self.horizon_len,), dtype=float) - if input_len < self.context_len: - num_front_pad = self.context_len - input_len - ts = np.concatenate( - [np.zeros(shape=(num_front_pad,), dtype=float), ts], axis=0 - ) - padding = np.concatenate( - [np.ones(shape=(num_front_pad,), dtype=float), padding], axis=0 - ) - elif input_len > self.context_len: - ts = ts[-self.context_len :] - padding = padding[-(self.context_len + self.horizon_len) :] - - input_ts.append(ts) - input_padding.append(padding) - inp_freq.append(freq[i]) - - # Padding the remainder batch. - for _ in range(pmap_pad): - input_ts.append(input_ts[-1]) - input_padding.append(input_padding[-1]) - inp_freq.append(inp_freq[-1]) - - return ( - np.stack(input_ts, axis=0), - np.stack(input_padding, axis=0), - np.array(inp_freq).astype(np.int32).reshape(-1, 1), - pmap_pad, - ) - - def forecast( - self, - inputs: Sequence[Any], - freq: Sequence[int] | None = None, - window_size: int | None = None, - forecast_context_len: int | None = None, - ) -> tuple[JTensor, JTensor]: - """Forecasts on a list of time series. - - Args: - inputs: list of time series forecast contexts. Each context time series - should be in a format convertible to JTensor by `jnp.array`. - freq: frequency of each context time series. 0 for high frequency - (default), 1 for medium, and 2 for low. Notice this is different from - the `freq` required by `forecast_on_df`. - window_size: window size of trend + residual decomposition. If None then - we do not do decomposition. - forecast_context_len: optional max context length. - - Returns: - A tuple for JTensors: - - the mean forecast of size (# inputs, # forecast horizon), - - the full forecast (mean + quantiles) of size - (# inputs, # forecast horizon, 1 + # quantiles). - - Raises: - ValueError: If the checkpoint is not properly loaded. - """ - if not self._train_state or not self._model: - raise ValueError( - "Checkpoint not loaded. Call `load_from_checkpoint` before" - " `forecast`." - ) - if forecast_context_len is None: - forecast_context_len = self.context_len - inputs = [np.array(ts)[-forecast_context_len:] for ts in inputs] - inp_min = np.min([np.min(ts) for ts in inputs]) - - if window_size is not None: - new_inputs = [] - for ts in inputs: - new_inputs.extend(moving_average(ts, window_size)) - inputs = new_inputs - - if freq is None: - logging.info("No frequency provided via `freq`. Default to high (0).") - freq = [0] * len(inputs) - - input_ts, input_padding, inp_freq, pmap_pad = self._preprocess(inputs, freq) - with base_layer.JaxContext.new_context(hparams=self._eval_context): - mean_outputs = [] - full_outputs = [] - assert input_ts.shape[0] % self.global_batch_size == 0 - for i in range(input_ts.shape[0] // self.global_batch_size): - input_ts_in = jnp.array( - input_ts[ - i * self.global_batch_size : (i + 1) * self.global_batch_size - ] - ) - input_padding_in = jnp.array( - input_padding[ - i * self.global_batch_size : (i + 1) * self.global_batch_size - ], - ) - inp_freq_in = jnp.array( - inp_freq[ - i * self.global_batch_size : (i + 1) * self.global_batch_size, : - ], - dtype=jnp.int32, - ) - pmapped_inputs = NestedMap({ - "input_ts": es.jax_einshape( - "(db)...->db...", - input_ts_in, - d=self.num_devices, - ), - "input_padding": es.jax_einshape( - "(db)...->db...", - input_padding_in, - d=self.num_devices, - ), - "date_features": None, - "freq": es.jax_einshape( - "(db)...->db...", - inp_freq_in, - d=self.num_devices, - ), - }) - mean_output, full_output = self._pmapped_decode(pmapped_inputs) - mean_output = es.jax_einshape( - "db...->(db)...", mean_output, d=self.num_devices - ) - full_output = es.jax_einshape( - "db...->(db)...", full_output, d=self.num_devices - ) - mean_output = np.array(mean_output) - full_output = np.array(full_output) - mean_outputs.append(mean_output) - full_outputs.append(full_output) - - mean_outputs = np.concatenate(mean_outputs, axis=0) - full_outputs = np.concatenate(full_outputs, axis=0) - - if pmap_pad > 0: - mean_outputs = mean_outputs[:-pmap_pad, ...] - full_outputs = full_outputs[:-pmap_pad, ...] - - if window_size is not None: - mean_outputs = mean_outputs[0::2, ...] + mean_outputs[1::2, ...] - full_outputs = full_outputs[0::2, ...] + full_outputs[1::2, ...] - if inp_min >= 0: - mean_outputs = np.maximum(mean_outputs, 0.0) - full_outputs = np.maximum(full_outputs, 0.0) - return mean_outputs, full_outputs - - def forecast_on_df( - self, - inputs: pd.DataFrame, - freq: str, - forecast_context_len: int = 0, - value_name: str = "values", - model_name: str = "timesfm", - window_size: int | None = None, - num_jobs: int = 1, - ) -> pd.DataFrame: - """Forecasts on a list of time series. - - Args: - inputs: A pd.DataFrame of all time series. The dataframe should have a - `unique_id` column for identifying the time series, a `ds` column for - timestamps and a value column for the time series values. - freq: string valued `freq` of data. Notice this is different from the - `freq` required by `forecast`. See `freq_map` for allowed values. - forecast_context_len: If provided none zero, we take the last - `forecast_context_len` time-points from each series as the forecast - context instead of the `context_len` set by the model. - value_name: The name of the value column. - model_name: name of the model to be written into future df. - window_size: window size of trend + residual decomposition. If None then - we do not do decomposition. - num_jobs: number of parallel processes to use for dataframe processing. - - Returns: - Future forecasts dataframe. - """ - if not ( - "unique_id" in inputs.columns - and "ds" in inputs.columns - and value_name in inputs.columns - ): - raise ValueError( - f"DataFrame must have unique_id, ds and {value_name} columns." - ) - if not forecast_context_len: - forecast_context_len = self.context_len - logging.info("Preprocessing dataframe.") - df_sorted = inputs.sort_values(by=["unique_id", "ds"]) - new_inputs = [] - uids = [] - if num_jobs == 1: - print("Processing dataframe with single process.") - for key, group in df_sorted.groupby("unique_id"): - inp, uid = process_group( - key, - group, - value_name, - forecast_context_len, - ) - new_inputs.append(inp) - uids.append(uid) - else: - if num_jobs == -1: - num_jobs = multiprocessing.cpu_count() - print("Processing dataframe with multiple processes.") - with multiprocessing.Pool(processes=num_jobs) as pool: - results = pool.starmap( - process_group, - [ - (key, group, value_name, forecast_context_len) - for key, group in df_sorted.groupby("unique_id") - ], - ) - new_inputs, uids = zip(*results) - print("Finished preprocessing dataframe.") - freq_inps = [freq_map(freq)] * len(new_inputs) - _, full_forecast = self.forecast( - new_inputs, freq=freq_inps, window_size=window_size - ) - print("Finished forecasting.") - fcst_df = make_future_dataframe( - uids=uids, - last_times=df_sorted.groupby("unique_id")["ds"].tail(1), - h=self.horizon_len, - freq=freq, - ) - fcst_df[model_name] = full_forecast[:, 0 : self.horizon_len, 0].reshape( - -1, 1 - ) - - if self._model.quantiles is not None: - for i, q in enumerate(self._model.quantiles): - q_col = f"{model_name}-q-{q}" - fcst_df[q_col] = full_forecast[:, 0 : self.horizon_len, 1 + i].reshape( - -1, 1 - ) - if q == 0.5: - fcst_df[model_name] = fcst_df[q_col] - logging.info("Finished creating output dataframe.") - return fcst_df