Fixed issue with offloading text encoder on ltx 2.5

This commit is contained in:
Jaret Burkett
2026-08-12 10:46:58 -06:00
parent 0fd3e61c4c
commit a1ddeeef13
2 changed files with 9 additions and 4 deletions

View File

@@ -1536,11 +1536,16 @@ class LTX25Model(LTX2Model):
self.model_config.layer_offloading self.model_config.layer_offloading
and self.model_config.layer_offloading_text_encoder_percent > 0 and self.model_config.layer_offloading_text_encoder_percent > 0
): ):
# layer_scalar is a bare tensor buffer on each decoder layer; the
# manager never enumerates it, so it must ride along explicitly
ignore_modules = [text_encoder.embed_tokens]
for layer in text_encoder.layers:
ignore_modules.append(layer.layer_scalar)
MemoryManager.attach( MemoryManager.attach(
text_encoder, text_encoder,
self.device_torch, self.device_torch,
offload_percent=self.model_config.layer_offloading_text_encoder_percent, offload_percent=self.model_config.layer_offloading_text_encoder_percent,
ignore_modules=[text_encoder.embed_tokens], ignore_modules=ignore_modules,
) )
text_encoder.to(self.device_torch) text_encoder.to(self.device_torch)

View File

@@ -56,8 +56,8 @@ class MemoryManager:
def memory_managed_to(self, *args, **kwargs): def memory_managed_to(self, *args, **kwargs):
# 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.nn.Parameter): if isinstance(module, torch.Tensor):
# Parameter cannot move this way # Parameters and bare tensor buffers cannot move this way
module.data = module.data.to(*args, **kwargs) module.data = module.data.to(*args, **kwargs)
else: else:
module.to(*args, **kwargs) module.to(*args, **kwargs)
@@ -181,7 +181,7 @@ class MemoryManager:
for unmanaged in module._memory_manager.unmanaged_modules: for unmanaged in module._memory_manager.unmanaged_modules:
try: try:
if isinstance(unmanaged, torch.nn.Parameter): if isinstance(unmanaged, torch.Tensor):
unmanaged.data = unmanaged.data.to('cpu') unmanaged.data = unmanaged.data.to('cpu')
else: else:
unmanaged.to('cpu') unmanaged.to('cpu')