Records the packing efficiency for each batch.
| 8 | |
| 9 | |
| 10 | class PackingEfficency(Callback): |
| 11 | """Records the packing efficiency for each batch.""" |
| 12 | |
| 13 | def __init__(self, log_interval: int = 100): |
| 14 | self.log_interval = log_interval |
| 15 | |
| 16 | def after_dataloader(self, state: State, logger: Logger) -> None: |
| 17 | if state.timestamp.batch.value % self.log_interval != 0: |
| 18 | return |
| 19 | logger.log_metrics( |
| 20 | { |
| 21 | "trainer/packing_efficiency": self._packing_efficiency(state), |
| 22 | } |
| 23 | ) |
| 24 | |
| 25 | def _packing_efficiency(self, state: State) -> float: |
| 26 | return state.batch["attention_mask"].sum().item() / state.batch["attention_mask"].numel() |