nemo_rl.data_plane.adapters.local#

Process-local data-plane adapter for colocated SFT loaders and policies.

Module Contents#

Classes#

_LocalPartition

LocalDataPlaneClient

Store a bounded set of complete SFT batches in the current process.

Functions#

is_local_batch_meta

Return whether metadata identifies a process-local partition version.

_value_batch_size

_is_identity_indices

_select_batch_value

_unwrap_local_value

local_batch_to_tensordict

Wrap a prepared local batch without flattening multimodal values.

materialize_local

Materialize local tensors and restore exact non-tensor field values.

Data#

API#

nemo_rl.data_plane.adapters.local._LOCAL_GENERATION_KEY#

‘local_partition_generation’

nemo_rl.data_plane.adapters.local._CHECKPOINT_FILENAME#

‘local_data_plane_state.pkl’

class nemo_rl.data_plane.adapters.local._LocalPartition#
fields: tuple[str, ...]#

None

num_samples: int#

None

consumer_tasks: tuple[str, ...]#

None

grpo_group_size: int | None#

None

enums: dict[str, list[str]]#

None

generation: int#

None

sample_ids: list[str]#

‘field(…)’

batch: dict[str, Any]#

‘field(…)’

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

None

consumed: dict[str, set[str]]#

‘field(…)’

nemo_rl.data_plane.adapters.local.is_local_batch_meta(
meta: nemo_rl.data_plane.interfaces.KVBatchMeta,
) → bool#

Return whether metadata identifies a process-local partition version.

nemo_rl.data_plane.adapters.local._value_batch_size(value: Any) → int#
nemo_rl.data_plane.adapters.local._is_identity_indices(indices: list[int], batch_size: int) → bool#
nemo_rl.data_plane.adapters.local._select_batch_value(
value: Any,
indices: list[int],
) → Any#
nemo_rl.data_plane.adapters.local._unwrap_local_value(value: Any) → Any#
nemo_rl.data_plane.adapters.local.local_batch_to_tensordict(
fields: collections.abc.Mapping[str, Any],
*,
batch_size: int,
) → tensordict.TensorDict#

Wrap a prepared local batch without flattening multimodal values.

Tensor leaves stay as tensors. Other leaves, including PackedTensor, use NonTensorData because this adapter never serializes them.

nemo_rl.data_plane.adapters.local.materialize_local(
td: tensordict.TensorDict,
layout: nemo_rl.data_plane.schema.Layout = 'padded',
pad_value_dict: dict[str, int | float] | None = None,
pad_to_seqlen: int = 0,
) → nemo_rl.distributed.batched_data_dict.BatchedDataDict[Any]#

Materialize local tensors and restore exact non-tensor field values.

Not interchangeable with :func:materialize. local_batch_to_tensordict stores each non-tensor column as one NonTensorData holding the whole column, which materialize reads as a single row. Batches written by this adapter must be read back through here.

class nemo_rl.data_plane.adapters.local.LocalDataPlaneClient(
cfg: nemo_rl.data_plane.interfaces.LocalDataPlaneConfig,
)#

Bases: nemo_rl.data_plane.interfaces.DataPlaneClient

Store a bounded set of complete SFT batches in the current process.

Initialization

_require_open() → None#
_partition(
partition_id: str,
) → nemo_rl.data_plane.adapters.local._LocalPartition#
_validate_meta(
meta: nemo_rl.data_plane.interfaces.KVBatchMeta,
) → nemo_rl.data_plane.adapters.local._LocalPartition#
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#
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#
get_data(
meta: nemo_rl.data_plane.interfaces.KVBatchMeta,
select_fields: list[str] | None = None,
) → tensordict.TensorDict#
check_consumption_status(
partition_id: str,
task_names: list[str],
) → bool#
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#
get_samples(
sample_ids: list[str],
partition_id: str,
select_fields: list[str],
) → tensordict.TensorDict#
list_sample_ids(partition_id: str) → list[str]#

List stored sample IDs without reading their batch payloads.

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

Persist the in-process partitions to checkpoint_dir.

State never leaves this process, so pickle is sufficient; callers must not load checkpoints from untrusted paths. The write goes to a sibling .tmp directory and is renamed into place so a crash mid-save leaves the previous checkpoint intact.

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

Restore partitions into a client that has not registered any yet.

close() → None#
static _sequence_lengths(
partition: nemo_rl.data_plane.adapters.local._LocalPartition,
indices: list[int],
) → list[int] | None#