Spaces:
Sleeping
Sleeping
| from __future__ import annotations | |
| import math | |
| import types | |
| from typing import Iterable, Sequence | |
| import torch | |
| _HEAD_ATTR_CANDIDATES = ("fc", "heads", "head", "classifier", "class_head", "brand_head", "year_head") | |
| _LORA_PARAM_TOKENS = (".lora_down.", ".lora_up.") | |
| _TIMM_VIT_LORA_TARGETS = { | |
| "qkv": ("attn", "qkv"), | |
| "proj": ("attn", "proj"), | |
| "fc1": ("mlp", "fc1"), | |
| "fc2": ("mlp", "fc2"), | |
| } | |
| class LoRALinear(torch.nn.Module): | |
| """Frozen linear layer with a trainable low-rank update.""" | |
| def __init__( | |
| self, | |
| base_layer: torch.nn.Linear, | |
| rank: int, | |
| alpha: float, | |
| dropout_p: float = 0.0, | |
| ) -> None: | |
| super().__init__() | |
| if rank <= 0: | |
| raise ValueError(f"LoRA rank must be positive, got {rank}") | |
| if not isinstance(base_layer, torch.nn.Linear): | |
| raise TypeError(f"LoRALinear expects torch.nn.Linear, got {type(base_layer)!r}") | |
| self.base_layer = base_layer | |
| self.rank = int(rank) | |
| self.alpha = float(alpha) | |
| self.scaling = self.alpha / self.rank | |
| self.lora_dropout = torch.nn.Dropout(p=dropout_p) if dropout_p > 0.0 else torch.nn.Identity() | |
| self.lora_down = torch.nn.Linear(base_layer.in_features, self.rank, bias=False) | |
| self.lora_up = torch.nn.Linear(self.rank, base_layer.out_features, bias=False) | |
| torch.nn.init.kaiming_uniform_(self.lora_down.weight, a=math.sqrt(5)) | |
| torch.nn.init.zeros_(self.lora_up.weight) | |
| for param in self.base_layer.parameters(): | |
| param.requires_grad = False | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| base_output = self.base_layer(x) | |
| lora_output = self.lora_up(self.lora_down(self.lora_dropout(x))) | |
| return base_output + (lora_output * self.scaling) | |
| def is_lora_parameter_name(name: str) -> bool: | |
| return any(token in name for token in _LORA_PARAM_TOKENS) | |
| def iter_head_parameter_prefixes(model: torch.nn.Module) -> Iterable[str]: | |
| for head_attr in _HEAD_ATTR_CANDIDATES: | |
| if hasattr(model, head_attr): | |
| yield head_attr | |
| def is_head_parameter_name(name: str, head_prefixes: Sequence[str] | None = None) -> bool: | |
| prefixes = tuple(head_prefixes) if head_prefixes is not None else tuple(_HEAD_ATTR_CANDIDATES) | |
| return any(name == prefix or name.startswith(f"{prefix}.") for prefix in prefixes) | |
| def resolve_lora_target_block_indices( | |
| num_blocks: int, | |
| last_n_blocks: int | None = None, | |
| block_indices: Sequence[int] | None = None, | |
| ) -> tuple[int, ...]: | |
| if num_blocks <= 0: | |
| raise ValueError(f"VisionTransformer must have at least one block, got {num_blocks}") | |
| if last_n_blocks is not None and block_indices is not None: | |
| raise ValueError("Specify either last_n_blocks or block_indices for LoRA, not both") | |
| if block_indices is not None: | |
| if len(block_indices) == 0: | |
| raise ValueError("block_indices must not be empty") | |
| normalized_indices = tuple(int(idx) for idx in block_indices) | |
| if len(set(normalized_indices)) != len(normalized_indices): | |
| raise ValueError(f"block_indices contains duplicates: {list(normalized_indices)}") | |
| invalid_indices = [idx for idx in normalized_indices if idx < 0 or idx >= num_blocks] | |
| if invalid_indices: | |
| raise ValueError( | |
| f"LoRA block_indices out of range for {num_blocks} blocks: {invalid_indices}" | |
| ) | |
| return tuple(sorted(normalized_indices)) | |
| if last_n_blocks is None: | |
| return tuple(range(num_blocks)) | |
| block_count = int(last_n_blocks) | |
| if block_count <= 0 or block_count > num_blocks: | |
| raise ValueError( | |
| f"last_n_blocks must be in [1, {num_blocks}] for this model, got {block_count}" | |
| ) | |
| return tuple(range(num_blocks - block_count, num_blocks)) | |
| def enable_lora_suffix_no_grad( | |
| model: torch.nn.Module, | |
| target_block_indices: Sequence[int], | |
| ) -> int: | |
| if not hasattr(model, "blocks"): | |
| raise ValueError("Frozen-prefix no_grad optimization requires VisionTransformer-style blocks") | |
| num_blocks = len(model.blocks) | |
| normalized_indices = tuple(int(idx) for idx in target_block_indices) | |
| if not normalized_indices: | |
| raise ValueError("target_block_indices must not be empty") | |
| first_trainable_block = normalized_indices[0] | |
| expected_suffix = tuple(range(first_trainable_block, num_blocks)) | |
| if normalized_indices != expected_suffix: | |
| raise ValueError( | |
| "Frozen-prefix no_grad optimization supports only contiguous suffix block selections. " | |
| f"Expected {list(expected_suffix)}, got {list(normalized_indices)}" | |
| ) | |
| if first_trainable_block == 0: | |
| model._lora_no_grad_prefix_blocks = () | |
| return 0 | |
| def forward_features_with_frozen_prefix( | |
| self, | |
| x: torch.Tensor, | |
| attn_mask: torch.Tensor | None = None, | |
| is_causal: bool = False, | |
| ) -> torch.Tensor: | |
| x = self.patch_embed(x) | |
| x = self._pos_embed(x) | |
| x = self.patch_drop(x) | |
| x = self.norm_pre(x) | |
| blocks = list(self.blocks) | |
| frozen_prefix = blocks[:first_trainable_block] | |
| trainable_suffix = blocks[first_trainable_block:] | |
| with torch.no_grad(): | |
| if attn_mask is not None or is_causal: | |
| for blk in frozen_prefix: | |
| x = blk(x, attn_mask=attn_mask, is_causal=is_causal) | |
| else: | |
| for blk in frozen_prefix: | |
| x = blk(x) | |
| # Prefix blocks are fully frozen, so their output can enter the trainable suffix detached. | |
| x = x.detach() | |
| if attn_mask is not None or is_causal: | |
| for blk in trainable_suffix: | |
| x = blk(x, attn_mask=attn_mask, is_causal=is_causal) | |
| else: | |
| for blk in trainable_suffix: | |
| x = blk(x) | |
| x = self.norm(x) | |
| return x | |
| model._lora_original_forward_features = getattr(model, "forward_features") | |
| model.forward_features = types.MethodType(forward_features_with_frozen_prefix, model) | |
| model._lora_no_grad_prefix_blocks = tuple(range(first_trainable_block)) | |
| return first_trainable_block | |
| def apply_lora_to_timm_attention( | |
| model: torch.nn.Module, | |
| rank: int, | |
| alpha: float, | |
| dropout_p: float, | |
| target_modules: Sequence[str], | |
| last_n_blocks: int | None = None, | |
| block_indices: Sequence[int] | None = None, | |
| ) -> int: | |
| if not hasattr(model, "blocks"): | |
| raise ValueError("LoRA injection currently supports timm VisionTransformer-style blocks only") | |
| supported_targets = tuple(_TIMM_VIT_LORA_TARGETS) | |
| invalid_targets = sorted(set(target_modules) - set(supported_targets)) | |
| if invalid_targets: | |
| raise ValueError(f"Unsupported LoRA target modules: {invalid_targets}. Supported: {list(supported_targets)}") | |
| target_block_indices = resolve_lora_target_block_indices( | |
| len(model.blocks), | |
| last_n_blocks=last_n_blocks, | |
| block_indices=block_indices, | |
| ) | |
| injected = 0 | |
| for block_idx in target_block_indices: | |
| block = model.blocks[block_idx] | |
| for target_name in target_modules: | |
| parent_attr, layer_attr = _TIMM_VIT_LORA_TARGETS[target_name] | |
| parent_module = getattr(block, parent_attr, None) | |
| if parent_module is None: | |
| raise ValueError(f"VisionTransformer block does not have submodule '{parent_attr}' for LoRA target '{target_name}'") | |
| layer = getattr(parent_module, layer_attr, None) | |
| if layer is None: | |
| raise ValueError(f"VisionTransformer block does not have target '{target_name}'") | |
| if not isinstance(layer, torch.nn.Linear): | |
| raise ValueError( | |
| f"LoRA target '{target_name}' must be torch.nn.Linear, got {type(layer)!r}" | |
| ) | |
| setattr(parent_module, layer_attr, LoRALinear(layer, rank=rank, alpha=alpha, dropout_p=dropout_p)) | |
| injected += 1 | |
| if injected == 0: | |
| raise ValueError("No VisionTransformer layers were patched with LoRA") | |
| model._lora_target_block_indices = target_block_indices | |
| return injected | |
| def freeze_all_parameters_except_lora_and_heads(model: torch.nn.Module) -> None: | |
| head_prefixes = tuple(iter_head_parameter_prefixes(model)) | |
| for name, param in model.named_parameters(): | |
| param.requires_grad = is_lora_parameter_name(name) or is_head_parameter_name(name, head_prefixes) | |
| def count_parameters(model: torch.nn.Module) -> dict[str, int]: | |
| total = 0 | |
| trainable = 0 | |
| lora = 0 | |
| head = 0 | |
| head_prefixes = tuple(iter_head_parameter_prefixes(model)) | |
| for name, param in model.named_parameters(): | |
| numel = param.numel() | |
| total += numel | |
| if param.requires_grad: | |
| trainable += numel | |
| if is_lora_parameter_name(name): | |
| lora += numel | |
| elif is_head_parameter_name(name, head_prefixes): | |
| head += numel | |
| return { | |
| "total": total, | |
| "trainable": trainable, | |
| "frozen": total - trainable, | |
| "lora": lora, | |
| "head": head, | |
| } | |