Add jupyter notebook
This commit is contained in:
@@ -10,7 +10,7 @@ import pandas as pd
|
||||
import torch
|
||||
import torch.multiprocessing as mp
|
||||
import yfinance as yf
|
||||
from finetuning_torch import FinetuningConfig, TimesFMFinetuner
|
||||
from timesfm.finetuning_torch import FinetuningConfig, TimesFMFinetuner
|
||||
from huggingface_hub import snapshot_download
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
@@ -0,0 +1,292 @@
|
||||
{
|
||||
"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, series: np.ndarray, context_length: int, horizon_length: int):\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",
|
||||
" \"\"\"\n",
|
||||
" self.series = series\n",
|
||||
" self.context_length = context_length\n",
|
||||
" self.horizon_length = horizon_length\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__(self, index: int) -> 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.zeros(1, dtype=torch.long)\n",
|
||||
"\n",
|
||||
" return x_context, input_padding, freq, x_future\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def prepare_datasets(\n",
|
||||
" series: np.ndarray, context_length: int, horizon_length: int, train_split: float = 0.8\n",
|
||||
") -> 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",
|
||||
" 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\n",
|
||||
" train_dataset = TimeSeriesDataset(train_data, context_length=context_length, horizon_length=horizon_length)\n",
|
||||
"\n",
|
||||
" val_dataset = TimeSeriesDataset(val_data, context_length=context_length, horizon_length=horizon_length)\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=192,\n",
|
||||
" )\n",
|
||||
" tfm = TimesFm(hparams=hparams, 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), context_vals, label=\"Historical Data\", color=\"blue\", 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), pred_vals, label=\"Prediction\", color=\"red\", 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",
|
||||
"\n",
|
||||
"def get_data(context_len: int, horizon_len: int) -> Tuple[Dataset, Dataset]:\n",
|
||||
" df = yf.download(\"AAPL\", start=\"2010-01-01\", end=\"2019-01-01\")\n",
|
||||
" time_series = df[\"Close\"].values\n",
|
||||
"\n",
|
||||
" train_dataset, val_dataset = prepare_datasets(\n",
|
||||
" series=time_series,\n",
|
||||
" context_length=context_len,\n",
|
||||
" horizon_length=horizon_len,\n",
|
||||
" train_split=0.8,\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" print(f\"Created datasets:\")\n",
|
||||
" print(f\"- Training samples: {len(train_dataset)}\")\n",
|
||||
" print(f\"- Validation samples: {len(val_dataset)}\")\n",
|
||||
" return train_dataset, val_dataset\n",
|
||||
"\n",
|
||||
"\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, num_epochs=5, learning_rate=1e-4, use_wandb=True)\n",
|
||||
"\n",
|
||||
" train_dataset, val_dataset = get_data(128, tfm_config.horizon_len)\n",
|
||||
" finetuner = TimesFMFinetuner(model, config)\n",
|
||||
"\n",
|
||||
" print(\"\\nStarting finetuning...\")\n",
|
||||
" results = finetuner.finetune(train_dataset=train_dataset, 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
|
||||
}
|
||||
@@ -1,340 +0,0 @@
|
||||
"""
|
||||
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
|
||||
|
||||
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.
|
||||
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.
|
||||
"""
|
||||
|
||||
batch_size: int = 32
|
||||
num_epochs: int = 20
|
||||
learning_rate: float = 1e-4
|
||||
weight_decay: float = 0.01
|
||||
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"
|
||||
|
||||
|
||||
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 _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))
|
||||
|
||||
return loss, predictions
|
||||
|
||||
def _train_epoch(self, train_loader: DataLoader, optimizer: torch.optim.Optimizer) -> float:
|
||||
"""Train for one epoch.
|
||||
|
||||
Args:
|
||||
train_loader: DataLoader for training data.
|
||||
optimizer: Optimizer instance.
|
||||
|
||||
Returns:
|
||||
Average training loss for the epoch.
|
||||
"""
|
||||
self.model.train()
|
||||
total_loss = 0.0
|
||||
|
||||
for batch in train_loader:
|
||||
loss, _ = self._process_batch(batch)
|
||||
|
||||
if self.config.distributed:
|
||||
losses = [torch.zeros_like(loss) for _ in range(dist.get_world_size())]
|
||||
dist.all_gather(losses, loss)
|
||||
|
||||
optimizer.zero_grad()
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
|
||||
total_loss += loss.item()
|
||||
|
||||
return total_loss / len(train_loader)
|
||||
|
||||
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
|
||||
|
||||
with torch.no_grad():
|
||||
for batch in val_loader:
|
||||
loss, _ = self._process_batch(batch)
|
||||
|
||||
if self.config.distributed:
|
||||
losses = [torch.zeros_like(loss) for _ in range(dist.get_world_size())]
|
||||
dist.all_gather(losses, loss)
|
||||
|
||||
total_loss += loss.item()
|
||||
|
||||
return total_loss / len(val_loader)
|
||||
|
||||
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}
|
||||
Reference in New Issue
Block a user