logits = model(input_batch)[:, -1, :] loss = torch.nn.functional.cross_entropy(logits, target_batch) __ __