nemo_automodel.components.speculative.streaming.refs
nemo_automodel.components.speculative.streaming.refs
Tensor-free control-plane contracts for the speculative-training stream.
Per the train-inference disaggregation RFC (issue #3062), every produced sample
is split into a tensor-free reference (this module) and the actual supervision
tensors, which live in a pluggable store. References hop between the
target-side producer and the draft-side consumer over queues, HTTP, and
checkpoint metadata — places where serializing a torch.Tensor would be
either expensive or wrong.
Three guarantees back this module:
assert_no_tensorsenforces “no tensors” recursively on every public object exposed from this package, so a producer that slipped a tensor into a queue or ref trips a validation error before it lands on a wire.SampleRefis a frozen dataclass. Once placed on a queue, a ref cannot be mutated under the holder, so the consumer is guaranteed to see the same feature keys the producer promised.FeatureSpeccarriesdtype+shape(and nothing else), the minimum metadata a consumer needs to preallocate the receive buffer before materializing the tensor — mirroring hownemo_automodel.components.speculative.eagle.remote.protocol.encode_nccl_metadataalready ships dtype + shape ahead of an NCCL recv so the client can allocate.
Module Contents
Classes
Functions
Data
API
Bases: enum.Enum
Which speculative-decoding draft family produced this sample.
The on-wire string value (eagle3 / dflash / dspark) is the
stable schema key the consumer matches against. New algorithms land by
adding a member here plus a schema row in the RFC’s “Feature schema per
algorithm” table; both moves are part of the same change.
Shape + dtype metadata for one named feature in a SampleRef.
Mirrors the {dtype_code, shape} pair the EAGLE-3 remote protocol
encodes in nemo_automodel.components.speculative.eagle.remote.protocol.encode_nccl_metadata
so the client can preallocate the NCCL receive buffer. Consumers here use
it to preallocate the receive tensor before calling
FeatureStore.get.
The torch.dtype the producer will store under
SampleRef.feature_keys[name]. Validation in
assert_no_tensors rejects stray tensors in the ref but
a torch.dtype object on this field is allowed.
Logical shape with the same axis order the producer used.
shape[0] is conventionally batch and shape[-1] the
feature dimension; algorithms with packed layouts (THD, BHSD)
document their conventions in the algorithm row of the
RFC’s “Feature schema per algorithm” table.
Tensor-free reference to one produced sample in a feature store.
Which draft family produced this sample; the consumer
matches it against a registered FeatureSpec registry
to validate the ref before materializing.
Same idea for the draft model’s weights, so a
consumer can refuse to train against a ref produced before its
own weight snapshot. Starts at "0"; PR 4 wires the resync.
Sum of every feature’s numel() * dtype_bytes.
The SampleRefQueue reads this against
StoreHealth.resident_bytes for watermark hysteresis.
Named features this sample contributes, mapped to
per-store keys (e.g. filenames for SharedDirFeatureStore in
PR 3, dict keys for LocalFeatureStore).
Per-feature FeatureSpec so the consumer can
allocate the receive buffer before calling FeatureStore.get.
Sum of attended tokens; used by the consumer for empty / short loss-mask neutralization and by the queue for backpressure.
Stable identifier across the whole training run. Same value on every ref of one run so producers and consumers can verify they are talking about the same run before materializing.
Stable identifier within run_id. The consumer uses it
to partition the stream across DP ranks (RFC §“Consumer-side DP
resharding”) and to ACK / FAIL the corresponding lease.
Bumped whenever the producer’s feature set or tensor
layout for algorithm changes incompatibly. The consumer uses
it as a hard gate on the ref.
scheme://location identifying the FeatureStore
back-end (e.g. mem://local for
LocalFeatureStore).
The store URI plus feature_keys form the lookup the consumer
makes; the store object itself is discovered through the URI at
materialization time, not stored on the ref.
Monotonically increasing identifier of the target-model weights that produced this sample. Used to reject refs from a stale target (relevant once train-with-decode / weight resync lands in PR 4).
Return the feature names in insertion order.
The order is significant because some algorithms size the draft’s
fc projection by the order of the aux hidden states, not just
their count; the producer chooses the order once, the ref freezes it,
and the consumer preserves it through FeatureStore.get.
Required feature-name set per algorithm (RFC “Feature schema per algorithm”).
Mirrors the “Required feature keys” column of the RFC’s schema table. PR 2 widens this into a typed registry that also validates per-key dtypes / shapes; PR 1 only checks that the producer named the right set.
Recursively validate that obj carries no tensors.
Walks dataclasses (any nested dataclass included), dict with str
keys, and list / tuple. Anything else must be a primitive
(str/int/float/bool/None/bytes) — a tensor, numpy
array, or duck-typed tensor-like (has data_ptr + is_cuda + numel) at
any depth raises ValueError.
The check is structural, not nominal: a third-party tensor type that quacks like one is rejected. PR 2’s registry / schema validators layer on top of this primitive guard, not instead of it.