| """PiSSA initialization: top-r SVD components of W absorbed into LoRA.""" | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| import torch | |
| 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.