Remote code: MLA/MTP attention is non-causal for unpadded input under the default (sdpa) load

#15
by bighead-liat - opened

The remote code in this repo (modeling_bailing_moe_v3.py, the same blob is used by Ling-3.0-tiny-base* and Ling-3.0-flash) runs the MLA attention layers and the MTP layer bidirectionally for unpadded input when the model is loaded with the default attn_implementation (sdpa) or with flash_attention_2:

  • BailingMoeV3Model.forward calls _prepare_4d_causal_attention_mask_for_sdpa, which by transformers convention returns None for unpadded input and expects the kernel to run with is_causal=True;
  • but BailingMoeV3Attention.forward hard-codes eager_attention_forward, which only adds a mask it is given, so nothing enforces causality.

generate() still looks fine (decode steps have query_length == 1) and SGLang / vLLM are unaffected, but every teacher-forced forward through transformers sees the future: hidden-state capture for EAGLE3 / DSpark / DFlash drafts, MTP-head or LoRA fine-tuning, perplexity. Measured on this checkpoint: top-1 agreement of a full-sequence forward with the model's own greedy output is 0.944 by default and 0.993 with an explicit causal mask; a padded batch is causal, an unpadded one is not.

Full write-up with a 10-line check and a one-block fix (dispatch through ALL_ATTENTION_FUNCTIONS, or build the mask inside the eager path): https://github.com/inclusionAI/Ling/issues/27

Workarounds until the file is updated: load with attn_implementation="eager", or pass an explicit 4-D additive causal mask to forward.

Sign up or log in to comment