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_lossterm). 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.