diff --git a/src/peft/tuners/lora/model.py b/src/peft/tuners/lora/model.py index f654d6c9c1..60e3ae4382 100644 --- a/src/peft/tuners/lora/model.py +++ b/src/peft/tuners/lora/model.py @@ -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 @@ -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 @@ -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 diff --git a/tests/test_custom_models.py b/tests/test_custom_models.py index 465779d4e0..80952a7b22 100644 --- a/tests/test_custom_models.py +++ b/tests/test_custom_models.py @@ -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