nemo_rl.data_plane.interfaces#

Stable boundary between NeMo-RL and data-plane implementations.

Wire shape adapters must support:

  • fields: TensorDict with tensor leaves AND optional NonTensorStack / NonTensorData leaves (TQ-native non-tensor passthrough). TQ’s storage backends handle encoding per backend (simple keeps Python objects; mooncake_client pickles internally).

  • tags: list[dict[str, Any]] per-sample primitives (kept separate from fields so non-tensor metadata like input_lengths doesn’t pollute the leaf-level schema).

  • keys: per-sample string uids.

  • partition_id: string-named address spaces with declared consumer_tasks and fields schemas.

All call sites in nemo_rl/algorithms, nemo_rl/experience and nemo_rl/models go through :class:DataPlaneClient — never import transfer_queue directly. This is what makes the implementation swappable.

See nemo_rl/data_plane/README.md for the full design.

Module Contents#

Classes#

SimpleStorageConfig

Sizing for backend="simple". Ignored by every other backend.

MooncakeCpuConfig

Sizing and RDMA knobs for backend="mooncake_cpu". Ignored otherwise.

DataPlaneConfig

Feature-gated config; defaults to disabled.

ObservabilityConfig

Optional middleware that records per-op metrics on the client.

LocalDataPlaneConfig

User configuration for process-local data transfer.

KVBatchMeta

Per-batch metadata for data-plane KV operations.

DataPlaneClient

Stable, swappable data-plane boundary.

Functions#

data_plane_supports_checkpointing

Return whether the configured backend supports complete save/load.

backend_config

Return the validated sizing block for cfg["backend"].

Data#

API#

nemo_rl.data_plane.interfaces.DATA_PLANE_CHECKPOINT_SCHEMA_VERSION#

2

class nemo_rl.data_plane.interfaces.SimpleStorageConfig#

Bases: pydantic.BaseModel

Sizing for backend="simple". Ignored by every other backend.

num_storage_units scales with the cluster: TQ round-robins storage units over Ray nodes and recommends >= 2 x the node count. No static default is correct across cluster sizes, so this is required rather than defaulted — a class field cannot see cluster.num_nodes, only the exemplar YAML can, via ${mul:2, ${cluster.num_nodes}}. Every recipe inherits that from the exemplar; set a plain int to pin it.

storage_capacity: int#

1000000

num_storage_units: int#

None

class nemo_rl.data_plane.interfaces.MooncakeCpuConfig#

Bases: pydantic.BaseModel

Sizing and RDMA knobs for backend="mooncake_cpu". Ignored otherwise.

global_segment_size / local_buffer_size are per client process (one per GPU), so a node pays gpus_per_node x (segment + buffer). Under RDMA that memory is pinned and resident from setup, so keep the per-node product in mind when raising them. Under-sizing surfaces as batch_get_tensor returned None.

reuse_registered_buffers keeps a pool of RDMA-registered buffers alive instead of registering a fresh one per transfer; set false to fall back to upstream’s per-call registration.

staging_buffer_size is that pool’s per-slot ceiling. It is a pooling threshold, not a size limit: a bigger payload still transfers, just with a transient registration. Slots ratchet — they grow to the largest payload admitted and never shrink — so raise it only when a per-key payload (one sample of one field) genuinely exceeds it, not for headroom.

use_gdr lets CUDA-initialized clients transfer through TransferQueue’s persistent GPU staging buffer. gdr_staging_buffer_mb is the positive HBM capacity of that buffer per active GDR client. CPU-only clients keep using the registered host-buffer path.

Every RDMA rail on the host is offered to mooncake (see rdma_devices). That is only safe with MC_ENABLE_DEST_DEVICE_AFFINITY=1, which pins each transfer’s peer rail to the local one by name; on a rail-isolated RoCE fabric a cross-rail pair has no route. It is set on RoCE-only hosts by nemo_rl.data_plane.adapters.transfer_queue_env.configure_engine_env.

global_segment_size: int#

68719476736

local_buffer_size: int#

4294967296

reuse_registered_buffers: bool#

True

staging_buffer_size: int#

268435456

use_gdr: bool#

False

gdr_staging_buffer_mb: pydantic.PositiveInt#

1024

class nemo_rl.data_plane.interfaces.DataPlaneConfig#

Bases: typing.TypedDict

Feature-gated config; defaults to disabled.

backend is the storage backend inside TransferQueue; it is owned by the TQ adapter, not by NeMo-RL. impl selects which adapter we go through.

Backend-specific knobs live under a block named for the backend that reads them — simple: and mooncake_cpu: — mirroring TransferQueue’s own config.yaml and the per-backend overlay :func:_init_tq builds. Only the block named by backend is consulted, so a config selecting simple never has to mention mooncake’s RDMA sizing at all. An absent mooncake_cpu: block means “use :class:MooncakeCpuConfig’s defaults” — but simple: is not optional: num_storage_units has no static default, since no single value is right across cluster sizes, so a simple run without the block fails validation.

Required keys (always set in the exemplar YAML): enabled, impl, backend, claim_meta_poll_interval_s.

storage_capacity / num_storage_units / global_segment_size / local_buffer_size used to sit at this level. A config still using that spelling is not rejected — the flat key is simply never read, and

Func:

backend_config resolves the nested block (or its defaults) as if it were absent. See there.

Initialization

Initialize self. See help(type(self)) for accurate signature.

enabled: bool#

None

impl: Literal[transfer_queue]#

None

backend: Literal[nemo_rl.data_plane.interfaces.DataPlaneConfig.simple, nemo_rl.data_plane.interfaces.DataPlaneConfig.mooncake_cpu]#

None

claim_meta_poll_interval_s: float#

None

simple: NotRequired[nemo_rl.data_plane.interfaces.SimpleStorageConfig]#

None

mooncake_cpu: NotRequired[nemo_rl.data_plane.interfaces.MooncakeCpuConfig]#

None

controller_address: NotRequired[str]#

None

ack_timeout_ms: NotRequired[int]#

None

observability: NotRequired[ObservabilityConfig]#

None

nemo_rl.data_plane.interfaces._CHECKPOINTABLE_BACKENDS: frozenset[str]#

‘frozenset(…)’

nemo_rl.data_plane.interfaces.data_plane_supports_checkpointing(
cfg: nemo_rl.data_plane.interfaces.DataPlaneConfig,
) → bool#

Return whether the configured backend supports complete save/load.

Simple and Mooncake support native TQ checkpoints. The existing checkpointing settings decide whether a run saves data-plane state; normal PUTs remain memory-only. An unrecognized future backend defaults to unsupported until its storage payload and controller metadata are both known to round-trip through a checkpoint.

nemo_rl.data_plane.interfaces._BACKEND_MODELS: dict[str, type[pydantic.BaseModel]]#

None

nemo_rl.data_plane.interfaces.backend_config(
cfg: nemo_rl.data_plane.interfaces.DataPlaneConfig,
) → Any#

Return the validated sizing block for cfg["backend"].

Reads the nested block and lets the model supply anything it omits, so no caller ever writes a fallback. Works whether cfg came through pydantic (block already coerced to a model) or as a plain dict from a test.

Sizing is read only from the nested block. A config still using the pre-nesting flat spelling gets this backend’s defaults, not its own values.

class nemo_rl.data_plane.interfaces.ObservabilityConfig#

Bases: typing.TypedDict

Optional middleware that records per-op metrics on the client.

Off by default. When enabled=True the factory wraps the chosen adapter with :class:MetricsDataPlaneClient. callback is injected programmatically (callables don’t round-trip through YAML) — set cfg["observability"]["callback"] = my_fn before

Func:

build_data_plane_client to plug into wandb / file / log. There is no default callback: per-step metrics reach the logger via get_step_metrics, so a per-op sink is opt-in.

verify_tensor_hash is a correctness check, not a metric: each put records a per-row torch.hash_tensor fold of the row’s values, mixed with the row’s dtype and shape, and each get re-checks it, so a value that changes between wire-in and wire-out is reported (hash/mismatches) instead of silently training on it. It reads every tensor element a second time on both sides — roughly 8 ms for a 107 MB batch — so leave it off outside of debugging. It does not detect a permutation within a row; see data_plane/README.md.

Initialization

Initialize self. See help(type(self)) for accurate signature.

enabled: bool#

None

callback: NotRequired[Callable[[dict[str, Any]], None]]#

None

verify_tensor_hash: NotRequired[bool]#

None

class nemo_rl.data_plane.interfaces.LocalDataPlaneConfig#

Bases: pydantic.BaseModel

User configuration for process-local data transfer.

max_partitions limits retained step batches. Set it to 1 for only the active step or 2 to retain one prefetched step as well.

enabled: Literal[True]#

True

impl: Literal[local]#

‘local’

max_partitions: Annotated[int, Field(ge=1)]#

2

observability: nemo_rl.data_plane.interfaces.ObservabilityConfig | None#

None

nemo_rl.data_plane.interfaces.DataPlaneRuntimeConfig#

None

class nemo_rl.data_plane.interfaces.KVBatchMeta#

Per-batch metadata for data-plane KV operations.

Carries the per-sample IDs (sample_ids) that address rows in the KV store plus per-row metadata (fields, sequence_lengths, tags) needed for downstream routing without fetching tensor data. Vocabulary is intentionally NeMo-RL-native rather than 1:1 with any specific backend — the adapter translates at the boundary.

Two roles:

  • Result type returned by :meth:DataPlaneClient.claim_meta — callers extract .sample_ids / .partition_id and pass them to

    meth:

    get_samples / :meth:get_data.

  • Argument type for the per-DP-rank fetch entrypoints. sequence_lengths lets the driver compute a balanced per-rank shard from metadata only (control plane), without ever materializing tensor data.

partition_id: str#

None

task_name: str | None#

None

sample_ids: list[str]#

None

fields: list[str] | None#

None

sequence_lengths: list[int] | None#

None

extra_info: dict[str, Any]#

‘field(…)’

tags: list[dict[str, Any]] | None#

None

__post_init__() → None#
property size: int#
stamp_tags(scalars: dict[str, Sequence[Any]]) → None#

Mirror per-row scalar columns onto :attr:tags.

Each entry in scalars is a length-size sequence (list, tensor, ndarray) whose elements are written to tags[i][name]. Initializes tags to a list of empty dicts if currently None.

_replace(
*,
sample_ids: list[str],
sequence_lengths: list[int] | None,
tags: list[dict[str, Any]] | None = None,
) → nemo_rl.data_plane.interfaces.KVBatchMeta#

Return a copy with new sample_ids/sequence_lengths/tags, same metadata otherwise.

subset(
indices: Sequence[int],
) → nemo_rl.data_plane.interfaces.KVBatchMeta#

Return a new meta with only the rows at indices (any order).

slice(
start: int,
stop: int,
) → nemo_rl.data_plane.interfaces.KVBatchMeta#

Return a new meta with rows in the contiguous range [start, stop).

concat(
*others: nemo_rl.data_plane.interfaces.KVBatchMeta,
) → nemo_rl.data_plane.interfaces.KVBatchMeta#

Append metadata from the same partition.

Sample IDs are concatenated in argument order, while fields are unioned in first-seen order. Sequence lengths and tags are retained only when every input provides them.

Parameters:

*others – Metadata batches whose partition_id matches this batch.

Returns:

A new metadata batch containing all input rows.

Raises:

ValueError – If any input has a different partition_id.

drop(indices: Sequence[int]) → KVBatchMeta | None#

Complement of :meth:subset. Returns None when all rows are dropped.

with_fields(
field_names: Sequence[str],
) → nemo_rl.data_plane.interfaces.KVBatchMeta#

Return a copy with field_names merged into fields (deduped, order-preserving).

class nemo_rl.data_plane.interfaces.DataPlaneClient#

Bases: abc.ABC

Stable, swappable data-plane boundary.

The methods are split into three groups by intent. Argument order mirrors the underlying transfer_queue API 1:1 so a future adapter (e.g. nv-dataplane) is a thin pass-through too.

A. Task-mediated — used by stages that wait for upstream production via the per-task consumer counter:

Meth:

register_partition, :meth:claim_meta, :meth:get_data,

meth:

check_consumption_status. B. Direct-by-key — used by stages that already know the exact uids (e.g. driver-side fan-out to DP ranks):

meth:

put_samples, :meth:get_samples, :meth:clear_samples. C. Lifecycle — :meth:save_checkpoint, :meth:load_checkpoint, and

meth:

close.

Stage-completion signal: there is intentionally no mark_consumed. The authoritative signal in TransferQueue is field production — when a stage calls :meth:put_samples for a new field, the controller flips production_status[sample, field] = 1. Downstream consumers waiting on that field only see those samples once produced.

abstractmethod register_partition(
partition_id: str,
fields: list[str],
num_samples: int,
consumer_tasks: list[str],
grpo_group_size: int | None = None,
enums: dict[str, list[str]] | None = None,
) → None#

Declare the partition schema and consumer tasks.

Parameters:
  • partition_id – Partition name.

  • fields – Superset of fields any producer may write here.

  • num_samples – Expected total samples; sizes controller arrays.

  • consumer_tasks – Named tasks; each gets its own consumption cursor.

  • grpo_group_size – Group size for GRPO balanced sampling.

  • enums – Per-field fixed-vocab string codec, shipped once at register.

abstractmethod claim_meta(
partition_id: str,
task_name: str,
required_fields: list[str],
batch_size: int,
dp_rank: int | None = None,
blocking: bool = True,
timeout_s: float = 60.0,
) → nemo_rl.data_plane.interfaces.KVBatchMeta#

Discover and claim up to batch_size ready samples.

Advances task_name’s per-sample consumption cursor (TQ’s mode='fetch'); claimed uids won’t be returned again. Samples stay readable via :meth:get_samples until :meth:clear_samples.

Parameters:
  • partition_id – Partition to claim from.

  • task_name – Consumer task whose cursor is advanced.

  • required_fields – Fields that must be produced for a sample to be claimable.

  • batch_size – Max samples to claim.

  • dp_rank – Reserved; driver-side balancing via :func:shard_meta_for_dp is used today.

  • blocking – Block until the batch can be claimed.

  • timeout_s – Max blocking time before raising.

Returns:

KVBatchMeta for the claimed batch; pass to :meth:get_data.

abstractmethod get_data(
meta: nemo_rl.data_plane.interfaces.KVBatchMeta,
select_fields: list[str] | None = None,
) → tensordict.TensorDict#

Resolve a meta to tensor data.

Field-set resolution: (1) explicit select_fields; (2) meta.fields if non-None; (3) fail loudly — never silently fetch all fields.

Parameters:
  • meta – From :meth:claim_meta or hand-built with explicit keys.

  • select_fields – Subset of fields to fetch.

Returns:

TensorDict keyed by field name, batched along meta.sample_ids.

abstractmethod check_consumption_status(
partition_id: str,
task_names: list[str],
) → bool#

True iff every task has consumed all samples in the partition.

Authoritative across workers — uses TQ’s controller-side counter, not the per-process client cache.

Parameters:
  • partition_id – Partition to check.

  • task_names – Tasks whose consumption cursors are inspected.

Returns:

True iff every task in task_names has consumed all samples.

abstractmethod put_samples(
sample_ids: list[str],
partition_id: str,
fields: tensordict.TensorDict | None = None,
tags: list[dict[str, Any]] | None = None,
) → nemo_rl.data_plane.interfaces.KVBatchMeta#

Write fields for sample_ids — the producer entrypoint.

Writing a field flips the controller’s production_status bit for (sample, field); that flip is the “stage finished” signal downstream consumers wait on. Tensor and NonTensorStack leaves both pass through to TQ; non-tensor encoding is per-backend.

Parameters:
  • sample_ids – Per-sample uids being written.

  • partition_id – Partition these samples belong to.

  • fields – Tensor / NonTensorStack leaves to write.

  • tags – Optional per-sample primitive metadata.

Returns:

KVBatchMeta covering sample_ids — usable for direct :meth:get_samples.

abstractmethod get_samples(
sample_ids: list[str],
partition_id: str,
select_fields: list[str],
) → tensordict.TensorDict#

Direct fetch by uids.

Used by per-DP-rank slice fetches. Does NOT advance any per-task consumption cursor — that only happens via :meth:claim_meta.

select_fields is required (no implicit “fetch every field” fallback): bulk schemas are wide and silent over-fetch is the most expensive shape the wire can take. Callers must name what they read.

Parameters:
  • sample_ids – Uids to fetch.

  • partition_id – Partition the samples live in.

  • select_fields – Subset of fields to fetch.

Returns:

TensorDict keyed by field name, batched along sample_ids.

abstractmethod list_sample_ids(partition_id: str) → list[str]#

List the sample IDs currently stored in a partition.

This metadata-only operation is intended for recovery validation and reconciliation. It must not fetch tensor payloads or advance consumer cursors.

Parameters:

partition_id – Partition whose stored keys should be listed.

Returns:

Stable, sorted sample IDs. An unknown or empty partition returns an empty list.

abstractmethod clear_samples(sample_ids: list[str] | None, partition_id: str) → None#

Drop key-value pairs.

Explicit form (sample_ids=[...]) drops exactly those uids and is the form callers should use whenever they have the meta in hand — both sync GRPO callers (driver passes meta.sample_ids) and future async-RL data-loader actors that don’t share a process-local registry with the producer.

Convenience form (sample_ids=None) drops “everything this process knows produced in this partition”. Adapters implement this via a local registry populated by :meth:put_samples, with a fallback query to the underlying store. Useful for step-end teardown when the caller is the producer (driver in sync GRPO). Workers / loader actors that didn’t produce the samples should pass explicit IDs — the None form may silently no-op for them, and adapters are expected to warn when that happens.

Parameters:
  • sample_ids – Uids to drop; None clears every uid this process produced in the partition.

  • partition_id – Partition the samples live in.

abstractmethod save_checkpoint(
checkpoint_dir: str | pathlib.Path,
*,
metadata: dict[str, Any] | None = None,
) → None#

Persist the complete data-plane state to checkpoint_dir.

The checkpoint must include both data and the implementation’s scheduling/consumption metadata. Callers must serialize checkpoint saves and prevent destructive operations such as clears until this method returns.

Parameters:
  • checkpoint_dir – New durable directory for this checkpoint.

  • metadata – Optional JSON-compatible recovery metadata.

abstractmethod load_checkpoint(
checkpoint_dir: str | pathlib.Path,
) → dict[str, Any]#

Restore a complete data-plane checkpoint.

The data-plane implementation must already be initialized, but no data operations may have run before restore. Implementations must reject a load after operations through the same client; callers must also ensure that no other client has modified shared data-plane state.

Parameters:

checkpoint_dir –

Directory previously written by

meth:

save_checkpoint.

Returns:

User metadata supplied to :meth:save_checkpoint. The caller may validate this metadata, but restoring data-plane state does not restore the surrounding controller or trainer state.

abstractmethod close() → None#

Release controller / storage handles. Idempotent.