simpler test

This commit is contained in:
wassname
2026-04-27 09:47:07 +08:00
parent b60a8c3f9b
commit 24ba8deb02
10 changed files with 566 additions and 41 deletions
+2
View File
@@ -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",
+13 -9
View File
@@ -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
View File
@@ -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
-1
View File
@@ -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
+137
View File
@@ -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)