~/wiki

Training Step Anatomy

Mis à jour le 2025-01-04Confiance : high
training-stepsforward-passbackward-passoptimizationmemory-patternsgpu-trainingneural-network-trainingpytorch-profilingmemory-dynamicsfirst-step-anomalyultra-scale-playbookpytorch-caching-allocatoroptimizer-statesactivation-clearingempirical-profilingmemory-allocation-patternsprogressive-memory-usagememory-componentshard-constraintmemory-as-bottleneckmemory-profiling-techniquesdynamic-memory-patternstraining-phasesmemory-component-hierarchycaching-allocator-preparationprogressive-activation-clearingactivation-memory-evolutiongradient-buildup-patternmemory-plateau-phenomenonthree-phase-trainingmemory-variancestep-by-step-analysispytorch-memory-managementoom-failure-patternsmemory-prediction-toolssingle-gpu-foundationmemory-usage-predictiontraining-initializationoptimizer-state-buildupmemory-step-progressiontraining-step-phasesactivation-gradient-lifecyclememory-optimization-foundationbatch-size-memory-relationshipsequence-length-impacttransformer-memory-patternsmemory-empirical-measurementprofiling-methodologymemory-variance-analysisstep-memory-dynamicsmemory-allocation-strategiescuda-kernel-memorymemory-fragmentation-impactprecision-format-impacttensor-shapes-memorymemory-usage-patternssystematic-memory-profilingtraining-step-memory-analysisempirical-memory-understandingmemory-component-evolutionstep-by-step-memory-tracking

The detailed breakdown of what happens during a single training step in neural network training, including memory usage patterns and computational phases. Understanding this anatomy is crucial for optimizing distributed training and memory management, as revealed through systematic PyTorch profiling in the ultra-scale-playbook.

Three Primary Training Phases

A complete training step consists of three sequential phases:

1. Forward Pass

  • Process: Input passes through model layers to generate outputs
  • Memory Pattern: Activations build up progressively as computation moves through layers
  • Storage: Intermediate activations stored for later gradient computation

2. Backward Pass

  • Process: Gradients computed using chain rule, propagating from output to input
  • Memory Pattern: Gradients accumulate while stored activations are progressively cleared
  • Optimization: Activations cleared as soon as corresponding gradients computed

3. Optimization Step

  • Process: Parameters updated using computed gradients
  • Memory Requirement: All gradients needed simultaneously
  • State Update: Optimizer states (momentum, variance) updated after parameter updates

Memory Dynamics Throughout Step

Progressive Memory Allocation

Memory usage follows predictable patterns during each phase:

  1. Forward Phase: Steady increase as activations accumulate
  2. Backward Phase: Peak usage as both activations and gradients exist temporarily
  3. Optimization Phase: Gradients maintained while optimizer states updated

Memory Component Evolution

Different memory components dominate at different phases:

  • Early Forward: Activations dominate memory usage
  • Peak Backward: Both activations and gradients compete for memory
  • Optimization: Optimizer states become prominent factor

First Step Anomaly

The initial training step exhibits distinctly different memory patterns:

Caching Allocator Preparation

  • PyTorch Caching Allocator: Performs extensive preparation work
  • Memory Plateau: Activations plateau longer than subsequent steps
  • Optimization: Pre-allocates memory blocks for faster subsequent allocations

Practical Implications

  • Training may succeed in first step but fail in subsequent steps
  • Optimizer states build up after first step, changing memory requirements
  • Memory predictions based on first step alone can be misleading

Empirical Memory Profiling

PyTorch Profiler Analysis

Systematic profiling reveals:

  • Dynamic Patterns: Memory usage varies significantly within single step
  • Step Variance: Different steps show different memory patterns
  • Component Dominance: Which memory components dominate varies by training phase

Memory Usage Variance

  • Consistent Phases: Forward/backward/optimization phases show consistent patterns
  • Variable Magnitude: Absolute memory usage can vary between steps
  • Progressive Clearing: Activations cleared in reverse order of computation

Memory Optimization Insights

Activation Management Strategy

Understanding step anatomy enables targeted optimization:

  • Gradient Checkpointing: Store activations only at strategic points
  • Progressive Clearing: Clear activations immediately after gradient computation
  • Recomputation: Trade computation for memory by recomputing rather than storing

Memory Prediction

Step anatomy analysis enables:

  • Accurate Estimation: Predict memory requirements based on step phase
  • OOM Prevention: Identify which phase most likely to cause memory issues
  • Optimization Strategy: Target optimizations to memory-dominant phases

Training Step Phases in Detail

Phase 1: Forward Pass Progression

  • Layer-by-layer computation
  • Activation accumulation
  • Memory steadily increases

Phase 2: Backward Pass Dynamics

  • Gradient computation from output to input
  • Simultaneous activation clearing
  • Memory peaks then decreases

Phase 3: Optimization Coordination

  • All gradients required simultaneously
  • Optimizer state updates
  • Memory pattern stabilizes

Scaling Implications

Step anatomy understanding becomes critical for:

  • Single GPU: Optimizing memory usage within device limits
  • Multi-GPU: Coordinating step phases across devices
  • Ultra-Scale: Managing step synchronization across thousands of GPUs

Profiling Methodology

Effective step anatomy analysis requires:

  • PyTorch Profiler: Built-in memory tracking capabilities
  • Step-by-Step Analysis: Individual step profiling rather than aggregate
  • Component Breakdown: Separate tracking of weights, gradients, optimizer states, activations
  • Multiple Steps: Analysis beyond first step to capture true patterns

Understanding training step anatomy provides the foundation for all memory optimization techniques and distributed training strategies.

See also