mirror of
https://github.com/wassname/lora-lite.git
synced 2026-08-02 12:50:47 +08:00
simpler test
This commit is contained in:
@@ -20,6 +20,7 @@ from .variants.dora import DoRAConfig
|
||||
from .variants.hra import HRAConfig
|
||||
from .variants.eva import EVAConfig
|
||||
from .variants.antipasto import AntiPaSTOConfig
|
||||
from .variants.road import RoadConfig
|
||||
|
||||
__all__ = [
|
||||
"AdapterConfig",
|
||||
@@ -32,6 +33,7 @@ __all__ = [
|
||||
"HRAConfig",
|
||||
"EVAConfig",
|
||||
"AntiPaSTOConfig",
|
||||
"RoadConfig",
|
||||
"attach",
|
||||
"detach",
|
||||
"save",
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
"""attach / detach / save / load. The whole runtime."""
|
||||
from __future__ import annotations
|
||||
import json
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.utils.hooks import RemovableHandle
|
||||
@@ -121,19 +122,22 @@ def save(model: nn.Module, path: str) -> None:
|
||||
if state is None:
|
||||
raise RuntimeError("no adapter attached; call attach() first")
|
||||
sd = {k: v.detach().cpu() for k, v in model.state_dict().items() if "lora_" in k}
|
||||
blob = {
|
||||
"cfg": state["cfg"].to_dict(),
|
||||
"state": sd,
|
||||
"base_fp": _base_weight_fingerprint(model),
|
||||
metadata = {
|
||||
"cfg": json.dumps(state["cfg"].to_dict()),
|
||||
"base_fp": json.dumps(_base_weight_fingerprint(model)),
|
||||
}
|
||||
torch.save(blob, path)
|
||||
from safetensors.torch import save_file
|
||||
save_file(sd, path, metadata=metadata)
|
||||
|
||||
|
||||
def load(model: nn.Module, path: str) -> list[RemovableHandle]:
|
||||
blob = torch.load(path, weights_only=True, map_location="cpu")
|
||||
cfg = AdapterConfig.from_dict(blob["cfg"])
|
||||
from safetensors.torch import load_file, safe_open
|
||||
with safe_open(path, framework="pt", device="cpu") as f:
|
||||
metadata = f.metadata()
|
||||
sd = load_file(path, device="cpu")
|
||||
cfg = AdapterConfig.from_dict(json.loads(metadata["cfg"]))
|
||||
handles = attach(model, cfg, _skip_group_init=True) # creates empty params; data-driven inits restored from state_dict
|
||||
missing, unexpected = model.load_state_dict(blob["state"], strict=False)
|
||||
missing, unexpected = model.load_state_dict(sd, strict=False)
|
||||
expected_lora = {k for k in model.state_dict() if "lora_" in k}
|
||||
missing_lora = sorted(expected_lora.intersection(missing))
|
||||
if missing_lora:
|
||||
@@ -141,7 +145,7 @@ def load(model: nn.Module, path: str) -> list[RemovableHandle]:
|
||||
unexpected_lora = [k for k in unexpected if "lora_" in k]
|
||||
if unexpected_lora:
|
||||
raise RuntimeError(f"unexpected lora keys in checkpoint: {unexpected_lora}")
|
||||
saved_fp = blob.get("base_fp", {})
|
||||
saved_fp = json.loads(metadata.get("base_fp", "{}"))
|
||||
if saved_fp:
|
||||
cur_fp = _base_weight_fingerprint(model)
|
||||
diffs = [k for k in saved_fp if saved_fp[k] != cur_fp.get(k)]
|
||||
|
||||
@@ -1 +1 @@
|
||||
from . import lora, pissa, delora, ia3, dora, hra, eva, antipasto # noqa: F401 side-effect: register
|
||||
from . import lora, pissa, delora, ia3, dora, hra, eva, antipasto, road # noqa: F401 side-effect: register
|
||||
|
||||
@@ -37,7 +37,6 @@ class LoRA:
|
||||
|
||||
@staticmethod
|
||||
def init(layer: nn.Module, cfg) -> None:
|
||||
# B is zeros => delta=0 at t=0; identity invariant holds.
|
||||
return
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
"""ROAD: Rotation ADaptation. https://arxiv.org/abs/2409.00119
|
||||
|
||||
ROAD applies a learned output-space block rotation/scaling after the frozen base
|
||||
layer:
|
||||
|
||||
y' = R y = R (W x + b)
|
||||
|
||||
This matches PEFT's unmerged forward path and fits lora-lite as a simple output
|
||||
hook. We implement the three PEFT variants (`road_1`, `road_2`, `road_4`) and
|
||||
skip merge/unmerge because this library keeps adapters as hooks.
|
||||
|
||||
Refs:
|
||||
- peft: https://github.com/huggingface/peft/blob/6030f9160ed2fc17220f6f41382a66f1257b6a93/src/peft/tuners/road/layer.py
|
||||
"""
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
|
||||
import torch
|
||||
from jaxtyping import Float
|
||||
from torch import nn, Tensor as T
|
||||
|
||||
from ..config import AdapterConfig, register_config
|
||||
from ..variant import ParamSpec, register
|
||||
|
||||
RoadVariant = Literal["road_1", "road_2", "road_4"]
|
||||
|
||||
|
||||
@register_config
|
||||
@dataclass
|
||||
class RoadConfig(AdapterConfig):
|
||||
variant: str = "road"
|
||||
road_variant: RoadVariant = "road_1"
|
||||
group_size: int = 64
|
||||
|
||||
|
||||
def _road_param_size(d_out: int, road_variant: str) -> int:
|
||||
if road_variant == "road_1":
|
||||
return d_out // 2
|
||||
if road_variant == "road_2":
|
||||
return d_out
|
||||
if road_variant == "road_4":
|
||||
return d_out * 2
|
||||
raise ValueError(f"road_variant must be 'road_1', 'road_2', or 'road_4', got {road_variant!r}")
|
||||
|
||||
|
||||
def _validate_group_geometry(d_out: int, group_size: int) -> None:
|
||||
if group_size <= 0 or group_size % 2 != 0:
|
||||
raise ValueError(f"ROAD group_size must be positive and even, got {group_size}")
|
||||
if d_out % group_size != 0:
|
||||
raise ValueError(f"ROAD d_out={d_out} must be divisible by group_size={group_size}")
|
||||
|
||||
|
||||
def _prepare_cols(
|
||||
road_variant: str,
|
||||
group_size: int,
|
||||
road_theta: torch.Tensor,
|
||||
road_alpha: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if road_variant == "road_1":
|
||||
# One θ/α per pair. Reuse it for both rows of each 2D rotation block.
|
||||
road_theta = road_theta.reshape(-1, group_size // 2).repeat_interleave(2, dim=0).flatten()
|
||||
road_alpha = road_alpha.reshape(-1, group_size // 2).repeat_interleave(2, dim=0).flatten()
|
||||
first_col = road_alpha * road_theta.cos()
|
||||
second_col = road_alpha * road_theta.sin()
|
||||
elif road_variant == "road_2":
|
||||
# One θ/α per output coordinate.
|
||||
first_col = road_alpha * road_theta.cos()
|
||||
second_col = road_alpha * road_theta.sin()
|
||||
elif road_variant == "road_4":
|
||||
# Independent θ/α for the first and second column contributions.
|
||||
road_theta = road_theta.reshape(-1, 2, group_size)
|
||||
road_alpha = road_alpha.reshape(-1, 2, group_size)
|
||||
first_col = road_alpha[:, 0, :].flatten() * road_theta[:, 0, :].cos().flatten()
|
||||
second_col = road_alpha[:, 1, :].flatten() * road_theta[:, 1, :].sin().flatten()
|
||||
else:
|
||||
raise ValueError(f"road_variant must be 'road_1', 'road_2', or 'road_4', got {road_variant!r}")
|
||||
return first_col, second_col
|
||||
|
||||
|
||||
def _apply_road(
|
||||
road_variant: str,
|
||||
group_size: int,
|
||||
road_theta: torch.Tensor,
|
||||
road_alpha: torch.Tensor,
|
||||
y: Float[T, '*B o'],
|
||||
) -> Float[T, '*B o']:
|
||||
first_col, second_col = _prepare_cols(road_variant, group_size, road_theta, road_alpha)
|
||||
y_grouped = y.reshape(-1, 2, group_size // 2)
|
||||
y1 = y_grouped[:, 0, :]
|
||||
y2 = y_grouped[:, 1, :]
|
||||
rotate_half_y = torch.stack((-y2, y1), dim=1).reshape(y.shape)
|
||||
return y * first_col + rotate_half_y * second_col
|
||||
|
||||
|
||||
def _road_matrix(
|
||||
road_variant: str,
|
||||
group_size: int,
|
||||
road_theta: torch.Tensor,
|
||||
road_alpha: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Explicit PEFT merge matrix. Used for tests and small-debug inspection."""
|
||||
first_col, second_col = _prepare_cols(road_variant, group_size, road_theta, road_alpha)
|
||||
size = second_col.shape[0]
|
||||
output = torch.diag(first_col)
|
||||
swapped_second_col = second_col.reshape(-1, 2, group_size // 2)[:, [1, 0], :].flatten()
|
||||
rotated_diag_second_col = torch.diag(swapped_second_col).reshape(-1, 2, group_size // 2, size)[:, [1, 0], :, :]
|
||||
rotated_diag_second_col[:, 0, :, :] *= -1
|
||||
return output + rotated_diag_second_col.reshape(size, size)
|
||||
|
||||
|
||||
@register
|
||||
class ROAD:
|
||||
name = "road"
|
||||
|
||||
@staticmethod
|
||||
def param_specs(d_in: int, d_out: int, cfg: RoadConfig) -> dict[str, ParamSpec]:
|
||||
_validate_group_geometry(d_out, cfg.group_size)
|
||||
size = _road_param_size(d_out, cfg.road_variant)
|
||||
return {
|
||||
"lora_road_theta": ParamSpec((size,), init="zeros", trainable=True),
|
||||
"lora_road_alpha": ParamSpec((size,), init="ones", trainable=True),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def init(layer: nn.Module, cfg: RoadConfig) -> None:
|
||||
return
|
||||
|
||||
@staticmethod
|
||||
def forward(
|
||||
layer: nn.Module,
|
||||
x: Float[T, '*B i'],
|
||||
y: Float[T, '*B o'],
|
||||
) -> Float[T, '*B o']:
|
||||
del x
|
||||
cfg = layer._lora_cfg
|
||||
y_cast = y.to(layer.lora_road_theta.dtype)
|
||||
return _apply_road(cfg.road_variant, cfg.group_size, layer.lora_road_theta, layer.lora_road_alpha, y_cast)
|
||||
Reference in New Issue
Block a user