MCPcopy Create free account
hub / github.com/PKU-Alignment/align-anything / PreferenceCollator

Class PreferenceCollator

align_anything/datasets/janus/preference.py:194–221  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

192
193
194class PreferenceCollator:
195
196 def __init__(self, pad_token_id: int) -> None:
197 """Initialize a collator."""
198 self.pad_token_id = pad_token_id
199
200 def __call__(self, samples: list[PreferenceSample]) -> PreferenceBatch:
201 return_dict = {}
202 current_device = get_current_device()
203
204 input_ids = [sample['better_input_ids'] for sample in samples] + [
205 sample['worse_input_ids'] for sample in samples
206 ] # size = (2 * B, L)
207 return_dict['input_ids'] = right_padding(input_ids, padding_value=self.pad_token_id).to(
208 current_device
209 ) # size = (2 * B, L)
210
211 if 'pixel_values' in samples[0].keys():
212 pixel_values = [sample['pixel_values'] for sample in samples]
213 return_dict['pixel_values'] = torch.cat(pixel_values + pixel_values, dim=0).to(current_device).unsqueeze(0)
214
215 if "better_images_emb_mask" in samples[0].keys():
216 better_images_emb_mask = [sample['better_images_emb_mask'] for sample in samples]
217 worse_images_emb_mask = [sample['worse_images_emb_mask'] for sample in samples]
218 return_dict['images_emb_mask'] = right_padding(better_images_emb_mask + worse_images_emb_mask, padding_value=0).to(current_device).detach()
219
220 return_dict['task'] = samples[0]['task']
221 return return_dict

Callers 2

get_collatorMethod · 0.70
get_collatorMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected