GPU support added, almost done
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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}
|
||||
|
||||
Reference in New Issue
Block a user