Method

DAMCHA stands for Data-Adaptive Mahalanobis Cross-Head Attention. It defines attention through a context-dependent bilinear metric.

Input-conditioned metric

For input X of shape [B, T, D], compute a mean context of shape [B, D]. An MLP maps this context to a flattened D × D cross-head matrix. With H heads and dh = D/H, reshape the matrix into [B, H, dh, H, dh].

The standard path learns every block directly. Both diagonal and off-diagonal scales equal one. The bilinear matrix has no symmetry or positive-definiteness constraint. Low-rank factorization and per-layer bias are available as ablations.

RowSum and head scores

Equation (4) defines mᵢ = Σⱼ Mᵢⱼ. Sum the block-column dimension to produce [B, H, dh, dh]. Appendix D uses the positional head slice Xᵢ = X[..., i*dh:(i+1)*dh] in the score equation:

Sᵢ = Xᵢ mᵢ Xᵢᵀ / √dh
Aᵢ = softmax(Sᵢ + mask)
Oᵢ = Aᵢ Vᵢ
Y  = concat(O₁, …, Oₕ) Wₒ

Value and output projections retain the standard multi-head interface. The implementation uses the per-head tensor dimensions specified in Appendix D for the quadratic form in Section 3.3.

Causal context

In an autoregressive decoder, query t conditions its metric on mean(X[:, :t+1]). This keeps the generator and attention scores causal. ViT uses a full-sequence mean. Padding tokens contribute to neither the context mean nor the attention key set.

The prefix convention is the repository's autoregressive implementation of input conditioning. The single-context complexity expression in Appendix D applies to the full-sequence mean; causal prefix generation has a context for each query.

Stack-wise sharing

With share_mlp=True, the decoder computes context metrics from the stack input and reuses them in every layer. With sharing disabled, each layer generates metrics from its own input. Generator parameters are registered at construction, so device conversion, optimizers and checkpoints include them before the first forward pass.

Efficient RowSum

For the standard MLP, its final affine projection and RowSum commute. The code sums the corresponding output weights and biases before applying the final projection. This yields exactly the same effective head metrics and gradients while avoiding a full [B, T, D, D] causal tensor. The full output parameterization remains trainable.

Code correspondence

Paper component Implementation
Input-dependent fθ(X), Eq. (3) StructuredMLPForM.head_metrics
Block-row sum, Eq. (4) Fused output projection or explicit block-row sum
Per-head quadratic scores, Eqs. (5)–(6) MBasedAttention.forward
Value aggregation and concatenation, Eqs. (7)–(8) MBasedAttention.forward
Stack-wise sharing, Section 3.4 Decoder.forward and VisionTransformer.forward
Static metric ablation use_M=True, use_mlp=False
Independent generators use_M=True, use_mlp=True, share_mlp=False

Ablations and API

The original linear_comb, bayesian, separate_B, compact-rank and layer-bias options remain available for experiments. Standard DAMCHA uses off_diag_mode='mlp', scales of 1.0, full rank, and no layer bias. --use_multihead_M is accepted for command compatibility; DAMCHA always uses RowSum per head.

Masks use True or integer 1 for visible entries. Floating-point attention masks add to the logits. padding_mask has shape [B, T] and marks real tokens. get_M_matrices() provides full metric matrices averaged across the last batch context for visualization; causal visualization uses the final prefix.