diff --git a/unsloth/context_parallel.py b/unsloth/context_parallel.py index 7fa21aea08..35cf018403 100644 --- a/unsloth/context_parallel.py +++ b/unsloth/context_parallel.py @@ -222,11 +222,6 @@ class ContextParallelManager: local_tokens = shift_labels.ne(-100).sum().float() global_tokens = local_tokens.clone() dist.all_reduce(global_tokens, op = dist.ReduceOp.SUM, group = self._cp_group) - - # Cache for reduce_grad_norm - self._cached_local_tokens = local_tokens.detach() - self._cached_global_tokens = global_tokens.detach() - return global_tokens def reduce_loss(self, loss): @@ -241,35 +236,14 @@ class ContextParallelManager: else: tensor = loss - # Each rank has loss = mean(local_CE). Average across ranks for reporting. - # This gives a good approximation of global mean when tokens are balanced. + # Sum across ranks and divide by cp_size for correct mean. global_loss = tensor.detach().clone() dist.all_reduce(global_loss, op = dist.ReduceOp.SUM, group = self._cp_group) global_loss = global_loss / self.settings.size self._set_report_loss(global_loss) - # Return unmodified loss for backward - gradients are correct with mean reduction return loss - def reduce_grad_norm(self, grad_norm: float) -> float: - """ - Reduce gradient norm across CP group. - - With mean reduction, each rank's gradients are scaled by 1/local_tokens. - Average local norms for consistent reporting with CP=1. - """ - if self._cp_group is None: - return grad_norm - - # Average local norms across CP ranks - local_norm = torch.tensor( - grad_norm, - dtype = torch.float32, - device = torch.device(DEVICE_TYPE_TORCH), - ) - dist.all_reduce(local_norm, op = dist.ReduceOp.SUM, group = self._cp_group) - return float(local_norm.item()) / self.settings.size - def patch_sft_config(): """Patch SFTConfig to add context_parallel_size and shuffle_dataset fields.""" @@ -319,7 +293,6 @@ def patch_sft_trainer() -> None: original_compute_loss = trainer_cls.compute_loss original_prediction_step = trainer_cls.prediction_step original_training_step = trainer_cls.training_step - original_log = trainer_cls.log original_get_train_sampler = getattr(trainer_cls, "_get_train_sampler", None) def _patch_train_sampler(original_fn): @@ -438,7 +411,9 @@ def patch_sft_trainer() -> None: use_cp = manager and isinstance(shift_labels, torch.Tensor) if use_cp: - # For CP, always use external loss to avoid fused CE token counting issues + global_tokens = manager.get_global_tokens(inputs) + + # For CP, use external loss to avoid fused CE token counting issues. saved_labels = inputs.pop("labels", None) local_shift_labels = inputs.pop("shift_labels", None) @@ -447,10 +422,13 @@ def patch_sft_trainer() -> None: from unsloth.kernels.cross_entropy_loss import fast_cross_entropy_loss - # Use mean reduction (n_items=None) - fast_cross_entropy_loss counts local tokens + # Scale by GA and cp_size to match gradient behavior of CP=1. + ga_steps = getattr(self.args, "gradient_accumulation_steps", 1) + cp_size = manager.settings.size loss = fast_cross_entropy_loss( logits = logits, labels = local_shift_labels, + n_items = global_tokens * ga_steps / cp_size, ) # Restore @@ -545,18 +523,10 @@ def patch_sft_trainer() -> None: return report_loss return loss - @functools.wraps(original_log) - def patched_log(self, logs, start_time = None): - manager = getattr(self, "_context_parallel_manager", None) - if manager and "grad_norm" in logs: - logs["grad_norm"] = manager.reduce_grad_norm(logs["grad_norm"]) - return original_log(self, logs, start_time) - trainer_cls.__init__ = patched_init trainer_cls.compute_loss = patched_compute_loss trainer_cls.prediction_step = patched_prediction_step trainer_cls.training_step = patched_training_step - trainer_cls.log = patched_log trainer_cls.__unsloth_context_parallel__ = True if original_get_train_sampler is not None: trainer_cls._get_train_sampler = _patch_train_sampler(