diff --git a/notebooks/finetuning_example.py b/notebooks/finetuning_example.py index 966c980..c42cf88 100644 --- a/notebooks/finetuning_example.py +++ b/notebooks/finetuning_example.py @@ -2,23 +2,21 @@ Example usage of the TimesFM Finetuning Framework. """ -import yfinance as yf -import torch from os import path -import numpy as np -from torch.utils.data import Dataset -from timesfm import TimesFm, TimesFmHparams, TimesFmCheckpoint -from timesfm.pytorch_patched_decoder import PatchedTimeSeriesDecoder -from finetuning_torch import FinetuningConfig, TimesFMFinetuner -from huggingface_hub import snapshot_download +from typing import Optional, Tuple + import numpy as np import pandas as pd -from torch.utils.data import Dataset import torch +import torch.multiprocessing as mp import yfinance as yf -from typing import Tuple, Optional +from finetuning_torch import FinetuningConfig, TimesFMFinetuner +from huggingface_hub import snapshot_download +from torch.utils.data import Dataset -from timesfm import TimesFm, TimesFmHparams +from timesfm import TimesFm, TimesFmCheckpoint, TimesFmHparams +from timesfm.pytorch_patched_decoder import PatchedTimeSeriesDecoder +import os class TimeSeriesDataset(Dataset): @@ -199,7 +197,7 @@ def get_data(context_len: int, horizon_len: int) -> Tuple[Dataset, Dataset]: return train_dataset, val_dataset -def basic_example(): +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) @@ -220,5 +218,76 @@ def basic_example(): ) +def setup_process(rank, world_size, model, config, train_dataset, val_dataset, return_dict): + """Initialize the distributed process.""" + # Set up the process group + os.environ["MASTER_ADDR"] = "localhost" + os.environ["MASTER_PORT"] = "12355" + + # Initialize the process group + torch.distributed.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=rank) + + # Set the device for this process + torch.cuda.set_device(rank) + + try: + finetuner = TimesFMFinetuner(model, config, rank=rank) + results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset) + + if rank == 0: # Only store results and plot from the main process + return_dict["results"] = results + plot_predictions( + model=model, + val_dataset=val_dataset, + save_path="timesfm_predictions.png", + ) + finally: + # Cleanup - important! + torch.distributed.destroy_process_group() + + +def multi_gpu_example(): + """Example of finetuning TimesFM using multiple GPUs.""" + # Define which GPUs to use + gpu_ids = [0] # Just using one GPU + world_size = len(gpu_ids) + + # Initialize model and config + model, hparams, tfm_config = get_model(load_weights=True) + config = FinetuningConfig( + batch_size=256, + num_epochs=5, + learning_rate=1e-4, + use_wandb=False, + distributed=True, + gpu_ids=gpu_ids, + ) + + # Get datasets + train_dataset, val_dataset = get_data(128, tfm_config.horizon_len) + + # Create a multiprocessing manager to share results between processes + 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, + ) + + # Get results from the main process + results = return_dict.get("results", None) + print("\nFinetuning completed!") + if results: + print(f"Training history: {len(results['history']['train_loss'])} epochs") + + return results + + if __name__ == "__main__": - basic_example() + # Use either single GPU or multi-GPU example + # basic_example() # Single GPU + multi_gpu_example() # Multi-GPU diff --git a/notebooks/finetuning_torch.py b/notebooks/finetuning_torch.py index e923136..a10d99e 100644 --- a/notebooks/finetuning_torch.py +++ b/notebooks/finetuning_torch.py @@ -26,17 +26,19 @@ Example usage: import abc import logging -from dataclasses import dataclass +import multiprocessing as mp +import os +from dataclasses import dataclass, field from pathlib import Path -from typing import Optional, Dict, Any +from typing import Any, Dict, List, Optional import torch -from torch.utils.data import Dataset, DataLoader +import torch.distributed as dist import torch.optim as optim from torch.nn.parallel import DistributedDataParallel as DDP -import wandb -import multiprocessing as mp +from torch.utils.data import DataLoader, Dataset +import wandb from timesfm import TimesFm @@ -59,49 +61,65 @@ class FinetuningConfig: use_wandb: bool = False wandb_project: str = "timesfm-finetuning" + gpu_ids: List[int] = field(default_factory=lambda: [0]) # List of GPU IDs to use + distributed: bool = False + master_port: str = "12355" + master_addr: str = "localhost" + class TimesFMFinetuner: - """Main class for finetuning TimesFM models.""" - def __init__( self, model: TimesFm, config: FinetuningConfig, + rank: int = 0, loss_fn: Optional[callable] = None, logger: Optional[logging.Logger] = None, ): - """ - Initialize TimesFM finetuner. - - Args: - model: TimesFM model to finetune - config: Finetuning configuration - logger: Optional logger instance - """ self.model = model self.config = config self.logger = logger or logging.getLogger(__name__) + self.rank = rank - self.device = torch.device(config.device) - self.loss_fn = loss_fn or (lambda x, y: torch.mean((x - y.squeeze(-1)) ** 2)) # MSELoss() + if config.distributed: + self._setup_distributed(rank) - if config.use_wandb: + 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 and rank == 0: # Only initialize wandb on main process self._setup_wandb() + def _setup_distributed(self, rank): + """Setup distributed training environment.""" + os.environ["MASTER_ADDR"] = self.config.master_addr + os.environ["MASTER_PORT"] = self.config.master_port + + if not dist.is_initialized(): + dist.init_process_group(backend="nccl", world_size=len(self.config.gpu_ids), rank=rank) + def _setup_wandb(self) -> None: """Initialize Weights & Biases logging.""" wandb.init(project=self.config.wandb_project, config=self.config.__dict__) def _create_dataloader(self, dataset: Dataset, name: str) -> DataLoader: """Create a dataloader from a dataset.""" + if self.config.distributed: + sampler = torch.utils.data.distributed.DistributedSampler( + dataset, num_replicas=len(self.config.gpu_ids), rank=dist.get_rank(), shuffle=name == "train" + ) + else: + sampler = None + return DataLoader( dataset, batch_size=self.config.batch_size, - shuffle=name == "train", - num_workers=mp.cpu_count(), + shuffle=(name == "train" and not self.config.distributed), + num_workers=mp.cpu_count() // len(self.config.gpu_ids), pin_memory=self.device.type == "cuda", persistent_workers=True, prefetch_factor=2, + sampler=sampler, ) def _train_epoch(self, train_loader: DataLoader, optimizer: torch.optim.Optimizer) -> float: @@ -157,13 +175,19 @@ class TimesFMFinetuner: """ self.model = self.model.to(self.device) + if self.config.distributed: + self.model = DDP( + self.model, + device_ids=[self.config.gpu_ids[dist.get_rank()]], + output_device=self.config.gpu_ids[dist.get_rank()], + ) + train_loader = self._create_dataloader(train_dataset, "train") val_loader = self._create_dataloader(val_dataset, "val") optimizer = 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...") @@ -196,4 +220,7 @@ class TimesFMFinetuner: except KeyboardInterrupt: self.logger.info("Training interrupted by user") + if self.config.distributed: + dist.destroy_process_group() + return {"history": history}