Task = POS tagging ```python def val_step(self, global_step: int, batch, device="cpu", encoder = None, encoder_kwargs={}): """ Can return multiple outputs. First output need not be loss. """ ... print(rels_predicted.shape) return label_loss, pointer_loss, rels_predicted, rels_labels ``` validation ptb_dep 3:: 0%| | 0/7 [00:00<?, ?it/s]torch.Size([1541]) torch.Size([1547]) torch.Size([1500]) torch.Size([1514]) torch.Size([1570]) torch.Size([1506]) torch.Size([1477]) torch.Size([1626]) validation ptb_dep 2:: 29%|█████████████████████████████████████████████████████████▏ | 2/7 [00:00<00:00, 30.67it/s] gathering validation ptb_dep 3:: 29%|█████████████████████████████████████████████████████████▏ | 2/7 [00:00<00:00, 29.47it/s] gathering validation ptb_dep 1:: 29%|█████████████████████████████████████████████████████████▏ | 2/7 [00:00<00:00, 28.46it/s] gathering validation ptb_dep 0:: 29%|█████████████████████████████████████████████████████████▏ | 2/7 [00:00<00:00, 27.57it/s] gathering