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: usizeIndex in the MODEL stack (0..num_hidden_layers). Diagnostics only.
attn_layer_idx: usizeIndex 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: f32FP8 latent-cache scale. Reads and writes must agree; the write takes 1/scale.
persist_bt: boolPersistent 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
impl Glm5NextDsaLayer
Sourcepub fn indexer_forward(
&self,
gpu: &dyn GpuBackend,
hidden: DevicePtr,
state: &mut Glm5NextDsaState,
pos_dev: Option<DevicePtr>,
stream: u64,
) -> Result<()>
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.
Sourcepub 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<()>
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.
Sourcepub 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<()>
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
impl TransformerLayer for Glm5NextDsaLayer
Source§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 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<()>
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<()>
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>>
fn alloc_state(&self, gpu: &dyn GpuBackend) -> Result<Box<dyn LayerState>>
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<()>
hidden in-place. Read moreSource§fn uses_local_mla_prefill(&self) -> bool
fn uses_local_mla_prefill(&self) -> bool
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>
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<()>
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)
decode_prestage).Source§fn decode_graph_unsupported(&self) -> bool
fn decode_graph_unsupported(&self) -> bool
Source§fn decode_multi_seq_unsupported(&self) -> bool
fn decode_multi_seq_unsupported(&self) -> bool
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
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>>>
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
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<()>
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
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<()>
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<()>
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<()>
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<()>
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<()>
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<()>
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<()>
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<()>
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>
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<()>
Source§fn is_ssm_layer(&self) -> bool
fn is_ssm_layer(&self) -> bool
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<()>
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<()>
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, )
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<()>
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<()>
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<()>
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<()>
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<()>
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