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#
Per-op-tag accumulation. |
|
Wire-in / wire-out fingerprint reconciliation. All zero unless enabled. |
|
Wrap a |
Functions#
Traffic totals derived from |
|
|
|
One digest per row: |
|
Leaf digests folded to one digest per top-level field. |
|
Materialize |
|
Wire bytes of one tensor leaf, rectangular or nested. |
|
Interpolated |
|
Approximate msgpack-encoded size of a non-tensor object. |
|
Extrapolate a |
|
Payload bytes of a TensorDict, as the wire will see them. |
|
The five series both step-metric paths report, identically. |
|
This step’s per-op detail, keyed by op, from two snapshots. |
|
This step’s hash-verification counters, or nothing if the guard is off. |
|
Bytes each op moved this step, in MB, per op that moved any. |
|
Every per-op series for one step, from two snapshots. |
|
Where this step’s data-plane time went, in percent. |
|
Whichever of :data: |
|
Fill in the derived per-op fields, in place. |
|
Combine per-process snapshots into one cluster-wide view. |
|
Per-step cluster metrics from two merged snapshots. |
|
One step’s metrics from two snapshots, cluster-wide or single-process. |
|
The subset of |
|
Reshape the flat per-op series into one row per op. |
|
Swallow anything the metrics panel raises, and say so. |
|
Emit one scope’s metrics: charted series, breakdown table, console line. |
|
Record an op’s outcome on its span. |
|
Whether |
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]#
fieldsfollowed by one mirror column each.Idempotent:
meta.fieldscomes back fromput_samplesalready carrying the mirrors, and suffixing those would ask fortokens_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,
One digest per row:
hash_tensor’s fold, with the row shape mixed in.torch.hash_tensoronly 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.mdrecords that.The fold being an XOR is also what lets one algorithm serve both layouts. XOR is associative and elementwise, so
hash_tensor(row)equalshash_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.boundsare leading-dim offsets,n_rows + 1of them: row i spansleaf[bounds[i] : bounds[i + 1]].row_shapemaps 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,
Leaf digests folded to one digest per top-level field.
select_fieldsnames top-level fields, so the mirror has to be per field rather than per leaf: a multimodalimagesarrives as several leaves and must reduce to a singleimages_hashthat 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_rowsinterpreted iterations there.int64arithmetic 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_idsonce;Nonepasses through._runconsumes 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
nbytesdispatches 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:
_valuesis a bound method on every dense tensor, so the first check is on the type, not for absence —buf is Nonewould 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 viaExtfor 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.budgetbounds the walk tomax_nodescontainer 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],
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,
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 singleitems()pass;keys()+get()would re-resolve every nested key from the root.leaves_only=Truewould hide the non-tensor entries entirely (NonTensorDatais not treated as a leaf), so this walks withleaves_only=Falseand skips container nodes itself.NonTensorDataandNonTensorStackare matched by type rather thanhasattr, and the distinction matters:NonTensorDataexposes BOTH.dataand.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_seencatch view-aliasing regressions.
- nemo_rl.data_plane.observability._step_deltas(
- snap: dict[str, Any],
- prev: dict[str, Any],
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-opmb– 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],
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_mscomes fromstep_max_ms, which the reader resets, rather than from the cumulativemax_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],
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_metricsprefers the cluster path whenever the fan-out reaches more than one process – which is every real run – so withverify_tensor_hashon,mismatchesnever reached the logger. A guard whose findings are not reported is not a guard.fields_skippedis 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_verifyblock.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_mbis the total and hides the asymmetry that matters: on a real stepgetmoved 20.8 MB againstput’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],
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]],
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 = 43reads “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 bystep/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_opanswers which call is expensive, and sums to 100 by construction.On the cluster path
wall_msis 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,
Whichever of :data:
_QUANTILESthis 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,
Fill in the derived per-op fields, in place.
Shared by :meth:
MetricsDataPlaneClient.snapshotand- Func:
merge_snapshotsso 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]],
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_outstandingis 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.snapshotper 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,
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
prevand passes it back.observability_overhead_msis the whole bill for measuring: every process’s wrapper time pluscollect_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,
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
metricsworth a time series.- Parameters:
metrics –
A flat dict from :func:
cluster_step_metricsor- meth:
MetricsDataPlaneClient.get_step_metrics.
- Returns:
Totals, time percentages, and hash counters – the per-op detail is dropped, since :func:
breakdown_tablepresents 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],
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_metricsor- 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,
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 bareKeyError: '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 > 0check 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,
Emit one scope’s metrics: charted series, breakdown table, console line.
The series and the table are derived from one
metricsdict, so they cannot disagree. A backend without a table type has no rows to log.- Parameters:
logger – Anything with
log_metricsandlog_table.metrics –
Output of :func:
cluster_step_metricsor- 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( ) None#
- nemo_rl.data_plane.observability._annotate(
- span: Any,
- n_keys: int,
- n_bytes: int,
- status: nemo_rl.data_plane.observability.EventStatus,
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.
statusdistinguishes 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_mscount every status.n_bytes/n_keyscount 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_unverifiedis as important asmismatches: 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.DataPlaneClientWrap a
DataPlaneClientwith 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_clientreports False, so the readers that poll every worker for data-plane stats stay off, asobservability.enabled: falseasked.
- 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_msis the aggregate data-plane cost;by_opbreaks it down per op tag with derivedmean_msandmb_per_sso 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_msand each op’sstep_max_msafter reading them, opening a fresh window. A maximum cannot be differenced out of a cumulative counter the waycallsandwall_mscan, 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,
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_metricsso trainers stay one line.frac_of_stepis the metric that decides whether optimising the data plane is worth anything:percent_of_dataplaneonly 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_samplescan subtract.Called after the underlying RPC succeeds so a failed put never leaves the accounting inflated.
n_bytesis 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 letsset.updatedo 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;
Nonemeans the whole partition was cleared.
- _release_cleared_samples(partition_id: str) None#
Reverse the put accounting for samples another process cleared.
_record_clearonly 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_idsis metadata-only and documented for reconciliation; diffing against it ties the accounting to the sample’s real lifetime.The stale uids go through
_record_clearso 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
monotonicper op on top of the two_runalready 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],
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.mdfor why the digest this replaced could not.Both layouts reduce to the same shape of work: hand
- Func:
_leaf_digeststhe 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;
Noneyields an empty result.sample_ids – Row i is attributed to
sample_ids[i], the ordering :meth:DataPlaneClient.get_samplespromises.
- Returns:
Field name -> one digest per row. A leaf that cannot be attributed per row is counted in
fields_skippedrather than silently dropped: a non-jaggednested layout, a leading dim that is notlen(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,
Return
fieldswith a<field>_hashcolumn beside each field.Never raises: a guard that cannot fold must not stop the put, so the original
fieldsgoes 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,
- _check_hashes(
- partition_id: str,
- sample_ids: list[str],
- out: Any,
Compare wire-out fingerprints against what was written. Never raises.
- _check_hashes_impl(
- partition_id: str,
- sample_ids: list[str],
- out: Any,
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
tokensand must never seetokens_hash.
- _run(
- op: str,
- partition_id: str,
- fn: Callable[[], Any],
- *,
- n_keys: int = 0,
- n_bytes: int = 0,
Run
fnand 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_byteswhen the return is aTensorDict.
- Returns:
Whatever
fnreturned.
- _emit(
- op: str,
- partition_id: str,
- n_keys: int,
- n_bytes: int,
- t0: float,
- status: nemo_rl.data_plane.observability.EventStatus,
- 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,
- load_checkpoint(
- checkpoint_dir: str | pathlib.Path,
- close() None#
- nemo_rl.data_plane.observability.is_metrics_client(
- client: Any,
Whether
clientcarries 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
snapshotattribute, which is not on the :class:DataPlaneClientABC.isinstance(None, ...)isFalse, 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.