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