From 9080939cb82ebdf4da2345b943ada472ea3b5b70 Mon Sep 17 00:00:00 2001 From: siriuz42 Date: Wed, 8 Oct 2025 16:32:41 +0000 Subject: [PATCH] Manually trigger jitting in Flax. --- src/timesfm/timesfm_2p5/timesfm_2p5_flax.py | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/src/timesfm/timesfm_2p5/timesfm_2p5_flax.py b/src/timesfm/timesfm_2p5/timesfm_2p5_flax.py index acd354b..51d6a3c 100644 --- a/src/timesfm/timesfm_2p5/timesfm_2p5_flax.py +++ b/src/timesfm/timesfm_2p5/timesfm_2p5_flax.py @@ -493,6 +493,8 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5): def compile(self, forecast_config: configs.ForecastConfig, **kwargs): # Acrobym used during validation. + print("Compiling model...") + fc = forecast_config if fc.max_context % self.model.p != 0: logging.info( @@ -581,3 +583,14 @@ class TimesFM_2p5_200M_flax(timesfm_2p5_base.TimesFM_2p5): self.compiled_decode = functools.partial( compiled_decode_kernel, self.forecast_config ) + + _ = self.compiled_decode( + self.forecast_config.max_horizon, + jnp.zeros( + (self.global_batch_size, self.forecast_config.max_context), dtype=jnp.float32 + ), + jnp.zeros( + (self.global_batch_size, self.forecast_config.max_context), dtype=jnp.bool + ), + ) + print("Compiling done.")