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 <noreply@anthropic.com>
This commit is contained in:
larsupb
2026-08-02 11:13:13 +02:00
co-authored by Claude Opus 5
parent d6b5ad87e6
commit 1ca216dcbb
12 changed files with 1106 additions and 56 deletions
+2
View File
@@ -0,0 +1,2 @@
__pycache__/
**/__pycache__/
+7
View File
@@ -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
+263
View File
@@ -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
+71 -21
View File
@@ -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})")
+20
View File
@@ -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)
+12 -16
View File
@@ -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(
+316 -19
View File
@@ -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, *]).
+30
View File
@@ -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),
])
+64
View File
@@ -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)])
+98
View File
@@ -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),
])
+119
View File
@@ -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)
+104
View File
@@ -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()