Torchmetric that grabs the precomputed z_loss value from the model outputs
| 112 | |
| 113 | @rename_class("ZLoss") |
| 114 | class EfficientZLoss(Metric): |
| 115 | """Torchmetric that grabs the precomputed z_loss value from the model outputs""" |
| 116 | |
| 117 | # Make torchmetrics call update only once |
| 118 | full_state_update = False |
| 119 | |
| 120 | def __init__(self, dist_sync_on_step: bool = False): |
| 121 | super().__init__(dist_sync_on_step=dist_sync_on_step) |
| 122 | self.add_state("sum_loss", default=torch.tensor(0.0), dist_reduce_fx="sum") |
| 123 | self.add_state("total_items", default=torch.tensor(0), dist_reduce_fx="sum") |
| 124 | |
| 125 | def update(self, loss: Tensor) -> None: |
| 126 | """Updates the internal state with results from a new batch. |
| 127 | |
| 128 | Args: |
| 129 | loss (~torch.Tensor): A Tensor of loss values to compare against. |
| 130 | """ |
| 131 | self.sum_loss += loss |
| 132 | self.total_items += 1 |
| 133 | |
| 134 | def compute(self) -> Tensor: |
| 135 | """Aggregate the state over all processes to compute the metric. |
| 136 | |
| 137 | Returns: |
| 138 | loss: The loss averaged across all batches as a :class:`~torch.Tensor`. |
| 139 | """ |
| 140 | # Return average loss over entire dataset |
| 141 | return self.sum_loss / self.total_items # type: ignore (third-party) |
| 142 | |
| 143 | |
| 144 | class EfficientHuggingFaceModel(HuggingFaceModel): |
no outgoing calls
no test coverage detected