nemo_rl.data_plane.observability#

Lean per-op metrics decorator for DataPlaneClient.

Wraps any DataPlaneClient and invokes a single user-provided callback on each operation. Each event is a flat dict::

{"op", "partition_id", "n_keys", "n_bytes", "wall_ms", "status"}

Plug wandb / file logging / debug print at the call site by passing on_event=<your function>. snapshot() returns cumulative totals plus live memory consumption: bytes_outstanding (sum of bytes currently held in TQ, i.e. put minus cleared) and peak_bytes_outstanding (high-water mark over the run lifetime).

Every method here runs on the hot path of a transfer, so nothing traverses a structure twice and nothing is allocated for a payload no callback reads.

verify_tensor_hash=True adds an opt-in correctness check: a per-row torch.hash_tensor fold over each row’s values, mixed with its dtype and shape, recorded at put and re-checked at get, so a tensor that changes between wire-in and wire-out is reported rather than trained on. It reads every tensor byte again on both sides, so it is a debugging tool, not a metric. See README.md for what it does and does not catch.

Module Contents#

Classes#

DataPlaneEvent

OpStats

Per-op-tag accumulation. calls/wall_ms count every status.

HashStats

Wire-in / wire-out fingerprint reconciliation. All zero unless enabled.

DataPlaneStats

MetricsDataPlaneClient

Wrap a DataPlaneClient with a per-op callback hook.

Functions#

_comm_volume

Traffic totals derived from by_op, so bytes have one source.

_hash_field

_with_mirrors

fields followed by one mirror column each.

_leaf_digests

One digest per row: hash_tensor’s fold, with the row shape mixed in.

_field_digests

Leaf digests folded to one digest per top-level field.

_as_list

Materialize sample_ids once; None passes through.

_tensor_bytes

Wire bytes of one tensor leaf, rectangular or nested.

_percentile_from_hist

Interpolated q-quantile (0-1) from bucket counts.

_estimate_encoded_bytes

Approximate msgpack-encoded size of a non-tensor object.

_nontensor_stack_bytes

Extrapolate a NonTensorStack’s payload from a strided row sample.

_td_bytes

Payload bytes of a TensorDict, as the wire will see them.

_step_deltas

The five series both step-metric paths report, identically.

_op_step_stats

This step’s per-op detail, keyed by op, from two snapshots.

_hash_deltas

This step’s hash-verification counters, or nothing if the guard is off.

_volume_mb

Bytes each op moved this step, in MB, per op that moved any.

_op_series

Every per-op series for one step, from two snapshots.

_percent_of_dataplane

Where this step’s data-plane time went, in percent.

_clamped_percentiles

Whichever of :data:_QUANTILES this sample can actually support.

_derive_op_metrics

Fill in the derived per-op fields, in place.

merge_snapshots

Combine per-process snapshots into one cluster-wide view.

cluster_step_metrics

Per-step cluster metrics from two merged snapshots.

_step_metrics

One step’s metrics from two snapshots, cluster-wide or single-process.

headline_series

The subset of metrics worth a time series.

breakdown_table

Reshape the flat per-op series into one row per op.

metrics_never_fail_the_step

Swallow anything the metrics panel raises, and say so.

log_step_metrics

Emit one scope’s metrics: charted series, breakdown table, console line.

log_event

_annotate

Record an op’s outcome on its span.

is_metrics_client

Whether client carries counters worth reading.

Data#

API#

nemo_rl.data_plane.observability.EventStatus#

None

class nemo_rl.data_plane.observability.DataPlaneEvent#

Bases: typing.TypedDict

op: str#

None

partition_id: str#

None

n_keys: int#

None

n_bytes: int#

None

wall_ms: float#

None

status: nemo_rl.data_plane.observability.EventStatus#

None

nemo_rl.data_plane.observability.logger#

‘getLogger(…)’

nemo_rl.data_plane.observability._OP_ATTR#

‘rl.data_plane.op’

nemo_rl.data_plane.observability._PARTITION_ATTR#

‘rl.data_plane.partition’

nemo_rl.data_plane.observability._KEYS_ATTR#

‘rl.data_plane.keys’

nemo_rl.data_plane.observability._BYTES_ATTR#

‘rl.data_plane.bytes’

nemo_rl.data_plane.observability._STATUS_ATTR#

‘rl.data_plane.status’

nemo_rl.data_plane.observability.LATENCY_BUCKETS_MS: tuple[float, ...]#

(0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 25.0, 50.0, 100.0, 250.0, 500.0, 1000.0, 2500.0, 5000.0)

nemo_rl.data_plane.observability._WRITE_OPS#

‘frozenset(…)’

nemo_rl.data_plane.observability._READ_OPS#

‘frozenset(…)’

nemo_rl.data_plane.observability._comm_volume(by_op: dict[str, Any]) → dict[str, int]#

Traffic totals derived from by_op, so bytes have one source.

Distinct from bytes_outstanding, which is occupancy (what is held) rather than traffic (what moved).

Parameters:

by_op – Per-op stats carrying n_bytes.

Returns:

bytes_written, bytes_read, and their sum.

nemo_rl.data_plane.observability._MAX_HASH_MISMATCH_LOGS#

20

nemo_rl.data_plane.observability._HASH_SUFFIX#

‘_hash’

nemo_rl.data_plane.observability._hash_field(name: str) → str#
nemo_rl.data_plane.observability._with_mirrors(fields: collections.abc.Sequence[str]) → list[str]#

fields followed by one mirror column each.

Idempotent: meta.fields comes back from put_samples already carrying the mirrors, and suffixing those would ask for tokens_hash_hash.

nemo_rl.data_plane.observability._RECONCILE_ROWS#

None

nemo_rl.data_plane.observability._HASH_FIELDS#

(‘rows_recorded’, ‘rows_checked’, ‘rows_unverified’, ‘mismatches’, ‘fields_skipped’, ‘guard_failures…

nemo_rl.data_plane.observability._QUANTILES#

((0.5, ‘p50_ms’, 20), (0.9, ‘p90_ms’, 40))

nemo_rl.data_plane.observability._INT_VIEW_BY_WIDTH#

None

nemo_rl.data_plane.observability._leaf_digests(
leaf: torch.Tensor,
bounds: collections.abc.Sequence[int],
row_shape: Callable[[int], tuple[int, ...]],
dtype: torch.dtype,
) → list[int]#

One digest per row: hash_tensor’s fold, with the row shape mixed in.

torch.hash_tensor only implements an XOR fold (mode=0), which is blind to a zero pad, a trailing-dim reshape and a dtype change, because none of those alter the multiset of element words. Mixing the row’s shape and dtype into the fold closes all three. It remains blind to a permutation within a row, which is the price of the fold being vectorized; README.md records that.

The fold being an XOR is also what lets one algorithm serve both layouts. XOR is associative and elementwise, so hash_tensor(row) equals hash_tensor(rect, dim=1)[i] – a rectangle reduces in one on-device call, a ragged leaf falls back to one call per row, and the two agree value-for-value. A field packed jagged and read back densified therefore reconciles without either side recording how the other reduced it.

bounds are leading-dim offsets, n_rows + 1 of them: row i spans leaf[bounds[i] : bounds[i + 1]]. row_shape maps a row’s length to the shape recorded for it, which is what keeps the jagged and dense views of one field agreeing: both must call a row (L, D).

Parameters:
  • leaf – The whole leaf – a dense tensor, or a jagged one’s values.

  • bounds – Leading-dim offsets delimiting each row.

  • row_shape – Row length -> the shape to record for that row.

  • dtype – Mixed in so a precision change diverges at equal byte width.

Returns:

One digest per row.

nemo_rl.data_plane.observability._field_digests(
leaf_digests: dict[str, list[int]],
n_rows: int,
) → dict[str, torch.Tensor]#

Leaf digests folded to one digest per top-level field.

select_fields names top-level fields, so the mirror has to be per field rather than per leaf: a multimodal images arrives as several leaves and must reduce to a single images_hash that the reader can recompute from the same leaves.

Sorted leaf order because dict order need not survive a round trip, and * 31 + rather than an XOR so two identical leaves do not cancel – the defect the row fold already has, which must not be repeated here.

Folded on tensors, not in Python: this runs on every put and every get, and a row loop per leaf costs n_leaves * n_rows interpreted iterations there. int64 arithmetic wraps two’s-complement, which is the same modular fold the scalar form spelled out.

nemo_rl.data_plane.observability._as_list(sample_ids: Any) → Any#

Materialize sample_ids once; None passes through.

_run consumes its lambda and the accounting needs the same sequence afterwards, so a generator would be exhausted by the time it is indexed.

nemo_rl.data_plane.observability._tensor_bytes(v: torch.Tensor) → int#

Wire bytes of one tensor leaf, rectangular or nested.

A nested tensor’s nbytes dispatches through __torch_function__; its packed values buffer answers the same question without dispatching, and every per-token field on this wire is nested.

Two guards, because a wrong byte count is worse than a slow one:

  • _values is a bound method on every dense tensor, so the first check is on the type, not for absence — buf is None would be wrong.

  • A buffer holding more elements than the offsets describe means the tensor views a larger allocation (torch.nested.narrow), where the buffer overcounts. Nothing here builds one, but this is handed whatever a caller passes.

nemo_rl.data_plane.observability._percentile_from_hist(hist: list[int], q: float) → float#

Interpolated q-quantile (0-1) from bucket counts.

Linear interpolation inside the containing bucket. A value landing in the overflow bucket returns the top edge as a lower bound – we know it exceeded 5 s but not by how much.

nemo_rl.data_plane.observability._estimate_encoded_bytes(obj: Any, budget: list[int]) → int#

Approximate msgpack-encoded size of a non-tensor object.

TQ encodes non-tensors with msgpack (serial_utils.batch_encode_into), falling back to pickle/cloudpickle via Ext for unknown types. Getting the exact size means running that encoder, which would double the serialisation work on the hot path – so this walks the structure and approximates instead. Container framing (1-5 bytes per element) is not modelled, so treat the result as a lower bound.

budget bounds the walk to max_nodes container elements. Only the container branches charge it: a leaf cannot itself expand the walk. Containers stop iterating once it is exhausted – summing a generator would otherwise keep walking every element while each recursive call returned 0, making the cost O(size) despite the budget.

nemo_rl.data_plane.observability._NONTENSOR_STACK_SAMPLES#

4

nemo_rl.data_plane.observability._nontensor_stack_bytes(
stack: tensordict.NonTensorStack,
budget: list[int],
) → int#

Extrapolate a NonTensorStack’s payload from a strided row sample.

nemo_rl.data_plane.observability._td_bytes(
td: tensordict.TensorDict | None,
max_nodes: int = 10000,
) → int#

Payload bytes of a TensorDict, as the wire will see them.

Tensor leaves count nbytes (see :func:_tensor_bytes), which is the size mooncake registers and sends. Non-tensor leaves are estimated with :func:_estimate_encoded_bytes, since TQ ships them over a separate msgpack path. Both kinds are counted in a single items() pass; keys() + get() would re-resolve every nested key from the root.

leaves_only=True would hide the non-tensor entries entirely (NonTensorData is not treated as a leaf), so this walks with leaves_only=False and skips container nodes itself.

NonTensorData and NonTensorStack are matched by type rather than hasattr, and the distinction matters: NonTensorData exposes BOTH .data and .tolist(), and its .tolist() broadcasts the single stored object across the batch dim (a 64-row batch reported 20x the real payload).

Aliased storage is counted per field: two keys viewing one buffer count twice, which is right for volume (both are serialised) and is what lets max_bytes_per_key_seen catch view-aliasing regressions.

nemo_rl.data_plane.observability._step_deltas(
snap: dict[str, Any],
prev: dict[str, Any],
) → dict[str, float]#

The five series both step-metric paths report, identically.

Shared so the single-process and cluster views cannot drift on series names – which is the whole point of the step//now/ convention they publish under.

Write and read volume are deliberately not here. They were computed and then dropped by :func:headline_series, charted by nobody, while the breakdown table already carries per-op mb – put’s is the write volume and get’s is the read volume, split finer than a global pair would be.

nemo_rl.data_plane.observability._op_step_stats(
by_op: dict[str, Any],
prev_ops: dict[str, Any],
) → dict[str, dict[str, float]]#

This step’s per-op detail, keyed by op, from two snapshots.

Shared by the single-process and cluster paths so the two cannot drift, and used for both the emitted percentages and the breakdown table – one computation, so a chart and the table beside it can never disagree.

max_ms comes from step_max_ms, which the reader resets, rather than from the cumulative max_ms: a maximum is not differenceable, so the cumulative one latches at the worst call ever seen and never comes back down. Ops with no calls this step are absent, not zero.

nemo_rl.data_plane.observability._hash_deltas(
hv: dict[str, int],
prev_hv: dict[str, int],
) → dict[str, float]#

This step’s hash-verification counters, or nothing if the guard is off.

Shared by both step-metric paths. It was emitted only on the driver path, but _log_data_plane_metrics prefers the cluster path whenever the fan-out reaches more than one process – which is every real run – so with verify_tensor_hash on, mismatches never reached the logger. A guard whose findings are not reported is not a guard.

fields_skipped is here for the same reason it exists at all: a guard that quietly stops covering a field still reports zero mismatches, so the abstention count has to be visible beside the finding count.

Parameters:
  • hv – This step’s cumulative hash_verify block.

  • prev_hv – The previous step’s, for differencing.

Returns:

step/hash/{counter} deltas, or {} when the guard never ran.

nemo_rl.data_plane.observability._volume_mb(per_op: dict[str, dict[str, float]]) → dict[str, float]#

Bytes each op moved this step, in MB, per op that moved any.

comm_volume_mb is the total and hides the asymmetry that matters: on a real step get moved 20.8 MB against put’s 2.7 MB, because every DP rank fetches its shard once for the logprob pass and again for the train pass. Those are separate transfers over the wire, not an accounting artifact, and the same is true of summing across processes – each rank pulls its own shard.

Ops that carry no payload (register, clear) are omitted rather than reported as zero, matching how the percentages treat an op that did not run.

nemo_rl.data_plane.observability._BY_OP#

‘step/by_op/’

nemo_rl.data_plane.observability._BY_OP_NAMESPACES#

None

nemo_rl.data_plane.observability._op_series(
by_op: dict[str, Any],
prev_ops: dict[str, Any],
) → dict[str, float]#

Every per-op series for one step, from two snapshots.

The two step-metric paths share this rather than each assembling the same keys: the helpers below exist so the single-process and cluster views cannot drift on series names, and duplicating the six lines that build those names one level up would have given the drift back.

nemo_rl.data_plane.observability._percent_of_dataplane(
per_op: dict[str, dict[str, float]],
) → dict[str, float]#

Where this step’s data-plane time went, in percent.

The name carries the denominator because that is the one thing a reader has to know before acting on the number: it is a percentage of the data plane, not of the step. by_op/put = 43 reads “43% of the time spent inside the data plane went to put”. Whether that time mattered at all against compute is a different question, answered by step/frac_of_step, which divides by the step’s own wall clock. A workload can be 43% put and still not be worth touching.

by_op answers which call is expensive, and sums to 100 by construction.

On the cluster path wall_ms is summed over processes that ran concurrently, so these are percentages of aggregate process-time rather than of elapsed time. That is the right denominator for “what should I optimise” and the wrong one for “what blocked the step”.

Parameters:

per_op – Per-op step detail from :func:_op_step_stats.

Returns:

step/percent_of_dataplane/by_op/{op} in percent. Empty when no op ran.

nemo_rl.data_plane.observability._clamped_percentiles(
hist: list[int],
max_ms: float,
) → dict[str, float]#

Whichever of :data:_QUANTILES this sample can actually support.

Two corrections, both needed wherever a percentile is taken off a coarse histogram. Each quantile is withheld until there are enough samples to resolve it: below that the interpolation returns bucket geometry rather than data – one sample in (100, 250] yields a p50 of 175 whatever the call took. And the interpolation spreads a bucket’s samples uniformly across it, so calls clustered low in a wide bucket read high, above the exact maximum measured beside them; the maximum is the tighter bound.

Returns a dict rather than a fixed pair so a caller emits only what the data supports. An absent series says “not enough calls”; a zero would read as a measurement.

nemo_rl.data_plane.observability._derive_op_metrics(
by_op: dict[str, Any],
total_wall_ms: float,
) → None#

Fill in the derived per-op fields, in place.

Shared by :meth:MetricsDataPlaneClient.snapshot and

Func:

merge_snapshots so a cluster-wide view is derived by exactly the same arithmetic as a single process – percentiles off the (summed) histogram, rates off the (summed) totals. Nothing derived is ever averaged across processes.

nemo_rl.data_plane.observability._SNAPSHOT_SUM#

(‘total_bytes’, ‘total_keys’, ‘total_ops’, ‘total_wall_ms’, ‘bytes_outstanding’, ‘peak_bytes_outstan…

nemo_rl.data_plane.observability._SNAPSHOT_MAX#

(‘max_bytes_per_key_seen’, ‘last_put_bytes_per_key’, ‘step_wall_ms’)

nemo_rl.data_plane.observability._OP_SUM#

(‘calls’, ‘errors’, ‘wall_ms’, ‘n_bytes’, ‘n_keys’)

nemo_rl.data_plane.observability._OP_MAX#

(‘max_ms’, ‘step_max_ms’)

nemo_rl.data_plane.observability.merge_snapshots(
snapshots: list[dict[str, Any]],
) → dict[str, Any]#

Combine per-process snapshots into one cluster-wide view.

This is what the accumulators were shaped for. Latency lives in fixed histogram buckets precisely so they add: summing 256 per-rank histograms gives the true cluster distribution, which averaging 256 per-rank percentiles cannot. Everything derived — percentiles, throughput — is recomputed from the merged totals, never averaged.

Counters sum. max_* fields take a maximum. peak_bytes_outstanding is the one approximation: summing per-process peaks assumes they coincided, so it is an upper bound on true cluster peak occupancy.

Parameters:

snapshots – One :meth:MetricsDataPlaneClient.snapshot per process.

Returns:

A snapshot-shaped dict covering every process, plus n_processes.

nemo_rl.data_plane.observability.cluster_step_metrics(
merged: dict[str, Any],
prev: dict[str, Any],
step_time_s: float,
collect_ms: float = 0.0,
) → dict[str, float]#

Per-step cluster metrics from two merged snapshots.

The single-process equivalent of this lives on the client, which owns its own previous reading. A cluster has no such owner, so the caller holds prev and passes it back.

observability_overhead_ms is the whole bill for measuring: every process’s wrapper time plus collect_ms, the fan-out that gathered the snapshots. The fan-out is the larger half; omitting it understates by an order of magnitude.

Parameters:
  • merged – Cluster-wide snapshot from :func:merge_snapshots.

  • prev – The previous merged snapshot, for differencing.

  • step_time_s – Step wall time, for frac_of_step.

  • collect_ms – Wall time the caller spent gathering and merging.

nemo_rl.data_plane.observability._step_metrics(
snap: dict[str, Any],
prev: dict[str, Any],
step_time_s: float,
collect_ms: float = 0.0,
) → dict[str, float]#

One step’s metrics from two snapshots, cluster-wide or single-process.

Both callers difference the same counters; only the scope (1 process off a single client) and collect_ms (0 when there was no fan-out to pay for) differ, so the arithmetic lives here once.

Parameters:
  • snap – This step’s snapshot, merged or per-client.

  • prev – The previous one, for differencing.

  • step_time_s – Step wall time, for frac_of_step.

  • collect_ms – Wall time spent gathering and merging, if any.

Returns:

The flat step/ metric dict, less any caller-specific keys.

nemo_rl.data_plane.observability._HEADLINE#

(‘step/wall_s’, ‘step/frac_of_step’, ‘step/comm_volume_mb’, ‘now/bytes_outstanding_mb’, ‘now/n_proce…

nemo_rl.data_plane.observability._HEADLINE_PREFIXES#

(‘step/percent_of_dataplane/’, ‘step/volume_mb/’, ‘step/hash/’, ‘step/self/’, ‘step/codec/’)

nemo_rl.data_plane.observability.headline_series(metrics: dict[str, float]) → dict[str, float]#

The subset of metrics worth a time series.

Parameters:

metrics –

A flat dict from :func:cluster_step_metrics or

meth:

MetricsDataPlaneClient.get_step_metrics.

Returns:

Totals, time percentages, and hash counters – the per-op detail is dropped, since :func:breakdown_table presents it better.

nemo_rl.data_plane.observability._BREAKDOWN_COLUMNS#

(‘percent_of_dataplane’, ‘calls’, ‘wall_ms’, ‘mean_ms’, ‘max_ms’, ‘p50_ms’, ‘p90_ms’, ‘mb’)

nemo_rl.data_plane.observability.breakdown_table(
metrics: dict[str, float],
) → tuple[list[str], list[list[Any]]]#

Reshape the flat per-op series into one row per op.

A stack of line charts answers “how did put’s wall time trend”; the question this feeds is “where did this step’s time go, across ops, at a glance” – which is a table, and reading it off eight separate charts is the wrong tool. Rows are ordered by their share of data-plane time, so the bottleneck is the first line read.

Built from the metrics dict that is logged rather than from the snapshot it came from, so the table and the series can never disagree: a value withheld from the series (a percentile below the sample gate) is absent from the table too.

Parameters:

metrics –

A flat step/{op}/{field} dict from

meth:

MetricsDataPlaneClient.get_step_metrics or

func:

cluster_step_metrics.

Returns:

(columns, rows) for :meth:Logger.log_table.

nemo_rl.data_plane.observability._panel_failures#

‘count(…)’

nemo_rl.data_plane.observability.metrics_never_fail_the_step(
step: int,
) → collections.abc.Iterator[None]#

Swallow anything the metrics panel raises, and say so.

Observability is on by default, so a fault here would otherwise take down every step of every recipe – a panel must never fail training.

Swallowing is why this has to be loud. A panel that raises every step logs nothing else, and no data_plane/* series reaches the dashboard at all: the symptom is an empty panel, which looks exactly like a data plane that cost nothing. So the first failure carries its traceback at ERROR – a bare KeyError: 'step_wall_ms' names the key but not the line that asked for it – and later ones carry a running count, which is what distinguishes “broken since step 1” from “flaked once”.

The nightly gate is the backstop: with the panel down, the suites’ rows_checked > 0 check reads an absent series and fails.

Parameters:

step – Step number, for the log line.

nemo_rl.data_plane.observability.log_step_metrics(
logger: Any,
metrics: dict[str, float],
step: int,
scope: str,
) → None#

Emit one scope’s metrics: charted series, breakdown table, console line.

The series and the table are derived from one metrics dict, so they cannot disagree. A backend without a table type has no rows to log.

Parameters:
  • logger – Anything with log_metrics and log_table.

  • metrics –

    Output of :func:cluster_step_metrics or

    meth:

    MetricsDataPlaneClient.get_step_metrics.

  • step – Step number to log against.

  • scope – "cluster" or "driver" – names the prefix, because the two differ by roughly the DP degree.

nemo_rl.data_plane.observability.log_event(
event: nemo_rl.data_plane.observability.DataPlaneEvent,
) → None#
nemo_rl.data_plane.observability._annotate(
span: Any,
n_keys: int,
n_bytes: int,
status: nemo_rl.data_plane.observability.EventStatus,
) → None#

Record an op’s outcome on its span.

Set after the call rather than at open because the byte and key counts are only known once the inner client has returned. status distinguishes a timeout from a generic error, which the exception the span already records does not.

class nemo_rl.data_plane.observability.OpStats#

Per-op-tag accumulation. calls/wall_ms count every status.

n_bytes/n_keys count successful calls only, matching the cumulative totals — a failed transfer moved no payload, but the time it burned is still time the data plane cost the step.

calls: int#

0

errors: int#

0

wall_ms: float#

0.0

n_bytes: int#

0

n_keys: int#

0

max_ms: float#

0.0

step_max_ms: float#

0.0

latency_hist: list[int]#

‘field(…)’

class nemo_rl.data_plane.observability.HashStats#

Wire-in / wire-out fingerprint reconciliation. All zero unless enabled.

rows_unverified is as important as mismatches: a run that reads back rows this process never wrote (the normal case for a consumer-side client, which sees only wire-out) verifies nothing, and a mismatch count of 0 would otherwise read as “checked and clean”.

rows_recorded: int#

0

rows_checked: int#

0

rows_unverified: int#

0

mismatches: int#

0

fields_skipped: int#

0

guard_failures: int#

0

class nemo_rl.data_plane.observability.DataPlaneStats#
total_bytes: int#

0

total_keys: int#

0

total_ops: int#

0

total_wall_ms: float#

0.0

step_wall_ms: float#

0.0

by_op: dict[str, nemo_rl.data_plane.observability.OpStats]#

‘field(…)’

bytes_outstanding: int#

0

peak_bytes_outstanding: int#

0

max_bytes_per_key_seen: int#

0

last_put_bytes_per_key: int#

0

self_ms: float#

0.0

pack_ms: float#

0.0

unpack_ms: float#

0.0

hash_verify: nemo_rl.data_plane.observability.HashStats#

‘field(…)’

class nemo_rl.data_plane.observability.MetricsDataPlaneClient(
inner: nemo_rl.data_plane.interfaces.DataPlaneClient,
on_event: Callable[[nemo_rl.data_plane.observability.DataPlaneEvent], None] | None = None,
verify_tensor_hash: bool = False,
observability_enabled: bool = True,
)#

Bases: nemo_rl.data_plane.interfaces.DataPlaneClient

Wrap a DataPlaneClient with a per-op callback hook.

Initialization

Wrap inner, accumulating per-op timing and volume.

Parameters:
  • inner – The client whose calls are measured.

  • on_event – Per-op callback. None (the default) skips building the event dict entirely — with metrics enabled but no sink, nothing is paid for a payload nobody reads.

  • verify_tensor_hash – Record a per-row fingerprint on put and re-check it on get. Debug aid, not a metric: it reads every tensor element again on both sides (~8 ms for a 107 MB batch of 1536 rows), so it is off unless the config asks.

  • observability_enabled – Whether the user asked for data-plane observability. False on a telemetry-only run, where the wrapper is installed for its spans alone: the counters stop being collected and :func:is_metrics_client reports False, so the readers that poll every worker for data-plane stats stay off, as observability.enabled: false asked.

property observability_enabled: bool#

Whether the counters this wrapper accumulates are worth reading.

snapshot(reset_step_window: bool = False) → dict[str, Any]#

Return cumulative totals plus live byte / key outstanding counts.

total_wall_ms is the aggregate data-plane cost; by_op breaks it down per op tag with derived mean_ms and mb_per_s so the backends can be compared without post-processing. Throughput is omitted for ops that move no payload (e.g. claim_meta, whose wall time is producer wait, not transfer).

Parameters:

reset_step_window – Zero step_wall_ms and each op’s step_max_ms after reading them, opening a fresh window. A maximum cannot be differenced out of a cumulative counter the way calls and wall_ms can, so the only way to scope one to a step is to reset it – and the reader that consumes it is the one that has to. Left off by default so an inspection snapshot never disturbs the step series.

get_step_metrics(
step_time_s: float,
snap: dict[str, Any] | None = None,
collect_ms: float = 0.0,
) → dict[str, float]#

Per-step data-plane metrics, as a ready-to-log flat dict.

Cumulative counters are differenced against the previous call, so this reports what the data plane cost this step. Mirrors VllmGeneration.get_step_metrics so trainers stay one line.

frac_of_step is the metric that decides whether optimising the data plane is worth anything: percent_of_dataplane only says where data-plane time went, never whether it mattered against compute.

Parameters:
  • step_time_s – Step wall time, for frac_of_step.

  • snap – A snapshot already taken by the caller. A caller that fans out has to read this client first and cannot read it twice – closing the step window a second time would zero every step/by_op/*/max_ms. Passing it here keeps the baseline in one place, this client, rather than a second copy on the caller.

  • collect_ms – Wall time the caller spent gathering, if any.

_record_put(partition_id: str, keys: list[str], n_bytes: int) → None#

Attribute put bytes per key so a later clear_samples can subtract.

Called after the underlying RPC succeeds so a failed put never leaves the accounting inflated.

n_bytes is a whole-batch figure, so there was never a per-key truth to keep: the old per-key dict stored an even split, and a subset clear released the mean either way. Holding one total and one key set says the same thing and lets set.update do the per-key work in C — 18.6 us to 3.0 us at 256 keys, which was the single largest remaining cost on the put path.

Parameters:
  • partition_id – Partition the keys were written to.

  • keys – Per-sample uids that were written.

  • n_bytes – Total bytes written; released pro rata on clear.

_record_clear(partition_id: str, keys: list[str] | None) → None#

Reverse the put accounting for keys.

Called after the underlying RPC succeeds so a failed clear keeps the accounting consistent with TQ’s actual state.

Bytes are released pro rata: the partition’s total times the share of its live keys being dropped. Clearing the last key releases the remainder exactly, so a partition always reconciles to zero however it is chopped up.

Parameters:
  • partition_id – Partition the keys were dropped from.

  • keys – Uids dropped; None means the whole partition was cleared.

_release_cleared_samples(partition_id: str) → None#

Reverse the put accounting for samples another process cleared.

_record_clear only fires in the process that issues the clear, which on the SC path is only ever SC: GenWorker and the value actor put through their own clients and never clear, so their accounting would keep every uid they ever wrote. list_sample_ids is metadata-only and documented for reconciliation; diffing against it ties the accounting to the sample’s real lifetime.

The stale uids go through _record_clear so both stores are released by the one rule a real clear uses.

_bill_self(entered: float) → None#

Charge this wrapper for the time it spent that was not the RPC.

One monotonic per op on top of the two _run already takes. Measuring the measurement is worth that: the alternative is asking a reader to trust a benchmark run on some other machine.

_row_fingerprints(
td: tensordict.TensorDict | None,
sample_ids: list[str],
) → dict[str, list[int]]#

A per-row digest of each tensor leaf, covering bytes, dtype and shape.

Every leaf is fingerprinted one row at a time, so every divergence names the sample that diverged – a genuinely ragged leaf included. See README.md for why the digest this replaced could not.

Both layouts reduce to the same shape of work: hand

Func:

_leaf_digests the leaf, the offsets delimiting its rows, and how to describe a row’s shape. It picks between one vectorized fold and a fold per row.

Parameters:
  • td – Leaves to fingerprint; None yields an empty result.

  • sample_ids – Row i is attributed to sample_ids[i], the ordering :meth:DataPlaneClient.get_samples promises.

Returns:

Field name -> one digest per row. A leaf that cannot be attributed per row is counted in fields_skipped rather than silently dropped: a non-jagged nested layout, a leading dim that is not len(sample_ids), or a leaf with no leading dim at all.

_hash_guard_failed(op: str, exc: Exception) → None#

Absorb a hash-guard failure: count it, log it, never re-raise.

The guard is a debug aid on a transfer that already succeeded, so a bug in it must not take the transfer down. Swallowing is only safe because the failure stays visible in step/hash/guard_failures – a guard that silently stopped checking would report zero mismatches.

_stamp_hashes(
sample_ids: list[str],
fields: tensordict.TensorDict | None,
) → tensordict.TensorDict | None#

Return fields with a <field>_hash column beside each field.

Never raises: a guard that cannot fold must not stop the put, so the original fields goes on the wire unstamped and the batch reads as unverified on the far side.

_stamp_hashes_impl(
sample_ids: list[str],
fields: tensordict.TensorDict | None,
) → tensordict.TensorDict | None#
_check_hashes(
partition_id: str,
sample_ids: list[str],
out: Any,
) → None#

Compare wire-out fingerprints against what was written. Never raises.

_check_hashes_impl(
partition_id: str,
sample_ids: list[str],
out: Any,
) → None#

Compare wire-out fingerprints against what was written.

The wire-in reading arrives with the row, so a shard read by a process that never wrote it reconciles the same as a same-process round trip. The mirror columns are stripped here: the caller asked for tokens and must never see tokens_hash.

_run(
op: str,
partition_id: str,
fn: Callable[[], Any],
*,
n_keys: int = 0,
n_bytes: int = 0,
) → Any#

Run fn and emit one observability event with wall-time and status.

Also opens one span per op, which is what puts transfer-queue traffic in the trace waterfall: on the single-controller path most of a step’s non-compute time is data-plane traffic, and without these spans that time showed up only as a gap between phases.

Parameters:
  • op – Operation tag ("put", "get", "clear", etc.).

  • partition_id – Partition the op targets.

  • fn – Zero-arg callable that invokes the inner client.

  • n_keys – Key count if known up front; otherwise inferred from the return value (KVBatchMeta.sample_ids).

  • n_bytes – Byte estimate; overridden by _td_bytes when the return is a TensorDict.

Returns:

Whatever fn returned.

_emit(
op: str,
partition_id: str,
n_keys: int,
n_bytes: int,
t0: float,
status: nemo_rl.data_plane.observability.EventStatus,
) → None#
register_partition(
partition_id,
fields,
num_samples,
consumer_tasks,
grpo_group_size=None,
enums=None,
)#
claim_meta(
partition_id,
task_name,
required_fields,
batch_size,
dp_rank=None,
blocking=True,
timeout_s=60.0,
)#
get_data(meta, select_fields=None)#
check_consumption_status(partition_id, task_names)#
put_samples(sample_ids, partition_id, fields=None, tags=None)#
get_samples(sample_ids, partition_id, select_fields)#
list_sample_ids(partition_id: str) → list[str]#
clear_samples(sample_ids, partition_id)#
save_checkpoint(
checkpoint_dir: str | pathlib.Path,
*,
metadata: dict[str, Any] | None = None,
) → None#
load_checkpoint(
checkpoint_dir: str | pathlib.Path,
) → dict[str, Any]#
close() → None#
nemo_rl.data_plane.observability.is_metrics_client(
client: Any,
) → TypeGuard[nemo_rl.data_plane.observability.MetricsDataPlaneClient]#

Whether client carries counters worth reading.

The one answer to “is observability on here”, replacing four call sites that asked it three ways – two by probing for a snapshot attribute, which is not on the :class:DataPlaneClient ABC. isinstance(None, ...) is False, so this covers “no client at all” too.

The type alone is not the answer: a telemetry-only run installs the wrapper for its spans with observability off, and the readers must not then poll every worker for stats the user switched off.