From 5c5d7de2af5dd7ef09907bcf4a78e563de314772 Mon Sep 17 00:00:00 2001 From: William collins Date: Thu, 12 Mar 2026 08:57:35 +0000 Subject: [PATCH] feat: Add native CGGR support for SFTTrainer [skip pre-commit.ci] --- unsloth/cggr/bridge.py | 49 ++++++++++++++++++++++++------------------ unsloth/cggr/router.py | 19 +++++++++++++++- 2 files changed, 46 insertions(+), 22 deletions(-) diff --git a/unsloth/cggr/bridge.py b/unsloth/cggr/bridge.py index e429fd953c..0701cd8d6e 100644 --- a/unsloth/cggr/bridge.py +++ b/unsloth/cggr/bridge.py @@ -76,10 +76,11 @@ class CGGRUnslothBridge: # Create truncated router for difficulty scoring self.router = create_truncated_router(model, num_layers=num_router_layers) - # Training state + # Training state (keep on device to avoid syncs) self.current_step = 0 - self.total_tokens_seen = 0 - self.hard_tokens_seen = 0 + self.device = next(model.parameters()).device + self.total_tokens_seen = torch.tensor(0, device=self.device, dtype=torch.long) + self.hard_tokens_seen = torch.tensor(0, device=self.device, dtype=torch.long) logger.info( f"Initialized CGGR Bridge: min_ratio={min_tokens_ratio}, " @@ -184,17 +185,22 @@ class CGGRUnslothBridge: # Get valid (non-ignored) token mask valid_mask = labels != -100 - # Compute ratio (avoid CPU sync by keeping on GPU) - ratio = self.min_tokens_ratio - if self.dynamic_threshold and valid_mask.any(): + # Compute ratio (stay on GPU to avoid sync) + ratio = torch.tensor(self.min_tokens_ratio, device=self.device, dtype=scores.dtype) + if self.dynamic_threshold: # More confident batch → keep fewer tokens + # We use a cautious approach: only adjust if we have valid scores + # Use where to avoid sync if possible, but a few ops are fine here valid_scores = scores.masked_select(valid_mask) - score_range = valid_scores.max() - valid_scores.min() + 1e-10 - mean_normalized = (valid_scores.mean() - valid_scores.min()) / score_range - # Lower mean score = more confident = keep fewer tokens - confidence = 1.0 - mean_normalized - ratio = self.min_tokens_ratio + (1.0 - self.min_tokens_ratio) * (1.0 - confidence) * 0.5 - ratio = max(self.min_tokens_ratio, min(ratio.item(), 1.0)) + if valid_scores.numel() > 0: + s_min = valid_scores.min() + s_max = valid_scores.max() + score_range = s_max - s_min + 1e-10 + mean_normalized = (valid_scores.mean() - s_min) / score_range + # Lower mean score = more confident = keep fewer tokens + confidence = 1.0 - mean_normalized + ratio = self.min_tokens_ratio + (1.0 - self.min_tokens_ratio) * (1.0 - confidence) * 0.5 + ratio = ratio.clamp(min=self.min_tokens_ratio, max=1.0) # Vectorized masking: compute per-sequence thresholds batch_size, seq_len = labels.shape @@ -226,11 +232,9 @@ class CGGRUnslothBridge: mask_tokens = below_threshold & valid_mask masked_labels.masked_fill_(mask_tokens, -100) - # Update statistics (use item() only once at end) - total = valid_mask.sum().item() - kept = (masked_labels != -100).sum().item() - self.total_tokens_seen += total - self.hard_tokens_seen += kept + # Update statistics (no .item() here - keeps computation on GPU) + self.total_tokens_seen += valid_mask.sum() + self.hard_tokens_seen += (masked_labels != -100).sum() return masked_labels @@ -239,13 +243,16 @@ class CGGRUnslothBridge: self.current_step += 1 def get_stats(self) -> Dict[str, float]: - """Get CGGR statistics for logging.""" - if self.total_tokens_seen == 0: + """Get CGGR statistics for logging (syncs here).""" + total = self.total_tokens_seen.item() + if total == 0: return {"cggr/hard_ratio": 0.0, "cggr/step": self.current_step} + + hard = self.hard_tokens_seen.item() return { - "cggr/hard_ratio": self.hard_tokens_seen / self.total_tokens_seen, + "cggr/hard_ratio": hard / total, "cggr/step": self.current_step, - "cggr/total_tokens": self.total_tokens_seen, + "cggr/total_tokens": total, } @classmethod diff --git a/unsloth/cggr/router.py b/unsloth/cggr/router.py index 52c8300028..5b644f2127 100644 --- a/unsloth/cggr/router.py +++ b/unsloth/cggr/router.py @@ -116,6 +116,7 @@ class TruncatedRouter(nn.Module): self, input_ids: torch.LongTensor, attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.LongTensor] = None, **kwargs, ) -> torch.Tensor: """ @@ -124,6 +125,7 @@ class TruncatedRouter(nn.Module): Args: input_ids: Input token IDs [batch, seq_len] attention_mask: Attention mask [batch, seq_len] + position_ids: Position IDs [batch, seq_len] Returns: logits: Output logits [batch, seq_len, vocab_size] @@ -131,11 +133,26 @@ class TruncatedRouter(nn.Module): # Embeddings hidden_states = self.embed_tokens(input_ids) + # Generate position_ids if not provided (needed for RoPE) + if position_ids is None: + position_ids = torch.arange( + input_ids.size(1), device=input_ids.device + ).unsqueeze(0).expand(input_ids.size(0), -1) + + # Simple expansion for 2D mask to 4D if needed by layers + mask_input = attention_mask + if attention_mask is not None and attention_mask.dim() == 2: + # Convert [batch, seq] to [batch, 1, 1, seq] + mask_input = attention_mask[:, None, None, :] + mask_input = mask_input.to(dtype=hidden_states.dtype) + mask_input = (1.0 - mask_input) * torch.finfo(hidden_states.dtype).min + # Pass through truncated layers for layer in self.layers: layer_outputs = layer( hidden_states, - attention_mask=attention_mask, + attention_mask=mask_input, + position_ids=position_ids, use_cache=False, ) hidden_states = layer_outputs[0] if isinstance(layer_outputs, tuple) else layer_outputs