| """ |
| Training Utilities — Shared helpers for additive delta-W fine-tuning. |
| |
| Instead of directly modifying model weights, we freeze all weights and |
| learn an additive ``delta_w`` per weight tensor. The effective weight is |
| ``W_original + delta_w``. The loss penalises delta_w for low rank |
| (nuclear norm) and small magnitude (Frobenius norm). |
| """ |
|
|
| from typing import Dict, List, Optional |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.utils.parametrize as parametrize |
|
|
| from sae.Training_Utils import freeze_model |
|
|
| from experiment.config.train_config import LayerSelectionConfig |
|
|
|
|
| |
| |
| |
|
|
| def get_object_token_ids(tokenizer, keywords: List[str]) -> torch.Tensor: |
| """Resolve a list of keywords into a unique set of token IDs. |
| |
| For each keyword, we encode it and collect all resulting token IDs. |
| Returns a 1-D tensor of unique token IDs to suppress. |
| """ |
| token_ids = set() |
| for kw in keywords: |
| ids = tokenizer.encode(kw, add_special_tokens=False) |
| token_ids.update(ids) |
| return torch.tensor(sorted(token_ids), dtype=torch.long) |
|
|
|
|
| |
| get_toilet_token_ids = get_object_token_ids |
|
|
|
|
| |
| |
| |
|
|
| def generate_clamped_target( |
| model: nn.Module, |
| inputs: dict, |
| object_token_ids: torch.Tensor, |
| ) -> torch.Tensor: |
| """Run a forward pass, clamp object token logits to -inf, return as target. |
| |
| After softmax the clamped tokens will have 0 probability, and the |
| remaining distribution is automatically renormalised by softmax. |
| |
| The model may already have delta parametrizations registered — that is |
| fine; the forward pass will use the current ``W_original + delta_w``. |
| """ |
| with torch.no_grad(): |
| outputs = model(**inputs) |
| logits = outputs.logits.clone() |
| logits[:, :, object_token_ids] = float("-inf") |
| return logits.detach() |
|
|
|
|
| |
| |
| |
|
|
| class WeightDelta(nn.Module): |
| """Additive delta registered via ``torch.nn.utils.parametrize``. |
| |
| On every forward pass the effective weight becomes ``W_original + delta_w``. |
| The delta is initialised to **zeros** so the model starts unchanged. |
| """ |
|
|
| def __init__(self, shape: tuple, device=None, dtype=None): |
| super().__init__() |
| |
| self.delta_w = nn.Parameter(torch.zeros(shape, device=device, dtype=torch.float32)) |
|
|
| def forward(self, X: torch.Tensor) -> torch.Tensor: |
| return X + self.delta_w |
|
|
|
|
| |
| |
| |
|
|
| _PATTERN_MAP = { |
| "all": None, |
| "text_only": ["language_model", "lm_head"], |
| "vision_only": ["vision_tower"], |
| "projector_only": ["multi_modal_projector"], |
| } |
|
|
|
|
| def get_target_param_names( |
| model: nn.Module, |
| layer_config: LayerSelectionConfig, |
| ) -> List[str]: |
| """Return the names of parameters that should be masked. |
| |
| Uses the same pattern-matching logic as the old ``select_trainable_params`` |
| but does **not** modify ``requires_grad`` on any parameter. |
| |
| If ``top_k_layers > 0``, only parameters in the last K transformer layers |
| (plus non-layer params like lm_head) are included. |
| """ |
| import re |
|
|
| patterns: Optional[List[str]] = None |
| if layer_config.mode == "specific_layers": |
| patterns = layer_config.specific_layer_patterns |
| else: |
| patterns = _PATTERN_MAP.get(layer_config.mode) |
|
|
| names = [] |
| for name, _param in model.named_parameters(): |
| if patterns is None: |
| names.append(name) |
| elif any(p in name for p in patterns): |
| names.append(name) |
|
|
| |
| if layer_config.top_k_layers > 0 and names: |
| |
| layer_indices = set() |
| layer_re = re.compile(r"\.layers\.(\d+)\.") |
| for name in names: |
| m = layer_re.search(name) |
| if m: |
| layer_indices.add(int(m.group(1))) |
| if layer_indices: |
| max_layer = max(layer_indices) |
| min_kept = max_layer - layer_config.top_k_layers + 1 |
| filtered = [] |
| for name in names: |
| m = layer_re.search(name) |
| if m: |
| idx = int(m.group(1)) |
| if idx >= min_kept: |
| filtered.append(name) |
| else: |
| |
| filtered.append(name) |
| names = filtered |
|
|
| return names |
|
|
|
|
| |
| |
| |
|
|
| def register_weight_deltas( |
| model: nn.Module, |
| target_param_names: List[str], |
| ) -> Dict[str, nn.Parameter]: |
| """Register a ``WeightDelta`` parametrization on every target parameter. |
| |
| 1. Register a ``WeightDelta`` so that ``module.attr`` now returns |
| ``W_original + delta_w`` during the forward pass. |
| 2. Freeze the original weight — only the delta is trainable. |
| |
| Returns: |
| deltas: ``{param_name: delta_w_parameter}`` |
| """ |
| deltas: Dict[str, nn.Parameter] = {} |
|
|
| for name in target_param_names: |
| parts = name.split(".") |
| attr = parts[-1] |
|
|
| |
| module = model |
| for part in parts[:-1]: |
| module = getattr(module, part) |
|
|
| |
| weight_tensor = getattr(module, attr) |
| delta_module = WeightDelta( |
| weight_tensor.shape, |
| device=weight_tensor.device, |
| dtype=weight_tensor.dtype, |
| ) |
| parametrize.register_parametrization(module, attr, delta_module) |
|
|
| |
| module.parametrizations[attr].original.requires_grad_(False) |
|
|
| deltas[name] = delta_module.delta_w |
|
|
| return deltas |
|
|
|
|
| |
| |
| |
|
|
| def collect_delta_w( |
| deltas: Dict[str, nn.Parameter], |
| ) -> Dict[str, torch.Tensor]: |
| """Return the delta_w dict for loss computation. |
| |
| Since deltas are learned directly, this simply returns them as-is. |
| The returned tensors remain **in the computation graph** so that the |
| Frobenius / nuclear-norm losses can back-propagate into the deltas. |
| """ |
| return {name: dw for name, dw in deltas.items()} |
|
|
|
|
| |
| |
| |
|
|
| def remove_deltas(model: nn.Module, target_param_names: List[str]) -> None: |
| """Remove all delta parametrizations, baking ``W + delta_w`` into each weight. |
| |
| After this call every ``module.attr`` is a plain ``nn.Parameter`` again |
| whose value equals ``W_original + delta_w``. |
| """ |
| for name in target_param_names: |
| parts = name.split(".") |
| attr = parts[-1] |
| module = model |
| for part in parts[:-1]: |
| module = getattr(module, part) |
| parametrize.remove_parametrizations(module, attr, leave_parametrized=True) |
|
|