From ca87a438a0341e25acdba7a0cfe658a441fd7850 Mon Sep 17 00:00:00 2001 From: misha-chertushkin Date: Sat, 1 Feb 2025 02:30:17 +0000 Subject: [PATCH] Update examples with new feedback --- notebooks/finetuning_torch.ipynb | 178 ++++++++++++++++++------------- 1 file changed, 104 insertions(+), 74 deletions(-) diff --git a/notebooks/finetuning_torch.ipynb b/notebooks/finetuning_torch.ipynb index 9515e7e..c50dca1 100644 --- a/notebooks/finetuning_torch.ipynb +++ b/notebooks/finetuning_torch.ipynb @@ -44,18 +44,27 @@ "class TimeSeriesDataset(Dataset):\n", " \"\"\"Dataset for time series data compatible with TimesFM.\"\"\"\n", "\n", - " def __init__(self, series: np.ndarray, context_length: int, horizon_length: int):\n", + " def __init__(self,\n", + " series: np.ndarray,\n", + " context_length: int,\n", + " horizon_length: int,\n", + " freq_type: int = 0):\n", " \"\"\"\n", - " Initialize dataset.\n", + " Initialize dataset.\n", + "\n", + " Args:\n", + " series: Time series data\n", + " context_length: Number of past timesteps to use as input\n", + " horizon_length: Number of future timesteps to predict\n", + " freq_type: Frequency type (0, 1, or 2)\n", + " \"\"\"\n", + " if freq_type not in [0, 1, 2]:\n", + " raise ValueError(\"freq_type must be 0, 1, or 2\")\n", "\n", - " Args:\n", - " series: Time series data\n", - " context_length: Number of past timesteps to use as input\n", - " horizon_length: Number of future timesteps to predict\n", - " \"\"\"\n", " self.series = series\n", " self.context_length = context_length\n", " self.horizon_length = horizon_length\n", + " self.freq_type = freq_type\n", " self._prepare_samples()\n", "\n", " def _prepare_samples(self) -> None:\n", @@ -66,47 +75,57 @@ " for start_idx in range(0, len(self.series) - total_length + 1):\n", " end_idx = start_idx + self.context_length\n", " x_context = self.series[start_idx:end_idx]\n", - " x_future = self.series[end_idx : end_idx + self.horizon_length]\n", + " x_future = self.series[end_idx:end_idx + self.horizon_length]\n", " self.samples.append((x_context, x_future))\n", "\n", " def __len__(self) -> int:\n", " return len(self.samples)\n", "\n", - " def __getitem__(self, index: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n", + " def __getitem__(\n", + " self, index: int\n", + " ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:\n", " x_context, x_future = self.samples[index]\n", "\n", " x_context = torch.tensor(x_context, dtype=torch.float32)\n", " x_future = torch.tensor(x_future, dtype=torch.float32)\n", "\n", " input_padding = torch.zeros_like(x_context)\n", - " freq = torch.zeros(1, dtype=torch.long)\n", + " freq = torch.tensor([self.freq_type], dtype=torch.long)\n", "\n", " return x_context, input_padding, freq, x_future\n", "\n", - "\n", - "def prepare_datasets(\n", - " series: np.ndarray, context_length: int, horizon_length: int, train_split: float = 0.8\n", - ") -> Tuple[Dataset, Dataset]:\n", + "def prepare_datasets(series: np.ndarray,\n", + " context_length: int,\n", + " horizon_length: int,\n", + " freq_type: int = 0,\n", + " train_split: float = 0.8) -> Tuple[Dataset, Dataset]:\n", " \"\"\"\n", - " Prepare training and validation datasets from time series data.\n", + " Prepare training and validation datasets from time series data.\n", "\n", - " Args:\n", - " series: Input time series data\n", - " context_length: Number of past timesteps to use\n", - " horizon_length: Number of future timesteps to predict\n", - " train_split: Fraction of data to use for training\n", + " Args:\n", + " series: Input time series data\n", + " context_length: Number of past timesteps to use\n", + " horizon_length: Number of future timesteps to predict\n", + " freq_type: Frequency type (0, 1, or 2)\n", + " train_split: Fraction of data to use for training\n", "\n", - " Returns:\n", - " Tuple of (train_dataset, val_dataset)\n", - " \"\"\"\n", + " Returns:\n", + " Tuple of (train_dataset, val_dataset)\n", + " \"\"\"\n", " train_size = int(len(series) * train_split)\n", " train_data = series[:train_size]\n", " val_data = series[train_size:]\n", "\n", - " # Create datasets\n", - " train_dataset = TimeSeriesDataset(train_data, context_length=context_length, horizon_length=horizon_length)\n", + " # Create datasets with specified frequency type\n", + " train_dataset = TimeSeriesDataset(train_data,\n", + " context_length=context_length,\n", + " horizon_length=horizon_length,\n", + " freq_type=freq_type)\n", "\n", - " val_dataset = TimeSeriesDataset(val_data, context_length=context_length, horizon_length=horizon_length)\n", + " val_dataset = TimeSeriesDataset(val_data,\n", + " context_length=context_length,\n", + " horizon_length=horizon_length,\n", + " freq_type=freq_type)\n", "\n", " return train_dataset, val_dataset\n" ] @@ -128,14 +147,16 @@ " device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n", " repo_id = \"google/timesfm-2.0-500m-pytorch\"\n", " hparams = TimesFmHparams(\n", - " backend=device,\n", - " per_core_batch_size=32,\n", - " horizon_len=128,\n", - " num_layers=50,\n", - " use_positional_embedding=False,\n", - " context_len=192,\n", + " backend=device,\n", + " per_core_batch_size=32,\n", + " horizon_len=128,\n", + " num_layers=50,\n", + " use_positional_embedding=False,\n", + " context_len=\n", + " 192, # Context length can be anything up to 2048 in multiples of 32\n", " )\n", - " tfm = TimesFm(hparams=hparams, checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))\n", + " tfm = TimesFm(hparams=hparams,\n", + " checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))\n", "\n", " model = PatchedTimeSeriesDecoder(tfm._model_config)\n", " if load_weights:\n", @@ -152,18 +173,18 @@ "outputs": [], "source": [ "def plot_predictions(\n", - " model: TimesFm,\n", - " val_dataset: Dataset,\n", - " save_path: Optional[str] = \"predictions.png\",\n", + " model: TimesFm,\n", + " val_dataset: Dataset,\n", + " save_path: Optional[str] = \"predictions.png\",\n", ") -> None:\n", " \"\"\"\n", - " Plot model predictions against ground truth for a batch of validation data.\n", + " Plot model predictions against ground truth for a batch of validation data.\n", "\n", - " Args:\n", - " model: Trained TimesFM model\n", - " val_dataset: Validation dataset\n", - " save_path: Path to save the plot\n", - " \"\"\"\n", + " Args:\n", + " model: Trained TimesFM model\n", + " val_dataset: Validation dataset\n", + " save_path: Path to save the plot\n", + " \"\"\"\n", " import matplotlib.pyplot as plt\n", "\n", " model.eval()\n", @@ -194,18 +215,26 @@ "\n", " plt.figure(figsize=(12, 6))\n", "\n", - " plt.plot(range(context_len), context_vals, label=\"Historical Data\", color=\"blue\", linewidth=2)\n", + " plt.plot(range(context_len),\n", + " context_vals,\n", + " label=\"Historical Data\",\n", + " color=\"blue\",\n", + " linewidth=2)\n", "\n", " plt.plot(\n", - " range(context_len, context_len + horizon_len),\n", - " future_vals,\n", - " label=\"Ground Truth\",\n", - " color=\"green\",\n", - " linestyle=\"--\",\n", - " linewidth=2,\n", + " range(context_len, context_len + horizon_len),\n", + " future_vals,\n", + " label=\"Ground Truth\",\n", + " color=\"green\",\n", + " linestyle=\"--\",\n", + " linewidth=2,\n", " )\n", "\n", - " plt.plot(range(context_len, context_len + horizon_len), pred_vals, label=\"Prediction\", color=\"red\", linewidth=2)\n", + " plt.plot(range(context_len, context_len + horizon_len),\n", + " pred_vals,\n", + " label=\"Prediction\",\n", + " color=\"red\",\n", + " linewidth=2)\n", "\n", " plt.xlabel(\"Time Step\")\n", " plt.ylabel(\"Value\")\n", @@ -218,43 +247,44 @@ " print(f\"Plot saved to {save_path}\")\n", "\n", " plt.close()\n", - "\n", - "\n", - "def get_data(context_len: int, horizon_len: int) -> Tuple[Dataset, Dataset]:\n", - " df = yf.download(\"AAPL\", start=\"2010-01-01\", end=\"2019-01-01\")\n", - " time_series = df[\"Close\"].values\n", - "\n", - " train_dataset, val_dataset = prepare_datasets(\n", - " series=time_series,\n", - " context_length=context_len,\n", - " horizon_length=horizon_len,\n", - " train_split=0.8,\n", - " )\n", - "\n", - " print(f\"Created datasets:\")\n", - " print(f\"- Training samples: {len(train_dataset)}\")\n", - " print(f\"- Validation samples: {len(val_dataset)}\")\n", - " return train_dataset, val_dataset\n", - "\n", + "\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ "\n", "def single_gpu_example():\n", " \"\"\"Basic example of finetuning TimesFM on stock data.\"\"\"\n", " model, hparams, tfm_config = get_model(load_weights=True)\n", - " config = FinetuningConfig(batch_size=256, num_epochs=5, learning_rate=1e-4, use_wandb=True)\n", + " config = FinetuningConfig(batch_size=256,\n", + " num_epochs=5,\n", + " learning_rate=1e-4,\n", + " use_wandb=True,\n", + " freq_type=1,\n", + " log_every_n_steps=10,\n", + " val_check_interval=0.5,\n", + " use_quantile_loss=True)\n", "\n", - " train_dataset, val_dataset = get_data(128, tfm_config.horizon_len)\n", + " train_dataset, val_dataset = get_data(128,\n", + " tfm_config.horizon_len,\n", + " freq_type=config.freq_type)\n", " finetuner = TimesFMFinetuner(model, config)\n", "\n", " print(\"\\nStarting finetuning...\")\n", - " results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset)\n", + " results = finetuner.finetune(train_dataset=train_dataset,\n", + " val_dataset=val_dataset)\n", "\n", " print(\"\\nFinetuning completed!\")\n", " print(f\"Training history: {len(results['history']['train_loss'])} epochs\")\n", "\n", " plot_predictions(\n", - " model=model,\n", - " val_dataset=val_dataset,\n", - " save_path=\"timesfm_predictions.png\",\n", + " model=model,\n", + " val_dataset=val_dataset,\n", + " save_path=\"timesfm_predictions.png\",\n", " )\n" ] },