[shardformer]: fix pooled token for left-padded inputs in pipeline sequence classification - #6458
Arthur031221 wants to merge 1 commit into
Conversation
…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.
|
Grep says # 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 lineThose five are not hit by the left-padding bug: on
Only the last row separates those five from 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. |
|
Thanks. Your last row is right, and it disproves the sentence in the PR body that said the |
Checklist before creating the PR
[doc/gemini/tensor/...]: A concise descriptionpip install pre-commit && pre-commit installIssue number
Fixes #6457
What does this PR do?
Anyone running
LlamaForSequenceClassification,Qwen2ForSequenceClassificationorOPTForSequenceClassificationwith 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 firstpad_token_idin the row (gpt2, bloom, falcon and gptj then take it moduloinput_ids.shape[-1]; mistral does not). That is not affected by the left-padding bug fixed here, but a row that containspad_token_idamong its real tokens, which can happen whenpad_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.pycompares each of the three pipeline forwards (one stage holding all layers) with the model's ownforward, on a left-padded and a right-padded batch. It passes with the change and fails with the old code in all three files withllama with left padding. With the old code in onlyqwen2.pyor onlyopt.pyit fails withqwen2 with left paddingandopt with left paddingrespectively. With more than one stage, the pipeline schedule passes the whole micro batch,input_idsincluded, 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
Do you enjoy contributing to Colossal-AI?