Remove columns that are populated exclusively by pad_token_id
(
input_ids, pad_token_id, attention_mask=None,
)
| 66 | |
| 67 | |
| 68 | def trim_batch( |
| 69 | input_ids, pad_token_id, attention_mask=None, |
| 70 | ): |
| 71 | """Remove columns that are populated exclusively by pad_token_id""" |
| 72 | keep_column_mask = input_ids.ne(pad_token_id).any(dim=0) |
| 73 | if attention_mask is None: |
| 74 | return input_ids[:, keep_column_mask] |
| 75 | else: |
| 76 | return (input_ids[:, keep_column_mask], attention_mask[:, keep_column_mask]) |
| 77 | |
| 78 | |
| 79 | class SummarizationDataset(Dataset): |
no outgoing calls
no test coverage detected