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