| 192 | |
| 193 | |
| 194 | class 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 |
no outgoing calls
no test coverage detected