You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

SageMaker中基于PyTorch的Temporal Fusion Transformer分布式训练示例及实践模式咨询

PyTorch Distributed Temporal Fusion Transformer (TFT) Training on Amazon SageMaker: Example & Best Practices

Got it, let's break this down for you—training a large Temporal Fusion Transformer (TFT) with PyTorch on SageMaker using distributed training is totally doable, and I’ve walked through this exact setup for a time-series forecasting project recently. Below is a concrete, runnable example plus key best practices to avoid common pitfalls.

Core Setup Overview

SageMaker’s PyTorch integration natively supports Distributed Data Parallel (DDP)—the standard for multi-GPU/multi-node PyTorch training. We’ll use the PyTorchEstimator to spin up distributed instances, and leverage environment variables SageMaker sets automatically to handle process grouping. For TFT, we’ll use pytorch-forecasting (a battle-tested library for time-series models) instead of building from scratch to save time.

Full Example Code

1. Training Script (train.py)

This script handles distributed init, data loading, model training, and checkpointing:

import os
import torch
import pytorch_lightning as pl
from pytorch_forecasting import TemporalFusionTransformer, TimeSeriesDataSet
from pytorch_forecasting.data import GroupNormalizer
from pytorch_forecasting.metrics import QuantileLoss
from pytorch_lightning.strategies import DDPStrategy

# SageMaker auto-set distributed training env vars
WORLD_SIZE = int(os.environ.get("WORLD_SIZE", 1))
LOCAL_RANK = int(os.environ.get("LOCAL_RANK", 0))
NODE_RANK = int(os.environ.get("NODE_RANK", 0))

def main():
    # ----------------------
    # 1. Load & Prep Time-Series Data
    # ----------------------
    # Replace this with your own S3-hosted dataset (use s3fs to load directly)
    data = TimeSeriesDataSet.generate_dummy_data()
    max_encoder_length = 30  # Past time steps to use for prediction
    max_decoder_length = 7   # Future time steps to predict

    training_dataset = TimeSeriesDataSet(
        data,
        time_idx="time_idx",
        target="value",
        group_ids=["series"],
        static_categoricals=["series"],
        time_varying_known_reals=["time_idx"],
        time_varying_unknown_reals=["value"],
        target_normalizer=GroupNormalizer(groups=["series"]),
        max_encoder_length=max_encoder_length,
        max_decoder_length=max_decoder_length,
    )

    # Create distributed-friendly dataloaders
    train_dataloader = training_dataset.to_dataloader(
        train=True, 
        batch_size=64, 
        num_workers=4,
        shuffle=True
    )

    # ----------------------
    # 2. Initialize TFT Model
    # ----------------------
    tft_model = TemporalFusionTransformer.from_dataset(
        training_dataset,
        learning_rate=0.03,
        hidden_size=128,  # Scale this for larger models
        attention_head_size=4,
        dropout=0.1,
        hidden_continuous_size=32,
        output_size=7,
        loss=QuantileLoss(),
        log_interval=10,
    )

    # ----------------------
    # 3. Distributed Training Setup
    # ----------------------
    strategy = DDPStrategy(
        accelerator="gpu",
        devices="auto",
        num_nodes=WORLD_SIZE,
        rank=NODE_RANK * torch.cuda.device_count() + LOCAL_RANK,
    )

    trainer = pl.Trainer(
        max_epochs=15,
        accelerator="gpu",
        devices="auto",
        strategy=strategy,
        enable_checkpointing=True,
        default_root_dir=os.environ.get("SM_MODEL_DIR", "./model"),
        logger=True,
        enable_progress_bar=LOCAL_RANK == 0,  # Only show progress on lead node
    )

    # ----------------------
    # 4. Start Training
    # ----------------------
    trainer.fit(tft_model, train_dataloaders=train_dataloader)

if __name__ == "__main__":
    main()

2. SageMaker Launch Script (sagemaker_launch.py)

This script configures the SageMaker estimator to spin up distributed nodes:

import sagemaker
from sagemaker.pytorch import PyTorchEstimator

# Initialize SageMaker session
sess = sagemaker.Session()
role = sagemaker.get_execution_role()

# Configure distributed training (DDP enabled)
distribution_config = {
    "torch_distributed": {
        "enabled": True
    }
}

# Define estimator
tft_estimator = PyTorchEstimator(
    entry_point="train.py",
    role=role,
    instance_count=2,  # Number of distributed nodes
    instance_type="ml.p3.8xlarge",  # Use ml.p4d.24xlarge for large models
    framework_version="2.1",
    py_version="py310",
    distribution=distribution_config,
    output_path="s3://your-bucket/tft-training-output",
    checkpoint_s3_uri="s3://your-bucket/tft-checkpoints",  # Resume training if nodes fail
    hyperparameters={
        "max_epochs": 15,
        "batch_size": 64
    },
    use_spot_instances=True,  # Save up to 70% on costs
    max_wait=3600  # Wait time for spot instances
)

# Start training job (point to your S3 training data)
tft_estimator.fit({"training": "s3://your-bucket/tft-training-data/"})

Key Best Practices

  • Instance Selection: For large TFT models (hidden size > 256), use ml.p4d.24xlarge (A100 GPUs) for 4x more memory and faster training. For smaller workloads, ml.p3.8xlarge is cost-effective.
  • Data Optimization:
    • Preprocess time-series data with SageMaker Data Wrangler to avoid runtime bottlenecks.
    • Use s3fs to load data directly from S3 without local downloads—reduces I/O lag.
  • Checkpointing: Always enable SageMaker’s checkpointing (as shown above) to resume training if a node fails. Store checkpoints in S3, not local instance storage.
  • Large Model Handling: For TFTs too big to fit in a single GPU, use Fully Sharded Data Parallel (FSDP) by adding "fsdp": True to the torch_distributed config. This splits model weights across GPUs.
  • Monitoring: Enable SageMaker Debugger to track GPU utilization and memory leaks. Integrate TensorBoard by passing tensorboard=True to the estimator.
  • Cost Savings: Use Spot Instances for non-critical training runs—you can save up to 70% on costs. Just add use_spot_instances=True and max_wait to the estimator.

Common Pitfalls to Avoid

  • Don’t hardcode GPU counts or ranks—always use SageMaker’s environment variables (WORLD_SIZE, LOCAL_RANK) to keep the script portable.
  • For TFT, ensure max_encoder_length and max_decoder_length match your time-series frequency (e.g., 90 days for daily data).
  • Avoid loading the entire dataset into memory—use PyTorch’s IterableDataset for large time-series datasets.

内容的提问来源于stack exchange,提问作者Philipp Schmid

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.27 14:22:37