Glm5NextDsaLayer

Struct Glm5NextDsaLayer 

Source
pub struct Glm5NextDsaLayer {
    pub cfg: Glm5NextDsaConfig,
    pub weights: Glm5NextDsaWeights,
    pub kernels: Glm5NextDsaLayerKernels,
    pub select_kernels: Glm5NextDsaKernels,
    pub decode_kernel: Glm5NextDsaDecodeKernel,
    pub workspace: Glm5NextDsaWorkspace,
    pub layer_idx: usize,
    pub attn_layer_idx: usize,
    pub rms_eps: f32,
    pub kv_scale: f32,
    pub persist_bt: bool,
}

Fields§

§cfg: Glm5NextDsaConfig§weights: Glm5NextDsaWeights§kernels: Glm5NextDsaLayerKernels§select_kernels: Glm5NextDsaKernels§decode_kernel: Glm5NextDsaDecodeKernel§workspace: Glm5NextDsaWorkspace§layer_idx: usize

Index in the MODEL stack (0..num_hidden_layers). Diagnostics only.

§attn_layer_idx: usize

Index in the KV POOL — the running ordinal over KV-cache-consuming layers, which for GLM-5.3 is 0..11 over the sparse layers, not 0..45.

🪤 These two are NOT interchangeable. The pool is sized to ModelConfig::num_attention_layers(); indexing it with layer_idx reads past the end of the allocation on every layer after the first.

§rms_eps: f32§kv_scale: f32

FP8 latent-cache scale. Reads and writes must agree; the write takes 1/scale.

§persist_bt: bool

Persistent block-table buffers instead of a gpu.alloc/gpu.free per DSA layer per token. ON by default; ATLAS_GLM_DSA_ALLOC_PER_STEP=1 restores the old path.

Implementations§

Source§

impl Glm5NextDsaLayer

Source

pub fn indexer_forward( &self, gpu: &dyn GpuBackend, hidden: DevicePtr, state: &mut Glm5NextDsaState, pos_dev: Option<DevicePtr>, stream: u64, ) -> Result<()>

Project hidden into the indexer cache at position pos, then advance.

Writes k_normed and gate directly into the state rows rather than through a staging buffer: the selector reads k[raw * D + d] over the whole context, so the cache is the natural destination and a copy would buy nothing.

Source

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

ONE drafter CONTEXT row: the KV latent and the indexer entry, with no query, no selection and no attend.

The MTP drafter’s context rows only have to EXIST in these two caches — their block output is discarded. Both caches are pure functions of the row’s own input, exactly as the Qwen drafter prefill exploits, so a context row costs kv_a + latent_write + the indexer’s wk, not a decode step. No MoE, no o_proj, no lm_head.

🪤 seq_len is BOTH the row’s KV slot and its RoPE position (the indexer takes its position from state.len()), so the drafter’s row space must stay DENSE — every pair key from 0 up must have been written. That is what prefill_drafter + the catch-up feed are for.

Source

pub fn decode_k( &self, hidden: DevicePtr, k: usize, state: &mut dyn LayerState, kv_cache: &mut PagedKvCache, seq_len: usize, block_table: &mut Vec<u32>, ctx: &ForwardContext<'_>, stream: u64, is_prefill: bool, ) -> Result<()>

K tokens of one sequence: the projections batched, selection and attention NOT.

The weight-heavy halves — q_a, the absorbed q_b, kv_a and the o_absorb output projection — sweep their weights ONCE for all K rows (1,290 MB/rank/token between them). Everything between them is a function of the individual token’s position: the paged KV slot, the indexer row, the selector geometry over [0, len) and the gather-attend.

🔴 Bit-identical to K serial TransformerLayer::decode calls, which is the requirement: an accepted draft must be the token the unspeculated engine would have emitted. ops::dense_mm_bf16 reproduces each row’s K-iteration order and reduction tree, and rms_norm_vanilla’s grid is the token axis.

🪤 REFUSES a SCALAR (num_seqs == 1) attn_metadata at k > 1. Those scalars — position, KV slot, seq len — describe ONE token, so K rows sharing them would write K queries into the same paged slot and select over the same position: a wrong answer with no shape error.

🔴 It ACCEPTS a K-ROW attn_metadata (num_seqs == k), which is what the graphed verify paths (verify_b/verify_c) already upload: positions [k] u32, slot [k] i64, seq_len [k] i32, block_table [k][max_blocks_per_seq] i32, all at stable device addresses written BEFORE capture or replay. Row r reads element r of each. Without this the layer fell through to its own per-row copy_h2d, and an H2D on a capturing stream fails with CUDA_ERROR_STREAM_CAPTURE_UNSUPPORTED — which is why the K-token verify was eager. At k == 1 every row offset is 0, so that path is unchanged byte for byte.

Trait Implementations§

Source§

impl TransformerLayer for Glm5NextDsaLayer

Source§

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

Release what alloc_state allocated — ANOMALIES A76. Reached by the non-composite paths that hold a bare Glm5NextDsaLayer; the composite Glm5NextLayer has its own, identical, override. Type-driven so a non-DSA state can never be freed here.

Source§

fn check_replay_room( &self, state: &dyn LayerState, seq_len: usize, k: usize, ) -> Result<()>

The replay’s writes end at seq_len + k; the buffer ends at capacity. A62.

Source§

fn sync_replayed_step( &self, state: &mut dyn LayerState, seq_len: usize, k: usize, ) -> Result<()>

The indexer cache length is the one thing this layer keeps on the host. A replayed graph writes the next row (the store kernel reads its position from device memory) but never calls decode, so the counter has to be advanced here or the NEXT eager step plans its selection over a stale length — and decode’s own lockstep check would fire.

Source§

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

Allocate per-sequence state for this layer. Read more
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 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 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<()>

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