etomoscow's picture
download
raw
7.04 kB
"""Load Kronecker Fisher factors and run eigendecompositions for LoRA init.
The factor library at
``../kronlingua/factors/e1_step2_llama31_base_12lang_gate/`` stores per-(language,
layer, module) safetensors files with keys ``A``, ``B``, ``count_a``, ``count_b``.
- ``A`` is the output-side second moment (gradient outer product avg),
shape ``[out_features, out_features]`` — for Llama-3.1-8B ``gate_proj`` that's
``[14336, 14336]``.
- ``B`` is the input-side second moment (activation outer product avg),
shape ``[in_features, in_features]`` — ``[4096, 4096]`` for ``gate_proj``.
Note: these existing factors are K-FAC ($S_X \otimes S_Y$). The
matrix-free MFF rank-1 Kron of $\mathbb{E}[gg^\top]$ from FisherKronecker is
not yet ported (see CHANGELOG.md "Deferred").
"""
from __future__ import annotations
import os
import sys
from dataclasses import dataclass
from pathlib import Path
import torch
from safetensors.torch import load_file
# Make kronlingua's mff package importable for downstream modules that need it
# (e.g., overlap analysis in E1.5). Idempotent.
_KRONLINGUA = Path(os.environ.get("MFFLORA_KRONLINGUA_ROOT", "external/kronlingua"))
if _KRONLINGUA.exists() and str(_KRONLINGUA) not in sys.path:
sys.path.insert(0, str(_KRONLINGUA))
DEFAULT_FACTOR_ROOT = _KRONLINGUA / "factors" / "e1_step2_llama31_base_12lang_gate"
@dataclass(frozen=True)
class FisherFactors:
"""Kronecker Fisher factors for one (language, layer, module) cell."""
A: torch.Tensor # [n, n], output-side
B: torch.Tensor # [m, m], input-side
layer_name: str
language: str
layer: int
module: str
count_a: int
count_b: int
def _layer_filename(layer: int, module: str) -> str:
return f"model__layers__{layer}__mlp__{module}.safetensors"
def load_factors(
language: str,
layer: int,
module: str = "gate_proj",
factor_dir: Path | str = DEFAULT_FACTOR_ROOT,
) -> FisherFactors:
"""Load Kronecker factors A, B for given (language, layer, module).
Returns:
FisherFactors with A, B as float32 tensors on CPU.
"""
factor_dir = Path(factor_dir)
path = factor_dir / language / _layer_filename(layer, module)
if not path.exists():
raise FileNotFoundError(f"Missing factor file: {path}")
tensors = load_file(str(path), device="cpu")
A = tensors["A"].to(torch.float32)
B = tensors["B"].to(torch.float32)
count_a = int(tensors["count_a"].item()) if "count_a" in tensors else 0
count_b = int(tensors["count_b"].item()) if "count_b" in tensors else 0
layer_name = f"model.layers.{layer}.mlp.{module}"
return FisherFactors(
A=A,
B=B,
layer_name=layer_name,
language=language,
layer=layer,
module=module,
count_a=count_a,
count_b=count_b,
)
def _symmetrize(matrix: torch.Tensor) -> torch.Tensor:
return (matrix + matrix.transpose(-1, -2)) * 0.5
def topk_eigvecs(
matrix: torch.Tensor,
k: int,
*,
ascending: bool = False,
method: str = "auto",
seed: int = 42,
oversampling: int = 16,
n_iter: int = 4,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Return ``(eigvecs, eigvals)`` for a symmetric PSD matrix.
- ``ascending=False`` → top-``k`` (largest eigenvalue first).
- ``ascending=True`` → bottom-``k`` (smallest eigenvalue first).
``method``:
- ``"exact"``: ``torch.linalg.eigh`` on the full matrix (cubic).
- ``"randomized"``: subspace iteration; only correct for the *top* end.
For ``ascending=True`` we silently fall back to exact (no reliable
randomized smallest-eigenvalue path here).
- ``"auto"`` (default): exact for ``n <= 4096``, randomized otherwise
when ``ascending=False``.
"""
if matrix.ndim != 2 or matrix.shape[0] != matrix.shape[1]:
raise ValueError(f"Expected square matrix, got {tuple(matrix.shape)}")
n = matrix.shape[0]
if k <= 0:
raise ValueError("k must be positive")
k = min(k, n)
original_dtype = matrix.dtype
sym = _symmetrize(matrix.to(torch.float32))
if method == "auto":
# eigh is O(n³); randomized is O(n k²). Switch when k << n.
# Even at n=4096, full eigh takes ~10 min on CPU. Use randomized
# whenever k < 10% of n and we want top eigenvectors.
use_random = (k < max(1, int(0.1 * n))) and (n >= 64) and not ascending
method = "randomized" if use_random else "exact"
if ascending:
# Smallest eigenvalues — no reliable randomized path; force exact.
method = "exact"
if method == "exact":
eigvals, eigvecs = torch.linalg.eigh(sym)
# eigh returns ascending eigenvalues.
if ascending:
sel_vals = eigvals[:k]
sel_vecs = eigvecs[:, :k]
else:
sel_vals = eigvals[-k:].flip(0)
sel_vecs = eigvecs[:, -k:].flip(1)
elif method == "randomized":
sel_vecs, sel_vals = _randomized_topk(
sym, k, oversampling=oversampling, n_iter=n_iter, seed=seed
)
else:
raise ValueError(f"Unknown method {method!r}")
return sel_vecs.to(original_dtype), sel_vals.to(original_dtype)
def _randomized_topk(
sym: torch.Tensor,
k: int,
*,
oversampling: int,
n_iter: int,
seed: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Randomized subspace iteration for top-k eigenpairs of symmetric matrix."""
n = sym.shape[0]
p = min(n, k + max(oversampling, 0))
generator = torch.Generator(device=sym.device).manual_seed(seed)
Q = torch.randn(n, p, dtype=torch.float32, device=sym.device, generator=generator)
Q, _ = torch.linalg.qr(Q, mode="reduced")
for _ in range(max(0, n_iter) + 1):
Z = sym @ Q
Q, _ = torch.linalg.qr(Z, mode="reduced")
small = _symmetrize(Q.T @ sym @ Q)
eigvals_small, eigvecs_small = torch.linalg.eigh(small)
# eigh: ascending; we want top-k → take last k and flip.
top_vals = eigvals_small[-k:].flip(0)
top_vecs = (Q @ eigvecs_small[:, -k:]).flip(1)
return top_vecs.contiguous(), top_vals.contiguous()
def fisher_energy_score(
eigvecs_A: torch.Tensor,
eigvecs_B: torch.Tensor,
A: torch.Tensor,
B: torch.Tensor,
) -> torch.Tensor:
"""FILet-style per-direction Fisher Energy on a paired basis.
For each ``i``: ``energy[i] = (v_i^T A v_i) * (u_i^T B u_i)`` where
``v_i = eigvecs_A[:, i]`` and ``u_i = eigvecs_B[:, i]``. Returns a 1-D
tensor of length ``min(eigvecs_A.shape[1], eigvecs_B.shape[1])``.
"""
if eigvecs_A.ndim != 2 or eigvecs_B.ndim != 2:
raise ValueError("Eigvec arguments must be 2-D")
A_f = A.to(torch.float32)
B_f = B.to(torch.float32)
Va = eigvecs_A.to(torch.float32)
Vb = eigvecs_B.to(torch.float32)
k = min(Va.shape[1], Vb.shape[1])
Va = Va[:, :k]
Vb = Vb[:, :k]
a_quad = ((A_f @ Va) * Va).sum(dim=0)
b_quad = ((B_f @ Vb) * Vb).sum(dim=0)
return (a_quad * b_quad).to(eigvecs_A.dtype)

Xet Storage Details

Size:
7.04 kB
·
Xet hash:
9b4c78952e7b860b110af31701dbae6778acf2857d7803a05086e6201ae0f360

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.