NemotronMoeLayer

Struct NemotronMoeLayer 

Source
pub struct NemotronMoeLayer { /* private fields */ }
Expand description

Nemotron-H standalone MoE FFN layer.

Implementations§

Source§

impl NemotronMoeLayer

Source

pub fn prepare_prefill_weights( &mut self, gpu: &dyn GpuBackend, config: &ModelConfig, )

Transpose expert weights for fast grouped GEMM prefill. Called from weight loader after construction. Skips expert transposition when memory is tight (Super 120B: 128 experts × 40 layers would OOM).

Source§

impl NemotronMoeLayer

Source

pub fn new( weights: NemotronMoeWeights, input_norm: DenseWeight, config: &ModelConfig, gpu: &dyn GpuBackend, moe_inter: usize, top_k: usize, ) -> Result<Self>

Trait Implementations§

Source§

impl TransformerLayer for NemotronMoeLayer

Source§

fn prefill( &self, hidden: DevicePtr, residual: DevicePtr, num_tokens: usize, _state: &mut dyn LayerState, _kv_cache: &mut PagedKvCache, _seq_len_start: usize, _block_table: &mut Vec<u32>, _disk_block_ids: &mut Vec<u32>, _disk_last_offloaded_per_layer: &mut Vec<u32>, _kv_write_start: usize, ctx: &ForwardContext<'_>, stream: u64, ) -> Result<()>

Batched MoE prefill: uses GEMM for gate/fc1/fc2/shared, per-token for routing + experts.

For Super 120B with 40 MoE layers, this replaces O(N * 7 kernel_launches) decode calls with O(4 GEMMs + N * 3 kernel_launches), cutting TTFT by 30-50%.

Source§

fn decode( &self, hidden: DevicePtr, residual: DevicePtr, _state: &mut dyn LayerState, _kv_cache: &mut PagedKvCache, _seq_len: usize, _block_table: &mut Vec<u32>, _disk_block_ids: &mut Vec<u32>, _disk_last_offloaded_per_layer: &mut Vec<u32>, ctx: &ForwardContext<'_>, stream: u64, ) -> Result<()>

Decode one token through this layer, modifying hidden in-place. Read more
Source§

fn alloc_state(&self, _gpu: &dyn GpuBackend) -> Result<Box<dyn LayerState>>

Allocate per-sequence state for this layer. Read more
Source§

fn uses_local_mla_prefill(&self) -> bool

True when this layer’s PREFILL attends only over the tokens it is handed, so a prefix-cache skip would hide the cached prefix from attention entirely. MLA layers on the paged path do; everything else reads the paged cache and is unaffected.
Source§

fn as_any_mut(&mut self) -> Option<&mut dyn Any>

&mut dyn Any downcast hook for post-construction weight overlays (e.g. the LoRA install walk). Default None; overlay-capable layers override.
Source§

fn fp8_calibration_frozen(&self) -> Option<bool>

Whether this layer’s ONLINE FP8-KV calibration has frozen its scale. None = this layer runs no online calibration (non-attention layer, static checkpoint scales, or a non-FP8 KV dtype). The scheduler’s graph-suppression gate keys off this rather than a token count: the scale freezes on the FIRST observe, so waiting calibration_tokens tokens would run ~256+ eager steps for a calibration that finished immediately.
Source§

fn decode_prestage( &self, _token: u32, _state: &mut dyn LayerState, _gpu: &dyn GpuBackend, _stream: u64, ) -> Result<()>

Hoisted per-step HOST work for layers that do host-side computation at decode (PLE: n-gram hash + NVMe fault-in + slot upload). The scheduler calls this every single-token decode step BEFORE any CUDA graph replay/capture — the same phasing as the token_ids upload — so the captured graph contains only kernels over stable device buffers. Layers with no host-side decode work keep the no-op default.
Source§

fn decode_prestage_rearm(&self, _state: &mut dyn LayerState)

Re-arm consumed prestage state so a failed CUDA-graph capture attempt can re-run the SAME step eagerly. Must be idempotent, and must not recompute (PLE’s history already advanced in decode_prestage).
Source§

fn decode_graph_unsupported(&self) -> bool

True when this layer’s decode can NEVER be captured into a CUDA graph — e.g. the QSA indexer’s host top-k round trip, whose captured dense fallback would silently replay WRONG attention once selection activates. The scheduler ORs this across layers once and keeps the whole model eager.
Source§

fn decode_multi_seq_unsupported(&self) -> bool

True when this layer cannot serve a BATCHED multi-sequence decode step — i.e. decode_multi_seq’s shared-ForwardContext loop would alias per-sequence state across rows rather than merely run slowly. Read more
Source§

fn decode_verify_multi_unsupported(&self) -> bool

True when this layer cannot serve a BATCHED multi-sequence VERIFY sweep (decode_verify_multi). Consumed by can_batch_verify_dispatch; a true layer falls back to the per-sequence verify loop, which is the sealed single-sequence path. Read more
Source§

fn snapshot_aux( &self, _state: &dyn LayerState, _gpu: &dyn GpuBackend, _stream: u64, ) -> Result<Option<Vec<u8>>>

Marconi aux state: host-serialized per-layer SEQUENCE state that must travel with an SSM snapshot for a prefix-cache hit to be complete — PLE’s n-gram history + conv state, QSA’s ingested indexer keys. Without these a restored prefix would silently serve the PREVIOUS request’s lexical state. Called at chunk-boundary snapshot saves; any D2H inside must be stream-ordered (copy_d2h_on_stream). Default: the layer carries no aux sequence state.
Source§

fn has_aux_state(&self) -> bool

True when this layer WOULD produce aux state — restore sites use it to decline snapshots that lack aux rather than restore a stale mix.
Source§

fn restore_aux( &self, _state: &mut dyn LayerState, _blob: &[u8], _gpu: &dyn GpuBackend, _stream: u64, ) -> Result<()>

Restore the aux state captured by Self::snapshot_aux on a prefix-cache hit, BEFORE the resumed prefill runs.
Source§

fn graph_stale_on_new_sequence(&self) -> bool

Prefill N tokens through this layer using GEMM-batched projections. Read more
Source§

fn sync_replayed_step( &self, _state: &mut dyn LayerState, _seq_len: usize, _k: usize, ) -> Result<()>

Reconcile whatever HOST-side per-sequence bookkeeping a step would have done, when that step was served by a replayed CUDA graph instead of being run. seq_len is the sequence length BEFORE this step’s k rows. Read more
Source§

fn check_replay_room( &self, _state: &dyn LayerState, _seq_len: usize, _k: usize, ) -> Result<()>

Refuse a step whose writes would land past a host-tracked cache — BEFORE the graph that performs them is replayed. Read more
Source§

fn prefill_phase1( &self, hidden: DevicePtr, residual: DevicePtr, num_tokens: usize, state: &mut dyn LayerState, kv_cache: &mut PagedKvCache, seq_len_start: usize, block_table: &mut Vec<u32>, disk_block_ids: &mut Vec<u32>, disk_last_offloaded_per_layer: &mut Vec<u32>, kv_write_start: usize, gdn_bufs: &GdnPrefillBuffers, token_offset: usize, ctx: &ForwardContext<'_>, stream: u64, ) -> Result<()>

Two-phase SSM prefill — Phase 1: projections and GDN input staging. Read more
Source§

fn prefill_phase1_proj_batched( &self, hidden_stacked: DevicePtr, residual_stacked: DevicePtr, total_tokens: usize, gdn_bufs: &GdnPrefillBuffers, ctx: &ForwardContext<'_>, stream: u64, ) -> Result<()>

M1 large-M batched Phase-1: token-parallel projections (RMS/QKVZ/BA-gates) over ALL stacked tokens in one large-M GEMM each. SSM-only; the caller runs prefill_phase1_conv1d_one per request then prefill_phase1_l2_batched.
Source§

fn prefill_phase1_conv1d_one( &self, state: &mut dyn LayerState, token_offset: usize, len: usize, gdn_bufs: &GdnPrefillBuffers, ctx: &ForwardContext<'_>, stream: u64, ) -> Result<()>

M1: per-request conv1d tail (advances per-request conv_state), reading the request’s slice of the stacked QKVZ scratch and writing into gdn_bufs.qkv.
Source§

fn prefill_phase1_l2_batched( &self, total_tokens: usize, gdn_bufs: &GdnPrefillBuffers, ctx: &ForwardContext<'_>, stream: u64, ) -> Result<()>

M1: batched L2 norm over the full stacked QKV buffer after all per-request conv1d tails have written their slices.
Source§

fn prefill_gdn_full( &self, _state: &mut dyn LayerState, _gdn_bufs: &GdnPrefillBuffers, _ctx: &ForwardContext<'_>, _stream: u64, ) -> Result<()>

Two-phase SSM prefill — Phase 2: GDN recurrence on the full sequence. Read more
Source§

fn prefill_inner_batched_q12( &self, _hidden_stacked: DevicePtr, _residual_stacked: DevicePtr, _num_tokens: usize, _kv_cache: &mut PagedKvCache, _seq_len_start: usize, _batched_meta: &BatchedAttnMetadata, _ctx: &ForwardContext<'_>, _stream: u64, ) -> Result<()>

Q12 Path B: batched attention prefill across N stacked-input streams. Read more
Source§

fn prefill_gdn_full_batched( &self, _h_state_ptrs: DevicePtr, _gdn_bufs: &GdnPrefillBuffers, _batch_size: u32, _chunk_len: u32, _ctx: &ForwardContext<'_>, _stream: u64, ) -> Result<()>

Q12 Path B: batched GDN recurrence across N streams. Read more
Source§

fn prefill_gdn_full_batched_fla_varlen( &self, _h_state_ptrs: DevicePtr, _gdn_bufs: &GdnPrefillBuffers, _batch_size: u32, _cu_seqlens: DevicePtr, _max_num_chunks: u32, _total_nt: usize, _max_seqlen: u32, _ctx: &ForwardContext<'_>, _stream: u64, ) -> Result<bool>

VARLEN batched GDN: process ragged co-dispatch lengths via cu_seqlens in ONE gdn_prefill_fla(batch=N, is_varlen) call (replaces the non-uniform per-request loop → fills chunk_delta_h’s 32→32N CTAs). Returns Ok(true) if it ran, Ok(false) if not eligible (caller falls back to the loop). Default (non-SSM layers, or FLA disabled): Ok(false).
Source§

fn prefill_phase3( &self, _hidden: DevicePtr, _residual: DevicePtr, _num_tokens: usize, _gdn_bufs: &GdnPrefillBuffers, _token_offset: usize, _ctx: &ForwardContext<'_>, _stream: u64, ) -> Result<()>

Two-phase SSM prefill — Phase 3: post-GDN processing. Read more
Source§

fn is_ssm_layer(&self) -> bool

Returns true if this layer is an SSM layer (supports two-phase prefill). Read more
Source§

fn transpose_moe_for_prefill( &mut self, _gpu: &dyn GpuBackend, _config: &ModelConfig, ) -> Result<()>

Allocate the transposed MoE expert weights used by the coalesced prefill GEMM kernels. Called as a post-load pass from factory::build after LM-head NVFP4 quantization has freed BF16 headroom, so memory-tight EP configurations (e.g. MiniMax M2.7-NVFP4 EP=2) can fit the transpose that layer-0 preflight would otherwise reject. Read more
Source§

fn transpose_moe_gate_up_for_prefill( &mut self, _gpu: &dyn GpuBackend, _config: &ModelConfig, ) -> Result<()>

Like transpose_moe_for_prefill but only transposes the gate+up projections (skips the down projection), reducing the transpose cost from 3× to 2× per expert. Used as a memory-tight fallback by the MiniMax loader when full transpose doesn’t fit.
Source§

fn set_moe_down_transpose_scratch( &mut self, _scratch_packed: DevicePtr, _scratch_scale: DevicePtr, _packed_ptrs_t: DevicePtr, _scale_ptrs_t: DevicePtr, )

Wire a shared per-prefill down_proj transpose scratch into this layer’s MoE block. Used as a memory-tight alternative to the persistent down transpose: factory allocates one shared scratch, every MoE layer reuses it layer-by-layer during sequential prefill. No-op for non-MoE layers and MoE layers that already have a persistent transposed down.
Source§

fn transpose_moe_for_prefill_unified( &mut self, _gpu: &dyn GpuBackend, _config: &ModelConfig, ) -> Result<()>

Phase 8a unified-layout MoE transpose: build persistent transposed gate/up/down for all experts and free the untransposed copies. Phased flow keeps memory budget tight enough for MiniMax M2.7 EP=2. After this call, the untransposed-layout decode kernels can no longer execute correctly — MoeLayer::use_t_layout_for_decode() must gate dispatch to the _t decode kernels. Default no-op.
Source§

fn transpose_moe_for_prefill_hybrid( &mut self, _gpu: &dyn GpuBackend, _config: &ModelConfig, ) -> Result<()>

Block C Path 2 hybrid-layout MoE transpose: build persistent transposed gate/up/down alongside the untransposed originals (no frees). Doubles MoE-weight memory but recovers the ~15 % decode regression of pure unified mode — decode + MTP verify dispatch keeps using the warp-reduction kernels on the originals while prefill (forward_batched) routes through transposed kernels. Caller must verify enough free memory before invocation. Default no-op for non-MoE layers.
Source§

fn decode_batched( &self, hidden: DevicePtr, residual: DevicePtr, num_tokens: usize, state: &mut dyn LayerState, kv_cache: &mut PagedKvCache, seq_len: usize, block_table: &mut Vec<u32>, disk_block_ids: &mut Vec<u32>, disk_last_offloaded_per_layer: &mut Vec<u32>, ctx: &ForwardContext<'_>, stream: u64, ) -> Result<()>

Decode K tokens through this layer using GEMM-batched projections. Read more
Source§

fn decode_multi_seq<'a, 'b: 'a>( &self, hidden: DevicePtr, residual: DevicePtr, num_seqs: usize, states: &'a mut [&'b mut (dyn LayerState + 'static)], kv_cache: &mut PagedKvCache, seq_lens: &[usize], block_tables: &[Vec<u32>], ctx: &ForwardContext<'_>, stream: u64, ) -> Result<()>

Decode N sequences through this layer in a single batched call. Read more
Source§

fn decode_verify_multi<'a, 'b: 'a>( &self, _hidden: DevicePtr, _residual: DevicePtr, _n_seqs: usize, _ks: &[usize], _states: &'a mut [&'b mut (dyn LayerState + 'static)], _kv_cache: &mut PagedKvCache, _wy_tables: DevicePtr, _ctx: &ForwardContext<'_>, _stream: u64, ) -> Result<()>

Batched MTP verify: n_seqs sequences × k tokens through this layer in ONE weight sweep (rows seq-major, r = i*k + j, contiguous in hidden/residual). Projections/FFN batch across all n_seqs*k rows; the stateful recurrence (conv/GDN) runs per-sequence against states[i] with row-offset buffer bases — per-sequence math is byte-identical to the single-sequence decode_batched K-token body. Read more
Source§

fn release_state( &self, _state: &mut dyn LayerState, _gpu: &dyn GpuBackend, ) -> Result<()>

Release the per-sequence state alloc_state produced, plus anything the layer attached to it lazily afterwards. Read more
Source§

fn uses_ssm_pool(&self) -> bool

Does this layer’s recurrent state live in the shared SSM pool? Read more

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

§

impl<T> Instrument for T

§

fn instrument(self, span: Span) -> Instrumented<Self>

Instruments this type with the provided [Span], returning an Instrumented wrapper. Read more
§

fn in_current_span(self) -> Instrumented<Self>

Instruments this type with the current Span, returning an Instrumented wrapper. Read more
Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = Infallible

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, <T as TryFrom<U>>::Error>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.
§

impl<V, T> VZip<V> for T
where V: MultiLane<T>,

§

fn vzip(self) -> V

§

impl<T> WithSubscriber for T

§

fn with_subscriber<S>(self, subscriber: S) -> WithDispatch<Self>
where S: Into<Dispatch>,

Attaches the provided Subscriber to this type, returning a [WithDispatch] wrapper. Read more
§

fn with_current_subscriber(self) -> WithDispatch<Self>

Attaches the current default Subscriber to this type, returning a [WithDispatch] wrapper. Read more