deleting old files
This commit is contained in:
-21
@@ -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
|
|
||||||
@@ -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
|
|
||||||
@@ -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,
|
|
||||||
)
|
|
||||||
@@ -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, :],
|
|
||||||
)
|
|
||||||
-602
@@ -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
|
|
||||||
Reference in New Issue
Block a user