From 74ed5fddb0be915e005aec7656117d4765ed881b Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Sun, 30 Aug 2026 11:01:52 -0600 Subject: [PATCH] Fix issue with ltx caching --- testing/test_model_loading.py | 6 +++- toolkit/memory_management/manager.py | 34 ++++++++++++++++++----- toolkit/models/v2/text_encoders/gemma3.py | 13 +++++---- toolkit/util/mixed_precision.py | 16 ++++++++++- 4 files changed, 54 insertions(+), 15 deletions(-) diff --git a/testing/test_model_loading.py b/testing/test_model_loading.py index 5553f15..43bae2c 100644 --- a/testing/test_model_loading.py +++ b/testing/test_model_loading.py @@ -21,6 +21,7 @@ import argparse import glob import json import os +import shutil import subprocess import sys import time @@ -235,7 +236,10 @@ def run_one( out_dir = os.path.join(out_dir, f"qtype_{qtype_override}") os.makedirs(out_dir, exist_ok=True) for old in glob.glob(os.path.join(out_dir, "*")): - os.remove(old) + if os.path.isdir(old): + shutil.rmtree(old) + else: + os.remove(old) from toolkit.config_modules import GenerateImageConfig, ModelConfig from toolkit.util.get_model import get_model_class diff --git a/toolkit/memory_management/manager.py b/toolkit/memory_management/manager.py index feaf449..3a89d16 100644 --- a/toolkit/memory_management/manager.py +++ b/toolkit/memory_management/manager.py @@ -55,13 +55,33 @@ class MemoryManager: self.unmanaged_modules: list[torch.nn.Module] = [] def memory_managed_to(self, *args, **kwargs): - # first move all the unmanaged modules - for module in self.unmanaged_modules: - if isinstance(module, torch.Tensor): - # Parameters and bare tensor buffers cannot move this way - module.data = module.data.to(*args, **kwargs) - else: - module.to(*args, **kwargs) + # the manager owns placement: the resident (unmanaged/ignore) set must + # live on the compute device for forwards to work. Legacy parking + # gestures (.to("cpu") between phases) would strand it there — the + # swapped .device property keeps reporting the compute device, so no + # holder heal ever brings it back. Honor device moves only TO the + # compute device; skip the device part of anything else (dtype + # handling below is unaffected). + target_device = kwargs.get("device", None) + for arg in args: + if isinstance(arg, (torch.device, str)) and not isinstance(arg, torch.dtype): + try: + target_device = torch.device(arg) + except (TypeError, RuntimeError): + pass + elif isinstance(arg, torch.device): + target_device = arg + move_resident = target_device is not None and ( + torch.device(target_device) == torch.device(self.process_device) + ) + if target_device is None or move_resident: + # first move all the unmanaged modules + for module in self.unmanaged_modules: + if isinstance(module, torch.Tensor): + # Parameters and bare tensor buffers cannot move this way + module.data = module.data.to(*args, **kwargs) + else: + module.to(*args, **kwargs) # check for a dtype argument dtype = None if "dtype" in kwargs: diff --git a/toolkit/models/v2/text_encoders/gemma3.py b/toolkit/models/v2/text_encoders/gemma3.py index 8e6bd1f..133c651 100644 --- a/toolkit/models/v2/text_encoders/gemma3.py +++ b/toolkit/models/v2/text_encoders/gemma3.py @@ -14,8 +14,10 @@ class Gemma3TextEncoder(Gemma3ForConditionalGeneration, OstrisTransformersMixin) # both layouts seen across transformers versions; missing paths skip return ["model.language_model.layers", "language_model.model.layers"] - def get_offload_ignore_modules(self): - return [self.model.language_model.base_model.embed_tokens] + # embed_tokens is NOT an ignore module: the manager's bouncing embedding + # keeps it cpu-resident with a cpu-side row gather (an ignore pin kept + # 2GB on the gpu and stranded it when legacy .to("cpu") gestures moved + # the resident set off-device) try: @@ -33,10 +35,9 @@ try: def get_offload_ignore_modules(self): # layer_scalar is a bare tensor buffer on each decoder layer; the - # manager never enumerates it, so it must ride along explicitly - return [self.embed_tokens] + [ - layer.layer_scalar for layer in self.layers - ] + # manager never enumerates it, so it must ride along explicitly. + # (embed_tokens is handled by the bouncing embedding manager.) + return [layer.layer_scalar for layer in self.layers] except ImportError: Gemma4TextEncoder = None diff --git a/toolkit/util/mixed_precision.py b/toolkit/util/mixed_precision.py index ae59618..cd16a8e 100644 --- a/toolkit/util/mixed_precision.py +++ b/toolkit/util/mixed_precision.py @@ -76,9 +76,23 @@ def pin_stored_fp32(root: nn.Module): if probe.dtype == torch.float32: return orig_apply(fn, *args, **kwargs) + # where fn sends a tensor depends on where it already is: a dtype-only + # cast keeps a cuda tensor on cuda (probing from cpu would wrongly + # report cpu and drag pinned tables off the gpu); an explicit device + # move relocates it. Probe per source device. + target_by_device = {} + + def _target_device(dev): + if dev not in target_by_device: + target_by_device[dev] = fn( + torch.zeros((), dtype=torch.float32, device=dev) + ).device + return target_by_device[dev] + def fn_pinned(t): if getattr(t, "_pin_dtype", False): - out = t if t.device == probe.device else t.to(probe.device) + target = _target_device(t.device) + out = t if t.device == target else t.to(target) out._pin_dtype = True return out return fn(t)