INNER CODE UNIT · Python
data_collator
princeton-nlp/WebShop · transfer/app.py:43
def data_collator(batch):
state_input_ids, state_attention_mask, action_input_ids, action_attention_mask, sizes, labels, images = [], [], [], [], [], [], []
for sample in batch:
state_input_ids.append(sample['state_input_ids'])
state_attention_mask.append(sample['state_attention_mask'])
action_input_ids.extend(sample['action_input_ids'])
action_attention_mask.extend(sample['action_attention_mask'])
sizes.append(sample['sizes'])
labels.append(sample['labels'])
images.append(sample['images'])
max_state_len = max(sum(x) for x in state_attention_mask)
max_action_len = max(sum(x) for x in action_attention_mask)
return {
'state_input_ids': torch.tensor(state_input_ids)[:, :max_state_len],
'state_attention_mask': torch.tensor(state_attention_mask)[:, :max_state_len],
'action_input_ids': torch.tensor(action_input_ids)[:, :max_action_len],
'action_attention_mask': torch.tensor(action_attention_mask)[:, :max_action_len],
'sizes': torch.tensor(sizes),