etomoscow's picture
download
raw
6.79 kB
"""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.