ML Platforms and Frameworks

ML Platforms and Frameworks #

Training at scale requires more than dividing data between GPUs. A platform must decide where model state lives, who computes gradients, how updates are combined, and how information moves between devices.

This page focuses on the Parameter Server model and the practical constraints that determine whether adding resources makes training faster.

Topics covered:

  • Parameter servers, workers, and independent scaling of storage and computation
  • Training-state memory and parameter sharding
  • Distributed loss, gradient aggregation, and the pull–compute–push cycle
  • Communication costs, RDMA, MPI, PCIe, and NVLink
  • Global batch size, learning-rate scaling, and warm-up
  • Monitoring, fault recovery, and sparse versus dense workloads

Learning Objectives #

By the end of this page, you should be able to:

  • distinguish parameter-server responsibilities from worker responsibilities
  • estimate model-state memory without confusing it with weights alone
  • write a consistent distributed SGD update
  • calculate ideal compute scaling and gradient-transfer time
  • explain why communication can erase the benefit of extra workers
  • adjust batch size and learning rate together, while checking convergence
  • choose an architecture using workload characteristics rather than GPU count alone

Big Picture #

A parameter server acts like a shared model-state service. Workers process different examples, calculate proposed changes, and send those changes to the servers that own the relevant parameters.

flowchart TD
    W1["Worker 1: data shard A"] <-->|"Pull parameters; push gradients"| S1["Server 1: parameter shard 1"]
    W1 <-->|"Pull parameters; push gradients"| S2["Server 2: parameter shard 2"]
    W2["Worker 2: data shard B"] <-->|"Pull parameters; push gradients"| S1
    W2 <-->|"Pull parameters; push gradients"| S2

    style W1 fill:#E1F5FE
    style W2 fill:#E1F5FE
    style S1 fill:#C8E6C9
    style S2 fill:#C8E6C9

Each worker may need parameters from every server, even though it processes only its own data shard. Data partitioning and parameter partitioning solve different problems.

For the broader choice between data, model, and pipeline parallelism, see Distributed Training Strategies.

1. Why Distribute Training? #

Two constraints often appear together:

  1. Computation: processing a large dataset takes too long on one worker.
  2. Memory: weights, gradients, and optimiser state exceed the available memory.

Let P be the number of model parameters, N the number of examples, and c the approximate computation per example.

Model-related storage grows with P, while processing the dataset requires approximately:

\[ W_{\text{compute}}=O(Nc) \]

Distributing examples addresses computation. Distributing model state addresses storage. A useful design must also account for the communication introduced by both.

Weights Alone Are Not the Training Memory #

FP16 weights use two bytes per parameter. An illustrative mixed-precision training configuration with Adam uses about 16 bytes per parameter for weights, gradients, master weights, and optimiser states combined.

\[ M_{\text{weights}}=2P\ \text{bytes}, \qquad M_{\text{training state}}\approx16P\ \text{bytes} \]

Using decimal units, where 1 GB = 10^9 bytes:

ParametersFP16 weights aloneTraining state at 16 bytes per parameter
1 billion2 GB16 GB
7 billion14 GB112 GB
70 billion140 GB1.12 TB
1 trillion2 TB16 TB

A seven-billion-parameter model can therefore have 14 GB of weights but approximately 112 GB of training state under this assumption.

The 16-byte estimate is configuration-dependent, not a universal constant. It excludes activations, input data, temporary workspaces, and allocator overhead. Weights fitting in memory does not imply that training fits.

2. Parameter Servers and Workers ☆ #

The architecture separates state management from gradient computation.

RoleMain responsibilitiesTypical pressure
Parameter serverOwn parameter shards, receive gradients, apply updates, return current parameters; often own corresponding optimiser stateMemory capacity, update throughput, network bandwidth
WorkerRead a data shard, obtain parameters, run forward and backward passes, send gradientsCompute throughput, local memory, input throughput

The parameter service is logically centralised, because it manages a shared model. It need not be one physical machine: several servers can each own a different shard.

A worker is also a logical role. A machine with several GPUs may run several workers; “one worker” does not necessarily mean “one machine”.

Independent Scaling #

  • Add workers when gradient computation is the bottleneck.
  • Add parameter servers when server-side memory or update service is the bottleneck.
  • Improve the communication path when workers spend most of their time waiting for parameters or sending gradients.

This separation allows resources to match their jobs. It does not guarantee that either side can scale without limits.

Parameter Sharding #

With S servers, partition the model into S disjoint parameter groups:

\[ \theta=[\theta_1,\theta_2,\ldots,\theta_S] \]

For balanced shards, each server owns approximately P/S parameters. If all of the estimated training state is sharded evenly:

\[ M_{\text{state per server}}\approx\frac{16P}{S}\ \text{bytes} \]

For the seven-billion-parameter example, four servers would each hold approximately 28 GB of that state, before additional overhead.

Sharding server-side state does not automatically shard each worker’s forward-pass model. In the basic data-parallel arrangement, workers still hold local model replicas. If a replica and its working memory do not fit, additional model or tensor partitioning is needed.

3. Distributed Loss and Gradient Updates ☆ #

Local and Global Loss #

Worker k owns a dataset shard containing N_k examples. Its local average loss is:

\[ L_k(\theta)=\frac{1}{N_k}\sum_{i\in D_k}\ell(x_i,y_i;\theta) \]

Here, D_k is the worker’s data shard and ℓ measures the error on one example. Training aims to find parameters that minimise the loss. For a global mean over all N examples:

\[ L(\theta)=\sum_{k=1}^{K}\frac{N_k}{N}L_k(\theta), \qquad N=\sum_{k=1}^{K}N_k \]

When all shards are equally sized, this becomes the mean of the K local losses.

Gradients from Mini-Batches #

During one synchronous step, each worker evaluates its local mini-batch using the same parameter version. If its mini-batch contains b_k examples, its mean gradient is:

\[ g_k=\frac{1}{b_k}\sum_{i\in B_k}\nabla_\theta\ell(x_i,y_i;\theta^{(t)}) \]

The combined mean gradient weights each worker by its batch size:

\[ g=\sum_{k=1}^{K}\frac{b_k}{B_{\text{global}}}g_k, \qquad B_{\text{global}}=\sum_{k=1}^{K}b_k \]

For equal local batch sizes:

\[ g=\frac{1}{K}\sum_{k=1}^{K}g_k, \qquad \theta^{(t+1)}=\theta^{(t)}-\eta g \]

The learning rate η controls the size of the update. Each server applies the component belonging to its own shard:

\[ \theta_s^{(t+1)}=\theta_s^{(t)}-\eta\frac{1}{K}\sum_{k=1}^{K}g_{k,s} \]

Here, g_{k,s} is the portion of worker k’s gradient associated with shard s.

Sum Versus Mean: Keep the Convention Consistent #

Distributed updates are also written using a sum of worker gradients. This is consistent when local objectives are defined as contributions to a sum, or when the learning rate includes the required normalisation.

If workers send local mean gradients for equal batches, summing them without dividing by K makes the update K times larger than averaging them at the same learning rate.

For example, with four workers, Σg_k = 4 × mean(g_k). Never switch silently between these conventions.

4. The Pull–Compute–Push Cycle #

One basic synchronous iteration proceeds as follows:

  1. Pull: workers obtain the current parameter shards and assemble their local model state.
  2. Compute: each worker processes its mini-batch and calculates gradients independently.
  3. Push: workers send gradient components to the servers that own the corresponding parameters.
  4. Aggregate and update: servers combine the required contributions and update their shards.
  5. Repeat: workers obtain the next parameter version and continue training.
flowchart TD
    A["Pull current parameters"] --> B["Compute local gradients"]
    B --> C["Push gradient shards"]
    C --> D["Aggregate and update"]
    D --> A

    style A fill:#E1F5FE
    style B fill:#C8E6C9
    style C fill:#FFF9C4
    style D fill:#EDE7F6

Synchronisation and Staleness #

In synchronous training, the update waits for the required workers. A slow worker can delay the group.

In asynchronous training, updates can be applied as they arrive. Workers may then calculate gradients using older parameters, producing stale gradients.

A system can instead use an explicit threshold, such as accepting 90 contributions out of 100 workers. This reduces waiting but changes the update policy. It must define which versions are accepted and how the selected gradients are normalised; it is not automatically equivalent to either full synchrony or unrestricted asynchronous SGD.

5. Compute Scaling Versus Communication ☆ #

Ideal Compute Scaling #

For a fixed total workload split evenly among K workers, the compute portion ideally falls to:

\[ T_{\text{compute}}(K)\approx\frac{T_{\text{compute}}(1)}{K}, \qquad W_{\text{per worker}}=O\!\left(\frac{Nc}{K}\right) \]

An illustrative workload requiring 64 days of compute on one worker would have these ideal times:

WorkersIdeal compute time
164 days
416 days
164 days
641 day

These figures exclude communication, synchronisation, imbalance, and input bottlenecks. They also do not describe a run where more workers simultaneously increase the total workload.

Total Step Time #

Without overlap, a simple model is:

\[ T_{\text{step}}\approx T_{\text{compute}}+T_{\text{communication}} \]

As compute becomes faster, communication may become the dominant term. Systems that overlap these activities require a more detailed timing model.

Communication Volume #

For a dense gradient with P entries and b_g bytes per entry, one worker’s gradient push contains:

\[ V_{\text{push}}=b_gP\ \text{bytes} \]

With K workers, a full push and pull per worker generates aggregate payload volume of approximately:

\[ V_{\text{total}}\approx K(b_g+b_\theta)P\ \text{bytes} \]

Here, b_θ is the number of bytes per transferred parameter. This excludes metadata, retries, and protocol overhead.

At fixed precision, traffic therefore grows roughly with worker count × parameter count. On a shared bottleneck link, more workers can increase waiting rather than useful throughput.

More server shards can increase coordination and the number of server contacts. A simplified O(S) contact count is not a universal wall-clock law: sharding can also distribute memory pressure and provide parallel network paths. Distinguish message count, transferred bytes, and elapsed time.

Worked Example: Transfer a 70-Billion-Parameter Gradient #

Assume a full FP16 gradient and decimal units:

\[ V=70\times10^9\times2=140\times10^9\ \text{bytes}=140\ \text{GB} \]

A 100 Gbit/s link has an ideal byte rate of:

\[ R=\frac{100}{8}=12.5\ \text{GB/s} \]

The ideal transfer time for one gradient push is:

\[ T_{\text{transfer}}=\frac{V}{R}=\frac{140}{12.5}=11.2\ \text{s} \]
Link rateIdeal byte rate7-billion-parameter FP16 gradient70-billion-parameter FP16 gradient
100 Gbit/s12.5 GB/s1.12 s11.2 s
400 Gbit/s50 GB/s0.28 s2.8 s
800 Gbit/s100 GB/s0.14 s1.4 s

These are payload-only lower bounds for one direction, assuming the whole stated link rate is available. They exclude the parameter pull, contention, and software overhead. Gbit/s is not GB/s, and a combined bidirectional bandwidth figure is not necessarily available to a one-way transfer.

6. Communication Tools and Hardware Paths #

RDMA, MPI, PCIe, and NVLink belong to different layers of the system. They are not interchangeable names for “a faster network”.

TechnologyWhat it isWhy it matters
RDMARemote Direct Memory Access, used for network transfers between registered memory regionsReduces CPU involvement and operating-system overhead in the transfer data path
MPIMessage Passing Interface: a standard programming interface for communication between processesProvides communication operations, including blocking and non-blocking forms
PCIePeripheral Component Interconnect Express: a general-purpose system interconnectConnects components such as GPUs, network cards, and SSDs to the host system
NVLinkA high-bandwidth GPU interconnectSupports fast communication between connected GPUs in supported topologies

Important Distinctions #

  • RDMA does not mean zero CPU use: setup, memory registration, and control still require system work.
  • An MPI non-blocking call does not imply asynchronous SGD: communication semantics and training-update policy are separate choices.
  • NVLink does not replace every inter-node network path: local GPU connectivity and cluster networking must both be considered.
  • PCIe lane allocation matters: several GPUs, network cards, and SSDs may compete for the host’s available connectivity.
  • Memory bandwidth matters as well as capacity: CPU workloads that repeatedly read large tensors can stall even when sufficient RAM is installed.

The useful design question is: which path carries each transfer, and what else shares that path? Faster GPUs cannot compensate indefinitely for a congested interconnect or slow input path.

7. Batch Size and Learning Rate ☆ #

Global Batch Size #

With K workers and equal local mini-batch size B_local:

\[ B_{\text{global}}=K B_{\text{local}} \]

Eight workers processing 32 examples each produce a global batch of 256. Increasing to 32 workers while keeping the local batch unchanged produces 1,024 examples per update.

This changes the optimisation process: each update uses more examples, and a fixed dataset requires fewer updates per pass.

Linear Learning-Rate Scaling #

A common starting heuristic scales the learning rate in proportion to global batch size:

\[ \eta_{\text{new}}=\eta_{\text{base}}\frac{B_{\text{new}}}{B_{\text{base}}} \]

For baseline batch size 256 and learning rate 0.1:

WorkersLocal batchGlobal batchLinearly scaled learning rate
8322560.1
32321,0240.4
128324,0961.6
256328,1923.2

For 32 workers:

\[ \eta_{\text{new}}=0.1\times\frac{1024}{256}=0.4 \]

The table illustrates the heuristic, not a guarantee that every listed rate will train stably. The model, optimiser, data, and gradient-normalisation convention all matter.

Warm-Up #

Warm-up begins with a smaller learning rate and gradually increases it towards the target rate. A later schedule may then reduce the rate.

This avoids immediately applying a large update before early training has stabilised. It is not the same as starting with the largest rate and decreasing it from the first step.

When scaling a run:

  1. calculate the new global batch
  2. choose an initial learning-rate rule
  3. introduce or adjust warm-up where appropriate
  4. compare training loss and validation quality
  5. retain the change only if throughput and convergence together improve

8. Monitoring and Fault Recovery #

Measure Where Time Goes #

A useful monitoring dashboard separates symptoms instead of relying on GPU utilisation alone.

ObservationWhat to investigate
GPUs idle between burstsParameter transfers, synchronisation, or input delivery
Servers heavily loaded while workers waitServer update throughput, memory access, or network contention
One worker consistently finishes lateUneven data, device performance, or local input bottlenecks
Faster steps but poorer validation qualityBatch size, learning rate, warm-up, or stale updates

Run short, comparable trials and record step time, throughput, network activity, memory use, and loss behaviour. The best configuration is not necessarily the largest one: seek a good time to acceptable model quality, with sensible resource cost.

Worker Failure Is Not Server Failure #

If a worker fails, a system may retry its work, restart it, or explicitly reconfigure the active worker set. A synchronous run must handle the missing contribution rather than wait forever.

If training continues with fewer workers, aggregation and global batch size may need to change. Simply dropping a worker’s gradient does not preserve the original update automatically.

If a parameter server fails, its shard and optimiser state must be recovered or reconstructed. Distribution alone does not provide that recovery: checkpoints, replication, or another explicit mechanism are required.

Operational Lesson #

A service can continue running on healthy GPUs while still sending new requests to a failed device. Fault detection must therefore connect to routing and scheduling, not merely to an alert.

Preserve a recoverable checkpoint or deployment version, validate changes separately, and confirm that work actually reaches healthy resources after recovery.

9. Where Parameter Servers Fit Best #

Sparse Embedding Workloads #

Recommendation systems can have very large embedding tables. A particular batch may touch only a small subset of the rows.

Rather than transferring an entire table, workers can request the relevant rows and send updates for the entries they used. Parameter servers are well suited to this combination of large shared state and sparse access.

Dense Neural-Network Workloads #

In dense training, a step may generate gradients for nearly every parameter. Repeatedly pushing and pulling full tensors through central services can create a communication bottleneck.

Collective communication such as all-reduce is therefore often useful for dense gradient aggregation. Sharded-state approaches can also distribute memory without requiring a conventional parameter-server service.

Workload propertyArchitectural implication
Huge embedding table; each batch touches few rowsSparse pull/push through parameter servers can be effective
Dense gradients; replicated model fitsCollective gradient aggregation is often attractive
Full model or training state does not fit locallyAdditional state sharding or model partitioning is needed

These are design tendencies, not absolute rules. Access patterns, hardware topology, consistency requirements, and implementation determine the final choice.

Preserve Useful Model Behaviour #

Further training for a narrow domain can improve specialised knowledge while weakening other abilities. Retain earlier checkpoints and evaluate both the new task and previously useful behaviour.

In some applications, a separate component can refine presentation or handle another task instead of repeatedly retraining the same model. Composition is a design option, not a guarantee that quality problems disappear.

10. Common Mistakes #

MistakeCorrection
Treat FP16 weight memory as total training memoryInclude gradients, optimiser state, activations, and overhead
Assume server sharding makes every worker’s model fitCheck the worker’s own memory and execution layout
Sum local mean gradients but use the averaging learning rate unchangedKeep loss normalisation, aggregation, and learning rate consistent
Divide by 100 GB/s for a 100 Gbit/s linkConvert bits to bytes first: 100/8 = 12.5 GB/s
Use total bidirectional bandwidth for one-way transfer timeUse the available rate for that direction and topology
Add workers without checking global batch sizeRecalculate batch size and validate the optimisation settings
Expect linear end-to-end speedupInclude communication, synchronisation, imbalance, and input costs
Assume distribution automatically handles every failureDefine recovery, checkpointing, and worker-membership policies

Practice Questions #

These are illustrative checks based on the concepts above.

  1. A model has seven billion parameters. Estimate FP16 weight memory and training-state memory at 16 bytes per parameter.
    • Answer: 14 GB of weights and 112 GB of training state, excluding activations and other overhead.
  2. Four parameter servers share that estimated training state evenly. How much does each hold?
    • Answer: 28 GB, before overhead. This does not determine worker memory.
  3. A fixed workload requires 64 days of compute on one worker. What is the ideal compute time on 16 workers?
    • Answer: 4 days, before communication and other overheads.
  4. How long does one 140 GB gradient push take over an ideal 100 Gbit/s link?
    • Answer: 11.2 seconds, because the ideal byte rate is 12.5 GB/s.
  5. Thirty-two workers use local batches of 32. What is the global batch? If the baseline is batch 256 at learning rate 0.1, what does linear scaling suggest?
    • Answer: global batch 1,024 and candidate learning rate 0.4; validate stability and quality.
  6. Why can adding workers make training slower?
    • Answer: communication and waiting can grow faster than the compute portion shrinks.
  7. Why are parameter servers attractive for large embedding tables?
    • Answer: they distribute shared state while allowing sparse access and updates to only the rows used.

Key Takeaways #

  • Parameter servers manage shared model state; workers calculate gradients.
  • Storage scaling and compute scaling are distinct from communication scaling.
  • A correct distributed update keeps gradient normalisation and learning rate consistent.
  • Transfer time depends on payload bytes and usable bandwidth, not a network label alone.
  • Larger worker counts change global batch size unless local batches are adjusted.
  • Monitor time to useful model quality and define recovery explicitly.

Understanding Checklist #

  • I can distinguish workers, parameter shards, and physical machines.
  • I can estimate weights-only memory and a stated training-state budget.
  • I can explain pull, compute, push, aggregate, and update.
  • I can distinguish summed gradients from averaged gradients.
  • I can convert Gbit/s to GB/s and calculate a one-way transfer lower bound.
  • I can calculate global batch size and a candidate scaled learning rate.
  • I can explain warm-up, stale gradients, and explicit fault recovery.
  • I can compare sparse embedding access with dense gradient communication.

Home | ML System Optimisation