diff --git a/notebooks/finetuning_torch.ipynb b/notebooks/finetuning_torch.ipynb new file mode 100644 index 0000000..c50dca1 --- /dev/null +++ b/notebooks/finetuning_torch.ipynb @@ -0,0 +1,322 @@ +{ + "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,\n", + " series: np.ndarray,\n", + " context_length: int,\n", + " horizon_length: int,\n", + " freq_type: int = 0):\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", + " 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", + " 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", + " \"\"\"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__(\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.tensor([self.freq_type], dtype=torch.long)\n", + "\n", + " return x_context, input_padding, freq, x_future\n", + "\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", + "\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", + " train_size = int(len(series) * train_split)\n", + " train_data = series[:train_size]\n", + " val_data = series[train_size:]\n", + "\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,\n", + " context_length=context_length,\n", + " horizon_length=horizon_length,\n", + " freq_type=freq_type)\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=\n", + " 192, # Context length can be anything up to 2048 in multiples of 32\n", + " )\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", + " 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),\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", + " )\n", + "\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", + " 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" + ] + }, + { + "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,\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,\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,\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", + " )\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/pyproject.toml b/pyproject.toml index f0238c0..e7163e4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -9,6 +9,7 @@ authors = [ "Abhimanyu Das ", "Petros Mol ", "Justin Güse ", + "Michael Chertushkin " ] readme = "README.md" keywords = ["time series", "timesfm", "forecast", "time series model"] diff --git a/src/finetuning/__init__.py b/src/finetuning/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/finetuning/finetuning_example.py b/src/finetuning/finetuning_example.py new file mode 100644 index 0000000..2c396d6 --- /dev/null +++ b/src/finetuning/finetuning_example.py @@ -0,0 +1,388 @@ +""" +Example usage of the TimesFM Finetuning Framework. + +For single GPU: +python script.py --training_mode=single + +For multiple GPUs: +python script.py --training_mode=multi --gpu_ids=0,1,2 +""" + +import os +from os import path +from typing import Optional, Tuple + +import numpy as np +import pandas as pd +import torch +import torch.multiprocessing as mp +import yfinance as yf +from absl import app, flags +from huggingface_hub import snapshot_download +from torch.utils.data import Dataset + +from finetuning.finetuning_torch import FinetuningConfig, TimesFMFinetuner +from timesfm import TimesFm, TimesFmCheckpoint, TimesFmHparams +from timesfm.pytorch_patched_decoder import PatchedTimeSeriesDecoder + +FLAGS = flags.FLAGS + +flags.DEFINE_enum( + "training_mode", + "single", + ["single", "multi"], + 'Training mode: "single" for single-GPU or "multi" for multi-GPU training.', +) + +flags.DEFINE_list( + "gpu_ids", ["0"], + "Comma-separated list of GPU IDs to use for multi-GPU training. Example: 0,1,2" +) + + +class TimeSeriesDataset(Dataset): + """Dataset for time series data compatible with TimesFM.""" + + def __init__(self, + series: np.ndarray, + context_length: int, + horizon_length: int, + freq_type: int = 0): + """ + Initialize dataset. + + Args: + series: Time series data + context_length: Number of past timesteps to use as input + horizon_length: Number of future timesteps to predict + freq_type: Frequency type (0, 1, or 2) + """ + if freq_type not in [0, 1, 2]: + raise ValueError("freq_type must be 0, 1, or 2") + + self.series = series + self.context_length = context_length + self.horizon_length = horizon_length + self.freq_type = freq_type + self._prepare_samples() + + def _prepare_samples(self) -> None: + """Prepare sliding window samples from the time series.""" + self.samples = [] + total_length = self.context_length + self.horizon_length + + for start_idx in range(0, len(self.series) - total_length + 1): + end_idx = start_idx + self.context_length + x_context = self.series[start_idx:end_idx] + x_future = self.series[end_idx:end_idx + self.horizon_length] + self.samples.append((x_context, x_future)) + + def __len__(self) -> int: + return len(self.samples) + + def __getitem__( + self, index: int + ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + x_context, x_future = self.samples[index] + + x_context = torch.tensor(x_context, dtype=torch.float32) + x_future = torch.tensor(x_future, dtype=torch.float32) + + input_padding = torch.zeros_like(x_context) + freq = torch.tensor([self.freq_type], dtype=torch.long) + + return x_context, input_padding, freq, x_future + + +def prepare_datasets(series: np.ndarray, + context_length: int, + horizon_length: int, + freq_type: int = 0, + train_split: float = 0.8) -> Tuple[Dataset, Dataset]: + """ + Prepare training and validation datasets from time series data. + + Args: + series: Input time series data + context_length: Number of past timesteps to use + horizon_length: Number of future timesteps to predict + freq_type: Frequency type (0, 1, or 2) + train_split: Fraction of data to use for training + + Returns: + Tuple of (train_dataset, val_dataset) + """ + train_size = int(len(series) * train_split) + train_data = series[:train_size] + val_data = series[train_size:] + + # Create datasets with specified frequency type + train_dataset = TimeSeriesDataset(train_data, + context_length=context_length, + horizon_length=horizon_length, + freq_type=freq_type) + + val_dataset = TimeSeriesDataset(val_data, + context_length=context_length, + horizon_length=horizon_length, + freq_type=freq_type) + + return train_dataset, val_dataset + + +def get_model(load_weights: bool = False): + device = "cuda" if torch.cuda.is_available() else "cpu" + repo_id = "google/timesfm-2.0-500m-pytorch" + hparams = TimesFmHparams( + backend=device, + per_core_batch_size=32, + horizon_len=128, + num_layers=50, + use_positional_embedding=False, + context_len= + 192, # Context length can be anything up to 2048 in multiples of 32 + ) + tfm = TimesFm(hparams=hparams, + checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id)) + + model = PatchedTimeSeriesDecoder(tfm._model_config) + if load_weights: + checkpoint_path = path.join(snapshot_download(repo_id), "torch_model.ckpt") + loaded_checkpoint = torch.load(checkpoint_path, weights_only=True) + model.load_state_dict(loaded_checkpoint) + return model, hparams, tfm._model_config + + +def plot_predictions( + model: TimesFm, + val_dataset: Dataset, + save_path: Optional[str] = "predictions.png", +) -> None: + """ + Plot model predictions against ground truth for a batch of validation data. + + Args: + model: Trained TimesFM model + val_dataset: Validation dataset + save_path: Path to save the plot + """ + import matplotlib.pyplot as plt + + model.eval() + + x_context, x_padding, freq, x_future = val_dataset[0] + x_context = x_context.unsqueeze(0) # Add batch dimension + x_padding = x_padding.unsqueeze(0) + freq = freq.unsqueeze(0) + x_future = x_future.unsqueeze(0) + + device = next(model.parameters()).device + x_context = x_context.to(device) + x_padding = x_padding.to(device) + freq = freq.to(device) + x_future = x_future.to(device) + + with torch.no_grad(): + predictions = model(x_context, x_padding.float(), freq) + predictions_mean = predictions[..., 0] # [B, N, horizon_len] + last_patch_pred = predictions_mean[:, -1, :] # [B, horizon_len] + + context_vals = x_context[0].cpu().numpy() + future_vals = x_future[0].cpu().numpy() + pred_vals = last_patch_pred[0].cpu().numpy() + + context_len = len(context_vals) + horizon_len = len(future_vals) + + plt.figure(figsize=(12, 6)) + + plt.plot(range(context_len), + context_vals, + label="Historical Data", + color="blue", + linewidth=2) + + plt.plot( + range(context_len, context_len + horizon_len), + future_vals, + label="Ground Truth", + color="green", + linestyle="--", + linewidth=2, + ) + + plt.plot(range(context_len, context_len + horizon_len), + pred_vals, + label="Prediction", + color="red", + linewidth=2) + + plt.xlabel("Time Step") + plt.ylabel("Value") + plt.title("TimesFM Predictions vs Ground Truth") + plt.legend() + plt.grid(True) + + if save_path: + plt.savefig(save_path) + print(f"Plot saved to {save_path}") + + plt.close() + + +def get_data(context_len: int, + horizon_len: int, + freq_type: int = 0) -> Tuple[Dataset, Dataset]: + df = yf.download("AAPL", start="2010-01-01", end="2019-01-01") + time_series = df["Close"].values + + train_dataset, val_dataset = prepare_datasets( + series=time_series, + context_length=context_len, + horizon_length=horizon_len, + freq_type=freq_type, + train_split=0.8, + ) + + print(f"Created datasets:") + print(f"- Training samples: {len(train_dataset)}") + print(f"- Validation samples: {len(val_dataset)}") + print(f"- Using frequency type: {freq_type}") + return train_dataset, val_dataset + + +def single_gpu_example(): + """Basic example of finetuning TimesFM on stock data.""" + model, hparams, tfm_config = get_model(load_weights=True) + config = FinetuningConfig(batch_size=256, + num_epochs=5, + learning_rate=1e-4, + use_wandb=True, + freq_type=1, + log_every_n_steps=10, + val_check_interval=0.5, + use_quantile_loss=True) + + train_dataset, val_dataset = get_data(128, + tfm_config.horizon_len, + freq_type=config.freq_type) + finetuner = TimesFMFinetuner(model, config) + + print("\nStarting finetuning...") + results = finetuner.finetune(train_dataset=train_dataset, + val_dataset=val_dataset) + + print("\nFinetuning completed!") + print(f"Training history: {len(results['history']['train_loss'])} epochs") + + plot_predictions( + model=model, + val_dataset=val_dataset, + save_path="timesfm_predictions.png", + ) + + +def setup_process(rank, world_size, model, config, train_dataset, val_dataset, + return_dict): + """Setup process function with optimized CUDA handling.""" + try: + if torch.cuda.is_available(): + torch.cuda.set_device(rank) + + os.environ["MASTER_ADDR"] = config.master_addr + os.environ["MASTER_PORT"] = config.master_port + if not torch.distributed.is_initialized(): + torch.distributed.init_process_group(backend="nccl", + world_size=world_size, + rank=rank) + + finetuner = TimesFMFinetuner(model, config, rank=rank) + + results = finetuner.finetune(train_dataset=train_dataset, + val_dataset=val_dataset) + + if rank == 0: + return_dict["results"] = results + plot_predictions( + model=model, + val_dataset=val_dataset, + save_path="timesfm_predictions.png", + ) + + except Exception as e: + print(f"Error in process {rank}: {str(e)}") + raise e + finally: + if torch.distributed.is_initialized(): + torch.distributed.destroy_process_group() + + +def multi_gpu_example(): + """Example of finetuning TimesFM using multiple GPUs with optimized spawn.""" + mp.set_start_method("spawn", force=True) + + gpu_ids = [0, 1] + world_size = len(gpu_ids) + + model, hparams, tfm_config = get_model(load_weights=True) + + # Create config + config = FinetuningConfig( + batch_size=256, + num_epochs=5, + learning_rate=3e-5, + use_wandb=True, + distributed=True, + gpu_ids=gpu_ids, + log_every_n_steps=50, + val_check_interval=0.5, + ) + train_dataset, val_dataset = get_data(128, tfm_config.horizon_len) + manager = mp.Manager() + return_dict = manager.dict() + + # Launch processes + mp.spawn( + setup_process, + args=(world_size, model, config, train_dataset, val_dataset, return_dict), + nprocs=world_size, + join=True, + ) + + results = return_dict.get("results", None) + print("\nFinetuning completed!") + return results + + +def main(argv): + """Main function that selects and runs the appropriate training mode.""" + + try: + if FLAGS.training_mode == "single": + print("\nStarting single-GPU training...") + single_gpu_example() + else: + gpu_ids = [int(id) for id in FLAGS.gpu_ids] + print(f"\nStarting multi-GPU training using GPUs: {gpu_ids}...") + + config = FinetuningConfig( + batch_size=256, + num_epochs=5, + learning_rate=3e-5, + use_wandb=True, + distributed=True, + gpu_ids=gpu_ids, + ) + + results = multi_gpu_example(config) + print("\nMulti-GPU training completed!") + + except Exception as e: + print(f"Training failed: {str(e)}") + finally: + if torch.distributed.is_initialized(): + torch.distributed.destroy_process_group() + + +if __name__ == "__main__": + app.run(main) diff --git a/src/finetuning/finetuning_torch.py b/src/finetuning/finetuning_torch.py new file mode 100644 index 0000000..5c2d8b3 --- /dev/null +++ b/src/finetuning/finetuning_torch.py @@ -0,0 +1,399 @@ +""" +TimesFM Finetuner: A flexible framework for finetuning TimesFM models on custom datasets. +""" + +import logging +import os +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from typing import Any, Callable, Dict, List, Optional + +import torch +import torch.distributed as dist +import torch.nn as nn +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.utils.data import DataLoader, Dataset +from timesfm.pytorch_patched_decoder import create_quantiles + +import wandb + + +class MetricsLogger(ABC): + """Abstract base class for logging metrics during training. + + This class defines the interface for logging metrics during model training. + Concrete implementations can log to different backends (e.g., WandB, TensorBoard). + """ + + @abstractmethod + def log_metrics(self, + metrics: Dict[str, Any], + step: Optional[int] = None) -> None: + """Log metrics to the specified backend. + + Args: + metrics: Dictionary containing metric names and values. + step: Optional step number or epoch for the metrics. + """ + pass + + @abstractmethod + def close(self) -> None: + """Clean up any resources used by the logger.""" + pass + + +class WandBLogger(MetricsLogger): + """Weights & Biases implementation of metrics logging. + + Args: + project: Name of the W&B project. + config: Configuration dictionary to log. + rank: Process rank in distributed training. + """ + + def __init__(self, project: str, config: Dict[str, Any], rank: int = 0): + self.rank = rank + if rank == 0: + wandb.init(project=project, config=config) + + def log_metrics(self, + metrics: Dict[str, Any], + step: Optional[int] = None) -> None: + """Log metrics to W&B if on the main process. + + Args: + metrics: Dictionary of metrics to log. + step: Current training step or epoch. + """ + if self.rank == 0: + wandb.log(metrics, step=step) + + def close(self) -> None: + """Finish the W&B run if on the main process.""" + if self.rank == 0: + wandb.finish() + + +class DistributedManager: + """Manages distributed training setup and cleanup. + + Args: + world_size: Total number of processes. + rank: Process rank. + master_addr: Address of the master process. + master_port: Port for distributed communication. + backend: PyTorch distributed backend to use. + """ + + def __init__( + self, + world_size: int, + rank: int, + master_addr: str = "localhost", + master_port: str = "12358", + backend: str = "nccl", + ): + self.world_size = world_size + self.rank = rank + self.master_addr = master_addr + self.master_port = master_port + self.backend = backend + + def setup(self) -> None: + """Initialize the distributed environment.""" + os.environ["MASTER_ADDR"] = self.master_addr + os.environ["MASTER_PORT"] = self.master_port + + if not dist.is_initialized(): + dist.init_process_group(backend=self.backend, + world_size=self.world_size, + rank=self.rank) + + def cleanup(self) -> None: + """Clean up the distributed environment.""" + if dist.is_initialized(): + dist.destroy_process_group() + + +@dataclass +class FinetuningConfig: + """Configuration for model training. + + Args: + batch_size: Number of samples per batch. + num_epochs: Number of training epochs. + learning_rate: Initial learning rate. + weight_decay: L2 regularization factor. + freq_type: Frequency, can be [0, 1, 2]. + use_quantile_loss: bool = False # Flag to enable/disable quantile loss + quantiles: Optional[List[float]] = None + device: Device to train on ('cuda' or 'cpu'). + distributed: Whether to use distributed training. + gpu_ids: List of GPU IDs to use. + master_port: Port for distributed training. + master_addr: Address for distributed training. + use_wandb: Whether to use Weights & Biases logging. + wandb_project: W&B project name. + log_every_n_steps: Log metrics every N steps (batches), this is inspired from Pytorch Lightning + val_check_interval: How often within one training epoch to check val metrics. (also from Pytorch Lightning) + Can be: float (0.0-1.0): fraction of epoch (e.g., 0.5 = validate twice per epoch) + int: validate every N batches + """ + + batch_size: int = 32 + num_epochs: int = 20 + learning_rate: float = 1e-4 + weight_decay: float = 0.01 + freq_type: int = 0 + use_quantile_loss: bool = False + quantiles: Optional[List[float]] = None + device: str = "cuda" if torch.cuda.is_available() else "cpu" + distributed: bool = False + gpu_ids: List[int] = field(default_factory=lambda: [0]) + master_port: str = "12358" + master_addr: str = "localhost" + use_wandb: bool = False + wandb_project: str = "timesfm-finetuning" + log_every_n_steps: int = 50 + val_check_interval: float = 0.5 + + +class TimesFMFinetuner: + """Handles model training and validation. + + Args: + model: PyTorch model to train. + config: Training configuration. + rank: Process rank for distributed training. + loss_fn: Loss function (defaults to MSE). + logger: Optional logging.Logger instance. + """ + + def __init__( + self, + model: nn.Module, + config: FinetuningConfig, + rank: int = 0, + loss_fn: Optional[Callable] = None, + logger: Optional[logging.Logger] = None, + ): + self.model = model + self.config = config + self.rank = rank + self.logger = logger or logging.getLogger(__name__) + self.device = torch.device( + f"cuda:{rank}" if torch.cuda.is_available() else "cpu") + self.loss_fn = loss_fn or (lambda x, y: torch.mean((x - y.squeeze(-1))**2)) + + if config.use_wandb: + self.metrics_logger = WandBLogger(config.wandb_project, config.__dict__, + rank) + + if config.distributed: + self.dist_manager = DistributedManager( + world_size=len(config.gpu_ids), + rank=rank, + master_addr=config.master_addr, + master_port=config.master_port, + ) + self.dist_manager.setup() + self.model = self._setup_distributed_model() + + def _setup_distributed_model(self) -> nn.Module: + """Configure model for distributed training.""" + self.model = self.model.to(self.device) + return DDP(self.model, + device_ids=[self.config.gpu_ids[self.rank]], + output_device=self.config.gpu_ids[self.rank]) + + def _create_dataloader(self, dataset: Dataset, is_train: bool) -> DataLoader: + """Create appropriate DataLoader based on training configuration. + + Args: + dataset: Dataset to create loader for. + is_train: Whether this is for training (affects shuffling). + + Returns: + DataLoader instance. + """ + if self.config.distributed: + sampler = torch.utils.data.distributed.DistributedSampler( + dataset, + num_replicas=len(self.config.gpu_ids), + rank=dist.get_rank(), + shuffle=is_train) + else: + sampler = None + + return DataLoader( + dataset, + batch_size=self.config.batch_size, + shuffle=(is_train and not self.config.distributed), + sampler=sampler, + ) + + def _quantile_loss(self, pred: torch.Tensor, actual: torch.Tensor, + quantile: float) -> torch.Tensor: + """Calculates quantile loss. + Args: + pred: Predicted values + actual: Actual values + quantile: Quantile at which loss is computed + Returns: + Quantile loss + """ + dev = actual - pred + loss_first = dev * quantile + loss_second = -dev * (1.0 - quantile) + return 2 * torch.where(loss_first >= 0, loss_first, loss_second) + + def _process_batch(self, batch: List[torch.Tensor]) -> tuple: + """Process a single batch of data. + + Args: + batch: List of input tensors. + + Returns: + Tuple of (loss, predictions). + """ + x_context, x_padding, freq, x_future = [ + t.to(self.device, non_blocking=True) for t in batch + ] + + predictions = self.model(x_context, x_padding.float(), freq) + predictions_mean = predictions[..., 0] + last_patch_pred = predictions_mean[:, -1, :] + + loss = self.loss_fn(last_patch_pred, x_future.squeeze(-1)) + if self.config.use_quantile_loss: + quantiles = self.config.quantiles or create_quantiles() + for i, quantile in enumerate(quantiles): + last_patch_quantile = predictions[:, -1, :, i + 1] + loss += torch.mean( + self._quantile_loss(last_patch_quantile, x_future.squeeze(-1), + quantile)) + + return loss, predictions + + def _train_epoch(self, train_loader: DataLoader, + optimizer: torch.optim.Optimizer) -> float: + """Train for one epoch in a distributed setting. + + Args: + train_loader: DataLoader for training data. + optimizer: Optimizer instance. + + Returns: + Average training loss for the epoch. + """ + self.model.train() + total_loss = 0.0 + num_batches = len(train_loader) + + for batch in train_loader: + loss, _ = self._process_batch(batch) + + optimizer.zero_grad() + loss.backward() + optimizer.step() + + total_loss += loss.item() + + avg_loss = total_loss / num_batches + + if self.config.distributed: + avg_loss_tensor = torch.tensor(avg_loss, device=self.device) + dist.all_reduce(avg_loss_tensor, op=dist.ReduceOp.SUM) + avg_loss = (avg_loss_tensor / dist.get_world_size()).item() + + return avg_loss + + def _validate(self, val_loader: DataLoader) -> float: + """Perform validation. + + Args: + val_loader: DataLoader for validation data. + + Returns: + Average validation loss. + """ + self.model.eval() + total_loss = 0.0 + num_batches = len(val_loader) + + with torch.no_grad(): + for batch in val_loader: + loss, _ = self._process_batch(batch) + total_loss += loss.item() + + avg_loss = total_loss / num_batches + + if self.config.distributed: + avg_loss_tensor = torch.tensor(avg_loss, device=self.device) + dist.all_reduce(avg_loss_tensor, op=dist.ReduceOp.SUM) + avg_loss = (avg_loss_tensor / dist.get_world_size()).item() + + return avg_loss + + def finetune(self, train_dataset: Dataset, + val_dataset: Dataset) -> Dict[str, Any]: + """Train the model. + + Args: + train_dataset: Training dataset. + val_dataset: Validation dataset. + + Returns: + Dictionary containing training history. + """ + self.model = self.model.to(self.device) + train_loader = self._create_dataloader(train_dataset, is_train=True) + val_loader = self._create_dataloader(val_dataset, is_train=False) + + optimizer = torch.optim.Adam(self.model.parameters(), + lr=self.config.learning_rate, + weight_decay=self.config.weight_decay) + + history = {"train_loss": [], "val_loss": [], "learning_rate": []} + + self.logger.info( + f"Starting training for {self.config.num_epochs} epochs...") + self.logger.info(f"Training samples: {len(train_dataset)}") + self.logger.info(f"Validation samples: {len(val_dataset)}") + + try: + for epoch in range(self.config.num_epochs): + train_loss = self._train_epoch(train_loader, optimizer) + val_loss = self._validate(val_loader) + current_lr = optimizer.param_groups[0]["lr"] + + metrics = { + "train_loss": train_loss, + "val_loss": val_loss, + "learning_rate": current_lr, + "epoch": epoch + 1, + } + + if self.config.use_wandb: + self.metrics_logger.log_metrics(metrics) + + history["train_loss"].append(train_loss) + history["val_loss"].append(val_loss) + history["learning_rate"].append(current_lr) + + if self.rank == 0: + self.logger.info( + f"[Epoch {epoch+1}] Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}" + ) + + except KeyboardInterrupt: + self.logger.info("Training interrupted by user") + + if self.config.distributed: + self.dist_manager.cleanup() + + if self.config.use_wandb: + self.metrics_logger.close() + + return {"history": history} diff --git a/src/timesfm/pytorch_patched_decoder.py b/src/timesfm/pytorch_patched_decoder.py index 67f6be4..15bf428 100644 --- a/src/timesfm/pytorch_patched_decoder.py +++ b/src/timesfm/pytorch_patched_decoder.py @@ -21,7 +21,7 @@ from torch import nn import torch.nn.functional as F -def _create_quantiles() -> list[float]: +def create_quantiles() -> list[float]: return [0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9] @@ -48,7 +48,7 @@ class TimesFMConfig: # Horizon length horizon_len: int = 128 # quantiles - quantiles: List[float] = dataclasses.field(default_factory=_create_quantiles) + quantiles: List[float] = dataclasses.field(default_factory=create_quantiles) # Padding value pad_val: float = 1123581321.0 # Tolerance