Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Project: prove multiworker receipt-training parity

Last updated: 7 Oct 20265 min read
project
AdvancedBy AITrove Editorial

Move a receipt-quality classifier from one worker to two while auditing duplicated sampler slots, unequal local batches, checkpoint recovery and metric denominators.

Pin a one-worker reference

Split receipt identities before any augmentation. Record the label map, model revision, random seeds, optimizer, precision and accepted examples per update. Train a tiny deterministic batch in full precision and save its gradients before stepping. Then make two model copies for a three-plus-one split; the global gradient should match the one-worker mean within numerical tolerance. The effective-batch rule is the reference, not the number of loader iterations.

Audit the sampler epoch

Use 47 receipt identities across three simulated ranks and show that padding to 48 slots repeats one identity. Decide whether training accepts that duplicate or uses a masked tail. Call set_epoch consistently on each rank, compare two epoch orders and log unique IDs seen. For validation, gather predictions with receipt IDs and count each identity once. The sampler lesson gives the audit query.

Execute a real two-worker smoke test

On a suitable multi-device runner, initialize a process group, put one model replica on each device and wrap it for synchronous gradient reduction. Give every rank the same number of update boundaries. Compare first-step gradients and parameters with the one-worker reference, then run an epoch and report accepted-example counts by rank. This page’s executable snippet simulates the gradient algebra on one process; it does not claim to start a distributed job.

Interrupt and recover

Save at a completed update with two worker shards and a manifest. Kill one rank halfway through the next save; restart from the last advertised generation and check the first resumed fixed-batch update. Change world size in a separate experiment and state the resulting data-order difference even if model state loads. The manifest rule decides which generation is safe to use.

Publish the measured trade

Compare one-worker and two-worker examples per second, synchronization time, peak memory, class-specific validation results and the number of sampler duplicates. Submit the small gradient-parity assertion, a duplicate-ID report, one injected-failure trace and a deduplicated validation metric. A faster run that silently doubles one rare-class receipt or weights the short rank equally with the long rank has not met the project contract.

Implementation

python
import copy
import torch
from torch import nn

torch.manual_seed(37)
reference_model = nn.Linear(3, 2)
worker_models = [copy.deepcopy(reference_model) for _ in range(2)]
receipt_features = torch.tensor([[0.7, 0.2, 0.1], [0.1, 0.8, 0.1],
                                 [0.2, 0.3, 0.5], [0.6, 0.1, 0.3]])
quality_labels = torch.tensor([0, 1, 1, 0])

reference_loss = nn.functional.cross_entropy(
    reference_model(receipt_features), quality_labels)
reference_loss.backward()
world_size = len(worker_models)
accepted_count = quality_labels.numel()
for worker_model, (first, last) in zip(worker_models, ((0, 3), (3, 4))):
    local_loss_sum = nn.functional.cross_entropy(
        worker_model(receipt_features[first:last]), quality_labels[first:last],
        reduction="sum")
    (local_loss_sum * world_size / accepted_count).backward()

for parameter_index, reference_parameter in enumerate(reference_model.parameters()):
    worker_parameters = [list(worker.parameters())[parameter_index]
                         for worker in worker_models]
    averaged_worker_gradient = sum(parameter.grad for parameter in worker_parameters) / world_size
    torch.testing.assert_close(averaged_worker_gradient, reference_parameter.grad,
                               rtol=1e-6, atol=1e-7)

Performance and operating cost

Each worker stores a full model and optimizer copy; activation memory follows its local microbatch. Synchronization transfers gradients of all parameters per update, and small batches can be communication-bound. The one-process parity check costs two extra model copies but catches loss-weight mistakes before a multi-device run. Validation deduplication uses stable receipt IDs and adds memory proportional to the number of evaluated examples or a sorted streaming merge.

Common Mistakes

  • Do not report the simulated gradient test as a completed multi-device run.
  • Do not divide each rank’s unequal-size mean by worker count and call it a global example mean.
  • Do not count padded sampler duplicates as independent validation receipts.

Read next

ai-data
deep-learning
Storage details