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)