diff --git a/notebooks/finetuning_example.py b/notebooks/finetuning_example.py index 07325b9..7ffb53d 100644 --- a/notebooks/finetuning_example.py +++ b/notebooks/finetuning_example.py @@ -20,275 +20,275 @@ import os 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): - """ - Initialize dataset. + def __init__(self, series: np.ndarray, context_length: int, horizon_length: int): + """ + Initialize dataset. - Args: - series: Time series data - context_length: Number of past timesteps to use as input - horizon_length: Number of future timesteps to predict - """ - self.series = series - self.context_length = context_length - self.horizon_length = horizon_length - self._prepare_samples() + Args: + series: Time series data + context_length: Number of past timesteps to use as input + horizon_length: Number of future timesteps to predict + """ + self.series = series + self.context_length = context_length + self.horizon_length = horizon_length + self._prepare_samples() - def _prepare_samples(self) -> None: - """Prepare sliding window samples from the time series.""" - self.samples = [] - total_length = self.context_length + self.horizon_length + def _prepare_samples(self) -> None: + """Prepare sliding window samples from the time series.""" + self.samples = [] + total_length = self.context_length + self.horizon_length - for start_idx in range(0, len(self.series) - total_length + 1): - end_idx = start_idx + self.context_length - x_context = self.series[start_idx:end_idx] - x_future = self.series[end_idx : end_idx + self.horizon_length] - self.samples.append((x_context, x_future)) + for start_idx in range(0, len(self.series) - total_length + 1): + end_idx = start_idx + self.context_length + x_context = self.series[start_idx:end_idx] + x_future = self.series[end_idx : end_idx + self.horizon_length] + self.samples.append((x_context, x_future)) - def __len__(self) -> int: - return len(self.samples) + def __len__(self) -> int: + return len(self.samples) - def __getitem__(self, index: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: - x_context, x_future = self.samples[index] + def __getitem__(self, index: int) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + x_context, x_future = self.samples[index] - x_context = torch.tensor(x_context, dtype=torch.float32) - x_future = torch.tensor(x_future, dtype=torch.float32) + x_context = torch.tensor(x_context, dtype=torch.float32) + x_future = torch.tensor(x_future, dtype=torch.float32) - input_padding = torch.zeros_like(x_context) - freq = torch.zeros(1, dtype=torch.long) + input_padding = torch.zeros_like(x_context) + 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( - 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]: - """ - Prepare training and validation datasets from time series data. + """ + Prepare training and validation datasets from time series data. - Args: - series: Input time series data - context_length: Number of past timesteps to use - horizon_length: Number of future timesteps to predict - train_split: Fraction of data to use for training + Args: + series: Input time series data + context_length: Number of past timesteps to use + horizon_length: Number of future timesteps to predict + train_split: Fraction of data to use for training - Returns: - Tuple of (train_dataset, val_dataset) - """ - train_size = int(len(series) * train_split) - train_data = series[:train_size] - val_data = series[train_size:] + Returns: + Tuple of (train_dataset, val_dataset) + """ + train_size = int(len(series) * train_split) + train_data = series[:train_size] + val_data = series[train_size:] - # Create datasets - train_dataset = TimeSeriesDataset(train_data, context_length=context_length, horizon_length=horizon_length) + # Create datasets + 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): - device = "cuda" if torch.cuda.is_available() else "cpu" - repo_id = "google/timesfm-2.0-500m-pytorch" - hparams = TimesFmHparams( - backend=device, - per_core_batch_size=32, - horizon_len=128, - num_layers=50, - use_positional_embedding=False, - context_len=192, - ) - tfm = TimesFm(hparams=hparams, checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id)) + device = "cuda" if torch.cuda.is_available() else "cpu" + repo_id = "google/timesfm-2.0-500m-pytorch" + hparams = TimesFmHparams( + backend=device, + per_core_batch_size=32, + horizon_len=128, + num_layers=50, + use_positional_embedding=False, + context_len=192, + ) + tfm = TimesFm(hparams=hparams, checkpoint=TimesFmCheckpoint(huggingface_repo_id=repo_id)) - model = PatchedTimeSeriesDecoder(tfm._model_config) - if load_weights: - checkpoint_path = path.join(snapshot_download(repo_id), "torch_model.ckpt") - loaded_checkpoint = torch.load(checkpoint_path, weights_only=True) - model.load_state_dict(loaded_checkpoint) - return model, hparams, tfm._model_config + model = PatchedTimeSeriesDecoder(tfm._model_config) + if load_weights: + checkpoint_path = path.join(snapshot_download(repo_id), "torch_model.ckpt") + loaded_checkpoint = torch.load(checkpoint_path, weights_only=True) + model.load_state_dict(loaded_checkpoint) + return model, hparams, tfm._model_config def plot_predictions( - model: TimesFm, - val_dataset: Dataset, - save_path: Optional[str] = "predictions.png", + model: TimesFm, + val_dataset: Dataset, + save_path: Optional[str] = "predictions.png", ) -> 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: - model: Trained TimesFM model - val_dataset: Validation dataset - save_path: Path to save the plot - """ - import matplotlib.pyplot as plt + Args: + model: Trained TimesFM model + val_dataset: Validation dataset + save_path: Path to save the plot + """ + import matplotlib.pyplot as plt - model.eval() + model.eval() - x_context, x_padding, freq, x_future = val_dataset[0] - x_context = x_context.unsqueeze(0) # Add batch dimension - x_padding = x_padding.unsqueeze(0) - freq = freq.unsqueeze(0) - x_future = x_future.unsqueeze(0) + x_context, x_padding, freq, x_future = val_dataset[0] + x_context = x_context.unsqueeze(0) # Add batch dimension + x_padding = x_padding.unsqueeze(0) + freq = freq.unsqueeze(0) + x_future = x_future.unsqueeze(0) - device = next(model.parameters()).device - x_context = x_context.to(device) - x_padding = x_padding.to(device) - freq = freq.to(device) - x_future = x_future.to(device) + device = next(model.parameters()).device + x_context = x_context.to(device) + x_padding = x_padding.to(device) + freq = freq.to(device) + x_future = x_future.to(device) - with torch.no_grad(): - predictions = model(x_context, x_padding.float(), freq) - predictions_mean = predictions[..., 0] # [B, N, horizon_len] - last_patch_pred = predictions_mean[:, -1, :] # [B, horizon_len] + with torch.no_grad(): + predictions = model(x_context, x_padding.float(), freq) + predictions_mean = predictions[..., 0] # [B, N, horizon_len] + last_patch_pred = predictions_mean[:, -1, :] # [B, horizon_len] - context_vals = x_context[0].cpu().numpy() - future_vals = x_future[0].cpu().numpy() - pred_vals = last_patch_pred[0].cpu().numpy() + context_vals = x_context[0].cpu().numpy() + future_vals = x_future[0].cpu().numpy() + pred_vals = last_patch_pred[0].cpu().numpy() - context_len = len(context_vals) - horizon_len = len(future_vals) + context_len = len(context_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( - range(context_len, context_len + horizon_len), - future_vals, - label="Ground Truth", - color="green", - linestyle="--", - linewidth=2, - ) + plt.plot( + range(context_len, context_len + horizon_len), + future_vals, + label="Ground Truth", + color="green", + linestyle="--", + 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.ylabel("Value") - plt.title("TimesFM Predictions vs Ground Truth") - plt.legend() - plt.grid(True) + plt.xlabel("Time Step") + plt.ylabel("Value") + plt.title("TimesFM Predictions vs Ground Truth") + plt.legend() + plt.grid(True) - if save_path: - plt.savefig(save_path) - print(f"Plot saved to {save_path}") + if save_path: + plt.savefig(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]: - df = yf.download("AAPL", start="2010-01-01", end="2019-01-01") - time_series = df["Close"].values + df = yf.download("AAPL", start="2010-01-01", end="2019-01-01") + time_series = df["Close"].values - train_dataset, val_dataset = prepare_datasets( - series=time_series, - context_length=context_len, - horizon_length=horizon_len, - train_split=0.8, - ) + train_dataset, val_dataset = prepare_datasets( + series=time_series, + context_length=context_len, + horizon_length=horizon_len, + train_split=0.8, + ) - print(f"Created datasets:") - print(f"- Training samples: {len(train_dataset)}") - print(f"- Validation samples: {len(val_dataset)}") - return train_dataset, val_dataset + print(f"Created datasets:") + print(f"- Training samples: {len(train_dataset)}") + print(f"- Validation samples: {len(val_dataset)}") + return train_dataset, val_dataset 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) + """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) - train_dataset, val_dataset = get_data(128, tfm_config.horizon_len) - finetuner = TimesFMFinetuner(model, config) + train_dataset, val_dataset = get_data(128, tfm_config.horizon_len) + finetuner = TimesFMFinetuner(model, config) - print("\nStarting finetuning...") - results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset) + print("\nStarting finetuning...") + results = finetuner.finetune(train_dataset=train_dataset, val_dataset=val_dataset) - print("\nFinetuning completed!") - print(f"Training history: {len(results['history']['train_loss'])} epochs") + print("\nFinetuning completed!") + print(f"Training history: {len(results['history']['train_loss'])} epochs") - plot_predictions( - model=model, - val_dataset=val_dataset, - save_path="timesfm_predictions.png", - ) + plot_predictions( + model=model, + val_dataset=val_dataset, + save_path="timesfm_predictions.png", + ) def setup_process(rank, world_size, model, config, train_dataset, val_dataset, return_dict): - """Setup process function with optimized CUDA handling.""" - try: - if torch.cuda.is_available(): - torch.cuda.set_device(rank) + """Setup process function with optimized CUDA handling.""" + try: + if torch.cuda.is_available(): + torch.cuda.set_device(rank) - os.environ["MASTER_ADDR"] = config.master_addr - os.environ["MASTER_PORT"] = config.master_port - if not torch.distributed.is_initialized(): - torch.distributed.init_process_group(backend="nccl", world_size=world_size, rank=rank) + os.environ["MASTER_ADDR"] = config.master_addr + os.environ["MASTER_PORT"] = config.master_port + if not torch.distributed.is_initialized(): + 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: - return_dict["results"] = results - plot_predictions( - model=model, - val_dataset=val_dataset, - save_path="timesfm_predictions.png", - ) + if rank == 0: + return_dict["results"] = results + plot_predictions( + model=model, + val_dataset=val_dataset, + save_path="timesfm_predictions.png", + ) - except Exception as e: - print(f"Error in process {rank}: {str(e)}") - raise e - finally: - if torch.distributed.is_initialized(): - torch.distributed.destroy_process_group() + except Exception as e: + print(f"Error in process {rank}: {str(e)}") + raise e + finally: + if torch.distributed.is_initialized(): + torch.distributed.destroy_process_group() def multi_gpu_example(): - """Example of finetuning TimesFM using multiple GPUs with optimized spawn.""" - mp.set_start_method("spawn", force=True) + """Example of finetuning TimesFM using multiple GPUs with optimized spawn.""" + mp.set_start_method("spawn", force=True) - gpu_ids = [0, 1] - world_size = len(gpu_ids) + gpu_ids = [0, 1] + 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 - config = FinetuningConfig( - batch_size=256, - num_epochs=5, - learning_rate=3e-5, - use_wandb=True, - distributed=True, - gpu_ids=gpu_ids, - ) - train_dataset, val_dataset = get_data(128, tfm_config.horizon_len) - manager = mp.Manager() - return_dict = manager.dict() + # Create config + config = FinetuningConfig( + batch_size=256, + num_epochs=5, + learning_rate=3e-5, + use_wandb=True, + distributed=True, + gpu_ids=gpu_ids, + ) + train_dataset, val_dataset = get_data(128, tfm_config.horizon_len) + 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, - ) + # Launch processes + mp.spawn( + setup_process, + args=(world_size, model, config, train_dataset, val_dataset, return_dict), + nprocs=world_size, + join=True, + ) - results = return_dict.get("results", None) - print("\nFinetuning completed!") - return results + results = return_dict.get("results", None) + print("\nFinetuning completed!") + return results if __name__ == "__main__": - try: - # single_gpu_example() # Single GPU - multi_gpu_example() # Multi-GPU - except Exception as e: - print(f"Training failed: {str(e)}") - finally: - if torch.distributed.is_initialized(): - torch.distributed.destroy_process_group() + try: + # single_gpu_example() # Single GPU + multi_gpu_example() # Multi-GPU + except Exception as e: + print(f"Training failed: {str(e)}") + finally: + if torch.distributed.is_initialized(): + torch.distributed.destroy_process_group() \ No newline at end of file diff --git a/notebooks/finetuning_torch.py b/notebooks/finetuning_torch.py index a3c7f3e..af2eadb 100644 --- a/notebooks/finetuning_torch.py +++ b/notebooks/finetuning_torch.py @@ -18,323 +18,323 @@ import wandb 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. - Concrete implementations can log to different backends (e.g., WandB, TensorBoard). + 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 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 + @abstractmethod + def close(self) -> None: + """Clean up any resources used by the logger.""" + pass 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: - project: Name of the W&B project. - config: Configuration dictionary to log. - rank: Process rank in distributed training. + metrics: Dictionary of metrics to log. + step: Current training step or epoch. """ + if self.rank == 0: + wandb.log(metrics, step=step) - 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() + 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. + """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. - """ + 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 __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 + 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) + 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() + def cleanup(self) -> None: + """Clean up the distributed environment.""" + if dist.is_initialized(): + dist.destroy_process_group() @dataclass class FinetuningConfig: - """Configuration for model training. + """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. - """ + 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" + 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. + """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: - 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. + 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 - 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)) + return DataLoader( + dataset, + batch_size=self.config.batch_size, + shuffle=(is_train and not self.config.distributed), + sampler=sampler, + ) - if config.use_wandb: - self.metrics_logger = WandBLogger(config.wandb_project, config.__dict__, rank) + def _process_batch(self, batch: List[torch.Tensor]) -> tuple: + """Process a single batch of data. - 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() + Args: + batch: List of input tensors. - 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] - ) + Returns: + Tuple of (loss, predictions). + """ + x_context, x_padding, freq, x_future = [t.to(self.device, non_blocking=True) for t in batch] - def _create_dataloader(self, dataset: Dataset, is_train: bool) -> DataLoader: - """Create appropriate DataLoader based on training configuration. + predictions = self.model(x_context, x_padding.float(), freq) + predictions_mean = predictions[..., 0] + last_patch_pred = predictions_mean[:, -1, :] - Args: - dataset: Dataset to create loader for. - is_train: Whether this is for training (affects shuffling). + loss = self.loss_fn(last_patch_pred, x_future.squeeze(-1)) - 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 loss, predictions - return DataLoader( - dataset, - batch_size=self.config.batch_size, - shuffle=(is_train and not self.config.distributed), - sampler=sampler, - ) + def _train_epoch(self, train_loader: DataLoader, optimizer: torch.optim.Optimizer) -> float: + """Train for one epoch. - def _process_batch(self, batch: List[torch.Tensor]) -> tuple: - """Process a single batch of data. + Args: + train_loader: DataLoader for training data. + optimizer: Optimizer instance. - Args: - batch: List of input tensors. + Returns: + Average training loss for the epoch. + """ + self.model.train() + total_loss = 0.0 - Returns: - Tuple of (loss, predictions). - """ - x_context, x_padding, freq, x_future = [t.to(self.device, non_blocking=True) for t in batch] + for batch in train_loader: + loss, _ = self._process_batch(batch) - predictions = self.model(x_context, x_padding.float(), freq) - predictions_mean = predictions[..., 0] - last_patch_pred = predictions_mean[:, -1, :] + if self.config.distributed: + losses = [torch.zeros_like(loss) for _ in range(dist.get_world_size())] + 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: - """Train for one epoch. + return total_loss / len(train_loader) - Args: - train_loader: DataLoader for training data. - optimizer: Optimizer instance. + def _validate(self, val_loader: DataLoader) -> float: + """Perform validation. - Returns: - Average training loss for the epoch. - """ - self.model.train() - total_loss = 0.0 + Args: + val_loader: DataLoader for validation data. - for batch in train_loader: - loss, _ = self._process_batch(batch) + Returns: + Average validation loss. + """ + self.model.eval() + total_loss = 0.0 - 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") + with torch.no_grad(): + for batch in val_loader: + loss, _ = self._process_batch(batch) 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: - 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} \ No newline at end of file