Merge pull request #178 from google-research/rajat_dev

Setting median outputs as default. Also minor changes to finetuning.
This commit is contained in:
Yichen Zhou
2024-11-05 16:31:23 -08:00
committed by GitHub
6 changed files with 492 additions and 421 deletions
+87 -15
View File
@@ -62,7 +62,7 @@ def freq_map(freq: str):
# Per time series normalization: forward.
def normalize(batch):
def _normalize(batch):
stats = [
(np.mean(x), np.where((w := np.std(x)) > _TOL, w, 1.0)) for x in batch
]
@@ -71,7 +71,7 @@ def normalize(batch):
# Per time series normalization: inverse.
def renormalize(batch, stats):
def _renormalize(batch, stats):
return [x * stat[1] + stat[0] for x, stat in zip(batch, stats)]
@@ -109,6 +109,8 @@ class TimesFmHparams:
per_core_batch_size: int = 32
backend: Literal["cpu", "gpu", "tpu"] = "cpu"
quantiles: Sequence[float] | None = DEFAULT_QUANTILES
# Hparams beyond the model.
point_forecast_mode: Literal["mean", "median"] = "median"
@dataclasses.dataclass(kw_only=True)
@@ -184,8 +186,9 @@ class TimesFmBase:
"""Loads a checkpoint and compiles the decoder."""
raise NotImplementedError("`load_from_checkpoint` is not implemented.")
def _preprocess(self, inputs: Sequence[np.array],
freq: Sequence[int]) -> tuple[np.array, np.array, int]:
def _preprocess(
self, inputs: Sequence[np.ndarray],
freq: Sequence[int]) -> tuple[np.ndarray, np.ndarray, np.ndarray, int]:
"""Formats and pads raw inputs to feed into the model.
This function both pads each time series to match the context length, and
@@ -200,6 +203,7 @@ class TimesFmBase:
A tuple of:
- the padded input time series to meet the model required context.
- the padding indicator.
- the frequency of each input time series.
- the number of padded examples for SPMD so that each core has the same
number (a multiple of `batch_size`) of examples.
"""
@@ -239,15 +243,14 @@ class TimesFmBase:
pmap_pad,
)
def forecast(
def _forecast(
self,
inputs: Sequence[Any],
freq: Sequence[int] | None = None,
window_size: int | None = None,
forecast_context_len: int | None = None,
return_forecast_on_context: bool = False,
truncate_negative: bool = False,
) -> tuple[np.array, np.array]:
) -> tuple[np.ndarray, np.ndarray]:
"""Forecasts on a list of time series.
Args:
@@ -261,11 +264,9 @@ class TimesFmBase:
forecast_context_len: optional max context length.
return_forecast_on_context: True to return the forecast on the context
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:
A tuple for JTensors:
A tuple for np.array:
- the mean forecast of size (# inputs, # forecast horizon),
- the full forecast (mean + quantiles) of size
(# inputs, # forecast horizon, 1 + # quantiles).
@@ -273,7 +274,78 @@ class TimesFmBase:
Raises:
ValueError: If the checkpoint is not properly loaded.
"""
raise NotImplementedError("`forecast` is not implemented.")
raise NotImplementedError("`_forecast` is not implemented.")
def forecast(
self,
inputs: Sequence[Any],
freq: Sequence[int] | None = None,
window_size: int | None = None,
forecast_context_len: int | None = None,
return_forecast_on_context: bool = False,
normalize: bool = False,
) -> tuple[np.ndarray, np.ndarray]:
"""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.
return_forecast_on_context: True to return the forecast on the context
when available, i.e. after the first input patch.
normalize: If True, then we normalize the inputs before forecasting and
the outputs are then renormalized to the original scale.
Returns:
A tuple for np.array:
- 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.
"""
stats = None
if normalize:
inputs, stats = _normalize(inputs)
mean_forecast, quantile_forecast = self._forecast(
inputs,
freq,
window_size,
forecast_context_len,
return_forecast_on_context,
)
if stats is not None:
stats = np.array(stats)
mu = stats[:, 0]
sigma = stats[:, 1]
mean_forecast = mean_forecast * sigma[:, None] + mu[:, None]
quantile_forecast = (quantile_forecast * sigma[:, None, None] +
mu[:, None, None])
if self.hparams.point_forecast_mode == "mean":
return mean_forecast, quantile_forecast
elif self.hparams.point_forecast_mode == "median":
if self._median_index == -1:
for i, quantile in enumerate(self.quantiles):
if quantile == 0.5:
self._median_index = i
break
if self._median_index == -1:
raise ValueError("Median (0.5) is not found in the model quantiles:"
f" {self.quantiles}. Please check the hparams.")
return (
quantile_forecast[:, :, 1 + self._median_index],
quantile_forecast,
)
else:
raise ValueError(
"Unsupported point forecast mode:"
f" {self.hparams.point_forecast_mode}. Use 'mean' or 'median'.")
def forecast_with_covariates(
self,
@@ -408,7 +480,7 @@ class TimesFmBase:
]
per_instance_stats = None
if normalize_xreg_target_per_input:
targets, per_instance_stats = normalize(targets)
targets, per_instance_stats = _normalize(targets)
xregs = xreg_lib.BatchedInContextXRegLinear(
targets=targets,
train_lens=train_lens,
@@ -431,7 +503,7 @@ class TimesFmBase:
assert_covariate_shapes=True,
)
if normalize_xreg_target_per_input:
xregs = renormalize(xregs, per_instance_stats)
xregs = _renormalize(xregs, per_instance_stats)
outputs = [
(mean_output[self._horizon_start:(self._horizon_start + test_len)] +
xreg)
@@ -446,7 +518,7 @@ class TimesFmBase:
]
per_instance_stats = None
if normalize_xreg_target_per_input:
targets, per_instance_stats = normalize(targets)
targets, per_instance_stats = _normalize(targets)
xregs, xregs_on_context, _, _, _ = xreg_lib.BatchedInContextXRegLinear(
targets=targets,
train_lens=train_lens,
@@ -484,7 +556,7 @@ class TimesFmBase:
for mean_output, test_len, xreg in zip(mean_outputs, test_lens, xregs)
]
if normalize_xreg_target_per_input:
outputs = renormalize(outputs, per_instance_stats)
outputs = _renormalize(outputs, per_instance_stats)
return outputs, xregs
+1 -8
View File
@@ -234,14 +234,13 @@ class TimesFmJax(timesfm_base.TimesFmBase):
}))
self._logging(f"Jitted decoding in {time.time() - start_time:.2f} seconds.")
def forecast(
def _forecast(
self,
inputs: Sequence[Any],
freq: Sequence[int] | None = None,
window_size: int | None = None,
forecast_context_len: int | None = None,
return_forecast_on_context: bool = False,
truncate_negative: bool = False,
) -> tuple[np.ndarray, np.ndarray]:
"""Forecasts on a list of time series.
@@ -256,8 +255,6 @@ class TimesFmJax(timesfm_base.TimesFmBase):
forecast_context_len: optional max context length.
return_forecast_on_context: True to return the forecast on the context
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:
A tuple for JTensors:
@@ -277,7 +274,6 @@ class TimesFmJax(timesfm_base.TimesFmBase):
else:
fcontext_len = forecast_context_len
inputs = [np.array(ts)[-fcontext_len:] for ts in inputs]
inp_min = np.min([np.min(ts) for ts in inputs])
if window_size is not None:
new_inputs = []
@@ -352,7 +348,4 @@ class TimesFmJax(timesfm_base.TimesFmBase):
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 and truncate_negative:
mean_outputs = np.maximum(mean_outputs, 0.0)
full_outputs = np.maximum(full_outputs, 0.0)
return mean_outputs, full_outputs
+2 -8
View File
@@ -46,6 +46,7 @@ class TimesFmTorch(timesfm_base.TimesFmBase):
self.global_batch_size = self.per_core_batch_size
self._device = torch.device("cuda:0" if (
torch.cuda.is_available() and self.backend == "gpu") else "cpu")
self._median_index = -1
def load_from_checkpoint(
self,
@@ -67,14 +68,13 @@ class TimesFmTorch(timesfm_base.TimesFmBase):
self._model.eval()
# TODO: add compilation.
def forecast(
def _forecast(
self,
inputs: Sequence[Any],
freq: Sequence[int] | None = None,
window_size: int | None = None,
forecast_context_len: int | None = None,
return_forecast_on_context: bool = False,
truncate_negative: bool = False,
) -> tuple[np.ndarray, np.ndarray]:
"""Forecasts on a list of time series.
@@ -89,8 +89,6 @@ class TimesFmTorch(timesfm_base.TimesFmBase):
forecast_context_len: optional max context length.
return_forecast_on_context: True to return the forecast on the context
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:
A tuple for JTensors:
@@ -110,7 +108,6 @@ class TimesFmTorch(timesfm_base.TimesFmBase):
else:
fcontext_len = forecast_context_len
inputs = [np.array(ts)[-fcontext_len:] for ts in inputs]
inp_min = np.min([np.min(ts) for ts in inputs])
if window_size is not None:
new_inputs = []
@@ -166,7 +163,4 @@ class TimesFmTorch(timesfm_base.TimesFmBase):
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 and truncate_negative:
mean_outputs = np.maximum(mean_outputs, 0.0)
full_outputs = np.maximum(full_outputs, 0.0)
return mean_outputs, full_outputs