Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
100 changes: 100 additions & 0 deletions tests/test_lora_fallback_effective_config.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,100 @@
import math
from copy import deepcopy

import pytest
import torch
import torch.nn as nn

from ultralytics.utils.lora import LoRAConfig
from ultralytics.utils.lora.fallback import (
FewShotLoRAConv,
ManualLoRAConv,
_collect_fallback_adapter_state,
_load_fallback_adapter_state,
_merge_fallback_modules,
apply_manual_lora,
)


def _manual_layer(*, use_rslora: bool) -> ManualLoRAConv:
return ManualLoRAConv(nn.Conv2d(4, 4, 1), r=8, alpha=16, use_rslora=use_rslora)


def test_fallback_honors_requested_rslora_scaling():
assert _manual_layer(use_rslora=True).scaling == pytest.approx(16 / math.sqrt(8))
assert _manual_layer(use_rslora=False).scaling == pytest.approx(16 / 8)
assert FewShotLoRAConv(nn.Conv2d(4, 4, 1), r=8, alpha=16, use_rslora=True).scaling == pytest.approx(
16 / math.sqrt(8)
)


def test_fallback_adapter_round_trip_preserves_effective_rslora(tmp_path):
base = nn.Sequential(nn.Conv2d(4, 4, 1))
restored_base = deepcopy(base)
source = apply_manual_lora(
base,
LoRAConfig(
r=8,
alpha=16,
dropout=0.0,
backend="fallback",
target_modules=["0"],
skip_stem=False,
use_rslora=True,
),
)
source[0].lora_B.data.normal_()
source.eval()
assert source.lora_runtime_metadata["requested_use_rslora"] is True
assert source.lora_runtime_metadata["effective_use_rslora"] is True
sample = torch.randn(2, 4, 3, 3)
expected = source(sample)
saved = _collect_fallback_adapter_state(source)
assert saved["modules"]["0"]["use_rslora"] is True
torch.save(saved, tmp_path / "fallback_adapter.pt")
payload = {"backend": "fallback", "weight_file": "fallback_adapter.pt"}

restored = _load_fallback_adapter_state(restored_base, tmp_path, payload)
restored.eval()

assert restored[0].use_rslora is True
assert restored[0].scaling == pytest.approx(16 / math.sqrt(8))
torch.testing.assert_close(restored(sample), expected)
assert _merge_fallback_modules(restored) == 1
torch.testing.assert_close(restored(sample), expected)


def test_few_shot_round_trip_preserves_dropconnect_and_adaptive_rank(tmp_path):
source = nn.Sequential(
FewShotLoRAConv(
nn.Conv2d(4, 4, 1),
r=8,
alpha=16,
dropconnect=0.27,
adaptive_rank=False,
)
)
saved = _collect_fallback_adapter_state(source)
torch.save(saved, tmp_path / "fallback_adapter.pt")
payload = {"backend": "fallback", "weight_file": "fallback_adapter.pt"}

restored = _load_fallback_adapter_state(nn.Sequential(nn.Conv2d(4, 4, 1)), tmp_path, payload)

assert isinstance(restored[0], FewShotLoRAConv)
assert restored[0].dropconnect_rate == pytest.approx(0.27)
assert restored[0].adaptive_rank is False
assert not hasattr(restored[0], "rank_mask")


def test_legacy_fallback_adapter_without_scaling_mode_keeps_lora_scaling(tmp_path):
source = nn.Sequential(_manual_layer(use_rslora=False))
saved = _collect_fallback_adapter_state(source)
for module_config in saved["modules"].values():
module_config.pop("use_rslora", None)
torch.save(saved, tmp_path / "fallback_adapter.pt")
payload = {"backend": "fallback", "weight_file": "fallback_adapter.pt"}

restored = _load_fallback_adapter_state(nn.Sequential(nn.Conv2d(4, 4, 1)), tmp_path, payload)

assert restored[0].use_rslora is False
assert restored[0].scaling == pytest.approx(16 / 8)
43 changes: 38 additions & 5 deletions ultralytics/utils/lora/fallback.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,14 @@
if TYPE_CHECKING:
from .config import LoRAConfig


def _fallback_lora_scaling(alpha: int, r: int, use_rslora: bool) -> float:
"""Return the effective fallback LoRA scale; direct wrappers default to historical standard LoRA."""
if r <= 0:
raise ValueError(f"Fallback LoRA rank must be positive, got r={r}.")
return alpha / math.sqrt(r) if use_rslora else alpha / r


class FewShotLoRAConv(nn.Module):
"""LoRA wrapper optimized for few-shot learning.

Expand All @@ -49,12 +57,14 @@ def __init__(self, conv: nn.Conv2d, r: int = 8, alpha: int = 16,
dropconnect_min: float = 0.0,
gradient_importance_weighted: bool = False,
variational_rank: bool = False,
rank_budget: float = 0.5):
rank_budget: float = 0.5,
use_rslora: bool = False):
super().__init__()
self.conv = conv
self.r = r
self.alpha = alpha
self.scaling = alpha / max(r, 1)
self.use_rslora = bool(use_rslora)
self.scaling = _fallback_lora_scaling(alpha, r, self.use_rslora)
self.lora_dropout = nn.Dropout(dropout) if dropout > 0 else None
self.dropconnect_rate = dropconnect
self.adaptive_rank = adaptive_rank
Expand Down Expand Up @@ -299,12 +309,20 @@ class ManualLoRAConv(nn.Module):
MUST be divisible by `groups`.
"""

def __init__(self, conv: nn.Conv2d, r: int = 8, alpha: int = 16, dropout: float = 0.0):
def __init__(
self,
conv: nn.Conv2d,
r: int = 8,
alpha: int = 16,
dropout: float = 0.0,
use_rslora: bool = False,
):
super().__init__()
self.conv = conv
self.r = r
self.alpha = alpha
self.scaling = alpha / max(r, 1)
self.use_rslora = bool(use_rslora)
self.scaling = _fallback_lora_scaling(alpha, r, self.use_rslora)
self.lora_dropout = nn.Dropout(dropout) if dropout > 0 else None

groups = conv.groups
Expand Down Expand Up @@ -534,6 +552,7 @@ def _replace_conv_with_manual_lora(module: nn.Module, config: "LoRAConfig", pref
lora_kwargs = {
"r": r, "alpha": max(r * 2, config.alpha),
"dropout": config.dropout,
"use_rslora": bool(getattr(config, "use_rslora", False)),
"dropconnect": getattr(config, "few_shot_dropconnect", 0.1),
"adaptive_rank": getattr(config, "few_shot_adaptive_rank", True),
"dropconnect_schedule": getattr(config, "few_shot_dropconnect_schedule", "constant"),
Expand All @@ -545,7 +564,12 @@ def _replace_conv_with_manual_lora(module: nn.Module, config: "LoRAConfig", pref
}
else:
lora_cls = ManualLoRAConv
lora_kwargs = {"r": r, "alpha": max(r * 2, config.alpha), "dropout": config.dropout}
lora_kwargs = {
"r": r,
"alpha": max(r * 2, config.alpha),
"dropout": config.dropout,
"use_rslora": bool(getattr(config, "use_rslora", False)),
}
setattr(module, name, lora_cls(child, **lora_kwargs))
replaced += 1
continue
Expand Down Expand Up @@ -628,6 +652,8 @@ def apply_manual_lora(model: nn.Module, config: "LoRAConfig", include_head: bool
effective_variant="lora",
requested_init_lora_weights=config.init_lora_weights,
effective_init_lora_weights=config.init_lora_weights,
requested_use_rslora=bool(getattr(config, "use_rslora", False)),
effective_use_rslora=bool(getattr(config, "use_rslora", False)),
include_head=include_head,
freeze_bn=bool(getattr(config, "freeze_bn", False)),
target_modules=model.lora_target_modules,
Expand Down Expand Up @@ -677,13 +703,16 @@ def _collect_fallback_adapter_state(model: nn.Module) -> Dict[str, Any]:
"r": int(module.r),
"alpha": int(module.alpha),
"dropout": float(module.lora_dropout.p if module.lora_dropout is not None else 0.0),
"use_rslora": bool(getattr(module, "use_rslora", False)),
}
state[name] = {
"lora_A": module.lora_A.detach().cpu(),
"lora_B": module.lora_B.detach().cpu(),
}
if isinstance(module, FewShotLoRAConv):
modules[name]["few_shot"] = True
modules[name]["dropconnect"] = float(module.dropconnect_rate)
modules[name]["adaptive_rank"] = bool(module.adaptive_rank)
modules[name]["dropconnect_schedule"] = getattr(module, "dropconnect_schedule", "constant")
modules[name]["dropconnect_max"] = getattr(module, "dropconnect_max", 0.3)
modules[name]["dropconnect_min"] = getattr(module, "dropconnect_min", 0.0)
Expand Down Expand Up @@ -711,6 +740,7 @@ def _load_fallback_adapter_state(model: nn.Module, path: Path, payload: Dict[str

for module_name, config in module_configs.items():
original = _get_module_by_name(target_root, module_name)
use_rslora = bool(config.get("use_rslora", False))
if isinstance(original, (ManualLoRAConv, FewShotLoRAConv)):
wrapped = original
else:
Expand All @@ -722,6 +752,7 @@ def _load_fallback_adapter_state(model: nn.Module, path: Path, payload: Dict[str
"r": int(config.get("r", 0)),
"alpha": int(config.get("alpha", 0)),
"dropout": float(config.get("dropout", 0.0)),
"use_rslora": use_rslora,
}
if is_few_shot:
lora_kwargs["dropconnect"] = float(config.get("dropconnect", 0.1))
Expand All @@ -735,6 +766,8 @@ def _load_fallback_adapter_state(model: nn.Module, path: Path, payload: Dict[str
wrapped = lora_cls(original, **lora_kwargs)
_set_module_by_name(target_root, module_name, wrapped)

wrapped.use_rslora = use_rslora
wrapped.scaling = _fallback_lora_scaling(wrapped.alpha, wrapped.r, use_rslora)
params = module_state.get(module_name, {})
wrapped.lora_A.data.copy_(params["lora_A"])
wrapped.lora_B.data.copy_(params["lora_B"])
Expand Down
Loading