etomoscow's picture
download
raw
1.17 kB
"""PiSSA initialization: top-r SVD components of W absorbed into LoRA."""
from __future__ import annotations
from dataclasses import dataclass
import torch
@dataclass(frozen=True)
class PiSSAResult:
U: torch.Tensor # [n, r]
V: torch.Tensor # [r, m]
residual: torch.Tensor # [n, m] = W - U V
def pissa_init(
layer_weight: torch.Tensor,
rank: int,
) -> PiSSAResult:
"""PiSSA: ``W = U_r diag(s_r) V_r^T``; LoRA holds ``U V`` with the top-r
components, and the LoRA-wrapped layer's residual weight is ``W - UV``.
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)
s_sqrt = s[:rank].clamp(min=0).sqrt()
U = u[:, :rank] * s_sqrt
V = s_sqrt.unsqueeze(1) * vh[:rank, :]
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.17 kB
·
Xet hash:
2705fed7ca997f8bbcf2cb7c1b82b24715a5c42817413f6c4ca4738e36199abc

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