| """LoRA random initialization strategies. | |
| Contains both the overscaled initializer (the failure-mode trigger studied in | |
| the paper) and the properly-scaled PEFT-standard initializer. | |
| """ | |
| from __future__ import annotations | |
| import math | |
| import torch | |
| def overscaled_random_init( | |
| layer_weight: torch.Tensor, | |
| rank: int, | |
| *, | |
| seed: int = 42, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| """Overscaled zero-preserving LoRA init (fan-in = rank). | |
| This initializer uses ``bound = 1/sqrt(rank)`` for the non-zero LoRA A/V | |
| factor, which is much larger than PEFT's standard fan-in scaling | |
| (``1/sqrt(in_features)``). For typical LoRA ranks (8–32) applied to | |
| layers with ``in_features`` in the thousands, this produces weight | |
| matrices whose rows are 10–40x larger than intended. | |
| We keep this initializer deliberately because it is the failure-mode | |
| trigger studied in our diagnostic paper. **Do not use this for production | |
| LoRA training.** Use :func:`kaiming_random_init` instead. | |
| ``V`` is uniform ``(-1/sqrt(rank), 1/sqrt(rank))``, ``U`` is zero. | |
| ``DeltaW = UV = 0`` at init. | |
| """ | |
| if rank <= 0: | |
| raise ValueError("rank must be positive") | |
| n, m = layer_weight.shape | |
| dtype = layer_weight.dtype | |
| device = layer_weight.device | |
| generator = torch.Generator(device=device).manual_seed(seed) | |
| bound = 1.0 / math.sqrt(rank) | |
| V = torch.empty(rank, m, dtype=dtype, device=device).uniform_(-bound, bound, generator=generator) | |
| U = torch.zeros(n, rank, dtype=dtype, device=device) | |
| return U, V | |
| # Backward compat alias — the old name was used in experiment runners. | |
| random_init = overscaled_random_init | |
| def kaiming_random_init( | |
| layer_weight: torch.Tensor, | |
| rank: int, | |
| *, | |
| seed: int = 42, | |
| ) -> tuple[torch.Tensor, torch.Tensor]: | |
| """PEFT-standard zero-preserving LoRA init (fan-in = in_features). | |
| ``V`` is Kaiming-uniform with fan-in = ``in_features`` (matching | |
| ``nn.Linear`` default and PEFT's built-in init), ``U`` is zero. | |
| ``DeltaW = UV = 0`` at init. | |
| """ | |
| if rank <= 0: | |
| raise ValueError("rank must be positive") | |
| n, m = layer_weight.shape | |
| dtype = layer_weight.dtype | |
| device = layer_weight.device | |
| generator = torch.Generator(device=device).manual_seed(seed) | |
| bound = 1.0 / math.sqrt(m) | |
| V = torch.empty(rank, m, dtype=dtype, device=device).uniform_(-bound, bound, generator=generator) | |
| U = torch.zeros(n, rank, dtype=dtype, device=device) | |
| return U, V | |
Xet Storage Details
- Size:
- 2.51 kB
- Xet hash:
- c2e81b7d787f616bbe276822d25236f6e3f3f618f6c6cb36c027427a3981377e
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.