pub struct NemotronMoeLayer { /* private fields */ }Expand description
Nemotron-H standalone MoE FFN layer.
Implementations§
Source§impl NemotronMoeLayer
impl NemotronMoeLayer
Sourcepub fn prepare_prefill_weights(
&mut self,
gpu: &dyn GpuBackend,
config: &ModelConfig,
)
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
impl NemotronMoeLayer
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
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<()>
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<()>
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 moreSource§fn alloc_state(&self, _gpu: &dyn GpuBackend) -> Result<Box<dyn LayerState>>
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
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>
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>
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<()>
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)
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
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
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 moreSource§fn decode_verify_multi_unsupported(&self) -> bool
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 moreSource§fn snapshot_aux(
&self,
_state: &dyn LayerState,
_gpu: &dyn GpuBackend,
_stream: u64,
) -> Result<Option<Vec<u8>>>
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
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<()>
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
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<()>
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 moreSource§fn check_replay_room(
&self,
_state: &dyn LayerState,
_seq_len: usize,
_k: usize,
) -> Result<()>
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<()>
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<()>
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<()>
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<()>
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<()>
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<()>
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<()>
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>
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<()>
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
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<()>
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 moreSource§fn transpose_moe_gate_up_for_prefill(
&mut self,
_gpu: &dyn GpuBackend,
_config: &ModelConfig,
) -> Result<()>
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,
)
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<()>
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<()>
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<()>
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<()>
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<()>
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 moreSource§fn release_state(
&self,
_state: &mut dyn LayerState,
_gpu: &dyn GpuBackend,
) -> Result<()>
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 moreSource§fn uses_ssm_pool(&self) -> bool
fn uses_ssm_pool(&self) -> bool
Does this layer’s recurrent state live in the shared SSM pool? Read more
Auto Trait Implementations§
impl Freeze for NemotronMoeLayer
impl RefUnwindSafe for NemotronMoeLayer
impl Send for NemotronMoeLayer
impl Sync for NemotronMoeLayer
impl Unpin for NemotronMoeLayer
impl UnwindSafe for NemotronMoeLayer
Blanket Implementations§
Source§impl<T> BorrowMut<T> for Twhere
T: ?Sized,
impl<T> BorrowMut<T> for Twhere
T: ?Sized,
Source§fn borrow_mut(&mut self) -> &mut T
fn borrow_mut(&mut self) -> &mut T
Mutably borrows from an owned value. Read more