From 1ca216dcbb01d9cc6e455511a2eef1b230f921da Mon Sep 17 00:00:00 2001 From: larsupb Date: Sun, 2 Aug 2026 11:13:13 +0200 Subject: [PATCH] feat(merge): delta-space interpolation, VRAM offload, and OOM-safe kernels Extends the delta-space merge path beyond the GTA family and hardens the large-layer code paths that OOM'd on 8 GB cards. - Merger node: route slerp/nuslerp/karcher/nearswap and SCE through the delta-space path (reconstruct each LoRA's full delta, merge, refactor), instead of merging up/down factors separately and injecting meaningless up_i @ down_j cross-terms. Serialize these on CUDA like GTA. - Merger node: `offload_models` toggle (default on) evicts resident DIT/CLIP/VAE from VRAM before a CUDA merge; ComfyUI reloads them lazily. - gta: add memory-frugal `karcher_delta_merge` and `sce_delta_merge` that consume the delta list in place instead of stacking [N, out, in]; the stock mergekit paths make ~3N full-tensor copies and OOM large layers. - gta: pick magnitude / magnitude_outliers / SCE-select thresholds from a GPU histogram above 16M elements, avoiding the int64 argsort + sort workspace (the OOM behind TIES/Breadcrumbs) without the ~300x host kthvalue penalty. Small tensors keep exact top-k for mergekit parity. - gta: cap della chunks by element budget as well as rows, and build per-row ranks with a single int32 `scatter_` instead of a second int64 argsort -- wide layers (e.g. KREA2 mlp [16384, 6144]) no longer need a ~1 GiB contiguous allocation. - lora_save: sanitize tensors before saving; refactoring produces transposed/sliced views that safetensors refuses to serialize. - Tests for interp fidelity/integration, save sanitization, VRAM offload and the new sparsify paths, plus RUN_TESTS.md and a .gitignore. Co-Authored-By: Claude Opus 5 --- .gitignore | 2 + README.md | 7 + RUN_TESTS.md | 263 +++++++++++++++++++++++ src/lora_mergekit_merge.py | 92 ++++++-- src/lora_save.py | 20 ++ src/merge/algorithms.py | 28 ++- src/merge/gta.py | 335 ++++++++++++++++++++++++++++-- tests/test_gta_sparsify.py | 30 +++ tests/test_interp_fidelity.py | 64 ++++++ tests/test_interp_integration.py | 98 +++++++++ tests/test_lora_save.py | 119 +++++++++++ tests/test_merger_vram_offload.py | 104 ++++++++++ 12 files changed, 1106 insertions(+), 56 deletions(-) create mode 100644 .gitignore create mode 100644 RUN_TESTS.md create mode 100644 tests/test_interp_fidelity.py create mode 100644 tests/test_interp_integration.py create mode 100644 tests/test_lora_save.py create mode 100644 tests/test_merger_vram_offload.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..6f9cf12 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +__pycache__/ +**/__pycache__/ \ No newline at end of file diff --git a/README.md b/README.md index 8ebc2a6..3cba131 100644 --- a/README.md +++ b/README.md @@ -241,6 +241,13 @@ Nearest neighbor parameter swapping. **Parameters:** - `distance_metric` ("cosine" or "euclidean"): Distance measure +> **Delta-space + `average_weights` (v2.2.5):** SLERP, NuSLERP, Karcher and NearSwap now merge in +> delta space (like the GTA family), reconstructing the full LoRA delta instead of merging the +> up/down factors separately. Each also exposes `average_weights` (default **OFF** = additive, so +> stacked LoRAs keep full magnitude; **ON** = normalized average/blend). For these interpolation +> methods, per-LoRA strength controls **magnitude only** — blend position comes from the method's +> own parameter (SLERP `t`, NearSwap threshold). + ### Utility Nodes #### PM LoRA Modifier diff --git a/RUN_TESTS.md b/RUN_TESTS.md new file mode 100644 index 0000000..28de3ca --- /dev/null +++ b/RUN_TESTS.md @@ -0,0 +1,263 @@ +# Running Tests + +Quick guide to running tests for the LoRA-Merger-ComfyUI custom nodes. + +## Integration Test + +### Quick Start + +```bash +cd /home/lars/SD/Apps/ComfyUI/custom_nodes/LoRA-Merger-ComfyUI +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 test_gradient_analyzer_integration.py +``` + +### What It Tests + +The integration test validates the **PM LoRA Semantic Analyzer (Gradient)** pipeline: + +1. ✅ Dependencies (PyTorch, diffusers, safetensors) +2. ✅ Node instantiation (all semantic merge nodes) +3. ✅ INPUT_TYPES configuration +4. ✅ Checkpoint dropdown population +5. ✅ Specification parsing + +### Expected Output + +``` +================================================================================ +PM LoRA Semantic Analyzer (Gradient) - Integration Test +================================================================================ + +[... file checks ...] + +================================================================================ +TEST SUMMARY +================================================================================ +✓ PASS Dependencies +✓ PASS LoRA Power Stacker +✓ PASS Semantic Analyzer Instantiation +✓ PASS Semantic Merge Spec +✓ PASS Semantic Merger Instantiation + +Results: 5/5 tests passed + +🎉 ALL TESTS PASSED! +``` + +### Exit Codes + +- `0` - All tests passed +- `1` - One or more tests failed +- `130` - Test interrupted (Ctrl+C) + +## Unit Tests + +### Run All Unit Tests + +```bash +cd /home/lars/SD/Apps/ComfyUI/custom_nodes/LoRA-Merger-ComfyUI +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m pytest tests/ +``` + +### Run Specific Test Files + +```bash +# Test types system +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m pytest tests/test_types.py -v + +# Test decomposition +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m pytest tests/test_decomposition.py -v + +# Test validation +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m pytest tests/test_validation.py -v + +# Test algorithms +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m pytest tests/test_algorithms.py -v +``` + +### Run with Coverage + +```bash +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m pytest --cov=src tests/ +``` + +## Syntax Validation + +### Check Python Syntax + +```bash +# Single file +python3 -m py_compile src/nodes_semantic_merge.py + +# All Python files +find . -name "*.py" -not -path "./venv/*" -exec python3 -m py_compile {} \; +``` + +### Check with Flake8 + +```bash +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m flake8 src/ --max-line-length=120 +``` + +## Test Files + +### Integration Tests +- `test_gradient_analyzer_integration.py` - Validates semantic analyzer pipeline + +### Unit Tests (in `tests/` directory) +- `test_types.py` - Type system and validators +- `test_algorithms.py` - Merge algorithms +- `test_decomposition.py` - SVD/QR decomposition +- `test_validation.py` - Input validation + +### Test Workflows +- `/home/lars/SD/Apps/ComfyUI/user/default/workflows/PM-Gradient-Analyzer-zImage.json` - Example workflow + +## Continuous Integration + +### GitHub Actions Example + +```yaml +name: Tests + +on: [push, pull_request] + +jobs: + test: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v2 + + - name: Set up Python + uses: actions/setup-python@v2 + with: + python-version: '3.10' + + - name: Install dependencies + run: | + pip install -r requirements.txt + pip install pytest pytest-cov + + - name: Run integration tests + run: python3 test_gradient_analyzer_integration.py + + - name: Run unit tests + run: pytest tests/ -v --cov=src +``` + +## Test Requirements + +### For Integration Tests + +**Required:** +- ComfyUI installed at `/home/lars/SD/Apps/ComfyUI` +- venv at `/home/lars/SD/Apps/ComfyUI/venv` +- Dependencies: `torch`, `diffusers>=0.30.0`, `safetensors` + +**Optional (for file checks):** +- zImage checkpoint in `models/diffusion_models/` +- LoRAs in `models/loras/` + +### For Unit Tests + +**Required:** +- PyTorch +- Standard Python libraries + +**Optional:** +- pytest +- pytest-cov (for coverage reports) + +## Quick Commands Reference + +```bash +# Integration test (basic validation) +./test_gradient_analyzer_integration.py + +# Unit tests (comprehensive) +pytest tests/ -v + +# Syntax check +python3 -m py_compile src/*.py + +# Coverage report +pytest --cov=src --cov-report=html tests/ + +# Specific test +pytest tests/test_types.py::test_validate_lora_tensors -v +``` + +## Debugging Test Failures + +### Enable Verbose Output + +```bash +# Integration test +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 test_gradient_analyzer_integration.py 2>&1 | tee test_output.log + +# Unit tests +pytest tests/ -vv --tb=long +``` + +### Check Imports + +```bash +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -c " +import sys +sys.path.insert(0, '/home/lars/SD/Apps/ComfyUI') +sys.path.insert(0, '/home/lars/SD/Apps/ComfyUI/custom_nodes/LoRA-Merger-ComfyUI') +from src.nodes_semantic_merge import PMLoRASemanticAnalyzerGradient +print('✓ Import successful') +" +``` + +### Test Individual Components + +```bash +# Test diffusers +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -c " +from diffusers.models import ZImageTransformer2DModel +print('✓ ZImageTransformer2DModel available') +" + +# Test folder_paths +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -c " +import sys +sys.path.insert(0, '/home/lars/SD/Apps/ComfyUI') +import folder_paths +print('✓ folder_paths available') +" +``` + +## Performance Testing + +### Time Individual Tests + +```bash +time /home/lars/SD/Apps/ComfyUI/venv/bin/python3 test_gradient_analyzer_integration.py +``` + +### Profile Tests + +```bash +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m cProfile -o test_profile.prof test_gradient_analyzer_integration.py + +# View profile +/home/lars/SD/Apps/ComfyUI/venv/bin/python3 -m pstats test_profile.prof +``` + +## Documentation + +- **TEST_GRADIENT_ANALYZER.md** - Detailed test documentation +- **QUICK_START_GRADIENT_ANALYZER.md** - User guide +- **ZIMAGE_LOADER_OPTIMIZATION.md** - Technical details +- **BUGFIX_ZIMAGE_IMPORT.md** - Import fix documentation + +## Support + +For issues or questions: +1. Check test output for error messages +2. Review documentation files +3. Verify all dependencies are installed +4. Check that files exist in expected locations +5. Open an issue on GitHub if problems persist diff --git a/src/lora_mergekit_merge.py b/src/lora_mergekit_merge.py index 09ee88b..edd61b7 100644 --- a/src/lora_mergekit_merge.py +++ b/src/lora_mergekit_merge.py @@ -23,6 +23,9 @@ from .merge import ( get_merge_method, prepare_method_args, simple_weighted_average, + sce_merge_deltas, + interp_delta_merge, + INTERP_MODES, ) from .mergekit_utils import load_on_device # Import centralized types @@ -98,6 +101,15 @@ class LoraMergerMergekit: ), }, ), + "offload_models": ("BOOLEAN", { + "default": True, + "tooltip": ( + "Evict resident models (DIT/CLIP/VAE) from VRAM before merging, " + "to avoid OOM on low-VRAM cards.\n" + "ComfyUI reloads them automatically when the sampler runs afterward.\n" + "Only acts when device=cuda; no effect on cpu merges." + ), + }), }, } @@ -129,7 +141,8 @@ Outputs: spectral_norm_scale: float = 0.0, merge_clip: bool = True, device=None, dtype=None, - refactor_method: str = "energy_rSVD"): + refactor_method: str = "energy_rSVD", + offload_models: bool = True): if components is None: raise Exception("No components provided for merging.") @@ -139,6 +152,21 @@ Outputs: device, dtype = map_device(device, dtype) + # Free VRAM held by resident models (DIT/CLIP/VAE) before the heavy + # CUDA merge, so low-VRAM cards don't OOM. ComfyUI reloads them lazily + # when the downstream sampler runs, so no manual reload is needed. + if offload_models and device.type == "cuda": + mm = comfy.model_management + _mib = lambda b: b / (1024 ** 2) + free_before = mm.get_free_memory(device) + logging.info("PM LoRA Merger: offloading resident models from VRAM before merge " + f"(free VRAM before: {_mib(free_before):.0f} MiB)") + mm.unload_all_models() + mm.soft_empty_cache() + free_after = mm.get_free_memory(device) + logging.info(f"PM LoRA Merger: free VRAM after offload: {_mib(free_after):.0f} MiB " + f"(reclaimed {_mib(free_after - free_before):.0f} MiB)") + # Use dispatcher to get merge method merge_method = get_merge_method(method['name']) @@ -268,8 +296,14 @@ Outputs: mode = method_args.get('mode') is_gta = (not is_clip) and (mode in GTA_MODES) + # SCE and the interpolation methods (slerp/nuslerp/karcher/nearswap) also + # merge in delta space: merging the up/down factors separately injects + # meaningless up_i @ down_j cross-terms. Reconstruct each LoRA's full + # delta, merge in delta space, then refactor back into a LoRA. + is_interp = (not is_clip) and (mode in INTERP_MODES) + is_delta_space = is_gta or is_interp or ((not is_clip) and mode == "sce") - if is_gta: + if is_delta_space: deltas = [] for (u, d, a) in lora_key_tuples.values(): u = u.to(device=device, dtype=torch.float32) @@ -283,17 +317,32 @@ Outputs: else: deltas.append(scale * (u @ d)) - merged = gta_merge( - deltas, - weights.to(torch.float32), - mode=mode, - normalize=method_args.get('normalize', True), - density=method_args.get('density', 1.0), - epsilon=method_args.get('epsilon', 0.0), - gamma=method_args.get('gamma', 0.0), - sign_consensus_algorithm=method_args.get('sign_consensus_algorithm', False), - rescale_norm=method_args.get('rescale_norm', 'default'), - ) + if is_gta: + merged = gta_merge( + deltas, + weights.to(torch.float32), + mode=mode, + normalize=method_args.get('normalize', True), + density=method_args.get('density', 1.0), + epsilon=method_args.get('epsilon', 0.0), + gamma=method_args.get('gamma', 0.0), + sign_consensus_algorithm=method_args.get('sign_consensus_algorithm', False), + rescale_norm=method_args.get('rescale_norm', 'default'), + ) + elif is_interp: + merged = interp_delta_merge( + method, deltas, weights.to(torch.float32), method_args, + key=key, normalize=method_args.get('normalize', False), + dtype=torch.float32, + ) + else: # SCE in delta space + merged = sce_merge_deltas( + deltas, + weights.to(torch.float32), + select_topk=method_args.get('select_topk', 0.5), + int8_mask=method_args.get('int8_mask', False), + normalize=method_args.get('normalize', False), + ) target_rank = max(u.shape[1] for (u, _, _) in lora_key_tuples.values()) up, down, alpha_out = merged_delta_to_lora( @@ -372,14 +421,15 @@ Outputs: update_frequency = max(1, len(keys_to_process) // 100) # Update at most 100 times completed_count = 0 - # The GTA delta-space path materializes full dense deltas (hundreds of MB - # each for large FLUX layers); running 8 concurrently multiplies peak VRAM - # ~8x and exhausts an 8 GB card with a resident model (OOM, or a hard - # segfault under the quantized-model CUDA context). Serialize it on CUDA -- - # GPU-serial is only ~11 s for 256 keys. The light non-GTA factored path - # keeps the 8-way pool. - is_gta = method_args.get('mode') in GTA_MODES - n_workers = 1 if (is_gta and getattr(device, "type", device) == "cuda") else 8 + # The delta-space path (GTA family + SCE) materializes full dense deltas + # (hundreds of MB each for large FLUX layers); running 8 concurrently + # multiplies peak VRAM ~8x and exhausts an 8 GB card with a resident model + # (OOM, or a hard segfault under the quantized-model CUDA context). + # Serialize it on CUDA -- GPU-serial is only ~11 s for 256 keys. The light + # non-GTA factored path keeps the 8-way pool. + _mode = method_args.get('mode') + is_delta_space = (_mode in GTA_MODES) or (_mode == "sce") or (_mode in INTERP_MODES) + n_workers = 1 if (is_delta_space and getattr(device, "type", device) == "cuda") else 8 logging.info(f"Merging with {n_workers} worker(s) " f"(mode={method_args.get('mode')}, device={device})") diff --git a/src/lora_save.py b/src/lora_save.py index ed7fefa..47067b5 100644 --- a/src/lora_save.py +++ b/src/lora_save.py @@ -2,10 +2,25 @@ import os import comfy import folder_paths +import torch from .architectures.sd_lora import convert_to_regular_lora +def sanitize_for_save(tensor): + """Return a tensor safetensors can serialize: CPU, dense, contiguous, no shared storage.""" + if not isinstance(tensor, torch.Tensor): + return tensor + tensor = tensor.detach() + if tensor.device.type != "cpu": + tensor = tensor.cpu() + if tensor.is_contiguous() and tensor.numel() == tensor.untyped_storage().nbytes() // tensor.element_size(): + return tensor + # `.contiguous()` is a no-op for a contiguous view into a larger storage, + # so clone those to drop the shared (oversized) storage. + return tensor.contiguous().clone() if tensor.is_contiguous() else tensor.contiguous() + + class LoraSave: def __init__(self): self.loaded_lora = None @@ -34,6 +49,11 @@ class LoraSave: # so we don't need to copy them from lora_raw anymore. # The merged CLIP weights are already in state_dict. + # Refactoring/merging produces transposed or sliced views (e.g. `(V * s).T`, + # `U[:, :r]`), which safetensors refuses to serialize. Detach the views from + # their backing storage before saving. + new_state_dict = {k: sanitize_for_save(v) for k, v in new_state_dict.items()} + print(f"Saving LoRA to {save_path}") comfy.utils.save_torch_file(new_state_dict, save_path) diff --git a/src/merge/algorithms.py b/src/merge/algorithms.py index 1b36d2b..bb1cc26 100644 --- a/src/merge/algorithms.py +++ b/src/merge/algorithms.py @@ -30,7 +30,7 @@ from mergekit.sparsify import RescaleNorm import mergekit.sparsify as sparsify_module from .utils import apply_weights_to_tensors, create_map, create_tensor_param -from .gta import sce_delta_merge +from .gta import sce_delta_merge, karcher_delta_merge def generalized_task_arithmetic_merge( tensors: Dict[ModelReference, torch.Tensor], @@ -350,22 +350,18 @@ def karcher_merge( """ method_args = method_args or {} - # Apply weights to tensors (Karcher uses equal weights internally, so pre-scale) - weighted_tensors = apply_weights_to_tensors(tensors, tensor_parameters) - - merge = KarcherMerge() - task = merge.make_task( - output_weight=weight_info, - tensors=gather_tensors, - base_model=None, # No base model for LoRA merging - parameters=ImmutableMap({ - "max_iter": method_args.get("max_iter", 10), - "tol": method_args.get("tol", 1e-5) - }), - tensor_parameters=tensor_parameters, + # Memory-frugal delta-space Karcher (see gta.karcher_delta_merge). The stock + # mergekit path makes ~3N full-tensor copies and OOMs large FLUX layers; this + # consumes the delta list in place instead. Karcher uses equal weights + # internally, so per-LoRA strength (applied as a magnitude post-scale in + # interp_delta_merge) does not enter here. + deltas = list(tensors.values()) + out = karcher_delta_merge( + deltas, + max_iter=method_args.get("max_iter", 10), + tol=method_args.get("tol", 1e-5), ) - - return task.execute(tensors=weighted_tensors) * method_args.get('lambda_', 1.0) + return out * method_args.get('lambda_', 1.0) def slerp_merge( diff --git a/src/merge/gta.py b/src/merge/gta.py index a20ce5b..8048181 100644 --- a/src/merge/gta.py +++ b/src/merge/gta.py @@ -5,6 +5,7 @@ repo's broken test env). Faithfully mirrors mergekit's sparsify + GTA math but operates directly on full LoRA deltas, avoiding the per-factor squaring and the meaningless factor-space sign vote of the old path. """ +import math from typing import List, Optional import torch @@ -32,16 +33,63 @@ def _rescaled_masked(tensor: torch.Tensor, mask: torch.Tensor, # ---------------------------------------------------------------- sparsify +# Above this many elements, magnitude / magnitude_outliers selection picks its +# threshold from a GPU histogram of |t| instead of a full argsort. A full argsort +# materializes an int64 permutation (8 bytes/elem == 2x an fp32 delta) plus sort +# workspace -- on a giant FLUX layer that single allocation OOMs an 8 GB card on +# top of the resident deltas (this is what broke TIES and Breadcrumbs). A host +# k-th-value avoids the OOM but is ~300x slower (a full PCIe copy + CPU select, +# per delta per key, serialized -- what made TIES/Breadcrumbs unusably slow). The +# histogram is a single O(numel) GPU pass, allocates only a few-KB histogram, and +# lands within ~1e-3 of the target density. Small tensors keep the exact GPU +# top-k so they stay bit-for-bit identical to mergekit (the parity tests rely on +# that). +_SPARSIFY_EXACT_ELEMS = 16 * 1024 * 1024 +_SPARSIFY_HIST_BINS = 8192 + + +def _topk_hist_threshold(w_abs: torch.Tensor, keep_k: int) -> float: + """Magnitude threshold ``tau`` such that ``w_abs >= tau`` keeps ~``keep_k`` of + the largest elements, found from a GPU histogram of ``w_abs`` (no sort, no host + transfer). ``w_abs`` is non-negative. Returns a python float.""" + numel = w_abs.numel() + if keep_k <= 0: + return float(w_abs.max()) + 1.0 # keep nothing + if keep_k >= numel: + return 0.0 # keep everything + hi = float(w_abs.max()) + if hi <= 0.0: + return 0.0 + hist = torch.histc(w_abs, bins=_SPARSIFY_HIST_BINS, min=0.0, max=hi) + # Cumulative count from the largest-magnitude bin downward; take the lower edge + # of the fewest top bins whose running count first reaches keep_k. + csum = torch.cumsum(hist.flip(0), 0) + j = int(torch.searchsorted(csum, torch.tensor(float(keep_k), device=w_abs.device))) + j = min(j, _SPARSIFY_HIST_BINS - 1) + return hi * (1.0 - (j + 1) / _SPARSIFY_HIST_BINS) + + def _magnitude(t, density, rescale_norm): if density >= 1: return t k = int(density * t.numel()) - mask = torch.zeros_like(t) - w = t.abs().view(-1) - if w.device.type == "cpu": - w = w.float() - topk = torch.argsort(w, descending=True)[:k] - mask.view(-1)[topk] = 1 + numel = t.numel() + if numel <= _SPARSIFY_EXACT_ELEMS: + mask = torch.zeros_like(t) + w = t.abs().view(-1) + if w.device.type == "cpu": + w = w.float() + topk = torch.argsort(w, descending=True)[:k] + mask.view(-1)[topk] = 1 + elif k <= 0: + mask = torch.zeros_like(t) + else: + # Keep the top-k by magnitude via a histogram threshold (~k up to bin + # resolution, negligible for merge quality). + w = t.abs() + tau = _topk_hist_threshold(w, k) + mask = (w >= tau).to(t.dtype) + del w return _rescaled_masked(t, mask, rescale_norm) @@ -55,12 +103,26 @@ def _magnitude_outliers(t, density, rescale_norm, gamma): if n_bot < 0: n_top += n_bot n_bot = 0 - w = t.abs().view(-1) - if w.device.type == "cpu": - w = w.float() - idx = torch.sort(w, descending=False).indices - mask = torch.zeros_like(t) - mask.view(-1)[idx[n_bot:-n_top]] = 1 + if n <= _SPARSIFY_EXACT_ELEMS: + w = t.abs().view(-1) + if w.device.type == "cpu": + w = w.float() + idx = torch.sort(w, descending=False).indices + mask = torch.zeros_like(t) + mask.view(-1)[idx[n_bot:-n_top]] = 1 + else: + # Middle band: drop the smallest n_bot and largest n_top by magnitude, via + # GPU-histogram thresholds (no sort workspace, no host round-trip). + # tau_low keeps the top (n - n_bot); tau_high keeps the top n_top (dropped). + w = t.abs() + mask = torch.ones_like(t) + if n_bot > 0: + tau_low = _topk_hist_threshold(w, n - n_bot) + mask *= (w >= tau_low).to(t.dtype) + if n_top > 0: + tau_high = _topk_hist_threshold(w, n_top) + mask *= (w < tau_high).to(t.dtype) + del w return _rescaled_masked(t, mask, rescale_norm) @@ -77,7 +139,13 @@ def _bernoulli(t, density, rescale_norm): # Rows above this are processed in blocks (below, in one shot). The whole-tensor # path is kept for small tensors so it stays bit-for-bit identical to mergekit's # della (same bernoulli RNG order) -- the parity unit test relies on that. -_DELLA_CHUNK_ROWS = 4096 +_DELLA_CHUNK_ROWS = 4096 # whole-vs-chunked threshold (rows) +# Cap the elements per chunk so della's temporaries (an int64 argsort + an int32 +# rank buffer) stay a small constant on *wide* layers -- e.g. KREA2 mlp deltas +# are [16384, 6144], where a 4096-row block needs a ~1 GiB contiguous int64 +# tensor that fails against a fragmented VRAM pool even with GiB free. Narrow +# layers (small `cols`) are unaffected: they still use the full row block. +_DELLA_CHUNK_ELEMS = 4 * 1024 * 1024 def _della_magprune(t, density, epsilon, rescale_norm): @@ -127,28 +195,42 @@ def _norm_value(t: torch.Tensor, norm: Optional[str]): raise ValueError(f"unknown rescale_norm {norm!r}") -def _della_chunked(x, density, epsilon, rescale_norm, chunk_rows=_DELLA_CHUNK_ROWS): - """Memory-frugal della for large layers (e.g. FLUX [21504, 3072]). +def _della_chunked(x, density, epsilon, rescale_norm, chunk_rows=None): + """Memory-frugal della for large layers (e.g. KREA2 mlp [16384, 6144]). della's ranking is per-row (``argsort(dim=1)``), so rows are independent and processing them in blocks is exact -- only the bernoulli draw order changes, and della is stochastic pruning anyway. The whole-tensor path peaks at ~11x the delta (two full-layer int64 argsorts + a stack of fp32 temporaries); this - keeps every temporary to a single ``chunk_rows`` block, so peak is a small - constant regardless of layer size. The mask is applied in place; the global - rescale (l1/l2/linf) is computed across the whole tensor to match semantics.""" + keeps every temporary to a single block, so peak is a small constant + regardless of layer size. The mask is applied in place; the global rescale + (l1/l2/linf) is computed across the whole tensor to match semantics. + + Two things bound peak VRAM: the block is capped by both a row count and an + element budget (so *wide* layers get short blocks), and the per-row rank is + built with a single ``scatter_`` into an int32 buffer instead of a second + int64 ``argsort`` (which would double the index memory and add sort + workspace -- the allocation that OOMs on an 8 GB card).""" rows, cols = x.shape + if chunk_rows is None: + chunk_rows = min(_DELLA_CHUNK_ROWS, max(1, _DELLA_CHUNK_ELEMS // cols)) in_place = x.dtype == torch.float32 out = x if in_place else x.to(torch.float32) # in-place on fp32 input; else one copy before = _norm_value(out, rescale_norm) # from original values, pre-mask denom = float(cols - 1) if cols > 1 else 1.0 two_eps = 2.0 * epsilon base = density - epsilon + # ascending rank positions 0..cols-1, reused for every block via broadcast + positions = torch.arange(cols, device=out.device, dtype=torch.int32).unsqueeze(0) for lo in range(0, rows, chunk_rows): hi = min(lo + chunk_rows, rows) blk = out[lo:hi] # view into out sorted_idx = torch.argsort(blk.abs(), dim=1, descending=False) - ranks = sorted_idx.argsort(dim=1) # 0..cols-1 per row + # per-row ascending rank of each element: rank[i, sorted_idx[i, j]] = j. + # Equivalent to sorted_idx.argsort(dim=1) but avoids a second int64 + # tensor + its sort workspace. + ranks = torch.empty((hi - lo, cols), dtype=torch.int32, device=out.device) + ranks.scatter_(1, sorted_idx, positions.expand(hi - lo, cols)) del sorted_idx probs = base + (ranks.to(torch.float32) / denom) * two_eps del ranks @@ -295,6 +377,221 @@ def _stream_merge(deltas: List[torch.Tensor], weights: torch.Tensor, *, return mixed +# Above this many elements, the top-k 'select' threshold is found from a GPU +# histogram rather than an exact top-k. torch.topk on CUDA allocates an O(numel) +# sort workspace (~4 bytes/elem) -- on a giant FLUX layer that single allocation is +# enough to OOM an 8 GB card on top of the resident deltas; a host kthvalue avoids +# the OOM but is ~300x slower (full PCIe copy + CPU select). The histogram (see +# _topk_hist_threshold) is a single O(numel) GPU pass, a few-KB buffer, and lands +# within ~1e-3 of the target density. Small tensors keep the exact GPU top-k. +_SCE_TOPK_GPU_LIMIT = 16 * 1024 * 1024 + + +def _sce_select_mask(var: torch.Tensor, select_topk: float) -> Optional[torch.Tensor]: + """0/1 mask keeping the ~`select_topk` fraction of highest-variance elements. + + Returns None if the whole tensor is kept, or a zero mask sentinel handling is + left to the caller (k==0 -> caller returns zeros). `var` is non-negative.""" + numel = var.numel() + nonzero = int(torch.count_nonzero(var)) + k = int(nonzero * select_topk) + if k <= 0: + return var.new_zeros(var.shape) # nothing selected -> zero delta + if numel <= _SCE_TOPK_GPU_LIMIT: + idx = torch.topk(var.view(-1), k=k, largest=True).indices + flat = var.new_zeros(numel) + flat[idx] = 1 + return flat.view(var.shape) + # Large layer: pick the top-k-variance threshold from a GPU histogram (no sort + # workspace, no host round-trip -- a host kthvalue here is ~300x slower). Ties + # may keep slightly >k, harmless for SCE's approximate variance selection. + tau = _topk_hist_threshold(var, k) + return (var >= tau).to(var.dtype) + + +def karcher_delta_merge(deltas: List[torch.Tensor], max_iter: int = 10, + tol: float = 1e-5) -> torch.Tensor: + """Memory-frugal Riemannian (Karcher) mean of a *list* of deltas. + + Numerically identical to mergekit's ``karcher_merge_tensors`` with equal + weights (verified rel < 1e-6), but it never makes the ~3N full-tensor copies + the stock path does (``apply_weights_to_tensors`` -> N, unit vectors -> N, a + per-iteration ``ui - dot*u`` transient), which OOMs a large FLUX layer on an + 8 GB card. Instead it **consumes** ``deltas`` (normalizes each entry in place + into its unit vector) and accumulates the tangent with scalar-alpha in-place + adds, so peak is the N unit tensors plus two ``[out, in]`` accumulators. + """ + n = len(deltas) + if n == 0: + raise ValueError("karcher_delta_merge requires at least one delta") + if n == 1: + return deltas[0] + dtype, device, shape = deltas[0].dtype, deltas[0].device, deltas[0].shape + + # Norms of the originals (for the final global scale), then normalize in place + # into unit vectors -- consuming the inputs so we never hold originals + units. + norms = [] + units = [] + for i in range(n): + t = deltas[i] + deltas[i] = None + nrm = torch.linalg.norm(t.float()).item() + norms.append(nrm) + if nrm > 0.0: + t.div_(nrm) + units.append(t) + # zero-norm deltas contribute nothing to the direction (and 0 to the scale) + if not units: + return torch.zeros(shape, dtype=dtype, device=device) + + m = len(units) + a = 1.0 / m # equal weights over the valid units + + # Initial guess: normalized arithmetic mean of the unit vectors. + u = torch.zeros_like(units[0]) + for ui in units: + u.add_(ui, alpha=a) + norm_u = torch.linalg.norm(u.float()).item() + if norm_u < tol: + u = units[0].clone() + else: + u.div_(norm_u) + + # Iterative Karcher mean on the hypersphere. + for _ in range(max_iter): + T = torch.zeros_like(u) + for ui in units: + dot = float(torch.clamp(torch.dot(u.flatten(), ui.flatten()), -1.0, 1.0)) + theta = math.acos(dot) + if theta < tol: + continue + coeff = a * (theta / math.sin(theta)) + # T += coeff * (ui - dot*u), done as two scalar-alpha in-place adds + # (no full-size ``ui - dot*u`` temporary). + T.add_(ui, alpha=coeff) + T.add_(u, alpha=-coeff * dot) + norm_T = torch.linalg.norm(T.float()).item() + if norm_T < tol: + break + # u = cos(||T||)*u + sin(||T||)*(T/||T||), in place. + u.mul_(math.cos(norm_T)).add_(T, alpha=math.sin(norm_T) / norm_T) + u_norm = torch.linalg.norm(u.float()).item() + if u_norm > tol: + u.div_(u_norm) + + # Global scale: equal-weight mean of the ORIGINAL norms (all n tensors). + s = sum(nrm for nrm in norms) / n + return u.mul_(s) + + +def sce_delta_merge(deltas: List[torch.Tensor], select_topk: float, + normalize: bool = True) -> torch.Tensor: + """Memory-frugal SCE (Select-Calculate-Erase) merge over a *list* of deltas. + + Faithful to mergekit's ``sce_merge`` with a zero base tensor, but it never + stacks the deltas into an ``[N, out, in]`` tensor and never allocates a + full-layer GPU sort workspace -- either of those, on top of the N resident + deltas, OOMs an 8 GB card on large FLUX layers. It keeps the same footprint as + the GTA stream merge: the N input deltas plus a small constant number of + ``[out, in]`` buffers (freed/consumed in place), independent of N. + + Steps (streamed): one pass for the per-element sum and sum-of-squares (variance + for the 'select' mask + elected sign), then -- for the merge -- a sign-consensus + accumulate over the selected task vectors. + + ``normalize`` picks the merge convention (both keep the same select + elected + sign): + * ``True`` -> mergekit's SCE: a normalized weighted AVERAGE (per-tensor + variance weights, divided by the surviving weight sum). A blend; magnitude + is bounded by a single delta regardless of how many LoRAs stack. + * ``False`` -> additive SUM of the sign-agreeing selected contributions, so + per-LoRA strengths act as gains and stacked LoRAs keep full magnitude + (matches ComfyUI stacking and the rest of the node suite). This is the + default the SCE node exposes. + + ``deltas`` are the already strength-weighted task vectors and are **consumed** + (masked/freed in place).""" + n = len(deltas) + if n == 0: + raise ValueError("sce_delta_merge requires at least one delta") + dtype, device, shape = deltas[0].dtype, deltas[0].device, deltas[0].shape + + # Single pass: sum (s1) and sum-of-squares (s2). Only two full-layer buffers + # are held here (plus one transient square per iter), vs. the previous + # s1+mean+var+var.abs() quartet. + s1 = torch.zeros(shape, dtype=dtype, device=device) + s2 = torch.zeros(shape, dtype=dtype, device=device) + for d in deltas: + s1 += d + s2 += d * d + # var = E[d^2] - E[d]^2 (unbiased=False). Computed in place into s2; the + # one-pass form can go slightly negative from cancellation, so clamp for the + # ranking (magnitude only affects which elements are 'most variable'). + s2.div_(n).sub_((s1 / n).square_()).clamp_(min=0) # s2 is now var + + # 'Select': top-`select_topk` fraction of highest-variance elements. + mask = None + if select_topk < 1: + mask = _sce_select_mask(s2, select_topk) + if int(torch.count_nonzero(mask)) == 0: + return torch.zeros(shape, dtype=dtype, device=device) + del s2 # free variance buffer + + # Elected per-element sign (TIES 'sum' method). Masking is a shared per-element + # factor, so sum_i(mask * d_i) == mask * s1 and the sign is unchanged where + # mask > 0; masked-out elements get sign +1 but are excluded below anyway. + if mask is not None: + s1 *= mask + majority_sign = (s1 >= 0).to(dtype) * 2 - 1 + del s1 + + # Apply the 'select' mask in place -> task vectors. + if mask is not None: + for i in range(n): + deltas[i] *= mask + del mask + + if not normalize: + # Additive convention: SUM the sign-agreeing selected contributions. No + # per-tensor weighting, no divide -- strengths (already baked into the + # deltas) act as gains. Accumulate in place; hold only {acc, majority_sign}. + acc = None + for i in range(n): + tv = deltas[i] + deltas[i] = None # free input as soon as consumed + agree = (torch.sign(tv) == majority_sign).to(dtype) # 0/1 + tv.mul_(agree) # keep only agreeing elements + acc = tv if acc is None else acc.add_(tv) + return acc + + # 'Calculate' (normalized average): per-tensor SCE weights = mean(tv_i**2), + # normalized over i. sum(tv_i**2) == ||tv_i||^2, so use the norm reduction (no + # full-layer square temporary). + denom = float(deltas[0].numel()) + tv_w = torch.empty(n, dtype=torch.float32, device=device) + for i in range(n): + tv_w[i] = deltas[i].float().norm() ** 2 / denom + wsum = float(tv_w.sum()) + tv_w = torch.ones_like(tv_w) / n if abs(wsum) < 1e-6 else tv_w / wsum + + # 'Erase' + merge: keep only contributions agreeing with the elected sign, + # weighted by the SCE weights, then normalize by the surviving weight sum. + # Accumulate in place to hold only {numerator, divisor, majority_sign} beyond + # the shrinking delta list. + numerator = None + divisor = None + for i in range(n): + tv = deltas[i] + deltas[i] = None # free input as soon as consumed + w_i = float(tv_w[i]) + agree = (torch.sign(tv) == majority_sign).to(dtype) # 0/1 + tv.mul_(agree).mul_(w_i) # tv <- w_i * agree * tv (in place) + numerator = tv if numerator is None else numerator.add_(tv) + agree.mul_(w_i) # agree <- w_i * agree (in place) + divisor = agree if divisor is None else divisor.add_(agree) + return numerator / divisor.clamp_(min=1e-6) + + # --------------------------------------------------------- sign + merge def elect_sign(weighted_deltas: torch.Tensor) -> torch.Tensor: """Per-element elected sign from stacked weighted deltas (shape [N, *]). diff --git a/tests/test_gta_sparsify.py b/tests/test_gta_sparsify.py index 7c310c6..7536fd1 100644 --- a/tests/test_gta_sparsify.py +++ b/tests/test_gta_sparsify.py @@ -110,6 +110,34 @@ def test_della_chunk_boundary_invariant_density(): assert abs(ka - kb) < 0.01 and abs(ka - 0.5) < 0.02 +def test_della_scatter_rank_equivalence(): + # The chunked path builds per-row ranks with scatter_ instead of a second + # argsort; the rank values must match exactly (only the memory differs). + torch.manual_seed(1) + mags = torch.randn(50, 40).abs() + sorted_idx = torch.argsort(mags, dim=1, descending=False) + ranks_argsort = sorted_idx.argsort(dim=1) # old method + positions = torch.arange(40, dtype=torch.int32).unsqueeze(0) + ranks_scatter = torch.empty((50, 40), dtype=torch.int32) + ranks_scatter.scatter_(1, sorted_idx, positions.expand(50, 40)) # new method + assert torch.equal(ranks_scatter.long(), ranks_argsort) + + +def test_della_wide_layer_preserves_density_and_monotonic(): + # Wide layer: rows>threshold AND large cols, so the element budget forces + # short blocks (the KREA2-style path). Statistics must still hold. + torch.manual_seed(0) + t = torch.randn(8000, 6000) + got = gta.sparsify(t, "della_magprune", density=0.6, epsilon=0.2) + kept = (got != 0).float().mean().item() + assert abs(kept - 0.6) < 0.02, f"kept {kept} != ~0.6" + mags = t.abs() + med = mags.median(dim=1, keepdim=True).values + hi = mags >= med + kept_mask = got != 0 + assert kept_mask[hi].float().mean().item() > kept_mask[~hi].float().mean().item() + 0.05 + + run([ ("magnitude", test_magnitude_matches_mergekit), ("magnitude_rescale_l1", test_magnitude_rescale_l1_matches), @@ -121,4 +149,6 @@ run([ ("della_large_l1", test_della_large_l1_preserved), ("della_large_monotonic", test_della_large_keep_prob_monotonic_in_magnitude), ("della_chunk_density", test_della_chunk_boundary_invariant_density), + ("della_scatter_rank_equivalence", test_della_scatter_rank_equivalence), + ("della_wide_density_monotonic", test_della_wide_layer_preserves_density_and_monotonic), ]) \ No newline at end of file diff --git a/tests/test_interp_fidelity.py b/tests/test_interp_fidelity.py new file mode 100644 index 0000000..b800fe8 --- /dev/null +++ b/tests/test_interp_fidelity.py @@ -0,0 +1,64 @@ +# tests/test_interp_fidelity.py +# Informational: delta-space blend should align better with the intended merge +# than the old factored path. Asserts a loose lower bound so it is not brittle. +import os, sys, traceback + +REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +PARENT = os.path.dirname(REPO) +COMFY_ROOT = os.path.dirname(PARENT) +sys.path.insert(0, COMFY_ROOT) +sys.path.insert(0, PARENT) + +import torch +import importlib.util +PKG = "LoRA_Merger_ComfyUI_test" +spec = importlib.util.spec_from_file_location( + PKG, os.path.join(REPO, "__init__.py"), submodule_search_locations=[REPO]) +pkg = importlib.util.module_from_spec(spec) +sys.modules[PKG] = pkg +spec.loader.exec_module(pkg) + +from LoRA_Merger_ComfyUI_test.src.merge.algorithms import ( + interp_delta_merge, slerp_merge, nuslerp_merge, karcher_merge) + + +def _delta(seed): + torch.manual_seed(seed) + return (torch.randn(96, 8) * 0.1) @ (torch.randn(8, 128) * 0.1) + + +def _cos(a, b): + return torch.nn.functional.cosine_similarity(a.flatten(), b.flatten(), dim=0).item() + + +def test_delta_space_blend_aligns_with_average(): + margs = { + "slerp": (slerp_merge, {"mode": "slerp", "t": 0.5, "lambda_": 1.0}), + "nuslerp": (nuslerp_merge, {"mode": "nuslerp", "nuslerp_flatten": True, + "nuslerp_row_wise": False, "lambda_": 1.0}), + "karcher": (karcher_merge, {"mode": "karcher", "max_iter": 10, "tol": 1e-5, + "lambda_": 1.0}), + } + for name, (fn, ma) in margs.items(): + deltas = [_delta(1), _delta(2)] + ref = 0.5 * (deltas[0] + deltas[1]) + out = interp_delta_merge(fn, [d.clone() for d in deltas], torch.tensor([1.0, 1.0]), + dict(ma), key="k", normalize=True) + c = _cos(out, ref) + print(f" {name}: cos(delta-space blend, mean-delta) = {c:+.3f}") + assert c > 0.6, f"{name}: unexpectedly low alignment {c}" + + +def run(tests): + failed = 0 + for name, fn in tests: + try: + fn(); print(f"PASS {name}") + except Exception: + failed += 1; print(f"FAIL {name}"); traceback.print_exc() + if failed: + print(f"\n{failed} FAILED"); sys.exit(1) + print(f"\nAll {len(tests)} passed") + + +run([("delta_space_blend_aligns", test_delta_space_blend_aligns_with_average)]) \ No newline at end of file diff --git a/tests/test_interp_integration.py b/tests/test_interp_integration.py new file mode 100644 index 0000000..8f76a89 --- /dev/null +++ b/tests/test_interp_integration.py @@ -0,0 +1,98 @@ +# tests/test_interp_integration.py +# Standalone end-to-end test: run LoraMergerMergekit.merge() for the interpolation +# modes on CPU and confirm they produce a valid, non-zero merged LoRA. +import os, sys, traceback + +REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +PARENT = os.path.dirname(REPO) +COMFY_ROOT = os.path.dirname(PARENT) +sys.path.insert(0, COMFY_ROOT) +sys.path.insert(0, PARENT) + +import torch +import importlib.util +PKG = "LoRA_Merger_ComfyUI_test" +spec = importlib.util.spec_from_file_location( + PKG, os.path.join(REPO, "__init__.py"), submodule_search_locations=[REPO]) +pkg = importlib.util.module_from_spec(spec) +sys.modules[PKG] = pkg +spec.loader.exec_module(pkg) + +from LoRA_Merger_ComfyUI_test.src.lora_mergekit_merge import LoraMergerMergekit +from LoRA_Merger_ComfyUI_test.src.merge import get_merge_method, prepare_method_args + + +def _lora(seed, rank=8, out=64, inn=96): + torch.manual_seed(seed) + up = torch.randn(out, rank) * 0.1 + down = torch.randn(rank, inn) * 0.1 + return (up, down, torch.tensor(float(rank))) + + +def _run(mode, settings, n_loras=2, key="lora_unet_test_layer"): + node = LoraMergerMergekit() + names = [f"loraA", f"loraB", f"loraC"][:n_loras] + node.components = {key: {nm: _lora(i) for i, nm in enumerate(names)}} + node.strengths = {nm: {"strength_model": 1.0, "strength_clip": 1.0} for nm in names} + method = get_merge_method(mode) + margs = prepare_method_args(mode, settings) + result = node.merge( + method=method, method_args=margs, lambda_=1.0, spectral_norm_scale=0.0, + merge_clip=False, device=torch.device("cpu"), dtype=torch.float32) + return result[0] + + +SETTINGS = { + "slerp": {"t": 0.5, "normalize": False}, + "nuslerp": {"nuslerp_flatten": True, "nuslerp_row_wise": False, "normalize": False}, + "karcher": {"max_iter": 10, "tol": 1e-5, "normalize": False}, + "nearswap": {"similarity_threshold": 0.001, "normalize": False}, +} + + +def test_each_mode_produces_nonzero_lora(): + for mode, st in SETTINGS.items(): + out = _run(mode, st) + adapters = out["lora"] + assert adapters, f"{mode}: empty adapter dict" + for k, adapter in adapters.items(): + up, down, alpha = adapter.weights[0], adapter.weights[1], adapter.weights[2] + assert torch.isfinite(up).all() and torch.isfinite(down).all(), f"{mode}: non-finite" + recon = up @ down + assert recon.norm().item() > 1e-6, f"{mode}: near-zero merge ({recon.norm()})" + + +def test_additive_stronger_than_average(): + def recon_norm(normalize): + st = dict(SETTINGS["slerp"]); st["normalize"] = normalize + out = _run("slerp", st) + a = next(iter(out["lora"].values())) + return (a.weights[0] @ a.weights[1]).norm().item() + off = recon_norm(False) + on = recon_norm(True) + assert off > on * 1.5, f"additive({off}) not > average({on})" + + +def test_single_owner_key_not_zero(): + out = _run("slerp", SETTINGS["slerp"], n_loras=1) + a = next(iter(out["lora"].values())) + assert (a.weights[0] @ a.weights[1]).norm().item() > 1e-6 + + +def run(tests): + failed = 0 + for name, fn in tests: + try: + fn(); print(f"PASS {name}") + except Exception: + failed += 1; print(f"FAIL {name}"); traceback.print_exc() + if failed: + print(f"\n{failed} FAILED"); sys.exit(1) + print(f"\nAll {len(tests)} passed") + + +run([ + ("each_mode_nonzero", test_each_mode_produces_nonzero_lora), + ("additive_stronger_than_average", test_additive_stronger_than_average), + ("single_owner_not_zero", test_single_owner_key_not_zero), +]) \ No newline at end of file diff --git a/tests/test_lora_save.py b/tests/test_lora_save.py new file mode 100644 index 0000000..7caf167 --- /dev/null +++ b/tests/test_lora_save.py @@ -0,0 +1,119 @@ +# tests/test_lora_save.py +# Standalone script test (repo pytest collection is broken). Loads the custom-node +# package under a synthetic name so the relative imports resolve, then checks that +# save-time sanitation turns the transposed / sliced views produced by the refactor +# paths into tensors safetensors will actually serialize. +import importlib.util, os, sys, tempfile, traceback + +import torch + +REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +PARENT = os.path.dirname(REPO) # .../custom_nodes +COMFY_ROOT = os.path.dirname(PARENT) # ComfyUI root, so `import comfy` resolves +PKG = "LoRA_Merger_ComfyUI_test" + +sys.path.insert(0, COMFY_ROOT) +sys.path.insert(0, PARENT) +spec = importlib.util.spec_from_file_location( + PKG, os.path.join(REPO, "__init__.py"), submodule_search_locations=[REPO]) +pkg = importlib.util.module_from_spec(spec) +sys.modules[PKG] = pkg +spec.loader.exec_module(pkg) + +from LoRA_Merger_ComfyUI_test.src.lora_save import sanitize_for_save + + +def storage_elems(t): + return t.untyped_storage().nbytes() // t.element_size() + + +def test_transposed_view_becomes_contiguous(): + # e.g. down = (V[:, :r] * s_sqrt).T + view = torch.randn(8, 4).T + assert not view.is_contiguous() + + out = sanitize_for_save(view) + + assert out.is_contiguous() + assert out.shape == view.shape + assert torch.equal(out, view) + + +def test_column_slice_drops_shared_storage(): + # e.g. up = U[:, :r] out of a q-column randomized SVD result + base = torch.randn(16, 8) + view = base[:, :3] + + out = sanitize_for_save(view) + + assert out.is_contiguous() + assert torch.equal(out, view) + assert storage_elems(out) == out.numel() + + +def test_row_slice_drops_shared_storage(): + # Contiguous view into a larger storage: .contiguous() alone is a no-op here, + # and safetensors rejects it as a shared-storage tensor. + base = torch.randn(16, 8) + view = base[:4] + assert view.is_contiguous() + + out = sanitize_for_save(view) + + assert torch.equal(out, view) + assert storage_elems(out) == out.numel() + assert out.data_ptr() != base.data_ptr() + + +def test_dense_tensor_is_passed_through_without_copy(): + t = torch.randn(4, 4) + assert sanitize_for_save(t).data_ptr() == t.data_ptr() + + +def test_non_tensor_values_pass_through(): + assert sanitize_for_save(1.0) == 1.0 + assert sanitize_for_save(None) is None + + +def test_sanitized_state_dict_is_serializable(): + import safetensors.torch as st + + base = torch.randn(16, 8) + state_dict = { + "lora_unet_blocks_4_mlp_down.lora_up.weight": base[:, :3], + "lora_unet_blocks_4_mlp_down.lora_down.weight": torch.randn(8, 3).T, + "lora_unet_blocks_4_mlp_down.alpha": torch.tensor(3.0), + } + + with tempfile.TemporaryDirectory() as tmp: + # Regression guard: the unsanitized dict is exactly what used to blow up. + raised = False + try: + st.save_file(state_dict, os.path.join(tmp, "raw.safetensors")) + except ValueError as e: + raised = "non contiguous" in str(e) or "shared" in str(e).lower() + assert raised, "expected safetensors to reject the raw views" + + path = os.path.join(tmp, "sanitized.safetensors") + st.save_file({k: sanitize_for_save(v) for k, v in state_dict.items()}, path) + + loaded = st.load_file(path) + assert set(loaded) == set(state_dict) + for k, v in state_dict.items(): + assert torch.equal(loaded[k], v) + + +if __name__ == "__main__": + failures = 0 + for name, fn in sorted(globals().items()): + if not name.startswith("test_") or not callable(fn): + continue + try: + fn() + print(f"PASS {name[5:]}") + except Exception: + failures += 1 + print(f"FAIL {name[5:]}") + traceback.print_exc() + print(f"\n{'All' if not failures else failures} " + ("passed" if not failures else "failed")) + sys.exit(1 if failures else 0) diff --git a/tests/test_merger_vram_offload.py b/tests/test_merger_vram_offload.py new file mode 100644 index 0000000..0f1b9c4 --- /dev/null +++ b/tests/test_merger_vram_offload.py @@ -0,0 +1,104 @@ +# tests/test_merger_vram_offload.py +# Verifies the PM LoRA Merger's `offload_models` widget: it evicts resident +# models from VRAM (comfy.model_management.unload_all_models) before the merge +# only when enabled AND the merge device is cuda. Uses the same standalone +# package-loader as test_merge_node_names so the node's relative imports resolve. +import importlib.util, os, sys, inspect +from unittest.mock import patch + +REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +PARENT = os.path.dirname(REPO) # .../custom_nodes +COMFY_ROOT = os.path.dirname(PARENT) # ComfyUI root, so `import comfy` resolves +PKG = "LoRA_Merger_ComfyUI_offload_test" + +sys.path.insert(0, COMFY_ROOT) +sys.path.insert(0, PARENT) +spec = importlib.util.spec_from_file_location( + PKG, os.path.join(REPO, "__init__.py"), submodule_search_locations=[REPO]) +pkg = importlib.util.module_from_spec(spec) +sys.modules[PKG] = pkg +spec.loader.exec_module(pkg) + +from LoRA_Merger_ComfyUI_offload_test.src import lora_mergekit_merge as mod +LoraMergerMergekit = mod.LoraMergerMergekit + + +class _StopAfterOffload(Exception): + """Sentinel raised in place of the merge so we test only the offload step.""" + + +def _run(offload_models, device): + """Call lora_mergekit far enough to execute the offload guard, then bail. + + Returns the unload_all_models mock so the caller can assert on it. We patch + `mod.comfy.model_management` (the exact object the node calls through) so the + test is valid whether comfy is the real module or a conftest MagicMock, and + never touches real CUDA/model state. + """ + mm = mod.comfy.model_management + with patch.object(mod, "get_merge_method", side_effect=_StopAfterOffload), \ + patch.object(mm, "unload_all_models") as unload, \ + patch.object(mm, "soft_empty_cache"): + try: + LoraMergerMergekit().lora_mergekit( + method={"name": "linear", "settings": {}}, + components={"dummy.layer": {}}, + strengths={}, + device=device, + dtype="float32", + offload_models=offload_models, + ) + except _StopAfterOffload: + pass + return unload + + +def test_offload_widget_present_default_true(): + req = LoraMergerMergekit.INPUT_TYPES()["required"] + assert "offload_models" in req, "offload_models widget missing" + spec_tuple = req["offload_models"] + assert spec_tuple[0] == "BOOLEAN", f"expected BOOLEAN, got {spec_tuple[0]!r}" + assert spec_tuple[1]["default"] is True, "offload_models should default to True" + + +def test_offload_param_default_true(): + params = inspect.signature(LoraMergerMergekit.lora_mergekit).parameters + assert "offload_models" in params, "offload_models param missing from signature" + assert params["offload_models"].default is True, "offload_models default should be True" + + +def test_unloads_on_cuda_when_enabled(): + assert _run(offload_models=True, device="cuda").called, \ + "unload_all_models should be called for cuda merge with offload_models=True" + + +def test_no_unload_when_disabled(): + assert not _run(offload_models=False, device="cuda").called, \ + "unload_all_models must not be called when offload_models=False" + + +def test_no_unload_on_cpu(): + assert not _run(offload_models=True, device="cpu").called, \ + "unload_all_models must not be called for a cpu merge" + + +def run(): + import traceback + tests = [ + ("offload_widget_present_default_true", test_offload_widget_present_default_true), + ("offload_param_default_true", test_offload_param_default_true), + ("unloads_on_cuda_when_enabled", test_unloads_on_cuda_when_enabled), + ("no_unload_when_disabled", test_no_unload_when_disabled), + ("no_unload_on_cpu", test_no_unload_on_cpu), + ] + failed = 0 + for name, fn in tests: + try: + fn(); print(f"PASS {name}") + except Exception: + failed += 1; print(f"FAIL {name}"); traceback.print_exc() + sys.exit(1 if failed else 0) + + +if __name__ == "__main__": + run()