From c0b93d4b6aa645a2397decba13927141c44306c8 Mon Sep 17 00:00:00 2001 From: periodLeo Date: Tue, 9 Jul 2024 20:16:43 +0900 Subject: [PATCH] added a verbose parameter to TimesFM.forecast_on_df(), in order to make it possible to desactivate print outputs. --- src/timesfm/timesfm.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/src/timesfm/timesfm.py b/src/timesfm/timesfm.py index 5dc7611..5bf3261 100644 --- a/src/timesfm/timesfm.py +++ b/src/timesfm/timesfm.py @@ -518,6 +518,7 @@ class TimesFm: model_name: str = "timesfm", window_size: int | None = None, num_jobs: int = 1, + verbose: bool = True, ) -> pd.DataFrame: """Forecasts on a list of time series. @@ -535,6 +536,7 @@ class TimesFm: 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. + verbose: output model states in terminal. Returns: Future forecasts dataframe. @@ -554,7 +556,8 @@ class TimesFm: new_inputs = [] uids = [] if num_jobs == 1: - print("Processing dataframe with single process.") + if verbose: + print("Processing dataframe with single process.") for key, group in df_sorted.groupby("unique_id"): inp, uid = process_group( key, @@ -567,7 +570,8 @@ class TimesFm: else: if num_jobs == -1: num_jobs = multiprocessing.cpu_count() - print("Processing dataframe with multiple processes.") + if verbose: + print("Processing dataframe with multiple processes.") with multiprocessing.Pool(processes=num_jobs) as pool: results = pool.starmap( process_group, @@ -577,12 +581,14 @@ class TimesFm: ], ) new_inputs, uids = zip(*results) - print("Finished preprocessing dataframe.") + if verbose: + 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.") + if verbose: + print("Finished forecasting.") fcst_df = make_future_dataframe( uids=uids, last_times=df_sorted.groupby("unique_id")["ds"].tail(1),