MCPcopy Create free account
hub / github.com/eric-mitchell/direct-preference-optimization / collate_fn

Function collate_fn

preference_datasets.py:186–210  ·  view source on GitHub ↗
(batch)

Source from the content-addressed store, hash-verified

184 ints [tokens] or strings [the original texts]) and returns a batch of examples,
185 PyTorch tensors padded to the maximum length. Strings are passed through."""
186 def collate_fn(batch):
187 # first, pad everything to the same length
188 padded_batch = {}
189 for k in batch[0].keys():
190 if k.endswith('_input_ids') or k.endswith('_attention_mask') or k.endswith('_labels'):
191 if 'prompt' in k: # adapted from https://stackoverflow.com/questions/73256206
192 to_pad = [torch.LongTensor(ex[k][::-1]) for ex in batch]
193 else:
194 to_pad = [torch.LongTensor(ex[k]) for ex in batch]
195 if k.endswith('_input_ids'):
196 padding_value = tokenizer.pad_token_id
197 elif k.endswith('_labels'):
198 padding_value = -100
199 elif k.endswith('_attention_mask'):
200 padding_value = 0
201 else:
202 raise ValueError(f"Unexpected key in batch '{k}'")
203
204 padded_batch[k] = pad_sequence(to_pad, batch_first=True, padding_value=padding_value)
205 if 'prompt' in k: # for the prompt, flip back so padding is on left side
206 padded_batch[k] = padded_batch[k].flip(dims=[1])
207 else:
208 padded_batch[k] = [ex[k] for ex in batch]
209
210 return padded_batch
211 return collate_fn
212
213

Callers 1

get_batch_iteratorFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected