| """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" | |
| 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.