From 458f2a93ef1dcdf4750514925a22532e135f309a Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Thu, 12 Mar 2026 08:57:42 +0000 Subject: [PATCH] [pre-commit.ci] auto fixes from pre-commit.com hooks --- unsloth/__init__.py | 2 +- unsloth/cggr/__init__.py | 17 ++-- unsloth/cggr/bridge.py | 185 +++++++++++++++++++++------------------ unsloth/cggr/router.py | 80 +++++++++-------- 4 files changed, 155 insertions(+), 129 deletions(-) diff --git a/unsloth/__init__.py b/unsloth/__init__.py index aaa4f44a20..b0da747cc2 100644 --- a/unsloth/__init__.py +++ b/unsloth/__init__.py @@ -283,6 +283,7 @@ from .trainer import * # CGGR (Confidence-Gated Gradient Routing) integration # Optional: requires `pip install cggr` for full functionality from .cggr import CGGR_AVAILABLE + if CGGR_AVAILABLE: from .cggr import CGGRUnslothBridge, patch_trainer_for_cggr, create_truncated_router @@ -299,4 +300,3 @@ from unsloth_zoo.rl_environments import ( # Patch TRL trainers for backwards compatibility _patch_trl_trainer() - diff --git a/unsloth/cggr/__init__.py b/unsloth/cggr/__init__.py index 8c30df714b..0268a29dc7 100644 --- a/unsloth/cggr/__init__.py +++ b/unsloth/cggr/__init__.py @@ -20,7 +20,7 @@ This module provides selective backpropagation via label masking, enabling Usage: from unsloth.cggr import CGGRUnslothBridge - + trainer = SFTTrainer(...) CGGRUnslothBridge.patch_trainer(trainer, min_tokens_ratio=0.25) trainer.train() @@ -42,6 +42,7 @@ __all__ = [ # Check if CGGR package is available try: import cggr + CGGR_AVAILABLE = True except ImportError: CGGR_AVAILABLE = False @@ -58,13 +59,17 @@ else: "CGGR is not installed. Install with: pip install cggr\n" "For CUDA acceleration: pip install cggr[cuda]" ) - + create_truncated_router = _cggr_not_available - CGGRUnslothBridge = type("CGGRUnslothBridge", (), { - "patch_trainer": staticmethod(_cggr_not_available), - }) + CGGRUnslothBridge = type( + "CGGRUnslothBridge", + (), + { + "patch_trainer": staticmethod(_cggr_not_available), + }, + ) patch_trainer_for_cggr = _cggr_not_available - + class TruncatedRouter: def __init__(self, *args, **kwargs): _cggr_not_available() diff --git a/unsloth/cggr/bridge.py b/unsloth/cggr/bridge.py index 0701cd8d6e..4b7a2e8143 100644 --- a/unsloth/cggr/bridge.py +++ b/unsloth/cggr/bridge.py @@ -35,17 +35,17 @@ __all__ = ["CGGRUnslothBridge", "patch_trainer_for_cggr"] class CGGRUnslothBridge: """ Bridge class for integrating CGGR with Unsloth trainers. - + Patches the trainer's compute_loss method to apply label masking before the forward pass, enabling selective gradient computation. - + Example: >>> from unsloth.cggr import CGGRUnslothBridge >>> trainer = SFTTrainer(...) >>> CGGRUnslothBridge.patch_trainer(trainer, min_tokens_ratio=0.25) >>> trainer.train() """ - + def __init__( self, model: torch.nn.Module, @@ -57,7 +57,7 @@ class CGGRUnslothBridge: ): """ Initialize CGGR bridge. - + Args: model: The model being trained min_tokens_ratio: Minimum fraction of tokens to keep gradients for (0.25 = top 25% hardest) @@ -72,21 +72,21 @@ class CGGRUnslothBridge: self.warmup_steps = warmup_steps self.scoring = scoring self.dynamic_threshold = dynamic_threshold - + # Create truncated router for difficulty scoring - self.router = create_truncated_router(model, num_layers=num_router_layers) - + self.router = create_truncated_router(model, num_layers = num_router_layers) + # Training state (keep on device to avoid syncs) self.current_step = 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) - + 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}, " f"router_layers={num_router_layers}, warmup={warmup_steps}" ) - + @torch.inference_mode() def compute_difficulty_scores( self, @@ -96,29 +96,29 @@ class CGGRUnslothBridge: ) -> torch.Tensor: """ Compute difficulty scores for each token using the truncated router. - + Args: input_ids: Input token IDs [batch, seq_len] labels: Target labels [batch, seq_len] attention_mask: Attention mask [batch, seq_len] - + Returns: difficulty_scores: Per-token difficulty [batch, seq_len] """ # Get logits from truncated router (fast forward pass) - logits = self.router(input_ids, attention_mask=attention_mask) - + logits = self.router(input_ids, attention_mask = attention_mask) + # Compute difficulty based on scoring strategy if self.scoring == "entropy": # High entropy = uncertain = hard # Use log_softmax for numerical stability (single fused kernel) - log_probs = F.log_softmax(logits, dim=-1) + log_probs = F.log_softmax(logits, dim = -1) probs = log_probs.exp() - scores = -torch.sum(probs * log_probs, dim=-1) + scores = -torch.sum(probs * log_probs, dim = -1) elif self.scoring == "margin": # Small margin between top-2 = hard # topk is efficient - only partial sort needed - top2 = torch.topk(logits, k=2, dim=-1).values + top2 = torch.topk(logits, k = 2, dim = -1).values scores = -(top2[..., 0] - top2[..., 1]) # Negative margin (high = hard) elif self.scoring == "loss": # High loss = hard - directly compute per-token loss @@ -127,34 +127,36 @@ class CGGRUnslothBridge: scores = F.cross_entropy( shift_logits.view(-1, shift_logits.size(-1)), shift_labels.view(-1), - reduction="none", - ignore_index=-100, + reduction = "none", + ignore_index = -100, ).view(shift_labels.shape) # Pad to match original sequence length - scores = F.pad(scores, (0, 1), value=0) + scores = F.pad(scores, (0, 1), value = 0) else: # combined - efficient fused computation # Compute log_softmax once (fused kernel) - log_probs = F.log_softmax(logits, dim=-1) + log_probs = F.log_softmax(logits, dim = -1) probs = log_probs.exp() - + # Entropy from log_probs (reuse computation) - entropy = -torch.sum(probs * log_probs, dim=-1) - + entropy = -torch.sum(probs * log_probs, dim = -1) + # Margin from topk - top2 = torch.topk(logits, k=2, dim=-1).values + top2 = torch.topk(logits, k = 2, dim = -1).values margin = top2[..., 0] - top2[..., 1] - + # Normalize and combine - use in-place operations where possible entropy_mean = entropy.mean() entropy_std = entropy.std() + 1e-10 margin_mean = margin.mean() margin_std = margin.std() + 1e-10 - + # Combined score: high entropy OR small margin = hard - scores = (entropy - entropy_mean) / entropy_std - (margin - margin_mean) / margin_std - + scores = (entropy - entropy_mean) / entropy_std - ( + margin - margin_mean + ) / margin_std + return scores - + def mask_easy_tokens( self, input_ids: torch.LongTensor, @@ -163,30 +165,32 @@ class CGGRUnslothBridge: ) -> torch.LongTensor: """ Mask easy tokens in labels with -100 to skip their gradients. - + Args: input_ids: Input token IDs [batch, seq_len] labels: Target labels [batch, seq_len] attention_mask: Attention mask [batch, seq_len] - + Returns: masked_labels: Labels with easy tokens set to -100 """ # During warmup, don't mask anything if self.current_step < self.warmup_steps: return labels - + # Clone labels to avoid modifying original masked_labels = labels.clone() - + # Compute difficulty scores scores = self.compute_difficulty_scores(input_ids, labels, attention_mask) - + # Get valid (non-ignored) token mask valid_mask = labels != -100 - + # Compute ratio (stay on GPU to avoid sync) - ratio = torch.tensor(self.min_tokens_ratio, device=self.device, dtype=scores.dtype) + 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 @@ -199,62 +203,67 @@ class CGGRUnslothBridge: 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) - + 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 - + # Set scores of invalid tokens to -inf so they're never selected as "hard" scores_for_threshold = scores.clone() - scores_for_threshold.masked_fill_(~valid_mask, float('-inf')) - + scores_for_threshold.masked_fill_(~valid_mask, float("-inf")) + # Count valid tokens per sequence - valid_counts = valid_mask.sum(dim=1) # [batch] - + valid_counts = valid_mask.sum(dim = 1) # [batch] + # Compute number to keep per sequence - num_keep = (valid_counts.float() * ratio).long().clamp(min=1) - + num_keep = (valid_counts.float() * ratio).long().clamp(min = 1) + # For each sequence, find the threshold score (k-th largest) # Use topk to find scores we should keep max_valid = valid_counts.max().item() if max_valid > 0: # Sort scores descending to find threshold - sorted_scores, _ = scores_for_threshold.sort(dim=1, descending=True) - + sorted_scores, _ = scores_for_threshold.sort(dim = 1, descending = True) + # Get threshold for each sequence (the num_keep-th highest score) # Clamp indices to valid range - threshold_indices = (num_keep - 1).clamp(min=0, max=seq_len - 1) - thresholds = sorted_scores.gather(1, threshold_indices.unsqueeze(1)).squeeze(1) # [batch] - + threshold_indices = (num_keep - 1).clamp(min = 0, max = seq_len - 1) + thresholds = sorted_scores.gather( + 1, threshold_indices.unsqueeze(1) + ).squeeze(1) # [batch] + # Mask tokens with scores below threshold below_threshold = scores < thresholds.unsqueeze(1) mask_tokens = below_threshold & valid_mask masked_labels.masked_fill_(mask_tokens, -100) - + # 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 - + def step(self): """Called after each training step to update internal state.""" self.current_step += 1 - + def get_stats(self) -> Dict[str, float]: """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": hard / total, "cggr/step": self.current_step, "cggr/total_tokens": total, } - + @classmethod def patch_trainer( cls, @@ -267,7 +276,7 @@ class CGGRUnslothBridge: ) -> "CGGRUnslothBridge": """ Patch a trainer to use CGGR selective backpropagation. - + Args: trainer: HuggingFace/TRL trainer instance min_tokens_ratio: Minimum fraction of tokens to keep (0.25 = 25% hardest) @@ -275,10 +284,10 @@ class CGGRUnslothBridge: warmup_steps: Train normally for this many steps first scoring: Scoring strategy ('entropy', 'margin', 'loss', 'combined') dynamic_threshold: Adjust ratio based on batch confidence - + Returns: CGGRUnslothBridge instance (for accessing stats) - + Example: >>> bridge = CGGRUnslothBridge.patch_trainer(trainer) >>> trainer.train() @@ -286,45 +295,49 @@ class CGGRUnslothBridge: """ # Create bridge instance bridge = cls( - model=trainer.model, - min_tokens_ratio=min_tokens_ratio, - num_router_layers=num_router_layers, - warmup_steps=warmup_steps, - scoring=scoring, - dynamic_threshold=dynamic_threshold, + model = trainer.model, + min_tokens_ratio = min_tokens_ratio, + num_router_layers = num_router_layers, + warmup_steps = warmup_steps, + scoring = scoring, + dynamic_threshold = dynamic_threshold, ) - + # Store reference on trainer trainer._cggr_bridge = bridge - + # Patch compute_loss to apply label masking original_compute_loss = trainer.compute_loss - + @wraps(original_compute_loss) def cggr_compute_loss(model, inputs, *args, **kwargs): # Apply CGGR label masking if "labels" in inputs and inputs["labels"] is not None: inputs = dict(inputs) # Don't modify original inputs["labels"] = bridge.mask_easy_tokens( - input_ids=inputs.get("input_ids"), - labels=inputs["labels"], - attention_mask=inputs.get("attention_mask"), + input_ids = inputs.get("input_ids"), + labels = inputs["labels"], + attention_mask = inputs.get("attention_mask"), ) - + # Call original compute_loss outputs = original_compute_loss(model, inputs, *args, **kwargs) - + # Update step counter bridge.step() - + return outputs - + trainer.compute_loss = cggr_compute_loss - + print(f"🦥 Unsloth + CGGR: Selective backpropagation enabled!") - print(f" → Keeping {min_tokens_ratio*100:.0f}% hardest tokens for gradient computation") - print(f" → Router uses {num_router_layers} layers, warmup={warmup_steps} steps") - + print( + f" → Keeping {min_tokens_ratio*100:.0f}% hardest tokens for gradient computation" + ) + print( + f" → Router uses {num_router_layers} layers, warmup={warmup_steps} steps" + ) + return bridge @@ -335,19 +348,19 @@ def patch_trainer_for_cggr( ) -> CGGRUnslothBridge: """ Convenience function to patch a trainer for CGGR. - + Equivalent to CGGRUnslothBridge.patch_trainer(). - + Args: trainer: Trainer instance to patch min_tokens_ratio: Fraction of tokens to keep (0.25 = 25% hardest) **kwargs: Additional arguments passed to CGGRUnslothBridge.patch_trainer() - + Returns: CGGRUnslothBridge instance """ return CGGRUnslothBridge.patch_trainer( trainer, - min_tokens_ratio=min_tokens_ratio, + min_tokens_ratio = min_tokens_ratio, **kwargs, ) diff --git a/unsloth/cggr/router.py b/unsloth/cggr/router.py index 5b644f2127..58d23dcce9 100644 --- a/unsloth/cggr/router.py +++ b/unsloth/cggr/router.py @@ -32,37 +32,37 @@ __all__ = ["TruncatedRouter", "create_truncated_router"] class TruncatedRouter(nn.Module): """ A truncated version of a language model using only the first N layers. - + Used for fast difficulty scoring in CGGR. Shares weights with the parent model, so uses zero additional memory. - + Args: model: The parent HuggingFace model num_layers: Number of decoder layers to use (default: 2) """ - + def __init__(self, model: nn.Module, num_layers: int = 2): super().__init__() self.num_layers = num_layers - + # Get the base model (handle PEFT wrapping) base_model = model if hasattr(model, "base_model"): base_model = model.base_model if hasattr(base_model, "model"): base_model = base_model.model - + # Store reference to model components (shares weights, no copy) self.embed_tokens = self._get_embed_tokens(base_model) self.layers = self._get_layers(base_model, num_layers) self.norm = self._get_norm(base_model) self.lm_head = self._get_lm_head(model, base_model) - + # Store config for reference self.config = getattr(base_model, "config", None) self.dtype = next(model.parameters()).dtype self.device = next(model.parameters()).device - + def _get_embed_tokens(self, model: nn.Module) -> nn.Module: """Extract embedding layer from model.""" if hasattr(model, "embed_tokens"): @@ -72,7 +72,7 @@ class TruncatedRouter(nn.Module): if hasattr(model, "transformer") and hasattr(model.transformer, "wte"): return model.transformer.wte # GPT-2 style raise ValueError(f"Cannot find embedding layer in model: {type(model)}") - + def _get_layers(self, model: nn.Module, num_layers: int) -> nn.ModuleList: """Extract first N decoder layers.""" layers = None @@ -84,13 +84,13 @@ class TruncatedRouter(nn.Module): layers = model.transformer.h # GPT-2 style elif hasattr(model, "encoder") and hasattr(model.encoder, "layer"): layers = model.encoder.layer # BERT style - + if layers is None: raise ValueError(f"Cannot find decoder layers in model: {type(model)}") - + # Return reference to first N layers (shares weights) return nn.ModuleList([layers[i] for i in range(min(num_layers, len(layers)))]) - + def _get_norm(self, model: nn.Module) -> Optional[nn.Module]: """Extract final normalization layer.""" if hasattr(model, "norm"): @@ -100,17 +100,21 @@ class TruncatedRouter(nn.Module): if hasattr(model, "transformer") and hasattr(model.transformer, "ln_f"): return model.transformer.ln_f # GPT-2 style return None - - def _get_lm_head(self, original_model: nn.Module, base_model: nn.Module) -> nn.Module: + + def _get_lm_head( + self, original_model: nn.Module, base_model: nn.Module + ) -> nn.Module: """Extract language model head.""" if hasattr(original_model, "lm_head"): return original_model.lm_head if hasattr(base_model, "lm_head"): return base_model.lm_head - if hasattr(original_model, "base_model") and hasattr(original_model.base_model, "lm_head"): + if hasattr(original_model, "base_model") and hasattr( + original_model.base_model, "lm_head" + ): return original_model.base_model.lm_head raise ValueError(f"Cannot find lm_head in model: {type(original_model)}") - + @torch.inference_mode() def forward( self, @@ -121,49 +125,53 @@ class TruncatedRouter(nn.Module): ) -> torch.Tensor: """ Forward pass through truncated model to get logits for scoring. - + 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] """ # 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) - + 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 = 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=mask_input, - position_ids=position_ids, - use_cache=False, + attention_mask = mask_input, + position_ids = position_ids, + use_cache = False, ) - hidden_states = layer_outputs[0] if isinstance(layer_outputs, tuple) else layer_outputs - + hidden_states = ( + layer_outputs[0] if isinstance(layer_outputs, tuple) else layer_outputs + ) + # Apply final norm if available if self.norm is not None: hidden_states = self.norm(hidden_states) - + # Get logits logits = self.lm_head(hidden_states) - + return logits @@ -173,23 +181,23 @@ def create_truncated_router( ) -> TruncatedRouter: """ Create a truncated router from a model for CGGR difficulty scoring. - + The router shares weights with the parent model, so uses zero additional GPU memory. It only runs the first N layers to quickly estimate token difficulty. - + Args: model: HuggingFace model (can be PEFT-wrapped) num_layers: Number of decoder layers to use (default: 2) - + Returns: TruncatedRouter instance - + Example: >>> from unsloth import FastLanguageModel >>> model, tokenizer = FastLanguageModel.from_pretrained(...) >>> router = create_truncated_router(model, num_layers=2) """ - router = TruncatedRouter(model, num_layers=num_layers) + router = TruncatedRouter(model, num_layers = num_layers) logger.info(f"Created truncated router with {num_layers} layers for CGGR scoring") return router