Skip to content

[shardformer]: fix pooled token for left-padded inputs in pipeline sequence classification - #6458

Open
Arthur031221 wants to merge 1 commit into
hpcaitech:mainfrom
Arthur031221:hotfix/seq-cls-left-padding
Open

Arthur031221 wants to merge 1 commit into
hpcaitech:mainfrom
Arthur031221:hotfix/seq-cls-left-padding

Conversation

@Arthur031221

@Arthur031221 Arthur031221 commented Sep 29, 2026 •

Copy link
Copy Markdown

Checklist before creating the PR

  • I have created an issue for this PR for traceability
  • The title follows the standard format: [doc/gemini/tensor/...]: A concise description
  • I have added relevant tags if possible for us to better distinguish different PRs
  • I have installed pre-commit: pip install pre-commit && pre-commit install

Issue number

Fixes #6457

What does this PR do?

Anyone running LlamaForSequenceClassification, Qwen2ForSequenceClassification or OPTForSequenceClassification with pipeline parallelism on left-padded batches gets logits and loss taken from the wrong token for every padded row, so the pipelined model trains and evaluates on different outputs than the same model without pipeline parallelism.

The pipeline forwards of these three classes found the pooled token with torch.ne(input_ids, pad_token_id).sum(-1) - 1, which is the last real token only for right-padded rows. This PR takes the rightmost non-pad token instead, the same expression transformers 4.51.3 and the Qwen3 pipeline forward in this repo already use, so both padding sides work.

The other pipeline forwards that pool one token (gpt2, bloom, falcon, gptj, mistral) use torch.eq(input_ids, pad_token_id).int().argmax(-1) - 1, which points one position before the first pad_token_id in the row (gpt2, bloom, falcon and gptj then take it modulo input_ids.shape[-1]; mistral does not). That is not affected by the left-padding bug fixed here, but a row that contains pad_token_id among its real tokens, which can happen when pad_token_id == eos_token_id, can be pooled at the wrong position. That is a separate issue, so these five files are unchanged here.

The script in the issue prints row matches model.forward: [False, True] on main and [True, True] with this change.

tests/test_shardformer/test_model/test_seq_cls_padding.py compares each of the three pipeline forwards (one stage holding all layers) with the model's own forward, on a left-padded and a right-padded batch. It passes with the change and fails with the old code in all three files with llama with left padding. With the old code in only qwen2.py or only opt.py it fails with qwen2 with left padding and opt with left padding respectively. With more than one stage, the pipeline schedule passes the whole micro batch, input_ids included, to every stage, so the last stage runs this same pooling code. The existing shardformer tests did not catch this because the Llama and Qwen2 sequence classification entries in the model zoo use unpadded inputs, and the OPT one is commented out.

Checklist before requesting a review

  • I have linked my PR to an issue (instruction)
  • My issue clearly describes the problem/feature/proposal, with diagrams/charts/table/code if possible
  • I have performed a self-review of my code
  • I have added thorough tests.
  • I have added docstrings for all the functions/methods I implemented

Do you enjoy contributing to Colossal-AI?

  • Yes, I do.
  • No, I don't.

…ification

The pipeline forwards of LlamaForSequenceClassification,
Qwen2ForSequenceClassification and OPTForSequenceClassification located
the pooled token with ne(input_ids, pad_token_id).sum(-1) - 1. That is
the last real token only for right-padded rows; for a left-padded row it
points into the middle of the sequence, so the logits and loss differ
from the model without pipeline parallelism.

Use the rightmost non-pad token, as transformers and the Qwen3 pipeline
forward already do.
@Arthur031221
Arthur031221 requested a review from a team as a code owner September 29, 2026 03:19
@mira687

mira687 commented Sep 30, 2026

Copy link
Copy Markdown

Grep says ne(input_ids, pad).sum(-1) - 1 appears in exactly these three files, so the scope is right for that pattern. But five of the other last-token poolers in shardformer/modeling are on a different legacy pattern and this PR doesn't reach them — falcon.py:517, gpt2.py:764, gptj.py:427, bloom.py:488, mistral.py:331:

# if no pad token found, use modulo instead of reverse indexing for ONNX compatibility
sequence_lengths = torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1
sequence_lengths = sequence_lengths % input_ids.shape[-1]   # mistral.py omits this second line

Those five are not hit by the left-padding bug: on [0, 0, 5, 6, 7, 8] the argmax is 0 and -1 wraps round to 5, which is what last_non_pad_token gives. They diverge somewhere else — when a pad id sits before the last real token, which is the ordinary case once pad_token_id == eos_token_id. Indices for pad 0, L = 6:

row ne.sum-1 eq.argmax-1 % L last_non_pad
[0,0,5,6,7,8] 3 5 5
[5,6,7,8,0,0] 3 3 3
[0,0,5,0,7,8] 2 5 5
[5,0,7,8,0,0] 2 0 3

Only the last row separates those five from last_non_pad_token. qwen3.py:434-436 is already on it, so this PR takes the directory to 4 of 9 aligned with the pinned transformers==4.51.3.

That is the index arithmetic computed on its own (torch 2.2.2, CPU) — I have not run those five pipeline forwards. And the modulo is there deliberately for ONNX, so it isn't a free swap; happy to leave the five if that's the intended scope.

@Arthur031221

Copy link
Copy Markdown
Author

Thanks. Your last row is right, and it disproves the sentence in the PR body that said the eq(...).argmax(-1) - 1 expression in the other five files was correct for left-padded, right-padded and unpadded rows. It points one position before the first pad_token_id in the row, so a row that contains pad_token_id among its real tokens can be pooled at the wrong position. I reproduced your four rows in plain Python and corrected that sentence in the body. I have only checked the index arithmetic, not those five forwards, so this PR stays on the three files with the left-padding bug.

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[BUG]: pipeline sequence classification picks the wrong token for left-padded inputs (Llama, Qwen2, OPT)

2 participants