etomoscow's picture
download
raw
2.51 kB
"""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.