From 86551f761af790d70e9869beb73939a313666c03 Mon Sep 17 00:00:00 2001 From: misha-chertushkin Date: Tue, 21 Jan 2025 19:57:50 +0000 Subject: [PATCH] Add jupyter notebook --- notebooks/finetuning_example.py | 2 +- notebooks/finetuning_torch.ipynb | 292 ++++++++++++++++++ .../timesfm}/finetuning_torch.py | 0 3 files changed, 293 insertions(+), 1 deletion(-) create mode 100644 notebooks/finetuning_torch.ipynb rename {notebooks => src/timesfm}/finetuning_torch.py (100%) diff --git a/notebooks/finetuning_example.py b/notebooks/finetuning_example.py index 7ffb53d..5f166b3 100644 --- a/notebooks/finetuning_example.py +++ b/notebooks/finetuning_example.py @@ -10,7 +10,7 @@ import pandas as pd import torch import torch.multiprocessing as mp import yfinance as yf -from finetuning_torch import FinetuningConfig, TimesFMFinetuner +from timesfm.finetuning_torch import FinetuningConfig, TimesFMFinetuner from huggingface_hub import snapshot_download from torch.utils.data import Dataset diff --git a/notebooks/finetuning_torch.ipynb b/notebooks/finetuning_torch.ipynb new file mode 100644 index 0000000..9515e7e --- /dev/null +++ b/notebooks/finetuning_torch.ipynb @@ -0,0 +1,292 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Introduction\n", + "This notebook shows how to use TimesFM with finetuning. \n", + "\n", + "In order to perform finetuning, you need to create the Pytorch Dataset in a proper format. The example of the Dataset is provided below.\n", + "The finetuning code can be found in timesfm.finetuning_torch.py. This notebook just imports the methods from finetuning" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Dataset Creation" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "from os import path\n", + "from typing import Optional, Tuple\n", + "\n", + "import numpy as np\n", + "import pandas as pd\n", + "import torch\n", + "import torch.multiprocessing as mp\n", + "import yfinance as yf\n", + "from timesfm.finetuning_torch import FinetuningConfig, TimesFMFinetuner\n", + "from huggingface_hub import snapshot_download\n", + "from torch.utils.data import Dataset\n", + "\n", + "from timesfm import TimesFm, TimesFmCheckpoint, TimesFmHparams\n", + "from timesfm.pytorch_patched_decoder import PatchedTimeSeriesDecoder\n", + "import os\n", + "\n", + "\n", + "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", + " \"\"\"\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", + " \"\"\"\n", + " self.series = series\n", + " self.context_length = context_length\n", + " self.horizon_length = horizon_length\n", + " self._prepare_samples()\n", + "\n", + " def _prepare_samples(self) -> None:\n", + " \"\"\"Prepare sliding window samples from the time series.\"\"\"\n", + " self.samples = []\n", + " total_length = self.context_length + self.horizon_length\n", + "\n", + " 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", + " 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", + " 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", + "\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", + " \"\"\"\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", + "\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", + "\n", + " val_dataset = TimeSeriesDataset(val_data, context_length=context_length, horizon_length=horizon_length)\n", + "\n", + " return train_dataset, val_dataset\n" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "### Model Creation" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "def get_model(load_weights: bool = False):\n", + " 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", + " )\n", + " tfm = TimesFm(hparams=hparams, checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))\n", + "\n", + " model = PatchedTimeSeriesDecoder(tfm._model_config)\n", + " if load_weights:\n", + " checkpoint_path = path.join(snapshot_download(repo_id), \"torch_model.ckpt\")\n", + " loaded_checkpoint = torch.load(checkpoint_path, weights_only=True)\n", + " model.load_state_dict(loaded_checkpoint)\n", + " return model, hparams, tfm._model_config\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "def plot_predictions(\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", + "\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", + "\n", + " x_context, x_padding, freq, x_future = val_dataset[0]\n", + " x_context = x_context.unsqueeze(0) # Add batch dimension\n", + " x_padding = x_padding.unsqueeze(0)\n", + " freq = freq.unsqueeze(0)\n", + " x_future = x_future.unsqueeze(0)\n", + "\n", + " device = next(model.parameters()).device\n", + " x_context = x_context.to(device)\n", + " x_padding = x_padding.to(device)\n", + " freq = freq.to(device)\n", + " x_future = x_future.to(device)\n", + "\n", + " with torch.no_grad():\n", + " predictions = model(x_context, x_padding.float(), freq)\n", + " predictions_mean = predictions[..., 0] # [B, N, horizon_len]\n", + " last_patch_pred = predictions_mean[:, -1, :] # [B, horizon_len]\n", + "\n", + " context_vals = x_context[0].cpu().numpy()\n", + " future_vals = x_future[0].cpu().numpy()\n", + " pred_vals = last_patch_pred[0].cpu().numpy()\n", + "\n", + " context_len = len(context_vals)\n", + " horizon_len = len(future_vals)\n", + "\n", + " plt.figure(figsize=(12, 6))\n", + "\n", + " plt.plot(range(context_len), context_vals, label=\"Historical Data\", color=\"blue\", 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", + " )\n", + "\n", + " plt.plot(range(context_len, context_len + horizon_len), pred_vals, label=\"Prediction\", color=\"red\", linewidth=2)\n", + "\n", + " plt.xlabel(\"Time Step\")\n", + " plt.ylabel(\"Value\")\n", + " plt.title(\"TimesFM Predictions vs Ground Truth\")\n", + " plt.legend()\n", + " plt.grid(True)\n", + "\n", + " if save_path:\n", + " plt.savefig(save_path)\n", + " 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", + "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", + "\n", + " train_dataset, val_dataset = get_data(128, tfm_config.horizon_len)\n", + " finetuner = TimesFMFinetuner(model, config)\n", + "\n", + " print(\"\\nStarting finetuning...\")\n", + " results = finetuner.finetune(train_dataset=train_dataset, 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", + " )\n" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "single_gpu_example()" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "timesfm-311", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.11.11" + } + }, + "nbformat": 4, + "nbformat_minor": 2 +} diff --git a/notebooks/finetuning_torch.py b/src/timesfm/finetuning_torch.py similarity index 100% rename from notebooks/finetuning_torch.py rename to src/timesfm/finetuning_torch.py