return float(reduced)
def _accumulate_batches(self, value: Any) -> None:
"""Accumulate batch metrics without retaining input tensor references."""
for batch in self._micro_batches(value):
token_count = self._batch_tokens(batch)
if callable(getattr(token_count, "detach", None)):
token_count = token_count.detach()
if self._local_step_tokens is None:
if callable(getattr(token_count, "clone", None)):
token_count = token_count.clone()
self._local_step_tokens = token_count
else:
self._local_step_tokens = self._local_step_tokens + token_count
self._local_step_samples += self._batch_samples(batch)
def _global_samples(self) -> int:
"""Reduce samples across DP+CP while removing CP replicas."""
cp_size = int(getattr(self.trainer.mesh, "cp_size", 1))
if cp_size < 1:
raise ValueError(f"mesh.cp_size must be positive, but got {cp_size}")
reduced_samples = self._reduce(self._local_step_samples, op="sum")
global_samples = reduced_samples / cp_size
if not global_samples.is_integer():
raise ValueError(
"Reduced sample count must be divisible by cp_size, "
f"but got reduced_samples={reduced_samples} and cp_size={cp_size}"
)
return int(global_samples)
def _current_lr(self) -> float:
"""Return the maximum learning rate across scheduler or optimizer groups."""
schedulers = self.trainer.lr_scheduler