etomoscow's picture
download
raw
4.18 kB
"""LoRA-GA: gradient-aligned initialization.
Approximates the first SGD update direction with a few-batch averaged gradient
of the full layer weight, then takes its top-r SVD.
Reference: LoRA-GA (arXiv 2407.05000). We follow the simplified formulation
``U = sqrt(rank) * U_g``, ``V = sqrt(rank) * S_g V_g^T`` with the standard
``s/r`` LoRA scaling baked into V (so the runner does not need a special
``alpha``). ΔW at init is rank-r and aligned with the negative gradient.
"""
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 _accumulate_layer_gradient(
model: nn.Module,
layer_name: str,
batches: Iterable[Any],
n_batches: int,
*,
loss_fn: Callable[[Any, Any], torch.Tensor] | None = None,
move_batch: Callable[[Any, torch.device], Any] | None = None,
) -> torch.Tensor:
"""Average the per-batch gradient of ``layer.weight`` over ``n_batches``."""
layer = resolve_linear_layer(model, layer_name)
device = next(model.parameters()).device
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)
layer.weight.requires_grad_(True)
grad_sum = torch.zeros_like(layer.weight, dtype=torch.float32)
count = 0
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()
if layer.weight.grad is None:
raise RuntimeError(f"No grad collected for {layer_name!r}")
grad_sum.add_(layer.weight.grad.detach().to(torch.float32))
count += 1
finally:
for p, rg in requires_grad_state:
p.requires_grad_(rg)
model.train(was_training)
model.zero_grad(set_to_none=True)
if count == 0:
raise RuntimeError("No gradient batches consumed")
return grad_sum / count
def lora_ga_init(
layer_weight: torch.Tensor,
rank: int,
*,
gradient: torch.Tensor | None = None,
model: nn.Module | None = None,
layer_name: str | None = None,
batches: Iterable[Any] | None = None,
n_batches: int = 8,
loss_fn: Callable[[Any, Any], torch.Tensor] | None = None,
) -> tuple[torch.Tensor, torch.Tensor]:
"""LoRA-GA: top-r SVD of the averaged layer-weight gradient.
Provide either ``gradient`` directly, or ``(model, layer_name, batches)``
so the function can accumulate it.
"""
if rank <= 0:
raise ValueError("rank must be positive")
if gradient is None:
if model is None or layer_name is None or batches is None:
raise ValueError(
"lora_ga_init needs either ``gradient`` or "
"``(model, layer_name, batches)``"
)
gradient = _accumulate_layer_gradient(
model, layer_name, batches, n_batches, loss_fn=loss_fn
)
work_dtype = torch.float32
target_dtype = layer_weight.dtype
G = gradient.to(work_dtype)
if G.shape != layer_weight.shape:
raise ValueError(
f"Gradient shape {tuple(G.shape)} != layer_weight {tuple(layer_weight.shape)}"
)
u, s, vh = torch.linalg.svd(G, full_matrices=False)
rank = min(rank, s.numel())
s_sqrt = s[:rank].clamp(min=0).sqrt()
U = u[:, :rank] * s_sqrt
V = s_sqrt.unsqueeze(1) * vh[:rank, :]
return U.to(target_dtype), V.to(target_dtype)

Xet Storage Details

Size:
4.18 kB
·
Xet hash:
f373e92139bedfa24d853bbc72464713548770f2ed6c764e2118441ea00ff6d3

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