Skip to content

Repository files navigation

AReaL-DTE

Delta Transfer Engine (dte) provides incremental weight synchronization for distributed reinforcement-learning training and inference.

English | 简体中文

Sparse, versioned weight deltas for RL training -> rollout sync, even when training and inference use different TP/PP/EP shard layouts.

dte is an incremental weight-transfer engine for online RL training. It keeps the rollout side in sync with the trainer by sending full weights only when a new base is needed, then sending sparse patches for later steps.

Why dte

Dense weight sync is easy to reason about, but it burns bandwidth every RL step. For bf16 weights, a sparse delta element costs 6 bytes: 4 bytes for the flat index and 2 bytes for the value. If 2% of a tensor changes, the raw payload is about 6% of the dense tensor before transport overhead. Less payload usually means shorter trainer-to-rollout handoff: the rollout side can start using the new policy sooner, and the RL loop spends less time waiting on weight sync. dte adds the machinery needed to make that safe in distributed training:

  • Faster weight handoff. Normal steps move sparse deltas instead of full tensors. When a tensor is not sparse enough, dte falls back to dense for that tensor, so the fast path does not become a penalty.
  • Less data on the wire. For bf16, 2% changed elements produce a raw sparse payload around 6% of dense size. The exact speedup depends on backend, topology, and scatter cost, but the transfer work scales with changed elements instead of total parameters on sparse steps.
  • Bitwise reconstruction. Change detection compares integer views of the floating-point storage. NaN payload bits and signed zero are preserved.
  • Shard-aware sparse transfer. Sparse indices are remapped through train_slices and inf_slices, so deltas can cross TP/PP/EP layout mismatches instead of assuming aligned shards.
  • Version-safe apply. Every delta carries base_version and payload_version; the receiver rejects a patch if it no longer has the right base.
  • Backend isolation. The delta algorithm is independent of NCCL, RDMA, awex, Mooncake, and runtime process management. Backends move bytes; dte defines what those bytes mean.

The transfer model is short:

  1. Send a full checkpoint once to create a receiver-side base.
  2. For later steps, detect which weight elements changed.
  3. Encode the changed flat indices and new values.
  4. Remap those indices when training and inference shards do not line up.
  5. Apply the patch only if the receiver still has the exact base version.

Architecture

Delta Transfer Engine architecture

Diagram source: docs/images/dte_architecture.svg

The diagram keeps only the main boundary: runtimes own model state, DTE owns delta semantics, and backends own byte movement. The smaller modules are:

Module Role
DeltaEngine Public orchestration layer. On the sender it chooses full vs delta, calls the tracker, and sends payloads. On the receiver it checks the version chain and writes reconstructed tensors into target params.
DeltaTracker Sender-side delta state. It seeds a base after full sync, computes bitwise masks from a snapshot or accepts external masks, and emits sparse/dense/unchanged entries.
Delta codec Flat named-tensor protocol: __awex_delta_header__, w@delta_idx, w@delta_val, and dense fallback w. This keeps the payload compatible with existing tensor grouping and IPC paths.
reconstruct_against_base Receiver-side reconstruction. It applies sparse patches on top of the stored base, adopts dense fallback tensors, and refreshes the base for the next version.
remap_delta_indices Converts sparse flat indices from a training shard into the corresponding inference-shard indices using train_slices and inf_slices.
delta_p2p / colocate_protocol Two-round protocol for variable-length sparse P2P: first exchange nnz, then exchange idx and val. This keeps cross-rank scheduling symmetric.
Transport, Plan, TransferOp, Payload Backend contract. Plan describes rank pairs and shard slices; Payload is what moves; Transport defines build_plan, send, and recv.
Backends loopback is the local CPU backend; awex_backend reuses awex reshard plans and NCCL scheduling; mooncake_backend provides the RDMA-oriented path.

For a normal delta step the data path is:

detect -> encode -> transport.send/recv -> decode -> reconstruct -> apply

With this split, bitwise diff, sparse encoding, version checks, and index remap can be tested without a cluster. NCCL, RDMA, MetaServer, and runtime integration stay inside backends.

Algorithm Overview

dte treats weight sync as a versioned state transition:

theta_v --Delta(v -> v+1)--> theta_{v+1}

The receiver stores a full base {name: tensor} for version v. A delta payload is valid only if its header says base_version == v; otherwise the receiver rejects it and the caller must run a full sync. This is the main correctness guard: a sparse patch is not a complete model state.

1. Detect changed elements

For each parameter w, the sender builds a change set:

C_w = { i | bits(theta_t[w][i]) != bits(theta_base[w][i]) }

The comparison is bitwise, not floating-point !=. NaN payload bits, signed zero, bf16 rounding artifacts, and fp32 router weights are handled by comparing same-width integer views of the tensor storage.

Delta mode accepts several change-detection inputs:

  • Snapshot diff: keep a CPU copy of the previous transferred weights and compare current weights against it. This is the standalone DeltaTracker fallback when encode(..., masks=None) is used.
  • AdamW inversion mask: reconstruct the pre-step weights from AdamW moments, then compare against the current weights. AReaL's DTE example launchers use this as their default delta configuration because it avoids a full extra snapshot. For decoupled AdamW:
u_t = (lr / (1 - beta1^step)) * m_t / (sqrt(v_t / (1 - beta2^step)) + eps)
theta_{t-1} = (theta_t + u_t) / (1 - lr * weight_decay)

The inversion path lets callers avoid storing a full snapshot inside DeltaTracker; they pass {name: bool_mask} into encode(...).

  • External masks and indices: callers may also pass boolean masks, sorted integer changed indices, or indices decoded from packed optimizer dirty bits.

DeltaEngine.mode remains either "delta" or "full". Snapshot, AdamW inversion, and dirty-bit detection are strategies used inside delta mode, not additional engine modes. The AReaL integration recommends AdamW inversion for normal runs and keeps snapshot diff available for fallback and verification, including MoE validation.

2. Choose sparse or dense per tensor

For a tensor with n elements, element size s, and k = |C_w| changed elements:

dense_bytes(w)  = n * s
sparse_bytes(w) = k * (4 + s)       # int32 flat index + value

DeltaTracker emits sparse entries only when:

k > 0
n < 2^31
sparse_bytes(w) <= sparse_bytes_ratio * dense_bytes(w)

Otherwise the tensor is sent dense for that step. This makes a payload heterogeneous by design:

__awex_delta_header__ -> [magic, codec, payload_version, base_version, counts]
w@delta_idx           -> int32[k]
w@delta_val           -> dtype[k]
w                     -> dense fallback tensor

Unchanged tensors are omitted. Dense fallback tensors refresh the receiver base just like a full sync for that tensor.

3. Remap sparse indices across shard layouts

Sparse indices are computed in the training shard's flat index space. If the rollout side uses a different TP/PP/EP layout, those indices must be projected through the transfer plan.

For each TransferOp, dte uses the source and destination overlap:

train_slices = op.train_slices
inf_slices   = op.inf_slices

The remap is:

p  = training-shard flat index
c  = unravel(p, train_shape)
keep c only if c is inside train_slices
c' = c - train_start + inf_start
p' = ravel(c', infer_shape)

The value tensor is filtered with the same mask, so the receiver scatters values[j] into p'[j] in its own inference shard. The remap function is pure tensor geometry; NCCL, RDMA, and IPC are still backend concerns.

4. Reconstruct on the receiver

After transport, the receiver decodes the payload and rebuilds full tensors against its base:

if tensor is dense:
    full = payload[name]
elif tensor is sparse:
    full = clone(base[name])
    full[delta_idx] = delta_val
else:
    full = clone(base[name])

The reconstructed tensors are written into live inference parameters, and the CPU base is refreshed to payload_version. In colocate NCCL paths, variable-size sparse patches use a two-round protocol: first exchange nnz for every op, then exchange (idx, val) buffers. Zero-nnz ops still participate so every rank walks the same schedule.

Repository Layout

  • src/dte/core/: delta algorithms: bitwise diff, AdamW inversion helper, sparse codec, reconstruction, shard index remap, and colocate P2P protocol logic.
  • src/dte/engine.py: DeltaEngine, the sender/receiver control flow.
  • src/dte/transport.py: backend-neutral Transport, Plan, TransferOp, and Payload types.
  • src/dte/backends/loopback.py: CPU in-memory backend used by tests.
  • src/dte/backends/awex_backend.py: awex adapter for transfer plans and NCCL scheduling.
  • src/dte/backends/mooncake_backend.py: Mooncake/RDMA-oriented packing path and transport interface.
  • docs/design.md: full design, equations, protocol invariants, and validation notes.
  • docs/awex-gpu-verification.md: GPU parity checklist for AwexTransport.

Install

Requirements:

  • Python >= 3.11 and < 3.13
  • PyTorch >= 2.9.1 and < 2.11 on the primary Linux stack

Install from a source checkout:

python -m pip install -e .

For development:

python -m pip install -e ".[dev]"

Optional transport backends:

python -m pip install -e ".[awex]"      # dingzhiqiang/asystem-awex + NCCL runtime
python -m pip install -e ".[mooncake]"  # Mooncake Transfer Engine

The default install is enough for the core algorithm and the loopback backend. Cluster backends need their runtime dependencies and hardware environment. After a PyPI release, the same extras can be installed from delta-transfer-engine[...].

Quickstart

from dte import DeltaEngine
from dte.backends import LoopbackTransport

engine = DeltaEngine(
    transport=LoopbackTransport(),
    mode="delta",
    anchor_interval=10,
)

# Trainer side.
engine.push(model.named_parameters(), version=step)

# Rollout side.
engine.pull(target_params, version=step)

mode="delta" seeds with a full sync, uses sparse deltas for normal steps, and can force periodic full anchors. mode="full" skips detection and sends dense weights every step. The standalone engine uses snapshot diff unless an integration supplies external masks; the AReaL examples supply AdamW-inversion masks by default.

Development

python -m pip install -e ".[dev]"
ruff check .
ruff format --check .
mdformat --check README.md README.zh-CN.md CONTRIBUTING.md docs
python -m pytest tests -q
python -m build

The CPU test suite covers the core algorithm, DeltaEngine, loopback transport, and backend contracts that do not require a live cluster. GPU parity for awex is tracked separately in docs/awex-gpu-verification.md.

Current Status

Area State
Core delta algorithm CPU tests cover snapshot diff, external masks and sorted indices, packed dirty bits, sparse/dense fallback, tied-storage snapshot dedup, remap fast paths, AdamW inversion, and version chains.
DeltaEngine CPU end-to-end tests cover full/delta sync, anchors, no-op deltas, broken chains, and transactional decode_for_live_apply / commit_live_apply.
loopback backend Local test backend.
awex backend Adapter uses dingzhiqiang/asystem-awex; the colocated sparse path supports coalesced two-round metadata/data exchange. Live cross-rank parity still needs NCCL + MetaServer + megatron/sglang.
mooncake backend Interface and pack/unpack path are present. RDMA execution needs a Mooncake runtime.

License

This project is licensed under the Apache License 2.0.

Releases

Packages

Contributors

Languages