GPU support added, almost done

This commit is contained in:
misha-chertushkin
2025-01-21 02:19:35 +00:00
parent fb6213ed59
commit a559718c66
2 changed files with 130 additions and 34 deletions
+82 -13
View File
@@ -2,23 +2,21 @@
Example usage of the TimesFM Finetuning Framework. Example usage of the TimesFM Finetuning Framework.
""" """
import yfinance as yf
import torch
from os import path from os import path
import numpy as np from typing import Optional, Tuple
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
import numpy as np import numpy as np
import pandas as pd import pandas as pd
from torch.utils.data import Dataset
import torch import torch
import torch.multiprocessing as mp
import yfinance as yf 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): class TimeSeriesDataset(Dataset):
@@ -199,7 +197,7 @@ def get_data(context_len: int, horizon_len: int) -> Tuple[Dataset, Dataset]:
return train_dataset, val_dataset return train_dataset, val_dataset
def basic_example(): def single_gpu_example():
"""Basic example of finetuning TimesFM on stock data.""" """Basic example of finetuning TimesFM on stock data."""
model, hparams, tfm_config = get_model(load_weights=True) model, hparams, tfm_config = get_model(load_weights=True)
config = FinetuningConfig(batch_size=256, num_epochs=5, learning_rate=1e-4, use_wandb=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__": if __name__ == "__main__":
basic_example() # Use either single GPU or multi-GPU example
# basic_example() # Single GPU
multi_gpu_example() # Multi-GPU
+48 -21
View File
@@ -26,17 +26,19 @@ Example usage:
import abc import abc
import logging import logging
from dataclasses import dataclass import multiprocessing as mp
import os
from dataclasses import dataclass, field
from pathlib import Path from pathlib import Path
from typing import Optional, Dict, Any from typing import Any, Dict, List, Optional
import torch import torch
from torch.utils.data import Dataset, DataLoader import torch.distributed as dist
import torch.optim as optim import torch.optim as optim
from torch.nn.parallel import DistributedDataParallel as DDP from torch.nn.parallel import DistributedDataParallel as DDP
import wandb from torch.utils.data import DataLoader, Dataset
import multiprocessing as mp
import wandb
from timesfm import TimesFm from timesfm import TimesFm
@@ -59,49 +61,65 @@ class FinetuningConfig:
use_wandb: bool = False use_wandb: bool = False
wandb_project: str = "timesfm-finetuning" 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: class TimesFMFinetuner:
"""Main class for finetuning TimesFM models."""
def __init__( def __init__(
self, self,
model: TimesFm, model: TimesFm,
config: FinetuningConfig, config: FinetuningConfig,
rank: int = 0,
loss_fn: Optional[callable] = None, loss_fn: Optional[callable] = None,
logger: Optional[logging.Logger] = 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.model = model
self.config = config self.config = config
self.logger = logger or logging.getLogger(__name__) self.logger = logger or logging.getLogger(__name__)
self.rank = rank
self.device = torch.device(config.device) if config.distributed:
self.loss_fn = loss_fn or (lambda x, y: torch.mean((x - y.squeeze(-1)) ** 2)) # MSELoss() 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() 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: def _setup_wandb(self) -> None:
"""Initialize Weights & Biases logging.""" """Initialize Weights & Biases logging."""
wandb.init(project=self.config.wandb_project, config=self.config.__dict__) wandb.init(project=self.config.wandb_project, config=self.config.__dict__)
def _create_dataloader(self, dataset: Dataset, name: str) -> DataLoader: def _create_dataloader(self, dataset: Dataset, name: str) -> DataLoader:
"""Create a dataloader from a dataset.""" """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( return DataLoader(
dataset, dataset,
batch_size=self.config.batch_size, batch_size=self.config.batch_size,
shuffle=name == "train", shuffle=(name == "train" and not self.config.distributed),
num_workers=mp.cpu_count(), num_workers=mp.cpu_count() // len(self.config.gpu_ids),
pin_memory=self.device.type == "cuda", pin_memory=self.device.type == "cuda",
persistent_workers=True, persistent_workers=True,
prefetch_factor=2, prefetch_factor=2,
sampler=sampler,
) )
def _train_epoch(self, train_loader: DataLoader, optimizer: torch.optim.Optimizer) -> float: 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) 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") train_loader = self._create_dataloader(train_dataset, "train")
val_loader = self._create_dataloader(val_dataset, "val") val_loader = self._create_dataloader(val_dataset, "val")
optimizer = optim.Adam( optimizer = optim.Adam(
self.model.parameters(), lr=self.config.learning_rate, weight_decay=self.config.weight_decay self.model.parameters(), lr=self.config.learning_rate, weight_decay=self.config.weight_decay
) )
history = {"train_loss": [], "val_loss": [], "learning_rate": []} history = {"train_loss": [], "val_loss": [], "learning_rate": []}
self.logger.info(f"Starting training for {self.config.num_epochs} epochs...") self.logger.info(f"Starting training for {self.config.num_epochs} epochs...")
@@ -196,4 +220,7 @@ class TimesFMFinetuner:
except KeyboardInterrupt: except KeyboardInterrupt:
self.logger.info("Training interrupted by user") self.logger.info("Training interrupted by user")
if self.config.distributed:
dist.destroy_process_group()
return {"history": history} return {"history": history}