| """FILet reproduction (arXiv 2605.01046). | |
| We use the (S_X, S_Y) interpretation: ``S_X`` is the input activation second | |
| moment and ``S_Y`` is the output gradient second moment, both estimated via | |
| forward and backward hooks on the target linear layer (the same machinery as | |
| K-FAC). FILet then ranks the right/left singular vectors of ``W`` by a | |
| per-direction Fisher Energy and *keeps the lowest-energy* directions. | |
| There is one ambiguous step in the paper: ``V_hat = normalize(W^T W)``. The | |
| default interpretation here is "right singular vectors of W (already | |
| orthonormal)". If reproduction misses by >0.5 we will swap to row-normalize | |
| or no-op as alternative interpretations (one-shot per Task 4.1). | |
| """ | |
| from __future__ import annotations | |
| from collections.abc import Callable, Iterable | |
| from typing import Any | |
| import torch | |
| from torch import nn | |
| from ._utils import move_batch_to_device, resolve_linear_layer | |
| def collect_filet_factors( | |
| model: nn.Module, | |
| layer_name: str, | |
| batches: Iterable[Any], | |
| n_batches: int = 8, | |
| *, | |
| loss_fn: Callable[[Any, Any], torch.Tensor] | None = None, | |
| move_batch: Callable[[Any, torch.device], Any] | None = None, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| """Returns ``(S_X, S_Y)`` on CPU as float32. | |
| - ``S_X = E[x x^T]`` with ``x`` the input to the layer (shape ``[m, m]``). | |
| - ``S_Y = E[g g^T]`` with ``g`` the gradient of loss wrt the layer output | |
| (shape ``[n, n]``). | |
| """ | |
| layer = resolve_linear_layer(model, layer_name) | |
| device = next(model.parameters()).device | |
| sx = torch.zeros(layer.in_features, layer.in_features, dtype=torch.float32) | |
| sy = torch.zeros(layer.out_features, layer.out_features, dtype=torch.float32) | |
| counts = {"x": 0, "y": 0} | |
| def fwd_hook(_mod, inputs, output): | |
| x = inputs[0].detach() | |
| x_flat = x.reshape(-1, x.shape[-1]).to(torch.float32).cpu() | |
| sx.add_(x_flat.T @ x_flat) | |
| counts["x"] += x_flat.shape[0] | |
| if torch.is_tensor(output) and not output.requires_grad: | |
| return output.detach().requires_grad_(True) | |
| return None | |
| def bwd_hook(_mod, _grad_input, grad_output): | |
| g = grad_output[0] | |
| if g is None: | |
| return | |
| g_flat = g.detach().reshape(-1, g.shape[-1]).to(torch.float32).cpu() | |
| sy.add_(g_flat.T @ g_flat) | |
| counts["y"] += g_flat.shape[0] | |
| handles = [ | |
| layer.register_forward_hook(fwd_hook), | |
| layer.register_full_backward_hook(bwd_hook), | |
| ] | |
| was_training = model.training | |
| requires_grad_state = [(p, p.requires_grad) for p in model.parameters()] | |
| for p, _ in requires_grad_state: | |
| p.requires_grad_(False) | |
| model.eval() | |
| try: | |
| for i, batch in enumerate(batches): | |
| if i >= n_batches: | |
| break | |
| if move_batch is not None: | |
| batch = move_batch(batch, device) | |
| else: | |
| batch = move_batch_to_device(batch, device) | |
| model.zero_grad(set_to_none=True) | |
| outputs = model(**batch) if isinstance(batch, dict) else model(batch) | |
| if loss_fn is not None: | |
| loss = loss_fn(outputs, batch) | |
| elif hasattr(outputs, "loss") and outputs.loss is not None: | |
| loss = outputs.loss | |
| else: | |
| raise ValueError("Provide loss_fn or use a model returning .loss") | |
| loss.backward() | |
| finally: | |
| for h in handles: | |
| h.remove() | |
| for p, rg in requires_grad_state: | |
| p.requires_grad_(rg) | |
| model.train(was_training) | |
| model.zero_grad(set_to_none=True) | |
| if counts["x"] == 0 or counts["y"] == 0: | |
| raise RuntimeError("FILet collection failed: zero counts") | |
| return sx / counts["x"], sy / counts["y"] | |
| def _normalize_columns(M: torch.Tensor, eps: float = 1e-12) -> torch.Tensor: | |
| norms = M.norm(dim=0, keepdim=True).clamp_min(eps) | |
| return M / norms | |
| def filet_init( | |
| layer_weight: torch.Tensor, | |
| rank: int, | |
| sx: torch.Tensor, | |
| sy: torch.Tensor, | |
| *, | |
| candidate_pool: int | None = None, | |
| normalize: str = "default", | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| """FILet init returning ``(U, V)`` with ``ΔW = U V`` at init nonzero. | |
| Args: | |
| layer_weight: ``[n, m]``. | |
| rank: LoRA rank. | |
| sx: ``[m, m]`` input activation second moment. | |
| sy: ``[n, n]`` output gradient second moment. | |
| candidate_pool: Pool of singular directions to score (defaults to | |
| ``min(n, m)``). | |
| normalize: One of {"default", "columns", "none"}. Controls how | |
| ``V_hat`` (and ``U_hat``) are post-processed; "default" leaves the | |
| SVD output untouched (right singular vectors are already | |
| orthonormal). The other modes are tried in Task 4.1 if reproduction | |
| misses. | |
| """ | |
| if rank <= 0: | |
| raise ValueError("rank must be positive") | |
| work_dtype = torch.float32 | |
| target_dtype = layer_weight.dtype | |
| W = layer_weight.to(work_dtype) | |
| n, m = W.shape | |
| if sx.shape != (m, m): | |
| raise ValueError(f"sx must be [m, m]; got {tuple(sx.shape)}") | |
| if sy.shape != (n, n): | |
| raise ValueError(f"sy must be [n, n]; got {tuple(sy.shape)}") | |
| pool = candidate_pool or min(n, m) | |
| pool = min(pool, n, m) | |
| if pool < rank: | |
| raise ValueError(f"candidate_pool={pool} < rank={rank}") | |
| # Step 1-2: V_hat from right SVD, U_hat = normalize(W V_hat). | |
| u_w, s_w, vh_w = torch.linalg.svd(W, full_matrices=False) | |
| V_hat = vh_w[:pool, :].T # [m, pool] | |
| U_hat = u_w[:, :pool] * s_w[:pool].unsqueeze(0) # [n, pool], == W @ V_hat | |
| if normalize == "columns": | |
| V_hat = _normalize_columns(V_hat) | |
| U_hat = _normalize_columns(U_hat) | |
| elif normalize == "none": | |
| pass | |
| elif normalize != "default": | |
| raise ValueError(f"Unknown normalize mode {normalize!r}") | |
| # After normalization, U_hat may not equal W V_hat exactly; recompute to be | |
| # consistent with FILet's intent ("U_hat = normalize(W V_hat)"). | |
| if normalize == "columns": | |
| U_hat = _normalize_columns(W @ V_hat) | |
| # Step 3: per-direction Fisher Energy. | |
| sx_f = sx.to(work_dtype) | |
| sy_f = sy.to(work_dtype) | |
| fisher = ((sx_f @ V_hat) * V_hat).sum(0) * ((sy_f @ U_hat) * U_hat).sum(0) | |
| # Step 4: select rank with LOWEST Fisher Energy. | |
| _, idx = torch.topk(-fisher, k=rank) | |
| idx = idx.sort().values | |
| V_sel = V_hat[:, idx] # [m, r] | |
| U_sel = U_hat[:, idx] # [n, r] | |
| sigma_sel = fisher[idx].clamp(min=0).sqrt() | |
| # Step 5-6: A = sqrt(sigma) V^T, B = U sqrt(sigma) → ΔW = U V. | |
| U_out = U_sel * sigma_sel.unsqueeze(0) | |
| V_out = (sigma_sel.unsqueeze(1) * V_sel.T) | |
| return U_out.to(target_dtype), V_out.to(target_dtype) | |
Xet Storage Details
- Size:
- 6.79 kB
- Xet hash:
- 541c08bbc025cade8a50eeacce1c03ab40108aac608a8f56d1487320e687f0e3
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.