Fix issue with ltx caching

This commit is contained in:
Jaret Burkett
2026-08-30 11:01:52 -06:00
parent 764b5064fb
commit 74ed5fddb0
4 changed files with 54 additions and 15 deletions

View File

@@ -21,6 +21,7 @@ import argparse
import glob import glob
import json import json
import os import os
import shutil
import subprocess import subprocess
import sys import sys
import time import time
@@ -235,6 +236,9 @@ def run_one(
out_dir = os.path.join(out_dir, f"qtype_{qtype_override}") out_dir = os.path.join(out_dir, f"qtype_{qtype_override}")
os.makedirs(out_dir, exist_ok=True) os.makedirs(out_dir, exist_ok=True)
for old in glob.glob(os.path.join(out_dir, "*")): for old in glob.glob(os.path.join(out_dir, "*")):
if os.path.isdir(old):
shutil.rmtree(old)
else:
os.remove(old) os.remove(old)
from toolkit.config_modules import GenerateImageConfig, ModelConfig from toolkit.config_modules import GenerateImageConfig, ModelConfig

View File

@@ -55,6 +55,26 @@ class MemoryManager:
self.unmanaged_modules: list[torch.nn.Module] = [] self.unmanaged_modules: list[torch.nn.Module] = []
def memory_managed_to(self, *args, **kwargs): def memory_managed_to(self, *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 # first move all the unmanaged modules
for module in self.unmanaged_modules: for module in self.unmanaged_modules:
if isinstance(module, torch.Tensor): if isinstance(module, torch.Tensor):

View File

@@ -14,8 +14,10 @@ class Gemma3TextEncoder(Gemma3ForConditionalGeneration, OstrisTransformersMixin)
# both layouts seen across transformers versions; missing paths skip # both layouts seen across transformers versions; missing paths skip
return ["model.language_model.layers", "language_model.model.layers"] return ["model.language_model.layers", "language_model.model.layers"]
def get_offload_ignore_modules(self): # embed_tokens is NOT an ignore module: the manager's bouncing embedding
return [self.model.language_model.base_model.embed_tokens] # 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: try:
@@ -33,10 +35,9 @@ try:
def get_offload_ignore_modules(self): def get_offload_ignore_modules(self):
# layer_scalar is a bare tensor buffer on each decoder layer; the # layer_scalar is a bare tensor buffer on each decoder layer; the
# manager never enumerates it, so it must ride along explicitly # manager never enumerates it, so it must ride along explicitly.
return [self.embed_tokens] + [ # (embed_tokens is handled by the bouncing embedding manager.)
layer.layer_scalar for layer in self.layers return [layer.layer_scalar for layer in self.layers]
]
except ImportError: except ImportError:
Gemma4TextEncoder = None Gemma4TextEncoder = None

View File

@@ -76,9 +76,23 @@ def pin_stored_fp32(root: nn.Module):
if probe.dtype == torch.float32: if probe.dtype == torch.float32:
return orig_apply(fn, *args, **kwargs) 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): def fn_pinned(t):
if getattr(t, "_pin_dtype", False): 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 out._pin_dtype = True
return out return out
return fn(t) return fn(t)