pub struct TpGdnDims {
pub tp_rank: usize,
pub tp_size: usize,
pub h: usize,
pub kd: usize,
pub vd: usize,
pub local_nk: usize,
pub full_nk: usize,
pub local_nv: usize,
pub full_nv: usize,
}Expand description
Pre-TP-shard GDN (linear-attention / SSM) dimensions reconstructed from
config.
Mirrors super::TpAttentionDims: topology.rs divides
linear_num_key_heads / linear_num_value_heads by tp_world_size at
startup, so by the time a loader runs config holds per-rank-local
head counts. The full_* fields multiply back up to the pre-shard sizes
that the segment slicers expect. Head dims (kd, vd) and the hidden
size h are never sharded.
Fields§
§tp_rank: usize§tp_size: usizetp_world_size clamped to >= 1. Loaders treat tp_size == 1 as the
no-shard fast path.
h: usizeHidden size (model embed dim) — never sharded.
kd: usizeKey head dim (linear_key_head_dim) — never sharded.
vd: usizeValue head dim (linear_value_head_dim) — never sharded.
local_nk: usizePer-rank key heads (Q and K share this count).
full_nk: usizeFull pre-shard key heads = local_nk * tp_size.
local_nv: usizePer-rank value heads.
full_nv: usizeFull pre-shard value heads = local_nv * tp_size.
Implementations§
Source§impl TpGdnDims
impl TpGdnDims
pub fn from_config(config: &ModelConfig) -> Self
Sourcepub fn full_key_dim(&self) -> usize
pub fn full_key_dim(&self) -> usize
Full (pre-shard) key projection width: full_nk * kd.
Sourcepub fn local_key_dim(&self) -> usize
pub fn local_key_dim(&self) -> usize
Local key projection width: local_nk * kd.
Sourcepub fn full_value_dim(&self) -> usize
pub fn full_value_dim(&self) -> usize
Full (pre-shard) value projection width: full_nv * vd.
Sourcepub fn local_value_dim(&self) -> usize
pub fn local_value_dim(&self) -> usize
Local value projection width: local_nv * vd.
Sourcepub fn full_conv_dim(&self) -> usize
pub fn full_conv_dim(&self) -> usize
Full conv / QKV width: 2*full_nk*kd + full_nv*vd.
Sourcepub fn local_conv_dim(&self) -> usize
pub fn local_conv_dim(&self) -> usize
Local conv / QKV width: 2*local_nk*kd + local_nv*vd.
Sourcepub fn full_qkvz_out(&self) -> usize
pub fn full_qkvz_out(&self) -> usize
Full QKVZ out dim: 2*full_nk*kd + 2*full_nv*vd.
Sourcepub fn local_qkvz_out(&self) -> usize
pub fn local_qkvz_out(&self) -> usize
Local QKVZ out dim: 2*local_nk*kd + 2*local_nv*vd.