etomoscow's picture
download
raw
1.1 kB
"""MiLoRA initialization: bottom-r SVD components of W."""
from __future__ import annotations
from .pissa import PiSSAResult
import torch
def milora_init(
layer_weight: torch.Tensor,
rank: int,
) -> PiSSAResult:
"""MiLoRA: take the *smallest* ``rank`` singular components of ``W``.
Same residual mechanic as PiSSA so the forward output is preserved at init.
"""
if rank <= 0:
raise ValueError("rank must be positive")
work_dtype = torch.float32
target_dtype = layer_weight.dtype
W = layer_weight.to(work_dtype)
n, m = W.shape
rank = min(rank, n, m)
u, s, vh = torch.linalg.svd(W, full_matrices=False)
# Take the trailing ``rank`` (smallest) components.
bottom_u = u[:, -rank:].flip(1)
bottom_s = s[-rank:].flip(0).clamp(min=0)
bottom_vh = vh[-rank:, :].flip(0)
s_sqrt = bottom_s.sqrt()
U = bottom_u * s_sqrt
V = s_sqrt.unsqueeze(1) * bottom_vh
residual = (W - U @ V).to(target_dtype)
return PiSSAResult(
U=U.to(target_dtype),
V=V.to(target_dtype),
residual=residual,
)

Xet Storage Details

Size:
1.1 kB
·
Xet hash:
ea99c8c7134640e17f1045ae9631eeb50581c462a8d70fbfe4e829e864b6154b

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