etomoscow's picture
download
raw
8.32 kB
"""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
@dataclass(frozen=True)
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.