MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / DataCollatorForSupervisedDataset

Class DataCollatorForSupervisedDataset

data/generation/generate.py:126–141  ·  view source on GitHub ↗

Collate examples for supervised fine-tuning.

Source from the content-addressed store, hash-verified

124
125@dataclass
126class DataCollatorForSupervisedDataset(object):
127 """Collate examples for supervised fine-tuning."""
128
129 tokenizer: transformers.PreTrainedTokenizer
130
131 def __call__(self, instances: Sequence[Dict]) -> Dict[str, torch.Tensor]:
132 input_ids, labels, ids = tuple([instance[key] for instance in instances] for key in ("input_ids", "labels", 'id'))
133 input_ids = padding(input_ids, self.tokenizer.pad_token_id, cutoff = 256)
134 labels = padding(labels, IGNORE_INDEX, cutoff = 256)
135
136 return dict(
137 input_ids=input_ids,
138 labels=labels,
139 id=torch.tensor(ids).to(input_ids.device),
140 attention_mask=input_ids.ne(self.tokenizer.pad_token_id),
141 )
142
143def make_supervised_data_module(tokenizer: transformers.PreTrainedTokenizer, dataset_name, max_sample=None) -> Dict:
144 """Make dataset and collator for supervised fine-tuning."""

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected