Skip to content
Open
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
15 changes: 6 additions & 9 deletions src/peft/tuners/lora/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@
from .gptq import dispatch_gptq
from .hqq import dispatch_hqq
from .inc import dispatch_inc
from .layer import Conv2d, LoraLayer, ParamWrapper, dispatch_default
from .layer import LoraLayer, ParamWrapper, _ConvNd, dispatch_default
from .te import dispatch_transformer_engine
from .torchao import dispatch_torchao
from .tp_layer import dispatch_megatron
Expand Down Expand Up @@ -980,13 +980,10 @@ def _svd_generalized_task_arithmetic_weighted_adapter(
else:
raise ValueError(f"Invalid value passed to combination type: {combination_type}")

conv2d = isinstance(target, Conv2d)
if conv2d:
conv2d_1x1 = target.weight.size()[2:4] == (1, 1)
if not conv2d_1x1:
delta_weight = delta_weight.flatten(start_dim=1)
else:
delta_weight = delta_weight.squeeze()
is_conv = isinstance(target, _ConvNd)
if is_conv:
# (out, in, *kernel) -> (out, in * prod(kernel)), works for Conv1d, Conv2d (incl. 1x1) and Conv3d
delta_weight = delta_weight.flatten(start_dim=1)
if (hasattr(target, "fan_in_fan_out") and target.fan_in_fan_out) or is_embedding:
delta_weight = delta_weight.T

Expand All @@ -1002,7 +999,7 @@ def _svd_generalized_task_arithmetic_weighted_adapter(
low_val = -hi_val
U = U.clamp(low_val, hi_val)
Vh = Vh.clamp(low_val, hi_val)
if conv2d:
if is_conv:
U = U.reshape(target_lora_B.data.shape)
Vh = Vh.reshape(target_lora_A.data.shape)
return Vh, U
Expand Down
40 changes: 40 additions & 0 deletions tests/test_custom_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -4670,6 +4670,46 @@ def test_add_weighted_adapter_with_different_scaling(self, weights, combination_
if max_mse is not None:
assert mse < max_mse

@pytest.mark.parametrize(
"conv_layer",
[
nn.Conv1d(10, 10, 3),
nn.Conv2d(10, 10, 3),
nn.Conv2d(10, 10, 1),
nn.Conv3d(10, 10, 3),
],
)
def test_add_weighted_adapter_svd_conv_layers(self, conv_layer):
# SVD based combination types used to flatten the delta weight only for Conv2d, so Conv1d and Conv3d failed
torch.manual_seed(0)

class ConvModel(nn.Module):
def __init__(self):
super().__init__()
self.conv = conv_layer

def forward(self, X):
return self.conv(X)

config = LoraConfig(target_modules=["conv"], r=4, init_lora_weights=False)
model = get_peft_model(ConvModel(), config, adapter_name="adapter1")
model.add_adapter("adapter2", config)
# with r1 + r2 <= svd_rank <= min(out_channels, in_channels * prod(kernel_size)) the SVD is exact and the
# merged delta weight must equal the weighted sum
model.add_weighted_adapter(
adapters=["adapter1", "adapter2"],
weights=[0.5, 0.5],
adapter_name="merged",
combination_type="svd",
svd_rank=8,
)

module = model.base_model.model.conv
expected = 0.5 * module.get_delta_weight("adapter1") + 0.5 * module.get_delta_weight("adapter2")
dw_merged = module.get_delta_weight("merged")
assert dw_merged.shape == module.base_layer.weight.shape
assert torch.allclose(dw_merged, expected, atol=1e-5, rtol=1e-5)

def test_multiple_adapters_no_needless_copy_modules_to_save(self):
# See 2206
# The problem was that we keep a "global" modules_to_save on the model which contains all possible
Expand Down
Loading