Skip to content

Ling-3.0 HF remote code: MLA/MTP attention is non-causal for unpadded input under the default (sdpa) attn_implementation #27

Description

@bighead-liat

Summary

With the HF remote code shared by every Ling-3.0 checkpoint (modeling_bailing_moe_v3.py, blob d3a2d415…, identical in Ling-3.0-tiny, Ling-3.0-tiny-base, -base-midtrain, -base-30T and Ling-3.0-flash), a plain model(input_ids) forward is not causal: the 6 MLA attention layers and the MTP layer attend to future tokens whenever the input has no padding and the model was loaded with the default attn_implementation (sdpa, since the class sets _supports_sdpa = True) or with flash_attention_2.

Serving in SGLang / vLLM is unaffected. generate() still produces sensible text (the last prompt position legitimately sees the whole prompt, and decode steps have query_length == 1), which is why the bug is easy to miss. It does corrupt every teacher-forced forward through transformers: hidden-state capture for draft models (EAGLE3 / DSpark / DFlash / MTP), MTP-head or LoRA fine-tuning, perplexity evaluation, and, more mildly, the prompt KV entries built during generate()'s prefill.

Mechanism

  1. BailingMoeV3Model.forward picks the mask by config._attn_implementation:
    • sdpa → _prepare_4d_causal_attention_mask_for_sdpa(...). Following the transformers convention, this returns None when attention_mask is None (or all ones) and key_value_length == query_length or query_length == 1, on the assumption that the attention kernel will be called with is_causal=True.
    • flash_attention_2 → the 2-D mask is passed through, i.e. None for unpadded input.
    • eager (the else branch) → _prepare_4d_causal_attention_mask(...), a real 4-D causal mask.
  2. BailingMoeV3Attention.forward ignores _attn_implementation for the actual computation: attention_interface: Callable = eager_attention_forward is hard-assigned (the flash/sdpa flags only toggle value padding around it), and eager_attention_forward adds the mask only if one was given. self.is_causal = True is set but never read.
  3. So for unpadded input under sdpa/flash_attention_2 the softmax runs over the full sequence. The KDA layers are recurrent and causal by construction, which is why the model still "works": only the MLA layers and the MTP layer leak.

Two consequences worth spelling out: a padded batch is causal (the sdpa helper materialises a 4-D mask when the 2-D mask has zeros) while an unpadded one is not, so results depend on batch composition; and the comment # tptest miss causal_mask = create_causal_mask( in forward is dead code, not the cause.

How it shows up

Measured on inclusionAI/Ling-3.0-tiny, bf16, transformers 4.57, full-sequence forward with use_cache=False, default load:

quantity default (mask dropped) explicit 4-D causal mask
top-1 agreement of the forward's logits with the model's own greedy generation on the same text 0.944 0.993
held-out next-token top-1 on the model's own outputs 0.917 0.951

A causal model re-scoring its own greedy output must reproduce it almost exactly; without the mask it does not, because later tokens change earlier logits. In our case a fine-tune of the MTP layer with the default forward reached 99.8 % train top-1 and then served worse than the untrained layer under SGLang NEXTN (acceptance length 2.22 vs 2.87 on MT-Bench): it had learned to read the future.

Minimal check

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

name = "inclusionAI/Ling-3.0-tiny"
tok = AutoTokenizer.from_pretrained(name, trust_remote_code=True)
model = AutoModelForCausalLM.from_pretrained(name, dtype=torch.bfloat16, trust_remote_code=True).cuda().eval()

ids = tok("The quick brown fox jumps over the lazy dog because", return_tensors="pt").input_ids.cuda()
k = 5
with torch.no_grad():
    a = model(ids[:, :k], use_cache=False).logits[0, -1]   # prefix only
    b = model(ids,        use_cache=False).logits[0, k - 1] # same position, future appended
print("max |Δlogit| at position k-1 when the future is appended:", (a - b).abs().max().item())
# causal model: ~0 (bf16 noise). Ling-3.0 remote code, default load: large.

Fix

Make the attention module honour what the model-level mask code assumes. Smallest change, in BailingMoeV3Attention.forward:

attention_interface = eager_attention_forward
if self.config._attn_implementation != "eager":
    attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation]
attn_output, attn_weights = attention_interface(self, query_states, key_states, value_states,
                                                attention_mask, dropout=..., scaling=self.scaling, **kwargs)

sdpa_attention_forward reads module.is_causal and passes is_causal=True when the mask is None, which is exactly the contract _prepare_4d_causal_attention_mask_for_sdpa relies on. Alternatively keep the eager path but build the causal mask in eager_attention_forward when attention_mask is None (a triu of finfo.min, sliced to key_states.shape[-2]), so the module is safe regardless of what the model passes down. The MTP layer's branch that turns a 4-D mask back into a 2-D int32 row for the linear-attention path already handles the 4-D additive form.

Workarounds for users until then (either is enough): load with attn_implementation="eager", which takes the else branch and builds a real mask; or pass an explicit 4-D additive causal mask to forward. All our numbers above use the second.

Environment

transformers 4.57.x (also reproduced on 5.12), torch 2.x, bf16, RTX 4090. Blob ids of modeling_bailing_moe_v3.py checked on the Hub on 2026-09-14: identical across the five Ling-3.0 repos listed above.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions