MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / EfficientZLoss

Class EfficientZLoss

src/flex_bert.py:114–141  ·  view source on GitHub ↗

Torchmetric that grabs the precomputed z_loss value from the model outputs

Source from the content-addressed store, hash-verified

112
113@rename_class("ZLoss")
114class 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
144class EfficientHuggingFaceModel(HuggingFaceModel):

Callers 1

create_flex_bert_mlmFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected