| """MFF-LoRA initialization (selection rule × strategy variants). | |
| Conventions | |
| ----------- | |
| We treat a linear layer as ``y = x W^T + b``, with ``W`` of shape ``[n, m]`` | |
| where ``n = out_features`` and ``m = in_features``. The LoRA delta is | |
| ``Δ = U V`` where ``U: [n, r]`` and ``V: [r, m]``, so that the effective | |
| forward becomes ``y = x (W + UV)^T + b``. | |
| Kronecker Fisher factors (per ``mfflora.factors``): | |
| - ``A: [n, n]`` — output-side (gradient covariance), aligns with ``U``. | |
| - ``B: [m, m]`` — input-side (activation covariance), aligns with ``V^T``. | |
| Strategy variants (``Δ`` at init time): | |
| - ``alpha`` (default): ``U = top-r eigvecs of A``, ``V = 0`` → ``Δ = 0``. | |
| - ``beta``: ``U = 0``, ``V = (top-r eigvecs of B)^T`` → ``Δ = 0``. | |
| - ``gamma``: both nonzero (PiSSA-style); the runner must subtract ``UV`` from | |
| the residual weight to preserve the forward output. Returns ``residual``. | |
| Selection rule (which eigenvectors of ``A`` and/or ``B`` to take): | |
| - ``"top"`` — largest eigenvalues (high-Fisher directions). | |
| - ``"bottom"`` — smallest eigenvalues (MiLoRA-style on the Fisher basis). | |
| - ``"energy"`` — Fisher Energy ranking on the joint basis (FILet-style scoring), | |
| with the lowest-energy ``r`` directions retained (matching FILet's choice). | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| import torch | |
| from ..factors import fisher_energy_score, topk_eigvecs | |
| class MFFLoraResult: | |
| """Return value of :func:`mff_lora_init`. | |
| ``residual`` is non-None only for strategy ``gamma`` and is the corrected | |
| weight ``W - UV`` that the LoRA-wrapped layer should hold. | |
| """ | |
| U: torch.Tensor | |
| V: torch.Tensor | |
| residual: torch.Tensor | None | |
| def _select_indices( | |
| A_eigvecs_top: torch.Tensor, | |
| B_eigvecs_top: torch.Tensor, | |
| A: torch.Tensor, | |
| B: torch.Tensor, | |
| rank: int, | |
| selection: str, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| """Return (chosen U-side eigvecs [n, r], chosen V-side eigvecs [m, r]).""" | |
| if selection in ("top", "bottom"): | |
| # Selection has already been performed by topk_eigvecs (top vs bottom) | |
| # for both A and B. We just take the leading ``rank`` columns. | |
| return A_eigvecs_top[:, :rank], B_eigvecs_top[:, :rank] | |
| if selection == "energy": | |
| # Compute Fisher Energy on the *top* eigvec basis of A and B; then | |
| # keep the ``rank`` columns with the LOWEST energy (FILet selection). | |
| # We evaluate on a larger candidate set than ``rank`` so the choice is | |
| # meaningful — passed in as the columns of A_eigvecs_top. | |
| candidates_A = A_eigvecs_top | |
| candidates_B = B_eigvecs_top | |
| k = min(candidates_A.shape[1], candidates_B.shape[1]) | |
| if k < rank: | |
| raise ValueError( | |
| f"Energy selection needs at least rank={rank} candidate " | |
| f"directions on each side; got {k}." | |
| ) | |
| energy = fisher_energy_score( | |
| candidates_A[:, :k], candidates_B[:, :k], A, B | |
| ) | |
| # Lowest energy first. | |
| _, idx = torch.topk(-energy, k=rank) | |
| idx = idx.sort().values | |
| return candidates_A[:, idx], candidates_B[:, idx] | |
| raise ValueError(f"Unknown selection rule {selection!r}") | |
| def mff_lora_init( | |
| layer_weight: torch.Tensor, | |
| A: torch.Tensor, | |
| B: torch.Tensor | None, | |
| rank: int, | |
| selection: str = "top", | |
| strategy: str = "alpha", | |
| *, | |
| candidate_oversample: int = 4, | |
| eig_method: str = "auto", | |
| seed: int = 42, | |
| ) -> MFFLoraResult: | |
| """Build LoRA factors ``(U, V)`` from MFF Kronecker Fisher factors. | |
| Args: | |
| layer_weight: ``[n, m]`` pretrained weight. Used by strategy ``gamma`` | |
| to compute the residual; ignored otherwise. | |
| A: ``[n, n]`` Fisher factor on the output side. | |
| B: ``[m, m]`` Fisher factor on the input side. May be ``None`` when | |
| ``strategy="alpha"`` and ``selection`` is ``"top"`` or ``"bottom"`` | |
| (B eigvecs are computed but discarded in those cases). | |
| rank: LoRA rank ``r``. | |
| selection: ``"top" | "bottom" | "energy"``. | |
| strategy: ``"alpha" | "beta" | "gamma"``. | |
| candidate_oversample: For ``selection="energy"``, look at | |
| ``rank * candidate_oversample`` top-Fisher directions when scoring. | |
| eig_method: Forwarded to :func:`topk_eigvecs`. | |
| seed: For randomized eigendecomposition. | |
| Returns: | |
| :class:`MFFLoraResult`. | |
| """ | |
| if layer_weight.ndim != 2: | |
| raise ValueError(f"layer_weight must be [n, m]; got {tuple(layer_weight.shape)}") | |
| n, m = layer_weight.shape | |
| if A.shape != (n, n): | |
| raise ValueError(f"A must be [n, n]={n,n}; got {tuple(A.shape)}") | |
| if B is not None and B.shape != (m, m): | |
| raise ValueError(f"B must be [m, m]={m,m}; got {tuple(B.shape)}") | |
| if B is None and selection == "energy": | |
| raise ValueError("B is required for selection='energy' (Fisher Energy scoring needs B).") | |
| if B is None and strategy in ("beta", "gamma"): | |
| raise ValueError(f"B is required for strategy='{strategy}'.") | |
| if rank <= 0 or rank > min(n, m): | |
| raise ValueError(f"rank must satisfy 0 < rank <= min(n,m)={min(n,m)}; got {rank}") | |
| work_dtype = torch.float32 | |
| target_dtype = layer_weight.dtype | |
| if selection == "energy": | |
| k_candidates = min(min(n, m), rank * candidate_oversample) | |
| ascending = False | |
| elif selection == "top": | |
| k_candidates = rank | |
| ascending = False | |
| elif selection == "bottom": | |
| k_candidates = rank | |
| ascending = True | |
| else: | |
| raise ValueError(f"Unknown selection {selection!r}") | |
| A_vecs, _ = topk_eigvecs( | |
| A.to(work_dtype), | |
| k_candidates, | |
| ascending=ascending, | |
| method=eig_method, | |
| seed=seed, | |
| ) | |
| # Fast-path: strategy α with top/bottom selection only uses A eigvecs. | |
| # Avoids materialising B eigvecs (and a large [m,m] placeholder) when B | |
| # is None or simply not needed, which would cost 820 MB per down_proj layer. | |
| if strategy == "alpha" and selection in ("top", "bottom"): | |
| chosen_A = A_vecs[:, :rank] | |
| U = chosen_A.to(target_dtype) | |
| V = torch.zeros(rank, m, dtype=target_dtype, device=layer_weight.device) | |
| return MFFLoraResult(U=U, V=V, residual=None) | |
| # B eigvecs are needed for energy selection and beta/gamma strategies. | |
| if B is None: | |
| raise ValueError("B is required for this (selection, strategy) combination.") | |
| B_vecs, _ = topk_eigvecs( | |
| B.to(work_dtype), | |
| k_candidates, | |
| ascending=ascending, | |
| method=eig_method, | |
| seed=seed + 1, | |
| ) | |
| chosen_A, chosen_B = _select_indices( | |
| A_vecs, B_vecs, | |
| A.to(work_dtype), B.to(work_dtype), | |
| rank, selection, | |
| ) | |
| # chosen_A: [n, r] (orthonormal cols), chosen_B: [m, r] (orthonormal cols). | |
| if strategy == "alpha": | |
| U = chosen_A.to(target_dtype) | |
| V = torch.zeros(rank, m, dtype=target_dtype, device=layer_weight.device) | |
| return MFFLoraResult(U=U, V=V, residual=None) | |
| if strategy == "beta": | |
| U = torch.zeros(n, rank, dtype=target_dtype, device=layer_weight.device) | |
| # ``V`` lives in [r, m]; eigvecs of B are columns of [m, r], so V = chosen_B.T. | |
| V = chosen_B.transpose(0, 1).contiguous().to(target_dtype) | |
| return MFFLoraResult(U=U, V=V, residual=None) | |
| if strategy == "gamma": | |
| # Project W into the chosen subspace and absorb scale via SVD on the | |
| # projected core, mirroring PiSSA's "subtract UV from W" trick. | |
| W = layer_weight.to(work_dtype) | |
| # Core of shape [r, r]: chosen_A.T @ W @ chosen_B. | |
| core = chosen_A.transpose(0, 1) @ W @ chosen_B | |
| u_c, s_c, vh_c = torch.linalg.svd(core, full_matrices=False) | |
| # U = chosen_A @ u_c * sqrt(s_c); V = sqrt(s_c) * vh_c @ chosen_B.T. | |
| s_sqrt = s_c.clamp(min=0).sqrt() | |
| U = (chosen_A @ u_c) * s_sqrt | |
| V = (s_sqrt.unsqueeze(1) * vh_c) @ chosen_B.transpose(0, 1) | |
| residual = (W - U @ V).to(target_dtype) | |
| return MFFLoraResult( | |
| U=U.to(target_dtype), | |
| V=V.to(target_dtype), | |
| residual=residual, | |
| ) | |
| raise ValueError(f"Unknown strategy {strategy!r}") | |
Xet Storage Details
- Size:
- 8.32 kB
- Xet hash:
- f8e9012de88f1538d82e1f2e860ff1d40508515bb46209b8c3da7a1fae4d137d
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.