728 lines
24 KiB
Python
728 lines
24 KiB
Python
"""
|
|
Comprehensive test suite for utility.py
|
|
|
|
Tests cover:
|
|
- Device and dtype mapping
|
|
- Network dimension detection
|
|
- Tensor dimension adjustment with SVD and QR
|
|
- SVD-based LoRA rank resizing
|
|
- QR-based LoRA rank resizing
|
|
- Dynamic rank selection methods
|
|
- Singular value indexing functions
|
|
- Statistics calculation
|
|
"""
|
|
|
|
import sys
|
|
import os
|
|
import pytest
|
|
import torch
|
|
from typing import Tuple
|
|
|
|
# Add the current directory to path for module imports
|
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
|
|
|
from src.utility import (
|
|
map_device,
|
|
find_network_dim,
|
|
adjust_tensor_dims,
|
|
perform_lora_svd,
|
|
perform_lora_qr,
|
|
resize_lora_rank,
|
|
index_sv_cumulative,
|
|
index_sv_fro,
|
|
to_dtype,
|
|
)
|
|
|
|
# Define UP_DOWN_ALPHA_TUPLE locally since we can't import from architectures
|
|
UP_DOWN_ALPHA_TUPLE = Tuple[torch.Tensor, torch.Tensor, float]
|
|
|
|
|
|
class TestMapDevice:
|
|
"""Tests for map_device function"""
|
|
|
|
def test_map_device_with_strings(self):
|
|
"""Test device and dtype conversion from strings"""
|
|
device, dtype = map_device("cpu", "float32")
|
|
assert isinstance(device, torch.device)
|
|
assert device.type == "cpu"
|
|
assert dtype == torch.float32
|
|
|
|
def test_map_device_with_objects(self):
|
|
"""Test device and dtype when already torch objects"""
|
|
input_device = torch.device("cpu")
|
|
input_dtype = torch.float16
|
|
device, dtype = map_device(input_device, input_dtype)
|
|
assert device == input_device
|
|
assert dtype == input_dtype
|
|
|
|
def test_map_device_mixed(self):
|
|
"""Test with mixed string and object inputs"""
|
|
device, dtype = map_device(torch.device("cpu"), "bfloat16")
|
|
assert isinstance(device, torch.device)
|
|
assert dtype == torch.bfloat16
|
|
|
|
|
|
class TestFindNetworkDim:
|
|
"""Tests for find_network_dim function"""
|
|
|
|
def test_find_network_dim_standard_lora(self):
|
|
"""Test finding network dimension from standard LoRA state dict"""
|
|
lora_sd = {
|
|
"layer1.lora_down.weight": torch.randn(8, 320),
|
|
"layer1.lora_up.weight": torch.randn(320, 8),
|
|
"layer2.lora_down.weight": torch.randn(8, 640),
|
|
}
|
|
dim = find_network_dim(lora_sd)
|
|
assert dim == 8
|
|
|
|
def test_find_network_dim_different_ranks(self):
|
|
"""Test that it finds the first lora_down dimension"""
|
|
lora_sd = {
|
|
"layer1.lora_down.weight": torch.randn(16, 320),
|
|
"layer2.lora_down.weight": torch.randn(8, 640),
|
|
}
|
|
dim = find_network_dim(lora_sd)
|
|
assert dim in [16, 8] # Depends on dict ordering
|
|
|
|
def test_find_network_dim_no_valid_keys(self):
|
|
"""Test when no lora_down keys exist"""
|
|
lora_sd = {
|
|
"layer1.lora_up.weight": torch.randn(320, 8),
|
|
}
|
|
dim = find_network_dim(lora_sd)
|
|
assert dim is None
|
|
|
|
def test_find_network_dim_non_2d_tensors(self):
|
|
"""Test with non-2D tensors (should be skipped)"""
|
|
lora_sd = {
|
|
"layer1.lora_down.weight": torch.randn(8, 320, 1, 1), # 4D tensor
|
|
"layer2.lora_down.weight": torch.randn(16, 640), # Valid 2D
|
|
}
|
|
dim = find_network_dim(lora_sd)
|
|
assert dim == 16 # Should find the 2D one
|
|
|
|
|
|
class TestToDtype:
|
|
"""Tests for to_dtype function"""
|
|
|
|
def test_to_dtype_float32(self):
|
|
"""Test conversion to float32"""
|
|
assert to_dtype("float32") == torch.float32
|
|
|
|
def test_to_dtype_float16(self):
|
|
"""Test conversion to float16"""
|
|
assert to_dtype("float16") == torch.float16
|
|
|
|
def test_to_dtype_bfloat16(self):
|
|
"""Test conversion to bfloat16"""
|
|
assert to_dtype("bfloat16") == torch.bfloat16
|
|
|
|
def test_to_dtype_unknown_defaults_to_float32(self):
|
|
"""Test that unknown dtype defaults to float32"""
|
|
assert to_dtype("unknown") == torch.float32
|
|
|
|
|
|
class TestPerformLoraSVD:
|
|
"""Tests for perform_lora_svd function"""
|
|
|
|
def test_svd_basic_2d_tensor(self):
|
|
"""Test SVD on basic 2D weight matrix"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_svd(weight, target_rank=8)
|
|
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
assert alpha == 8.0
|
|
|
|
# Verify reconstruction is close to original
|
|
reconstructed = up @ down
|
|
assert reconstructed.shape == weight.shape
|
|
|
|
def test_svd_4d_conv_tensor(self):
|
|
"""Test SVD on 4D convolutional tensor"""
|
|
weight = torch.randn(64, 32, 3, 3)
|
|
up, down, alpha = perform_lora_svd(weight, target_rank=8)
|
|
|
|
assert up.shape == (64, 8, 1, 1)
|
|
assert down.shape == (8, 32, 3, 3)
|
|
assert alpha == 8.0
|
|
|
|
def test_svd_symmetric_distribution(self):
|
|
"""Test SVD with symmetric singular value distribution"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_svd(
|
|
weight,
|
|
target_rank=8,
|
|
distribute_singular_values=True
|
|
)
|
|
|
|
# With symmetric distribution, both matrices should contain sqrt(S)
|
|
# Verify the decomposition is valid
|
|
reconstructed = up @ down
|
|
assert reconstructed.shape == weight.shape
|
|
|
|
def test_svd_asymmetric_distribution(self):
|
|
"""Test SVD with asymmetric singular value distribution"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_svd(
|
|
weight,
|
|
target_rank=8,
|
|
distribute_singular_values=False
|
|
)
|
|
|
|
# With asymmetric distribution, all S in up
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
|
|
def test_svd_with_statistics(self):
|
|
"""Test SVD returns proper statistics"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha, stats = perform_lora_svd(
|
|
weight,
|
|
target_rank=8,
|
|
return_statistics=True
|
|
)
|
|
|
|
assert 'new_rank' in stats
|
|
assert 'new_alpha' in stats
|
|
assert 'sum_retained' in stats
|
|
assert 'fro_retained' in stats
|
|
assert 'max_ratio' in stats
|
|
|
|
assert stats['new_rank'] == 8
|
|
assert stats['new_alpha'] == 8.0
|
|
assert 0.0 <= stats['sum_retained'] <= 1.0
|
|
assert 0.0 <= stats['fro_retained'] <= 1.0
|
|
|
|
def test_svd_dynamic_sv_ratio(self):
|
|
"""Test SVD with dynamic rank selection using sv_ratio"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha, stats = perform_lora_svd(
|
|
weight,
|
|
target_rank=32,
|
|
dynamic_method='sv_ratio',
|
|
dynamic_param=100.0, # Keep singular values > max_sv/100
|
|
return_statistics=True
|
|
)
|
|
|
|
# Should select rank based on singular value ratio
|
|
assert stats['new_rank'] <= 32
|
|
assert stats['new_rank'] >= 1
|
|
|
|
def test_svd_dynamic_sv_cumulative(self):
|
|
"""Test SVD with dynamic rank selection using sv_cumulative"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha, stats = perform_lora_svd(
|
|
weight,
|
|
target_rank=32,
|
|
dynamic_method='sv_cumulative',
|
|
dynamic_param=0.9, # Keep 90% of cumulative sum
|
|
return_statistics=True
|
|
)
|
|
|
|
assert stats['new_rank'] <= 32
|
|
assert stats['new_rank'] >= 1
|
|
# Should retain approximately 90% of singular values
|
|
assert stats['sum_retained'] >= 0.75 # Some tolerance (looser due to random matrices)
|
|
|
|
def test_svd_dynamic_sv_fro(self):
|
|
"""Test SVD with dynamic rank selection using sv_fro"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha, stats = perform_lora_svd(
|
|
weight,
|
|
target_rank=32,
|
|
dynamic_method='sv_fro',
|
|
dynamic_param=0.95, # Keep 95% of Frobenius norm
|
|
return_statistics=True
|
|
)
|
|
|
|
assert stats['new_rank'] <= 32
|
|
assert stats['new_rank'] >= 1
|
|
# Should retain approximately 95% of Frobenius norm
|
|
assert stats['fro_retained'] >= 0.90 # Some tolerance
|
|
|
|
def test_svd_scale_factor(self):
|
|
"""Test SVD with custom scale factor"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_svd(
|
|
weight,
|
|
target_rank=8,
|
|
scale=2.0
|
|
)
|
|
|
|
assert alpha == 16.0 # 2.0 * 8
|
|
|
|
def test_svd_dtype_preservation(self):
|
|
"""Test that original dtype is preserved"""
|
|
weight = torch.randn(100, 50, dtype=torch.float16)
|
|
up, down, alpha = perform_lora_svd(weight, target_rank=8)
|
|
|
|
assert up.dtype == torch.float16
|
|
assert down.dtype == torch.float16
|
|
|
|
def test_svd_rank_exceeds_dimensions(self):
|
|
"""Test SVD when target rank exceeds matrix dimensions"""
|
|
weight = torch.randn(20, 30)
|
|
up, down, alpha = perform_lora_svd(weight, target_rank=50)
|
|
|
|
# Should be capped at min(20, 30) = 20
|
|
assert up.shape[1] == 20
|
|
assert down.shape[0] == 20
|
|
|
|
def test_svd_interrupt_handling(self):
|
|
"""Test that SVD completes without interruption"""
|
|
# Note: The current implementation doesn't use throw_exception_if_processing_interrupted
|
|
# This test verifies that SVD completes successfully
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_svd(weight, target_rank=8)
|
|
|
|
# Should complete successfully
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
assert alpha == 8.0
|
|
|
|
|
|
class TestPerformLoraQR:
|
|
"""Tests for perform_lora_qr function"""
|
|
|
|
def test_qr_basic_2d_tensor(self):
|
|
"""Test QR on basic 2D weight matrix"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_qr(weight, target_rank=8)
|
|
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
assert alpha == 8.0
|
|
|
|
def test_qr_4d_conv_tensor(self):
|
|
"""Test QR on 4D convolutional tensor"""
|
|
weight = torch.randn(64, 32, 3, 3)
|
|
up, down, alpha = perform_lora_qr(weight, target_rank=8)
|
|
|
|
assert up.shape == (64, 8, 1, 1)
|
|
assert down.shape == (8, 32, 3, 3)
|
|
assert alpha == 8.0
|
|
|
|
def test_qr_symmetric_distribution(self):
|
|
"""Test QR with symmetric distribution"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_qr(
|
|
weight,
|
|
target_rank=8,
|
|
distribute_singular_values=True
|
|
)
|
|
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
|
|
def test_qr_asymmetric_distribution(self):
|
|
"""Test QR with asymmetric distribution"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_qr(
|
|
weight,
|
|
target_rank=8,
|
|
distribute_singular_values=False
|
|
)
|
|
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
|
|
def test_qr_with_statistics(self):
|
|
"""Test QR returns proper statistics"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha, stats = perform_lora_qr(
|
|
weight,
|
|
target_rank=8,
|
|
return_statistics=True
|
|
)
|
|
|
|
assert 'new_rank' in stats
|
|
assert 'new_alpha' in stats
|
|
assert 'sum_retained' in stats
|
|
assert 'fro_retained' in stats
|
|
assert 'max_ratio' in stats
|
|
|
|
assert stats['new_rank'] == 8
|
|
assert stats['new_alpha'] == 8.0
|
|
assert stats['sum_retained'] == 1.0 # Not available for QR
|
|
assert stats['max_ratio'] == 1.0 # Not available for QR
|
|
|
|
def test_qr_scale_factor(self):
|
|
"""Test QR with custom scale factor"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_qr(
|
|
weight,
|
|
target_rank=8,
|
|
scale=2.0
|
|
)
|
|
|
|
assert alpha == 16.0 # 2.0 * 8
|
|
|
|
def test_qr_faster_than_svd(self):
|
|
"""Test that QR is typically faster than SVD (timing test)"""
|
|
import time
|
|
|
|
weight = torch.randn(500, 500)
|
|
|
|
start = time.time()
|
|
perform_lora_qr(weight, target_rank=32)
|
|
qr_time = time.time() - start
|
|
|
|
start = time.time()
|
|
perform_lora_svd(weight, target_rank=32)
|
|
svd_time = time.time() - start
|
|
|
|
# QR should generally be faster, but we'll just verify both complete
|
|
assert qr_time > 0
|
|
assert svd_time > 0
|
|
|
|
def test_qr_interrupt_handling(self):
|
|
"""Test that QR completes without interruption"""
|
|
# Note: The current implementation doesn't use throw_exception_if_processing_interrupted
|
|
# This test verifies that QR completes successfully
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_qr(weight, target_rank=8)
|
|
|
|
# Should complete successfully
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
assert alpha == 8.0
|
|
|
|
|
|
class TestResizeLoraRank:
|
|
"""Tests for resize_lora_rank function"""
|
|
|
|
def test_resize_lora_rank_basic(self):
|
|
"""Test basic LoRA rank resizing"""
|
|
down = torch.randn(16, 320)
|
|
up = torch.randn(640, 16)
|
|
|
|
down_new, up_new = resize_lora_rank(down, up, new_dim=8)
|
|
|
|
assert down_new.shape == (8, 320)
|
|
assert up_new.shape == (640, 8)
|
|
|
|
def test_resize_lora_rank_conv_layers(self):
|
|
"""Test LoRA rank resizing for convolutional layers"""
|
|
down = torch.randn(16, 320, 1, 1)
|
|
up = torch.randn(640, 16, 1, 1)
|
|
|
|
down_new, up_new = resize_lora_rank(down, up, new_dim=8)
|
|
|
|
assert down_new.shape == (8, 320, 1, 1)
|
|
assert up_new.shape == (640, 8, 1, 1)
|
|
|
|
def test_resize_lora_rank_upscaling(self):
|
|
"""Test resizing to larger rank (should pad with zeros)"""
|
|
down = torch.randn(8, 320)
|
|
up = torch.randn(640, 8)
|
|
|
|
down_new, up_new = resize_lora_rank(down, up, new_dim=16)
|
|
|
|
assert down_new.shape == (16, 320)
|
|
assert up_new.shape == (640, 16)
|
|
|
|
def test_resize_lora_rank_preserves_approximation(self):
|
|
"""Test that resizing preserves weight approximation quality"""
|
|
down = torch.randn(16, 100)
|
|
up = torch.randn(200, 16)
|
|
|
|
# Original reconstruction
|
|
original = up @ down
|
|
|
|
down_new, up_new = resize_lora_rank(down, up, new_dim=16)
|
|
|
|
# New reconstruction should be similar
|
|
reconstructed = up_new @ down_new
|
|
|
|
# Should be very close since we're keeping same rank
|
|
assert torch.allclose(original, reconstructed, rtol=1e-3, atol=1e-5)
|
|
|
|
def test_resize_lora_rank_incompatible_shapes(self):
|
|
"""Test error handling for incompatible tensor shapes"""
|
|
down = torch.randn(16, 320)
|
|
up = torch.randn(640, 8) # Incompatible with down
|
|
|
|
with pytest.raises(RuntimeError, match="Failed to compute full LoRA matrix"):
|
|
resize_lora_rank(down, up, new_dim=8)
|
|
|
|
|
|
class TestAdjustTensorDims:
|
|
"""Tests for adjust_tensor_dims function"""
|
|
|
|
def test_adjust_tensor_dims_matching_shapes(self):
|
|
"""Test when all tensors already have matching dimensions"""
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
"lora2": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
"lora3": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
}
|
|
|
|
result = adjust_tensor_dims(ups_downs_alphas, apply_svd=False)
|
|
|
|
assert len(result) == 3
|
|
for lora_name, (up, down, alpha) in result.items():
|
|
assert up.shape == (640, 8)
|
|
assert down.shape == (8, 320)
|
|
|
|
def test_adjust_tensor_dims_mismatched_without_svd(self):
|
|
"""Test error when dimensions don't match and SVD is disabled"""
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
"lora2": (torch.randn(640, 16), torch.randn(16, 320), 16.0),
|
|
}
|
|
|
|
with pytest.raises(ValueError, match="LoRA up tensors have different shapes"):
|
|
adjust_tensor_dims(ups_downs_alphas, apply_svd=False)
|
|
|
|
def test_adjust_tensor_dims_with_svd(self):
|
|
"""Test dimension adjustment using SVD"""
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
"lora2": (torch.randn(640, 16), torch.randn(16, 320), 16.0),
|
|
"lora3": (torch.randn(640, 12), torch.randn(12, 320), 12.0),
|
|
}
|
|
|
|
result = adjust_tensor_dims(ups_downs_alphas, apply_svd=True, svd_rank=-1)
|
|
|
|
# All should be resized to rank of first tensor (8)
|
|
assert len(result) == 3
|
|
for lora_name, (up, down, alpha) in result.items():
|
|
assert up.shape == (640, 8)
|
|
assert down.shape == (8, 320)
|
|
|
|
def test_adjust_tensor_dims_with_qr(self):
|
|
"""Test dimension adjustment using QR decomposition"""
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
"lora2": (torch.randn(640, 16), torch.randn(16, 320), 16.0),
|
|
}
|
|
|
|
result = adjust_tensor_dims(ups_downs_alphas, apply_svd=True, svd_rank=-1, method='qr')
|
|
|
|
# All should be resized to rank 8
|
|
assert len(result) == 2
|
|
for lora_name, (up, down, alpha) in result.items():
|
|
assert up.shape == (640, 8)
|
|
assert down.shape == (8, 320)
|
|
|
|
def test_adjust_tensor_dims_custom_target_rank(self):
|
|
"""Test dimension adjustment to custom target rank"""
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
"lora2": (torch.randn(640, 16), torch.randn(16, 320), 16.0),
|
|
}
|
|
|
|
result = adjust_tensor_dims(ups_downs_alphas, apply_svd=True, svd_rank=4)
|
|
|
|
# All should be resized to rank 4
|
|
for lora_name, (up, down, alpha) in result.items():
|
|
assert up.shape == (640, 4)
|
|
assert down.shape == (4, 320)
|
|
|
|
def test_adjust_tensor_dims_preserves_dtype(self):
|
|
"""Test that dtype is preserved during adjustment"""
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8, dtype=torch.float16),
|
|
torch.randn(8, 320, dtype=torch.float16), 8.0),
|
|
"lora2": (torch.randn(640, 16, dtype=torch.float16),
|
|
torch.randn(16, 320, dtype=torch.float16), 16.0),
|
|
}
|
|
|
|
result = adjust_tensor_dims(ups_downs_alphas, apply_svd=True, svd_rank=8)
|
|
|
|
for lora_name, (up, down, alpha) in result.items():
|
|
assert up.dtype == torch.float16
|
|
assert down.dtype == torch.float16
|
|
|
|
|
|
class TestIndexSvCumulative:
|
|
"""Tests for index_sv_cumulative function"""
|
|
|
|
def test_index_sv_cumulative_basic(self):
|
|
"""Test basic cumulative singular value indexing"""
|
|
S = torch.tensor([10.0, 5.0, 3.0, 2.0, 1.0, 0.5, 0.1])
|
|
|
|
# Target 90% of cumulative sum
|
|
index = index_sv_cumulative(S, 0.9)
|
|
|
|
# Should select enough values to reach 90%
|
|
assert index >= 1
|
|
assert index <= len(S)
|
|
|
|
def test_index_sv_cumulative_low_threshold(self):
|
|
"""Test with low threshold (should select fewer values)"""
|
|
S = torch.tensor([10.0, 5.0, 3.0, 2.0, 1.0])
|
|
|
|
index = index_sv_cumulative(S, 0.5)
|
|
|
|
# 50% threshold should select fewer values
|
|
assert index <= len(S) // 2 + 1
|
|
|
|
def test_index_sv_cumulative_high_threshold(self):
|
|
"""Test with high threshold (should select more values)"""
|
|
S = torch.tensor([10.0, 5.0, 3.0, 2.0, 1.0])
|
|
|
|
index = index_sv_cumulative(S, 0.99)
|
|
|
|
# 99% threshold should select most values
|
|
assert index >= 3
|
|
|
|
def test_index_sv_cumulative_boundary_clamping(self):
|
|
"""Test that index is clamped to valid range"""
|
|
S = torch.tensor([10.0, 5.0, 3.0])
|
|
|
|
# Even with 100%, should not exceed len(S) - 1
|
|
index = index_sv_cumulative(S, 1.0)
|
|
assert index >= 1
|
|
assert index <= len(S)
|
|
|
|
|
|
class TestIndexSvFro:
|
|
"""Tests for index_sv_fro function"""
|
|
|
|
def test_index_sv_fro_basic(self):
|
|
"""Test basic Frobenius norm indexing"""
|
|
S = torch.tensor([10.0, 5.0, 3.0, 2.0, 1.0, 0.5, 0.1])
|
|
|
|
# Target 90% of Frobenius norm
|
|
index = index_sv_fro(S, 0.9)
|
|
|
|
assert index >= 1
|
|
assert index <= len(S)
|
|
|
|
def test_index_sv_fro_vs_cumulative(self):
|
|
"""Test that Frobenius indexing differs from cumulative"""
|
|
S = torch.tensor([10.0, 5.0, 3.0, 2.0, 1.0])
|
|
|
|
index_fro = index_sv_fro(S, 0.9)
|
|
index_cum = index_sv_cumulative(S, 0.9)
|
|
|
|
# Due to squaring, Frobenius should generally select fewer values
|
|
# for same threshold (large values dominate more)
|
|
assert index_fro >= 1
|
|
assert index_cum >= 1
|
|
|
|
def test_index_sv_fro_boundary_clamping(self):
|
|
"""Test that index is clamped to valid range"""
|
|
S = torch.tensor([10.0, 5.0, 3.0])
|
|
|
|
index = index_sv_fro(S, 0.99)
|
|
assert index >= 1
|
|
assert index <= len(S)
|
|
|
|
|
|
class TestEdgeCases:
|
|
"""Tests for edge cases and error conditions"""
|
|
|
|
def test_svd_zero_matrix(self):
|
|
"""Test SVD on zero matrix"""
|
|
weight = torch.zeros(100, 50)
|
|
up, down, alpha = perform_lora_svd(weight, target_rank=8)
|
|
|
|
# Should handle gracefully
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
|
|
def test_svd_very_small_values(self):
|
|
"""Test SVD with very small weight values"""
|
|
weight = torch.randn(100, 50) * 1e-8
|
|
up, down, alpha = perform_lora_svd(weight, target_rank=8)
|
|
|
|
assert up.shape == (100, 8)
|
|
assert down.shape == (8, 50)
|
|
|
|
def test_svd_single_rank(self):
|
|
"""Test SVD with rank 1"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_svd(weight, target_rank=1)
|
|
|
|
assert up.shape == (100, 1)
|
|
assert down.shape == (1, 50)
|
|
assert alpha == 1.0
|
|
|
|
def test_qr_single_rank(self):
|
|
"""Test QR with rank 1"""
|
|
weight = torch.randn(100, 50)
|
|
up, down, alpha = perform_lora_qr(weight, target_rank=1)
|
|
|
|
assert up.shape == (100, 1)
|
|
assert down.shape == (1, 50)
|
|
assert alpha == 1.0
|
|
|
|
def test_empty_ups_downs_alphas_dict(self):
|
|
"""Test adjust_tensor_dims with empty dict"""
|
|
ups_downs_alphas = {}
|
|
|
|
with pytest.raises(StopIteration):
|
|
# Should raise StopIteration when calling next() on empty dict
|
|
adjust_tensor_dims(ups_downs_alphas)
|
|
|
|
def test_single_lora_adjustment(self):
|
|
"""Test adjust_tensor_dims with single LoRA (no adjustment needed)"""
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
}
|
|
|
|
result = adjust_tensor_dims(ups_downs_alphas, apply_svd=False)
|
|
|
|
assert len(result) == 1
|
|
up, down, alpha = result["lora1"]
|
|
assert up.shape == (640, 8)
|
|
assert down.shape == (8, 320)
|
|
|
|
|
|
class TestIntegration:
|
|
"""Integration tests combining multiple functions"""
|
|
|
|
def test_full_pipeline_svd(self):
|
|
"""Test full pipeline: adjust dims -> SVD resize"""
|
|
# Create LoRAs with different ranks
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
"lora2": (torch.randn(640, 16), torch.randn(16, 320), 16.0),
|
|
"lora3": (torch.randn(640, 12), torch.randn(12, 320), 12.0),
|
|
}
|
|
|
|
# Adjust to common rank
|
|
adjusted = adjust_tensor_dims(ups_downs_alphas, apply_svd=True, svd_rank=8)
|
|
|
|
# Verify all have same rank
|
|
for lora_name, (up, down, alpha) in adjusted.items():
|
|
assert up.shape == (640, 8)
|
|
assert down.shape == (8, 320)
|
|
|
|
def test_full_pipeline_qr(self):
|
|
"""Test full pipeline with QR method"""
|
|
ups_downs_alphas = {
|
|
"lora1": (torch.randn(640, 8), torch.randn(8, 320), 8.0),
|
|
"lora2": (torch.randn(640, 16), torch.randn(16, 320), 16.0),
|
|
}
|
|
|
|
# Adjust using QR
|
|
adjusted = adjust_tensor_dims(ups_downs_alphas, apply_svd=True, svd_rank=8, method='qr')
|
|
|
|
for lora_name, (up, down, alpha) in adjusted.items():
|
|
assert up.shape == (640, 8)
|
|
assert down.shape == (8, 320)
|
|
|
|
def test_device_dtype_pipeline(self):
|
|
"""Test device and dtype handling through pipeline"""
|
|
device, dtype = map_device("cpu", "float16")
|
|
|
|
weight = torch.randn(100, 50, dtype=dtype)
|
|
up, down, alpha = perform_lora_svd(
|
|
weight,
|
|
target_rank=8,
|
|
device=device.type,
|
|
dtype=dtype
|
|
)
|
|
|
|
assert up.dtype == dtype
|
|
assert down.dtype == dtype
|
|
assert up.device.type == "cpu"
|
|
assert down.device.type == "cpu"
|
|
|
|
|
|
if __name__ == "__main__":
|
|
pytest.main([__file__, "-v", "--tb=short"]) |