nemo_automodel.components.speculative.streaming.refs

View as Markdown

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:

  1. assert_no_tensors enforces “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.
  2. SampleRef is 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.
  3. FeatureSpec carries dtype + shape (and nothing else), the minimum metadata a consumer needs to preallocate the receive buffer before materializing the tensor — mirroring how nemo_automodel.components.speculative.eagle.remote.protocol.encode_nccl_metadata already ships dtype + shape ahead of an NCCL recv so the client can allocate.

Module Contents

Classes

NameDescription
FeatureAlgorithmWhich speculative-decoding draft family produced this sample.
FeatureSpecShape + dtype metadata for one named feature in a SampleRef.
SampleRefTensor-free reference to one produced sample in a feature store.

Functions

NameDescription
_algorithm_required_featuresRequired feature-name set per algorithm (RFC “Feature schema per algorithm”).
_is_numpy_array-
_is_torch_tensor-
_reject_tensor-
assert_no_tensorsRecursively validate that obj carries no tensors.

Data

_PRIMITIVE_TYPES

logger

API

class nemo_automodel.components.speculative.streaming.refs.FeatureAlgorithm

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.

DFLASH
= 'dflash'
DSPARK
= 'dspark'
EAGLE3
= 'eagle3'
class nemo_automodel.components.speculative.streaming.refs.FeatureSpec(
shape: tuple[int, ...],
dtype: torch.dtype
)
Dataclass

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.

dtype
dtype

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.

shape
tuple[int, ...]

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.

nemo_automodel.components.speculative.streaming.refs.FeatureSpec.__post_init__() -> None
class nemo_automodel.components.speculative.streaming.refs.SampleRef(
sample_id: str,
run_id: str,
store_uri: str,
feature_keys: dict[str, str],
schema_version: int,
num_tokens: int,
estimated_bytes: int,
target_model_version: str,
draft_weight_version: str
)
Dataclass

Tensor-free reference to one produced sample in a feature store.

algorithm
FeatureAlgorithm

Which draft family produced this sample; the consumer matches it against a registered FeatureSpec registry to validate the ref before materializing.

draft_weight_version
str

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.

estimated_bytes
int

Sum of every feature’s numel() * dtype_bytes. The SampleRefQueue reads this against StoreHealth.resident_bytes for watermark hysteresis.

feature_keys
dict[str, str]

Named features this sample contributes, mapped to per-store keys (e.g. filenames for SharedDirFeatureStore in PR 3, dict keys for LocalFeatureStore).

feature_specs
dict[str, FeatureSpec]

Per-feature FeatureSpec so the consumer can allocate the receive buffer before calling FeatureStore.get.

num_tokens
int

Sum of attended tokens; used by the consumer for empty / short loss-mask neutralization and by the queue for backpressure.

run_id
str

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.

sample_id
str

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.

schema_version
int

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.

store_uri
str

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.

target_model_version
str

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).

nemo_automodel.components.speculative.streaming.refs.SampleRef.__post_init__() -> None
nemo_automodel.components.speculative.streaming.refs.SampleRef.feature_names() -> tuple[str, ...]

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.

nemo_automodel.components.speculative.streaming.refs._algorithm_required_features(
) -> frozenset[str]

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.

nemo_automodel.components.speculative.streaming.refs._is_numpy_array(
obj: typing.Any
) -> bool
nemo_automodel.components.speculative.streaming.refs._is_torch_tensor(
obj: typing.Any
) -> bool
nemo_automodel.components.speculative.streaming.refs._reject_tensor(
obj: typing.Any,
path: str
) -> typing.NoReturn
nemo_automodel.components.speculative.streaming.refs.assert_no_tensors(
obj: typing.Any,
path: str = 'ref'
) -> None

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.

nemo_automodel.components.speculative.streaming.refs._PRIMITIVE_TYPES = (str, int, float, bool, type(None), bytes)
nemo_automodel.components.speculative.streaming.refs.logger = logging.getLogger(__name__)