INNER CODE UNIT · Python
max_state_len
princeton-nlp/WebShop · transfer/app.py:53
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),
'images': torch.tensor(images),
'labels': torch.tensor(labels),
}
def bart_predict(input):
input_ids = bart_tokenizer(input)['input_ids']
input_ids = torch.tensor(input_ids).unsqueeze(0)
output = bart_model.generate(input_ids, max_length=512, num_return_sequences=5, num_beams=5)
return bart_tokenizer.batch_decode(output.tolist(), skip_special_tokens=True)[0]