Style fix

This commit is contained in:
misha-chertushkin
2025-01-21 19:29:42 +00:00
parent 4f4649c212
commit f84366e3d1
2 changed files with 467 additions and 467 deletions
+201 -201
View File
@@ -20,275 +20,275 @@ import os
class TimeSeriesDataset(Dataset): class TimeSeriesDataset(Dataset):
"""Dataset for time series data compatible with TimesFM.""" """Dataset for time series data compatible with TimesFM."""
def __init__(self, series: np.ndarray, context_length: int, horizon_length: int): def __init__(self, series: np.ndarray, context_length: int, horizon_length: int):
""" """
Initialize dataset. Initialize dataset.
Args: Args:
series: Time series data series: Time series data
context_length: Number of past timesteps to use as input context_length: Number of past timesteps to use as input
horizon_length: Number of future timesteps to predict horizon_length: Number of future timesteps to predict
""" """
self.series = series self.series = series
self.context_length = context_length self.context_length = context_length
self.horizon_length = horizon_length self.horizon_length = horizon_length
self._prepare_samples() self._prepare_samples()
def _prepare_samples(self) -> None: def _prepare_samples(self) -> None:
"""Prepare sliding window samples from the time series.""" """Prepare sliding window samples from the time series."""
self.samples = [] self.samples = []
total_length = self.context_length + self.horizon_length total_length = self.context_length + self.horizon_length
for start_idx in range(0, len(self.series) - total_length + 1): for start_idx in range(0, len(self.series) - total_length + 1):
end_idx = start_idx + self.context_length end_idx = start_idx + self.context_length
x_context = self.series[start_idx:end_idx] x_context = self.series[start_idx:end_idx]
x_future = self.series[end_idx : end_idx + self.horizon_length] x_future = self.series[end_idx : end_idx + self.horizon_length]
self.samples.append((x_context, x_future)) self.samples.append((x_context, x_future))
def __len__(self) -> int: def __len__(self) -> int:
return len(self.samples) return len(self.samples)
def __getitem__(self, index: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: def __getitem__(self, index: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
x_context, x_future = self.samples[index] x_context, x_future = self.samples[index]
x_context = torch.tensor(x_context, dtype=torch.float32) x_context = torch.tensor(x_context, dtype=torch.float32)
x_future = torch.tensor(x_future, dtype=torch.float32) x_future = torch.tensor(x_future, dtype=torch.float32)
input_padding = torch.zeros_like(x_context) input_padding = torch.zeros_like(x_context)
freq = torch.zeros(1, dtype=torch.long) freq = torch.zeros(1, dtype=torch.long)
return x_context, input_padding, freq, x_future return x_context, input_padding, freq, x_future
def prepare_datasets( def prepare_datasets(
series: np.ndarray, context_length: int, horizon_length: int, train_split: float = 0.8 series: np.ndarray, context_length: int, horizon_length: int, train_split: float = 0.8
) -> Tuple[Dataset, Dataset]: ) -> Tuple[Dataset, Dataset]:
""" """
Prepare training and validation datasets from time series data. Prepare training and validation datasets from time series data.
Args: Args:
series: Input time series data series: Input time series data
context_length: Number of past timesteps to use context_length: Number of past timesteps to use
horizon_length: Number of future timesteps to predict horizon_length: Number of future timesteps to predict
train_split: Fraction of data to use for training train_split: Fraction of data to use for training
Returns: Returns:
Tuple of (train_dataset, val_dataset) Tuple of (train_dataset, val_dataset)
""" """
train_size = int(len(series) * train_split) train_size = int(len(series) * train_split)
train_data = series[:train_size] train_data = series[:train_size]
val_data = series[train_size:] val_data = series[train_size:]
# Create datasets # Create datasets
train_dataset = TimeSeriesDataset(train_data, context_length=context_length, horizon_length=horizon_length) train_dataset = TimeSeriesDataset(train_data, context_length=context_length, horizon_length=horizon_length)
val_dataset = TimeSeriesDataset(val_data, context_length=context_length, horizon_length=horizon_length) val_dataset = TimeSeriesDataset(val_data, context_length=context_length, horizon_length=horizon_length)
return train_dataset, val_dataset return train_dataset, val_dataset
def get_model(load_weights: bool = False): def get_model(load_weights: bool = False):
device = "cuda" if torch.cuda.is_available() else "cpu" device = "cuda" if torch.cuda.is_available() else "cpu"
repo_id = "google/timesfm-2.0-500m-pytorch" repo_id = "google/timesfm-2.0-500m-pytorch"
hparams = TimesFmHparams( hparams = TimesFmHparams(
backend=device, backend=device,
per_core_batch_size=32, per_core_batch_size=32,
horizon_len=128, horizon_len=128,
num_layers=50, num_layers=50,
use_positional_embedding=False, use_positional_embedding=False,
context_len=192, context_len=192,
) )
tfm = TimesFm(hparams=hparams, checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id)) tfm = TimesFm(hparams=hparams, checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id))
model = PatchedTimeSeriesDecoder(tfm._model_config) model = PatchedTimeSeriesDecoder(tfm._model_config)
if load_weights: if load_weights:
checkpoint_path = path.join(snapshot_download(repo_id), "torch_model.ckpt") checkpoint_path = path.join(snapshot_download(repo_id), "torch_model.ckpt")
loaded_checkpoint = torch.load(checkpoint_path, weights_only=True) loaded_checkpoint = torch.load(checkpoint_path, weights_only=True)
model.load_state_dict(loaded_checkpoint) model.load_state_dict(loaded_checkpoint)
return model, hparams, tfm._model_config return model, hparams, tfm._model_config
def plot_predictions( def plot_predictions(
model: TimesFm, model: TimesFm,
val_dataset: Dataset, val_dataset: Dataset,
save_path: Optional[str] = "predictions.png", save_path: Optional[str] = "predictions.png",
) -> None: ) -> None:
""" """
Plot model predictions against ground truth for a batch of validation data. Plot model predictions against ground truth for a batch of validation data.
Args: Args:
model: Trained TimesFM model model: Trained TimesFM model
val_dataset: Validation dataset val_dataset: Validation dataset
save_path: Path to save the plot save_path: Path to save the plot
""" """
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
model.eval() model.eval()
x_context, x_padding, freq, x_future = val_dataset[0] x_context, x_padding, freq, x_future = val_dataset[0]
x_context = x_context.unsqueeze(0) # Add batch dimension x_context = x_context.unsqueeze(0) # Add batch dimension
x_padding = x_padding.unsqueeze(0) x_padding = x_padding.unsqueeze(0)
freq = freq.unsqueeze(0) freq = freq.unsqueeze(0)
x_future = x_future.unsqueeze(0) x_future = x_future.unsqueeze(0)
device = next(model.parameters()).device device = next(model.parameters()).device
x_context = x_context.to(device) x_context = x_context.to(device)
x_padding = x_padding.to(device) x_padding = x_padding.to(device)
freq = freq.to(device) freq = freq.to(device)
x_future = x_future.to(device) x_future = x_future.to(device)
with torch.no_grad(): with torch.no_grad():
predictions = model(x_context, x_padding.float(), freq) predictions = model(x_context, x_padding.float(), freq)
predictions_mean = predictions[..., 0] # [B, N, horizon_len] predictions_mean = predictions[..., 0] # [B, N, horizon_len]
last_patch_pred = predictions_mean[:, -1, :] # [B, horizon_len] last_patch_pred = predictions_mean[:, -1, :] # [B, horizon_len]
context_vals = x_context[0].cpu().numpy() context_vals = x_context[0].cpu().numpy()
future_vals = x_future[0].cpu().numpy() future_vals = x_future[0].cpu().numpy()
pred_vals = last_patch_pred[0].cpu().numpy() pred_vals = last_patch_pred[0].cpu().numpy()
context_len = len(context_vals) context_len = len(context_vals)
horizon_len = len(future_vals) horizon_len = len(future_vals)
plt.figure(figsize=(12, 6)) plt.figure(figsize=(12, 6))
plt.plot(range(context_len), context_vals, label="Historical Data", color="blue", linewidth=2) plt.plot(range(context_len), context_vals, label="Historical Data", color="blue", linewidth=2)
plt.plot( plt.plot(
range(context_len, context_len + horizon_len), range(context_len, context_len + horizon_len),
future_vals, future_vals,
label="Ground Truth", label="Ground Truth",
color="green", color="green",
linestyle="--", linestyle="--",
linewidth=2, linewidth=2,
) )
plt.plot(range(context_len, context_len + horizon_len), pred_vals, label="Prediction", color="red", linewidth=2) plt.plot(range(context_len, context_len + horizon_len), pred_vals, label="Prediction", color="red", linewidth=2)
plt.xlabel("Time Step") plt.xlabel("Time Step")
plt.ylabel("Value") plt.ylabel("Value")
plt.title("TimesFM Predictions vs Ground Truth") plt.title("TimesFM Predictions vs Ground Truth")
plt.legend() plt.legend()
plt.grid(True) plt.grid(True)
if save_path: if save_path:
plt.savefig(save_path) plt.savefig(save_path)
print(f"Plot saved to {save_path}") print(f"Plot saved to {save_path}")
plt.close() plt.close()
def get_data(context_len: int, horizon_len: int) -> Tuple[Dataset, Dataset]: def get_data(context_len: int, horizon_len: int) -> Tuple[Dataset, Dataset]:
df = yf.download("AAPL", start="2010-01-01", end="2019-01-01") df = yf.download("AAPL", start="2010-01-01", end="2019-01-01")
time_series = df["Close"].values time_series = df["Close"].values
train_dataset, val_dataset = prepare_datasets( train_dataset, val_dataset = prepare_datasets(
series=time_series, series=time_series,
context_length=context_len, context_length=context_len,
horizon_length=horizon_len, horizon_length=horizon_len,
train_split=0.8, train_split=0.8,
) )
print(f"Created datasets:") print(f"Created datasets:")
print(f"- Training samples: {len(train_dataset)}") print(f"- Training samples: {len(train_dataset)}")
print(f"- Validation samples: {len(val_dataset)}") print(f"- Validation samples: {len(val_dataset)}")
return train_dataset, val_dataset return train_dataset, val_dataset
def single_gpu_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)
train_dataset, val_dataset = get_data(128, tfm_config.horizon_len) train_dataset, val_dataset = get_data(128, tfm_config.horizon_len)
finetuner = TimesFMFinetuner(model, config) finetuner = TimesFMFinetuner(model, config)
print("\nStarting finetuning...") print("\nStarting finetuning...")
results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset) results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset)
print("\nFinetuning completed!") print("\nFinetuning completed!")
print(f"Training history: {len(results['history']['train_loss'])} epochs") print(f"Training history: {len(results['history']['train_loss'])} epochs")
plot_predictions( plot_predictions(
model=model, model=model,
val_dataset=val_dataset, val_dataset=val_dataset,
save_path="timesfm_predictions.png", save_path="timesfm_predictions.png",
) )
def setup_process(rank, world_size, model, config, train_dataset, val_dataset, return_dict): def setup_process(rank, world_size, model, config, train_dataset, val_dataset, return_dict):
"""Setup process function with optimized CUDA handling.""" """Setup process function with optimized CUDA handling."""
try: try:
if torch.cuda.is_available(): if torch.cuda.is_available():
torch.cuda.set_device(rank) torch.cuda.set_device(rank)
os.environ["MASTER_ADDR"] = config.master_addr os.environ["MASTER_ADDR"] = config.master_addr
os.environ["MASTER_PORT"] = config.master_port os.environ["MASTER_PORT"] = config.master_port
if not torch.distributed.is_initialized(): if not torch.distributed.is_initialized():
torch.distributed.init_process_group(backend="nccl", world_size=world_size, rank=rank) torch.distributed.init_process_group(backend="nccl", world_size=world_size, rank=rank)
finetuner = TimesFMFinetuner(model, config, rank=rank) finetuner = TimesFMFinetuner(model, config, rank=rank)
results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset) results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset)
if rank == 0: if rank == 0:
return_dict["results"] = results return_dict["results"] = results
plot_predictions( plot_predictions(
model=model, model=model,
val_dataset=val_dataset, val_dataset=val_dataset,
save_path="timesfm_predictions.png", save_path="timesfm_predictions.png",
) )
except Exception as e: except Exception as e:
print(f"Error in process {rank}: {str(e)}") print(f"Error in process {rank}: {str(e)}")
raise e raise e
finally: finally:
if torch.distributed.is_initialized(): if torch.distributed.is_initialized():
torch.distributed.destroy_process_group() torch.distributed.destroy_process_group()
def multi_gpu_example(): def multi_gpu_example():
"""Example of finetuning TimesFM using multiple GPUs with optimized spawn.""" """Example of finetuning TimesFM using multiple GPUs with optimized spawn."""
mp.set_start_method("spawn", force=True) mp.set_start_method("spawn", force=True)
gpu_ids = [0, 1] gpu_ids = [0, 1]
world_size = len(gpu_ids) world_size = len(gpu_ids)
model, hparams, tfm_config = get_model(load_weights=True) model, hparams, tfm_config = get_model(load_weights=True)
# Create config # Create config
config = FinetuningConfig( config = FinetuningConfig(
batch_size=256, batch_size=256,
num_epochs=5, num_epochs=5,
learning_rate=3e-5, learning_rate=3e-5,
use_wandb=True, use_wandb=True,
distributed=True, distributed=True,
gpu_ids=gpu_ids, gpu_ids=gpu_ids,
) )
train_dataset, val_dataset = get_data(128, tfm_config.horizon_len) train_dataset, val_dataset = get_data(128, tfm_config.horizon_len)
manager = mp.Manager() manager = mp.Manager()
return_dict = manager.dict() return_dict = manager.dict()
# Launch processes # Launch processes
mp.spawn( mp.spawn(
setup_process, setup_process,
args=(world_size, model, config, train_dataset, val_dataset, return_dict), args=(world_size, model, config, train_dataset, val_dataset, return_dict),
nprocs=world_size, nprocs=world_size,
join=True, join=True,
) )
results = return_dict.get("results", None) results = return_dict.get("results", None)
print("\nFinetuning completed!") print("\nFinetuning completed!")
return results return results
if __name__ == "__main__": if __name__ == "__main__":
try: try:
# single_gpu_example() # Single GPU # single_gpu_example() # Single GPU
multi_gpu_example() # Multi-GPU multi_gpu_example() # Multi-GPU
except Exception as e: except Exception as e:
print(f"Training failed: {str(e)}") print(f"Training failed: {str(e)}")
finally: finally:
if torch.distributed.is_initialized(): if torch.distributed.is_initialized():
torch.distributed.destroy_process_group() torch.distributed.destroy_process_group()
+266 -266
View File
@@ -18,323 +18,323 @@ import wandb
class MetricsLogger(ABC): class MetricsLogger(ABC):
"""Abstract base class for logging metrics during training. """Abstract base class for logging metrics during training.
This class defines the interface for logging metrics during model training. This class defines the interface for logging metrics during model training.
Concrete implementations can log to different backends (e.g., WandB, TensorBoard). 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 @abstractmethod
def log_metrics(self, metrics: Dict[str, Any], step: Optional[int] = None) -> None: def close(self) -> None:
"""Log metrics to the specified backend. """Clean up any resources used by the logger."""
pass
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): class WandBLogger(MetricsLogger):
"""Weights & Biases implementation of metrics logging. """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: Args:
project: Name of the W&B project. metrics: Dictionary of metrics to log.
config: Configuration dictionary to log. step: Current training step or epoch.
rank: Process rank in distributed training.
""" """
if self.rank == 0:
wandb.log(metrics, step=step)
def __init__(self, project: str, config: Dict[str, Any], rank: int = 0): def close(self) -> None:
self.rank = rank """Finish the W&B run if on the main process."""
if rank == 0: if self.rank == 0:
wandb.init(project=project, config=config) wandb.finish()
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: class DistributedManager:
"""Manages distributed training setup and cleanup. """Manages distributed training setup and cleanup.
Args: Args:
world_size: Total number of processes. world_size: Total number of processes.
rank: Process rank. rank: Process rank.
master_addr: Address of the master process. master_addr: Address of the master process.
master_port: Port for distributed communication. master_port: Port for distributed communication.
backend: PyTorch distributed backend to use. backend: PyTorch distributed backend to use.
""" """
def __init__( def __init__(
self, self,
world_size: int, world_size: int,
rank: int, rank: int,
master_addr: str = "localhost", master_addr: str = "localhost",
master_port: str = "12358", master_port: str = "12358",
backend: str = "nccl", backend: str = "nccl",
): ):
self.world_size = world_size self.world_size = world_size
self.rank = rank self.rank = rank
self.master_addr = master_addr self.master_addr = master_addr
self.master_port = master_port self.master_port = master_port
self.backend = backend self.backend = backend
def setup(self) -> None: def setup(self) -> None:
"""Initialize the distributed environment.""" """Initialize the distributed environment."""
os.environ["MASTER_ADDR"] = self.master_addr os.environ["MASTER_ADDR"] = self.master_addr
os.environ["MASTER_PORT"] = self.master_port os.environ["MASTER_PORT"] = self.master_port
if not dist.is_initialized(): if not dist.is_initialized():
dist.init_process_group(backend=self.backend, world_size=self.world_size, rank=self.rank) dist.init_process_group(backend=self.backend, world_size=self.world_size, rank=self.rank)
def cleanup(self) -> None: def cleanup(self) -> None:
"""Clean up the distributed environment.""" """Clean up the distributed environment."""
if dist.is_initialized(): if dist.is_initialized():
dist.destroy_process_group() dist.destroy_process_group()
@dataclass @dataclass
class FinetuningConfig: class FinetuningConfig:
"""Configuration for model training. """Configuration for model training.
Args: Args:
batch_size: Number of samples per batch. batch_size: Number of samples per batch.
num_epochs: Number of training epochs. num_epochs: Number of training epochs.
learning_rate: Initial learning rate. learning_rate: Initial learning rate.
weight_decay: L2 regularization factor. weight_decay: L2 regularization factor.
device: Device to train on ('cuda' or 'cpu'). device: Device to train on ('cuda' or 'cpu').
distributed: Whether to use distributed training. distributed: Whether to use distributed training.
gpu_ids: List of GPU IDs to use. gpu_ids: List of GPU IDs to use.
master_port: Port for distributed training. master_port: Port for distributed training.
master_addr: Address for distributed training. master_addr: Address for distributed training.
use_wandb: Whether to use Weights & Biases logging. use_wandb: Whether to use Weights & Biases logging.
wandb_project: W&B project name. wandb_project: W&B project name.
""" """
batch_size: int = 32 batch_size: int = 32
num_epochs: int = 20 num_epochs: int = 20
learning_rate: float = 1e-4 learning_rate: float = 1e-4
weight_decay: float = 0.01 weight_decay: float = 0.01
device: str = "cuda" if torch.cuda.is_available() else "cpu" device: str = "cuda" if torch.cuda.is_available() else "cpu"
distributed: bool = False distributed: bool = False
gpu_ids: List[int] = field(default_factory=lambda: [0]) gpu_ids: List[int] = field(default_factory=lambda: [0])
master_port: str = "12358" master_port: str = "12358"
master_addr: str = "localhost" master_addr: str = "localhost"
use_wandb: bool = False use_wandb: bool = False
wandb_project: str = "timesfm-finetuning" wandb_project: str = "timesfm-finetuning"
class TimesFMFinetuner: class TimesFMFinetuner:
"""Handles model training and validation. """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: Args:
model: PyTorch model to train. dataset: Dataset to create loader for.
config: Training configuration. is_train: Whether this is for training (affects shuffling).
rank: Process rank for distributed training.
loss_fn: Loss function (defaults to MSE). Returns:
logger: Optional logging.Logger instance. 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
def __init__( return DataLoader(
self, dataset,
model: nn.Module, batch_size=self.config.batch_size,
config: FinetuningConfig, shuffle=(is_train and not self.config.distributed),
rank: int = 0, sampler=sampler,
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: def _process_batch(self, batch: List[torch.Tensor]) -> tuple:
self.metrics_logger = WandBLogger(config.wandb_project, config.__dict__, rank) """Process a single batch of data.
if config.distributed: Args:
self.dist_manager = DistributedManager( batch: List of input tensors.
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: Returns:
"""Configure model for distributed training.""" Tuple of (loss, predictions).
self.model = self.model.to(self.device) """
return DDP( x_context, x_padding, freq, x_future = [t.to(self.device, non_blocking=True) for t in batch]
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: predictions = self.model(x_context, x_padding.float(), freq)
"""Create appropriate DataLoader based on training configuration. predictions_mean = predictions[..., 0]
last_patch_pred = predictions_mean[:, -1, :]
Args: loss = self.loss_fn(last_patch_pred, x_future.squeeze(-1))
dataset: Dataset to create loader for.
is_train: Whether this is for training (affects shuffling).
Returns: return loss, predictions
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( def _train_epoch(self, train_loader: DataLoader, optimizer: torch.optim.Optimizer) -> float:
dataset, """Train for one epoch.
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: Args:
"""Process a single batch of data. train_loader: DataLoader for training data.
optimizer: Optimizer instance.
Args: Returns:
batch: List of input tensors. Average training loss for the epoch.
"""
self.model.train()
total_loss = 0.0
Returns: for batch in train_loader:
Tuple of (loss, predictions). loss, _ = self._process_batch(batch)
"""
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) if self.config.distributed:
predictions_mean = predictions[..., 0] losses = [torch.zeros_like(loss) for _ in range(dist.get_world_size())]
last_patch_pred = predictions_mean[:, -1, :] dist.all_gather(losses, loss)
loss = self.loss_fn(last_patch_pred, x_future.squeeze(-1)) optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss, predictions total_loss += loss.item()
def _train_epoch(self, train_loader: DataLoader, optimizer: torch.optim.Optimizer) -> float: return total_loss / len(train_loader)
"""Train for one epoch.
Args: def _validate(self, val_loader: DataLoader) -> float:
train_loader: DataLoader for training data. """Perform validation.
optimizer: Optimizer instance.
Returns: Args:
Average training loss for the epoch. val_loader: DataLoader for validation data.
"""
self.model.train()
total_loss = 0.0
for batch in train_loader: Returns:
loss, _ = self._process_batch(batch) Average validation loss.
"""
self.model.eval()
total_loss = 0.0
if self.config.distributed: with torch.no_grad():
losses = [torch.zeros_like(loss) for _ in range(dist.get_world_size())] for batch in val_loader:
dist.all_gather(losses, loss) loss, _ = self._process_batch(batch)
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: if self.config.distributed:
self.dist_manager.cleanup() 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: if self.config.use_wandb:
self.metrics_logger.close() self.metrics_logger.log_metrics(metrics)
return {"history": history} 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}