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