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