| """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.