Files
WildAi 2aaef8228b feat: native lowvram streaming for oversized models (v2.6.0)
- comfy_stream: convert the transformers tree to comfy-cast streaming modules; direct-param containers (Block1D gammas) self-relocate to the compute device on attribute access
- GGUF/ConvRot quant residents stream natively, dtype-preserving
- patcher: drop force-placement workaround, pure core delegation
- vendored generate: derive stage devices from activations, not parameter residency (lies under partial offload)
- generation: place inputs on the runtime compute device
- tests: lowvram protocol, streaming parity, residents, e2e
2026-08-27 17:52:17 +03:00

99 lines
3.8 KiB
Python

"""Phase 0: characterization of core's lowvram protocol on FOREIGN trees.
These tests drive the REAL ``comfy.model_patcher.ModelPatcher.load`` /
``partially_unload`` on a small foreign (plain nn.Module) tree using CPU
devices only. They encode the protocol facts that motivate the streaming
integration (plan 2026-08-26, §1):
- modules beyond a tight lowvram budget are silently SKIPPED by ``load()``
(never placed, never flagged);
- placed modules are flagged ``comfy_patched_weights`` regardless of
castability;
- ``partially_unload`` strips flagged modules (flag cleared) with no
streaming hook to bring them back — the production bug's precondition.
Device differentials are covered by the manual GPU gate (plan §5.4); here we
pin flags, ordering, and accounting deterministically.
"""
import pytest
import torch
from torch import nn
import comfy.model_management as comfy_mm
from comfy.model_patcher import ModelPatcher
def _foreign_tree():
class _Tree(nn.Module):
def __init__(self):
super().__init__()
self.hot = nn.Linear(128, 128) # ~66 KB bf16 / 66K elems f32
self.cold_tail = nn.Linear(4, 4) # tiny
return _Tree()
def _budget_for(tree, which="hot"):
"""A lowvram budget that admits `which` module but not the other."""
mods = {n: m for n, m in tree.named_modules()}
hot = comfy_mm.module_size(mods["hot"])
tail = comfy_mm.module_size(mods["cold_tail"])
# Sorted descending by size: hot loads first; tail must NOT fit.
return hot + max(hot // 8, 1)
@pytest.fixture
def patcher():
tree = _foreign_tree()
mp = ModelPatcher(
tree,
load_device=torch.device("cpu"),
offload_device=torch.device("cpu"),
size=comfy_mm.module_size(tree),
)
yield mp, tree
class TestCoreLowvramProtocolOnForeignTree:
"""Core sorts modules LARGEST-FIRST; under a tight budget the biggest
uncastable modules fail `lowvram_fits` and are silently skipped."""
def test_tight_budget_skips_all_uncastable_modules(self, patcher):
"""Uncastable + over-budget: the FIRST failing module sets an
offload_buffer it can never repay, so every later module fails
`lowvram_fits` too -> NOTHING is placed, nothing streamed, no flags.
On GPU this is exactly where CPU strays come from at first load."""
mp, tree = patcher
hot_mem = comfy_mm.module_size(tree.hot)
tail_mem = comfy_mm.module_size(tree.cold_tail)
budget = tail_mem + max(tail_mem // 8, 1)
mp.patch_model(device_to=torch.device("cpu"),
lowvram_model_memory=budget)
for m in (tree.hot, tree.cold_tail):
assert not hasattr(m, "comfy_patched_weights"), m
assert not hasattr(m, "weight_function"), m
assert mp.model.model_loaded_weight_memory == 0
def test_full_load_flags_everything(self, patcher):
mp, tree = patcher
mp.patch_model(device_to=torch.device("cpu"),
lowvram_model_memory=0) # 0 => full_load
assert getattr(tree.hot, "comfy_patched_weights", False) is True
assert getattr(tree.cold_tail, "comfy_patched_weights", False) is True
def test_partially_unload_strips_flagged_foreign_modules(self, patcher):
mp, tree = patcher
mp.patch_model(device_to=torch.device("cpu"),
lowvram_model_memory=0)
assert getattr(tree.hot, "comfy_patched_weights", False) is True
freed = mp.partially_unload(torch.device("cpu"), memory_to_free=1)
assert freed > 0
# Flag cleared: core considers it "offloaded"; a foreign tree has NO
# mechanism to pull the weights back during forward — this is the
# production crash precondition (plan §1 step 4).
assert getattr(tree.hot, "comfy_patched_weights", False) is False