MCPcopy Create free account
hub / github.com/huggingface/transformers / trim_batch

Function trim_batch

examples/seq2seq/utils.py:68–76  ·  view source on GitHub ↗

Remove columns that are populated exclusively by pad_token_id

(
    input_ids, pad_token_id, attention_mask=None,
)

Source from the content-addressed store, hash-verified

66
67
68def 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
79class SummarizationDataset(Dataset):

Callers 3

trim_seq2seq_batchMethod · 0.85
collate_fnMethod · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected