add support for inference sans positional encoding

This commit is contained in:
Yichen Zhou
2024-12-12 11:04:21 -08:00
parent 02bc2f2212
commit 27f40371c2
4 changed files with 16 additions and 10 deletions
+2
View File
@@ -237,6 +237,7 @@ class PatchedTimeSeriesDecoder(base_layer.BaseLayer):
stacked_transformer_params_tpl: LayerTpl = template_field( stacked_transformer_params_tpl: LayerTpl = template_field(
transformers.StackedTransformer) transformers.StackedTransformer)
use_freq: bool = True use_freq: bool = True
use_pos_emb: bool = True
def setup(self) -> None: def setup(self) -> None:
"""Construct the model.""" """Construct the model."""
@@ -333,6 +334,7 @@ class PatchedTimeSeriesDecoder(base_layer.BaseLayer):
# A patch should not be padded even if there is at least one zero. # A patch should not be padded even if there is at least one zero.
patched_padding = jnp.min(patched_pads, axis=-1) patched_padding = jnp.min(patched_pads, axis=-1)
if self.use_pos_emb:
if pos_emb is None: if pos_emb is None:
position_emb = self.position_emb(seq_length=model_input.shape[1]) position_emb = self.position_emb(seq_length=model_input.shape[1])
else: else:
+2
View File
@@ -109,6 +109,7 @@ class TimesFmHparams:
per_core_batch_size: int = 32 per_core_batch_size: int = 32
backend: Literal["cpu", "gpu", "tpu"] = "cpu" backend: Literal["cpu", "gpu", "tpu"] = "cpu"
quantiles: Sequence[float] | None = DEFAULT_QUANTILES quantiles: Sequence[float] | None = DEFAULT_QUANTILES
use_positional_embedding: bool = True
# Hparams beyond the model. # Hparams beyond the model.
point_forecast_mode: Literal["mean", "median"] = "median" point_forecast_mode: Literal["mean", "median"] = "median"
@@ -172,6 +173,7 @@ class TimesFmBase:
self.backend = hparams.backend self.backend = hparams.backend
self.quantiles = hparams.quantiles self.quantiles = hparams.quantiles
self.num_heads = hparams.num_heads self.num_heads = hparams.num_heads
self.use_pos_emb = hparams.use_positional_embedding
# Rewrite these values in __post_init__ for SPMD. # Rewrite these values in __post_init__ for SPMD.
self.num_cores = 1 self.num_cores = 1
+1
View File
@@ -117,6 +117,7 @@ class TimesFmJax(timesfm_base.TimesFmBase):
residual_block_tpl=pax_fiddle.Config(patched_decoder.ResidualBlock), residual_block_tpl=pax_fiddle.Config(patched_decoder.ResidualBlock),
quantiles=self.quantiles, quantiles=self.quantiles,
use_freq=True, use_freq=True,
use_pos_emb=self.use_pos_emb,
stacked_transformer_params_tpl=pax_fiddle.Config( stacked_transformer_params_tpl=pax_fiddle.Config(
transformers.StackedTransformer, transformers.StackedTransformer,
num_heads=self.num_heads, num_heads=self.num_heads,
+1
View File
@@ -40,6 +40,7 @@ class TimesFmTorch(timesfm_base.TimesFmBase):
horizon_len=self.output_patch_len, horizon_len=self.output_patch_len,
head_dim=self.model_dims // self.num_heads, head_dim=self.model_dims // self.num_heads,
quantiles=self.quantiles, quantiles=self.quantiles,
use_positional_embedding=self.use_pos_emb,
) )
self._model = None self._model = None
self.num_cores = 1 self.num_cores = 1