deleting old files
This commit is contained in:
@@ -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