etomoscow's picture
download
raw
3.18 kB
"""EVA initialization: top-r eigenvectors of the input activation covariance.
Reference: EVA (arXiv 2410.07170). We use a forward-only collection of input
activation second moments and pick the top-r eigenvectors of that matrix as
the LoRA right factor; the left factor is zero so ``ΔW = 0`` at init.
"""
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_eva_factor(
model: nn.Module,
layer_name: str,
batches: Iterable[Any],
n_batches: int = 8,
*,
move_batch: Callable[[Any, torch.device], Any] | None = None,
) -> torch.Tensor:
"""Estimate ``B = E[x x^T]`` for a layer's input over a few batches."""
layer = resolve_linear_layer(model, layer_name)
device = next(model.parameters()).device
factor_b = torch.zeros(layer.in_features, layer.in_features, dtype=torch.float32)
count = 0
def hook(_mod, inputs, _output):
nonlocal count
x = inputs[0].detach()
x = x.reshape(-1, x.shape[-1]).to(torch.float32).cpu()
factor_b.add_(x.T @ x)
count += x.shape[0]
handle = layer.register_forward_hook(hook)
was_training = model.training
model.eval()
try:
with torch.no_grad():
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(**batch) if isinstance(batch, dict) else model(batch)
finally:
handle.remove()
model.train(was_training)
if count == 0:
raise RuntimeError("No activations collected")
return factor_b / count
def eva_init(
layer_weight: torch.Tensor,
rank: int,
*,
activation_cov: torch.Tensor | None = None,
model: nn.Module | None = None,
layer_name: str | None = None,
batches: Iterable[Any] | None = None,
n_batches: int = 8,
) -> tuple[torch.Tensor, torch.Tensor]:
"""EVA: ``V = top-r eigvecs(B)^T``, ``U = 0``. ``ΔW = 0`` at init."""
if rank <= 0:
raise ValueError("rank must be positive")
n, m = layer_weight.shape
if activation_cov is None:
if model is None or layer_name is None or batches is None:
raise ValueError(
"eva_init needs either ``activation_cov`` or "
"``(model, layer_name, batches)``"
)
activation_cov = collect_eva_factor(model, layer_name, batches, n_batches)
work_dtype = torch.float32
target_dtype = layer_weight.dtype
B = activation_cov.to(work_dtype)
if B.shape != (m, m):
raise ValueError(f"activation_cov must be [m, m]={m,m}; got {tuple(B.shape)}")
sym = (B + B.T) * 0.5
_, eigvecs = torch.linalg.eigh(sym)
top = eigvecs[:, -rank:].flip(1)
V = top.transpose(0, 1).contiguous().to(target_dtype)
U = torch.zeros(n, rank, dtype=target_dtype, device=layer_weight.device)
return U, V

Xet Storage Details

Size:
3.18 kB
·
Xet hash:
901bd6fd5ec3c1bb5ad105b0932f15a4b29323dbac6fb06d6a43dad6d673a040

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