Fix issue with ltx caching
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user