Merge pull request #126 from google-research/rajat_dev
have an option for truncating negative values
This commit is contained in:
+147
-187
@@ -11,7 +11,6 @@
|
|||||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
# See the License for the specific language governing permissions and
|
# See the License for the specific language governing permissions and
|
||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
|
|
||||||
"""TimesFM forecast API for inference."""
|
"""TimesFM forecast API for inference."""
|
||||||
|
|
||||||
import collections
|
import collections
|
||||||
@@ -60,23 +59,16 @@ def moving_average(arr, window_size):
|
|||||||
"""Calculates the moving average using NumPy's convolution function."""
|
"""Calculates the moving average using NumPy's convolution function."""
|
||||||
# Pad with zeros to handle initial window positions
|
# Pad with zeros to handle initial window positions
|
||||||
arr_padded = np.pad(arr, (window_size - 1, 0), "constant")
|
arr_padded = np.pad(arr, (window_size - 1, 0), "constant")
|
||||||
smoothed_arr = (
|
smoothed_arr = (np.convolve(arr_padded, np.ones(window_size), "valid") /
|
||||||
np.convolve(arr_padded, np.ones(window_size), "valid") / window_size
|
window_size)
|
||||||
)
|
|
||||||
return [smoothed_arr, arr - smoothed_arr]
|
return [smoothed_arr, arr - smoothed_arr]
|
||||||
|
|
||||||
|
|
||||||
def freq_map(freq: str):
|
def freq_map(freq: str):
|
||||||
"""Returns the frequency map for the given frequency string."""
|
"""Returns the frequency map for the given frequency string."""
|
||||||
freq = str.upper(freq)
|
freq = str.upper(freq)
|
||||||
if (
|
if (freq.endswith("H") or freq.endswith("T") or freq.endswith("MIN") or
|
||||||
freq.endswith("H")
|
freq.endswith("D") or freq.endswith("B") or freq.endswith("U")):
|
||||||
or freq.endswith("T")
|
|
||||||
or freq.endswith("MIN")
|
|
||||||
or freq.endswith("D")
|
|
||||||
or freq.endswith("B")
|
|
||||||
or freq.endswith("U")
|
|
||||||
):
|
|
||||||
return 0
|
return 0
|
||||||
elif freq.endswith(("W", "M", "MS")):
|
elif freq.endswith(("W", "M", "MS")):
|
||||||
return 1
|
return 1
|
||||||
@@ -202,9 +194,7 @@ class TimesFm:
|
|||||||
num_layers=num_layers,
|
num_layers=num_layers,
|
||||||
transformer_layer_params_tpl=pax_fiddle.Config(
|
transformer_layer_params_tpl=pax_fiddle.Config(
|
||||||
transformers.Transformer,
|
transformers.Transformer,
|
||||||
ln_tpl=pax_fiddle.Config(
|
ln_tpl=pax_fiddle.Config(normalizations.RmsNorm,),
|
||||||
normalizations.RmsNorm,
|
|
||||||
),
|
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -222,34 +212,38 @@ class TimesFm:
|
|||||||
|
|
||||||
def _get_sample_inputs(self):
|
def _get_sample_inputs(self):
|
||||||
return {
|
return {
|
||||||
"input_ts": jnp.zeros(
|
"input_ts":
|
||||||
(
|
jnp.zeros(
|
||||||
self.per_core_batch_size,
|
(
|
||||||
self.context_len + self.output_patch_len,
|
self.per_core_batch_size,
|
||||||
|
self.context_len + self.output_patch_len,
|
||||||
|
),
|
||||||
|
dtype=jnp.float32,
|
||||||
),
|
),
|
||||||
dtype=jnp.float32,
|
"input_padding":
|
||||||
),
|
jnp.zeros(
|
||||||
"input_padding": jnp.zeros(
|
(
|
||||||
(
|
self.per_core_batch_size,
|
||||||
self.per_core_batch_size,
|
self.context_len + self.output_patch_len,
|
||||||
self.context_len + self.output_patch_len,
|
),
|
||||||
|
dtype=jnp.float32,
|
||||||
),
|
),
|
||||||
dtype=jnp.float32,
|
"freq":
|
||||||
),
|
jnp.zeros(
|
||||||
"freq": jnp.zeros(
|
(
|
||||||
(
|
self.per_core_batch_size,
|
||||||
self.per_core_batch_size,
|
1,
|
||||||
1,
|
),
|
||||||
|
dtype=jnp.int32,
|
||||||
),
|
),
|
||||||
dtype=jnp.int32,
|
|
||||||
),
|
|
||||||
}
|
}
|
||||||
|
|
||||||
def load_from_checkpoint(
|
def load_from_checkpoint(
|
||||||
self,
|
self,
|
||||||
checkpoint_path: Optional[str] = None,
|
checkpoint_path: Optional[str] = None,
|
||||||
repo_id: str = "google/timesfm-1.0-200m",
|
repo_id: str = "google/timesfm-1.0-200m",
|
||||||
checkpoint_type: checkpoints.CheckpointType = checkpoints.CheckpointType.FLAX,
|
checkpoint_type: checkpoints.CheckpointType = checkpoints.CheckpointType.
|
||||||
|
FLAX,
|
||||||
step: int | None = None,
|
step: int | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Loads a checkpoint and compiles the decoder.
|
"""Loads a checkpoint and compiles the decoder.
|
||||||
@@ -269,8 +263,7 @@ class TimesFm:
|
|||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
self._model = instantiate(self.model_p)
|
self._model = instantiate(self.model_p)
|
||||||
var_weight_hparams = self._model.abstract_init_with_metadata(
|
var_weight_hparams = self._model.abstract_init_with_metadata(
|
||||||
self._get_sample_inputs(), do_eval=True
|
self._get_sample_inputs(), do_eval=True)
|
||||||
)
|
|
||||||
train_state_partition_specs = tasks_lib.create_state_partition_specs(
|
train_state_partition_specs = tasks_lib.create_state_partition_specs(
|
||||||
var_weight_hparams,
|
var_weight_hparams,
|
||||||
mesh_shape=self.mesh_shape,
|
mesh_shape=self.mesh_shape,
|
||||||
@@ -284,8 +277,7 @@ class TimesFm:
|
|||||||
learners=None,
|
learners=None,
|
||||||
)
|
)
|
||||||
self._logging(
|
self._logging(
|
||||||
f"Constructed model weights in {time.time() - start_time:.2f} seconds."
|
f"Constructed model weights in {time.time() - start_time:.2f} seconds.")
|
||||||
)
|
|
||||||
|
|
||||||
# Load the model weights.
|
# Load the model weights.
|
||||||
self._logging(f"Restoring checkpoint from {checkpoint_path}.")
|
self._logging(f"Restoring checkpoint from {checkpoint_path}.")
|
||||||
@@ -298,8 +290,7 @@ class TimesFm:
|
|||||||
step=step,
|
step=step,
|
||||||
)
|
)
|
||||||
self._logging(
|
self._logging(
|
||||||
f"Restored checkpoint in {time.time() - start_time:.2f} seconds."
|
f"Restored checkpoint in {time.time() - start_time:.2f} seconds.")
|
||||||
)
|
|
||||||
self.jit_decode()
|
self.jit_decode()
|
||||||
|
|
||||||
def jit_decode(self):
|
def jit_decode(self):
|
||||||
@@ -335,34 +326,36 @@ class TimesFm:
|
|||||||
with base_layer.JaxContext.new_context(hparams=self._eval_context):
|
with base_layer.JaxContext.new_context(hparams=self._eval_context):
|
||||||
_ = self._pmapped_decode(
|
_ = self._pmapped_decode(
|
||||||
NestedMap({
|
NestedMap({
|
||||||
"input_ts": jnp.zeros(
|
"input_ts":
|
||||||
(
|
jnp.zeros(
|
||||||
self.num_devices,
|
(
|
||||||
self.per_core_batch_size,
|
self.num_devices,
|
||||||
self.context_len,
|
self.per_core_batch_size,
|
||||||
|
self.context_len,
|
||||||
|
),
|
||||||
|
dtype=jnp.float32,
|
||||||
),
|
),
|
||||||
dtype=jnp.float32,
|
"input_padding":
|
||||||
),
|
jnp.zeros(
|
||||||
"input_padding": jnp.zeros(
|
(
|
||||||
(
|
self.num_devices,
|
||||||
self.num_devices,
|
self.per_core_batch_size,
|
||||||
self.per_core_batch_size,
|
self.context_len + self.horizon_len,
|
||||||
self.context_len + self.horizon_len,
|
),
|
||||||
|
dtype=jnp.float32,
|
||||||
),
|
),
|
||||||
dtype=jnp.float32,
|
"date_features":
|
||||||
),
|
None,
|
||||||
"date_features": None,
|
"freq":
|
||||||
"freq": jnp.zeros(
|
jnp.zeros(
|
||||||
(self.num_devices, self.per_core_batch_size, 1),
|
(self.num_devices, self.per_core_batch_size, 1),
|
||||||
dtype=jnp.int32,
|
dtype=jnp.int32,
|
||||||
),
|
),
|
||||||
})
|
}))
|
||||||
)
|
|
||||||
self._logging(f"Jitted decoding in {time.time() - start_time:.2f} seconds.")
|
self._logging(f"Jitted decoding in {time.time() - start_time:.2f} seconds.")
|
||||||
|
|
||||||
def _preprocess(
|
def _preprocess(self, inputs: Sequence[np.array],
|
||||||
self, inputs: Sequence[np.array], freq: Sequence[int]
|
freq: Sequence[int]) -> tuple[np.array, np.array, int]:
|
||||||
) -> tuple[np.array, np.array, int]:
|
|
||||||
"""Formats and pads raw inputs to feed into the model.
|
"""Formats and pads raw inputs to feed into the model.
|
||||||
|
|
||||||
This function both pads each time series to match the context length, and
|
This function both pads each time series to match the context length, and
|
||||||
@@ -383,24 +376,21 @@ class TimesFm:
|
|||||||
|
|
||||||
input_ts, input_padding, inp_freq = [], [], []
|
input_ts, input_padding, inp_freq = [], [], []
|
||||||
|
|
||||||
pmap_pad = (
|
pmap_pad = ((len(inputs) - 1) // self.global_batch_size +
|
||||||
(len(inputs) - 1) // self.global_batch_size + 1
|
1) * self.global_batch_size - len(inputs)
|
||||||
) * self.global_batch_size - len(inputs)
|
|
||||||
|
|
||||||
for i, ts in enumerate(inputs):
|
for i, ts in enumerate(inputs):
|
||||||
input_len = ts.shape[0]
|
input_len = ts.shape[0]
|
||||||
padding = np.zeros(shape=(input_len + self.horizon_len,), dtype=float)
|
padding = np.zeros(shape=(input_len + self.horizon_len,), dtype=float)
|
||||||
if input_len < self.context_len:
|
if input_len < self.context_len:
|
||||||
num_front_pad = self.context_len - input_len
|
num_front_pad = self.context_len - input_len
|
||||||
ts = np.concatenate(
|
ts = np.concatenate([np.zeros(shape=(num_front_pad,), dtype=float), ts],
|
||||||
[np.zeros(shape=(num_front_pad,), dtype=float), ts], axis=0
|
axis=0)
|
||||||
)
|
|
||||||
padding = np.concatenate(
|
padding = np.concatenate(
|
||||||
[np.ones(shape=(num_front_pad,), dtype=float), padding], axis=0
|
[np.ones(shape=(num_front_pad,), dtype=float), padding], axis=0)
|
||||||
)
|
|
||||||
elif input_len > self.context_len:
|
elif input_len > self.context_len:
|
||||||
ts = ts[-self.context_len :]
|
ts = ts[-self.context_len:]
|
||||||
padding = padding[-(self.context_len + self.horizon_len) :]
|
padding = padding[-(self.context_len + self.horizon_len):]
|
||||||
|
|
||||||
input_ts.append(ts)
|
input_ts.append(ts)
|
||||||
input_padding.append(padding)
|
input_padding.append(padding)
|
||||||
@@ -426,6 +416,7 @@ class TimesFm:
|
|||||||
window_size: int | None = None,
|
window_size: int | None = None,
|
||||||
forecast_context_len: int | None = None,
|
forecast_context_len: int | None = None,
|
||||||
return_forecast_on_context: bool = False,
|
return_forecast_on_context: bool = False,
|
||||||
|
truncate_negative: bool = False,
|
||||||
) -> tuple[JTensor, JTensor]:
|
) -> tuple[JTensor, JTensor]:
|
||||||
"""Forecasts on a list of time series.
|
"""Forecasts on a list of time series.
|
||||||
|
|
||||||
@@ -440,6 +431,8 @@ class TimesFm:
|
|||||||
forecast_context_len: optional max context length.
|
forecast_context_len: optional max context length.
|
||||||
return_forecast_on_context: True to return the forecast on the context
|
return_forecast_on_context: True to return the forecast on the context
|
||||||
when available, i.e. after the first input patch.
|
when available, i.e. after the first input patch.
|
||||||
|
truncate_negative: truncate to only non-negative values if all the contexts
|
||||||
|
have non-negative values.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
A tuple for JTensors:
|
A tuple for JTensors:
|
||||||
@@ -453,8 +446,7 @@ class TimesFm:
|
|||||||
if not self._train_state or not self._model:
|
if not self._train_state or not self._model:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Checkpoint not loaded. Call `load_from_checkpoint` before"
|
"Checkpoint not loaded. Call `load_from_checkpoint` before"
|
||||||
" `forecast`."
|
" `forecast`.")
|
||||||
)
|
|
||||||
if forecast_context_len is None:
|
if forecast_context_len is None:
|
||||||
forecast_context_len = self.context_len
|
forecast_context_len = self.context_len
|
||||||
inputs = [np.array(ts)[-forecast_context_len:] for ts in inputs]
|
inputs = [np.array(ts)[-forecast_context_len:] for ts in inputs]
|
||||||
@@ -476,50 +468,48 @@ class TimesFm:
|
|||||||
full_outputs = []
|
full_outputs = []
|
||||||
assert input_ts.shape[0] % self.global_batch_size == 0
|
assert input_ts.shape[0] % self.global_batch_size == 0
|
||||||
for i in range(input_ts.shape[0] // self.global_batch_size):
|
for i in range(input_ts.shape[0] // self.global_batch_size):
|
||||||
input_ts_in = jnp.array(
|
input_ts_in = jnp.array(input_ts[i * self.global_batch_size:(i + 1) *
|
||||||
input_ts[
|
self.global_batch_size])
|
||||||
i * self.global_batch_size : (i + 1) * self.global_batch_size
|
|
||||||
]
|
|
||||||
)
|
|
||||||
input_padding_in = jnp.array(
|
input_padding_in = jnp.array(
|
||||||
input_padding[
|
input_padding[i * self.global_batch_size:(i + 1) *
|
||||||
i * self.global_batch_size : (i + 1) * self.global_batch_size
|
self.global_batch_size],)
|
||||||
],
|
|
||||||
)
|
|
||||||
inp_freq_in = jnp.array(
|
inp_freq_in = jnp.array(
|
||||||
inp_freq[
|
inp_freq[i * self.global_batch_size:(i + 1) *
|
||||||
i * self.global_batch_size : (i + 1) * self.global_batch_size, :
|
self.global_batch_size, :],
|
||||||
],
|
|
||||||
dtype=jnp.int32,
|
dtype=jnp.int32,
|
||||||
)
|
)
|
||||||
pmapped_inputs = NestedMap({
|
pmapped_inputs = NestedMap({
|
||||||
"input_ts": es.jax_einshape(
|
"input_ts":
|
||||||
"(db)...->db...",
|
es.jax_einshape(
|
||||||
input_ts_in,
|
"(db)...->db...",
|
||||||
d=self.num_devices,
|
input_ts_in,
|
||||||
),
|
d=self.num_devices,
|
||||||
"input_padding": es.jax_einshape(
|
),
|
||||||
"(db)...->db...",
|
"input_padding":
|
||||||
input_padding_in,
|
es.jax_einshape(
|
||||||
d=self.num_devices,
|
"(db)...->db...",
|
||||||
),
|
input_padding_in,
|
||||||
"date_features": None,
|
d=self.num_devices,
|
||||||
"freq": es.jax_einshape(
|
),
|
||||||
"(db)...->db...",
|
"date_features":
|
||||||
inp_freq_in,
|
None,
|
||||||
d=self.num_devices,
|
"freq":
|
||||||
),
|
es.jax_einshape(
|
||||||
|
"(db)...->db...",
|
||||||
|
inp_freq_in,
|
||||||
|
d=self.num_devices,
|
||||||
|
),
|
||||||
})
|
})
|
||||||
mean_output, full_output = self._pmapped_decode(pmapped_inputs)
|
mean_output, full_output = self._pmapped_decode(pmapped_inputs)
|
||||||
if not return_forecast_on_context:
|
if not return_forecast_on_context:
|
||||||
mean_output = mean_output[:, :, self._horizon_start :, ...]
|
mean_output = mean_output[:, :, self._horizon_start:, ...]
|
||||||
full_output = full_output[:, :, self._horizon_start :, ...]
|
full_output = full_output[:, :, self._horizon_start:, ...]
|
||||||
mean_output = es.jax_einshape(
|
mean_output = es.jax_einshape("db...->(db)...",
|
||||||
"db...->(db)...", mean_output, d=self.num_devices
|
mean_output,
|
||||||
)
|
d=self.num_devices)
|
||||||
full_output = es.jax_einshape(
|
full_output = es.jax_einshape("db...->(db)...",
|
||||||
"db...->(db)...", full_output, d=self.num_devices
|
full_output,
|
||||||
)
|
d=self.num_devices)
|
||||||
mean_output = np.array(mean_output)
|
mean_output = np.array(mean_output)
|
||||||
full_output = np.array(full_output)
|
full_output = np.array(full_output)
|
||||||
mean_outputs.append(mean_output)
|
mean_outputs.append(mean_output)
|
||||||
@@ -535,7 +525,7 @@ class TimesFm:
|
|||||||
if window_size is not None:
|
if window_size is not None:
|
||||||
mean_outputs = mean_outputs[0::2, ...] + mean_outputs[1::2, ...]
|
mean_outputs = mean_outputs[0::2, ...] + mean_outputs[1::2, ...]
|
||||||
full_outputs = full_outputs[0::2, ...] + full_outputs[1::2, ...]
|
full_outputs = full_outputs[0::2, ...] + full_outputs[1::2, ...]
|
||||||
if inp_min >= 0:
|
if inp_min >= 0 and truncate_negative:
|
||||||
mean_outputs = np.maximum(mean_outputs, 0.0)
|
mean_outputs = np.maximum(mean_outputs, 0.0)
|
||||||
full_outputs = np.maximum(full_outputs, 0.0)
|
full_outputs = np.maximum(full_outputs, 0.0)
|
||||||
return mean_outputs, full_outputs
|
return mean_outputs, full_outputs
|
||||||
@@ -543,16 +533,13 @@ class TimesFm:
|
|||||||
def forecast_with_covariates(
|
def forecast_with_covariates(
|
||||||
self,
|
self,
|
||||||
inputs: list[Sequence[float]],
|
inputs: list[Sequence[float]],
|
||||||
dynamic_numerical_covariates: (
|
dynamic_numerical_covariates: (dict[str, Sequence[Sequence[float]]] |
|
||||||
dict[str, Sequence[Sequence[float]]] | None
|
None) = None,
|
||||||
) = None,
|
dynamic_categorical_covariates: (dict[str, Sequence[Sequence[Category]]] |
|
||||||
dynamic_categorical_covariates: (
|
None) = None,
|
||||||
dict[str, Sequence[Sequence[Category]]] | None
|
|
||||||
) = None,
|
|
||||||
static_numerical_covariates: dict[str, Sequence[float]] | None = None,
|
static_numerical_covariates: dict[str, Sequence[float]] | None = None,
|
||||||
static_categorical_covariates: (
|
static_categorical_covariates: (dict[str, Sequence[Category]] |
|
||||||
dict[str, Sequence[Category]] | None
|
None) = None,
|
||||||
) = None,
|
|
||||||
freq: Sequence[int] | None = None,
|
freq: Sequence[int] | None = None,
|
||||||
window_size: int | None = None,
|
window_size: int | None = None,
|
||||||
forecast_context_len: int | None = None,
|
forecast_context_len: int | None = None,
|
||||||
@@ -594,17 +581,12 @@ class TimesFm:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
# Verify and bookkeep covariates.
|
# Verify and bookkeep covariates.
|
||||||
if not (
|
if not (dynamic_numerical_covariates or dynamic_categorical_covariates or
|
||||||
dynamic_numerical_covariates
|
static_numerical_covariates or static_categorical_covariates):
|
||||||
or dynamic_categorical_covariates
|
|
||||||
or static_numerical_covariates
|
|
||||||
or static_categorical_covariates
|
|
||||||
):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"At least one of dynamic_numerical_covariates,"
|
"At least one of dynamic_numerical_covariates,"
|
||||||
" dynamic_categorical_covariates, static_numerical_covariates,"
|
" dynamic_categorical_covariates, static_numerical_covariates,"
|
||||||
" static_categorical_covariates must be set."
|
" static_categorical_covariates must be set.")
|
||||||
)
|
|
||||||
|
|
||||||
# Track the lengths of (1) each input, (2) the part that can be used in the
|
# Track the lengths of (1) each input, (2) the part that can be used in the
|
||||||
# linear model, and (3) the horizon.
|
# linear model, and (3) the horizon.
|
||||||
@@ -624,20 +606,18 @@ class TimesFm:
|
|||||||
|
|
||||||
if dynamic_numerical_covariates:
|
if dynamic_numerical_covariates:
|
||||||
test_lens.append(
|
test_lens.append(
|
||||||
len(list(dynamic_numerical_covariates.values())[0][i]) - input_len
|
len(list(dynamic_numerical_covariates.values())[0][i]) - input_len)
|
||||||
)
|
|
||||||
elif dynamic_categorical_covariates:
|
elif dynamic_categorical_covariates:
|
||||||
test_lens.append(
|
test_lens.append(
|
||||||
len(list(dynamic_categorical_covariates.values())[0][i]) - input_len
|
len(list(dynamic_categorical_covariates.values())[0][i]) -
|
||||||
)
|
input_len)
|
||||||
else:
|
else:
|
||||||
test_lens.append(self.horizon_len)
|
test_lens.append(self.horizon_len)
|
||||||
|
|
||||||
if test_lens[-1] > self.horizon_len:
|
if test_lens[-1] > self.horizon_len:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Forecast requested longer horizon than the model definition "
|
"Forecast requested longer horizon than the model definition "
|
||||||
f"supports: {test_lens[-1]} vs {self.horizon_len}."
|
f"supports: {test_lens[-1]} vs {self.horizon_len}.")
|
||||||
)
|
|
||||||
|
|
||||||
# Prepare the covariates into train and test.
|
# Prepare the covariates into train and test.
|
||||||
train_dynamic_numerical_covariates = collections.defaultdict(list)
|
train_dynamic_numerical_covariates = collections.defaultdict(list)
|
||||||
@@ -660,11 +640,9 @@ class TimesFm:
|
|||||||
continue
|
continue
|
||||||
for covariate_name, covariate_values in covariates.items():
|
for covariate_name, covariate_values in covariates.items():
|
||||||
for input_len, train_len, covariate_value in zip(
|
for input_len, train_len, covariate_value in zip(
|
||||||
input_lens, train_lens, covariate_values
|
input_lens, train_lens, covariate_values):
|
||||||
):
|
|
||||||
train_covariates[covariate_name].append(
|
train_covariates[covariate_name].append(
|
||||||
covariate_value[(input_len - train_len) : input_len]
|
covariate_value[(input_len - train_len):input_len])
|
||||||
)
|
|
||||||
test_covariates[covariate_name].append(covariate_value[input_len:])
|
test_covariates[covariate_name].append(covariate_value[input_len:])
|
||||||
|
|
||||||
# Fit models.
|
# Fit models.
|
||||||
@@ -678,15 +656,10 @@ class TimesFm:
|
|||||||
return_forecast_on_context=True,
|
return_forecast_on_context=True,
|
||||||
)
|
)
|
||||||
targets = [
|
targets = [
|
||||||
(
|
(np.array(input_ts)[-train_len:] -
|
||||||
np.array(input_ts)[-train_len:]
|
mean_output[(self._horizon_start - train_len):self._horizon_start])
|
||||||
- mean_output[
|
for input_ts, mean_output, train_len in zip(inputs, mean_outputs,
|
||||||
(self._horizon_start - train_len) : self._horizon_start
|
train_lens)
|
||||||
]
|
|
||||||
)
|
|
||||||
for input_ts, mean_output, train_len in zip(
|
|
||||||
inputs, mean_outputs, train_lens
|
|
||||||
)
|
|
||||||
]
|
]
|
||||||
per_instance_stats = None
|
per_instance_stats = None
|
||||||
if normalize_xreg_target_per_input:
|
if normalize_xreg_target_per_input:
|
||||||
@@ -697,8 +670,10 @@ class TimesFm:
|
|||||||
test_lens=test_lens,
|
test_lens=test_lens,
|
||||||
train_dynamic_numerical_covariates=train_dynamic_numerical_covariates,
|
train_dynamic_numerical_covariates=train_dynamic_numerical_covariates,
|
||||||
test_dynamic_numerical_covariates=test_dynamic_numerical_covariates,
|
test_dynamic_numerical_covariates=test_dynamic_numerical_covariates,
|
||||||
train_dynamic_categorical_covariates=train_dynamic_categorical_covariates,
|
train_dynamic_categorical_covariates=
|
||||||
test_dynamic_categorical_covariates=test_dynamic_categorical_covariates,
|
train_dynamic_categorical_covariates,
|
||||||
|
test_dynamic_categorical_covariates=
|
||||||
|
test_dynamic_categorical_covariates,
|
||||||
static_numerical_covariates=static_numerical_covariates,
|
static_numerical_covariates=static_numerical_covariates,
|
||||||
static_categorical_covariates=static_categorical_covariates,
|
static_categorical_covariates=static_categorical_covariates,
|
||||||
).fit(
|
).fit(
|
||||||
@@ -713,12 +688,8 @@ class TimesFm:
|
|||||||
if normalize_xreg_target_per_input:
|
if normalize_xreg_target_per_input:
|
||||||
xregs = _renormalize(xregs, per_instance_stats)
|
xregs = _renormalize(xregs, per_instance_stats)
|
||||||
outputs = [
|
outputs = [
|
||||||
(
|
(mean_output[self._horizon_start:(self._horizon_start + test_len)] +
|
||||||
mean_output[
|
xreg)
|
||||||
self._horizon_start : (self._horizon_start + test_len)
|
|
||||||
]
|
|
||||||
+ xreg
|
|
||||||
)
|
|
||||||
for mean_output, test_len, xreg in zip(mean_outputs, test_lens, xregs)
|
for mean_output, test_len, xreg in zip(mean_outputs, test_lens, xregs)
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -737,8 +708,10 @@ class TimesFm:
|
|||||||
test_lens=test_lens,
|
test_lens=test_lens,
|
||||||
train_dynamic_numerical_covariates=train_dynamic_numerical_covariates,
|
train_dynamic_numerical_covariates=train_dynamic_numerical_covariates,
|
||||||
test_dynamic_numerical_covariates=test_dynamic_numerical_covariates,
|
test_dynamic_numerical_covariates=test_dynamic_numerical_covariates,
|
||||||
train_dynamic_categorical_covariates=train_dynamic_categorical_covariates,
|
train_dynamic_categorical_covariates=
|
||||||
test_dynamic_categorical_covariates=test_dynamic_categorical_covariates,
|
train_dynamic_categorical_covariates,
|
||||||
|
test_dynamic_categorical_covariates=
|
||||||
|
test_dynamic_categorical_covariates,
|
||||||
static_numerical_covariates=static_numerical_covariates,
|
static_numerical_covariates=static_numerical_covariates,
|
||||||
static_categorical_covariates=static_categorical_covariates,
|
static_categorical_covariates=static_categorical_covariates,
|
||||||
).fit(
|
).fit(
|
||||||
@@ -761,12 +734,8 @@ class TimesFm:
|
|||||||
return_forecast_on_context=True,
|
return_forecast_on_context=True,
|
||||||
)
|
)
|
||||||
outputs = [
|
outputs = [
|
||||||
(
|
(mean_output[self._horizon_start:(self._horizon_start + test_len)] +
|
||||||
mean_output[
|
xreg)
|
||||||
self._horizon_start : (self._horizon_start + test_len)
|
|
||||||
]
|
|
||||||
+ xreg
|
|
||||||
)
|
|
||||||
for mean_output, test_len, xreg in zip(mean_outputs, test_lens, xregs)
|
for mean_output, test_len, xreg in zip(mean_outputs, test_lens, xregs)
|
||||||
]
|
]
|
||||||
if normalize_xreg_target_per_input:
|
if normalize_xreg_target_per_input:
|
||||||
@@ -806,14 +775,10 @@ class TimesFm:
|
|||||||
Returns:
|
Returns:
|
||||||
Future forecasts dataframe.
|
Future forecasts dataframe.
|
||||||
"""
|
"""
|
||||||
if not (
|
if not ("unique_id" in inputs.columns and "ds" in inputs.columns and
|
||||||
"unique_id" in inputs.columns
|
value_name in inputs.columns):
|
||||||
and "ds" in inputs.columns
|
|
||||||
and value_name in inputs.columns
|
|
||||||
):
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"DataFrame must have unique_id, ds and {value_name} columns."
|
f"DataFrame must have unique_id, ds and {value_name} columns.")
|
||||||
)
|
|
||||||
if not forecast_context_len:
|
if not forecast_context_len:
|
||||||
forecast_context_len = self.context_len
|
forecast_context_len = self.context_len
|
||||||
logging.info("Preprocessing dataframe.")
|
logging.info("Preprocessing dataframe.")
|
||||||
@@ -840,36 +805,31 @@ class TimesFm:
|
|||||||
with multiprocessing.Pool(processes=num_jobs) as pool:
|
with multiprocessing.Pool(processes=num_jobs) as pool:
|
||||||
results = pool.starmap(
|
results = pool.starmap(
|
||||||
process_group,
|
process_group,
|
||||||
[
|
[(key, group, value_name, forecast_context_len)
|
||||||
(key, group, value_name, forecast_context_len)
|
for key, group in df_sorted.groupby("unique_id")],
|
||||||
for key, group in df_sorted.groupby("unique_id")
|
|
||||||
],
|
|
||||||
)
|
)
|
||||||
new_inputs, uids = zip(*results)
|
new_inputs, uids = zip(*results)
|
||||||
if verbose:
|
if verbose:
|
||||||
print("Finished preprocessing dataframe.")
|
print("Finished preprocessing dataframe.")
|
||||||
freq_inps = [freq_map(freq)] * len(new_inputs)
|
freq_inps = [freq_map(freq)] * len(new_inputs)
|
||||||
_, full_forecast = self.forecast(
|
_, full_forecast = self.forecast(new_inputs,
|
||||||
new_inputs, freq=freq_inps, window_size=window_size
|
freq=freq_inps,
|
||||||
)
|
window_size=window_size)
|
||||||
if verbose:
|
if verbose:
|
||||||
print("Finished forecasting.")
|
print("Finished forecasting.")
|
||||||
fcst_df = make_future_dataframe(
|
fcst_df = make_future_dataframe(
|
||||||
uids=uids,
|
uids=uids,
|
||||||
last_times=df_sorted.groupby("unique_id")["ds"].tail(1),
|
last_times=df_sorted.groupby("unique_id")["ds"].tail(1),
|
||||||
h=self.horizon_len,
|
h=self.horizon_len,
|
||||||
freq=freq,
|
freq=freq,
|
||||||
)
|
)
|
||||||
fcst_df[model_name] = full_forecast[:, 0 : self.horizon_len, 0].reshape(
|
fcst_df[model_name] = full_forecast[:, 0:self.horizon_len, 0].reshape(-1, 1)
|
||||||
-1, 1
|
|
||||||
)
|
|
||||||
|
|
||||||
if self._model.quantiles is not None:
|
if self._model.quantiles is not None:
|
||||||
for i, q in enumerate(self._model.quantiles):
|
for i, q in enumerate(self._model.quantiles):
|
||||||
q_col = f"{model_name}-q-{q}"
|
q_col = f"{model_name}-q-{q}"
|
||||||
fcst_df[q_col] = full_forecast[:, 0 : self.horizon_len, 1 + i].reshape(
|
fcst_df[q_col] = full_forecast[:, 0:self.horizon_len,
|
||||||
-1, 1
|
1 + i].reshape(-1, 1)
|
||||||
)
|
|
||||||
if q == 0.5:
|
if q == 0.5:
|
||||||
fcst_df[model_name] = fcst_df[q_col]
|
fcst_df[model_name] = fcst_df[q_col]
|
||||||
logging.info("Finished creating output dataframe.")
|
logging.info("Finished creating output dataframe.")
|
||||||
|
|||||||
Reference in New Issue
Block a user