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 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

View File

@@ -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:

View File

@@ -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

View File

@@ -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)