Garfield.modules.compute_contrastive_clusterloss

Garfield.modules.compute_contrastive_clusterloss(c_i, c_j, class_num, temperature, include_entropy=True)[source]

Cluster loss function.

Args:

c_i (torch.Tensor): First set of cluster probabilities. c_j (torch.Tensor): Second set of cluster probabilities. class_num (int): Number of classes. temperature (float): Temperature scaling factor. include_entropy (bool): If False, drop the cluster-assignment entropy

regularizer H(Y) (the ne_loss term). Used for the H(Y) ablation requested by reviewers. Default True preserves original behavior.

device (torch.device): The device to perform computations on.

Returns:

torch.Tensor: The computed loss value.