# num_pos and num_neg below contain the number of positive and negative pairs # computed relative to the i'th input. In an actual setting, this number should # be the same for every input element, but we let it vary here for maximum # flexibility. num_pos = target.sum(dim=1) num_neg = target.size(0) - num_pos