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:
@@ -0,0 +1,2 @@
|
||||
__pycache__/
|
||||
**/__pycache__/
|
||||
@@ -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
@@ -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
@@ -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})")
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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, *]).
|
||||
|
||||
@@ -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),
|
||||
])
|
||||
@@ -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)])
|
||||
@@ -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),
|
||||
])
|
||||
@@ -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)
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user