Fixed embedding scale for offloading for ltx 2.3. Was a new bug added today.
Some checks failed
Close Stale Issues and PRs / close-stale (push) Has been cancelled

This commit is contained in:
Jaret Burkett
2026-08-31 15:53:58 -06:00
parent 6940ebf533
commit 9d6a9a0803
2 changed files with 8 additions and 10 deletions

View File

@@ -653,6 +653,10 @@ class EmbeddingLayerMemoryManager(BaseLayerMemoryManager):
# cpu-resident weight; no pinning — the weight never crosses the bus # cpu-resident weight; no pinning — the weight never crosses the bus
module.weight.data = module.weight.data.to("cpu") module.weight.data = module.weight.data.to("cpu")
# subclass buffers (gemma's embed_scale) must join the cpu-side math
for buf_name, buf in module._buffers.items():
if buf is not None:
module._buffers[buf_name] = buf.to("cpu")
self._original_forward = module.forward self._original_forward = module.forward
@@ -664,15 +668,9 @@ class EmbeddingLayerMemoryManager(BaseLayerMemoryManager):
if input_ids.device.type == "cuda" if input_ids.device.type == "cuda"
else self.manager.process_device else self.manager.process_device
) )
out = F.embedding( # the original forward preserves subclass behavior (scaled word
input_ids.to("cpu"), # embeddings multiply by embed_scale; raw F.embedding would not)
self.module.weight, out = self._original_forward(input_ids.to("cpu"))
self.module.padding_idx,
self.module.max_norm,
self.module.norm_type,
self.module.scale_grad_by_freq,
self.module.sparse,
)
return out.to(out_device, non_blocking=True) return out.to(out_device, non_blocking=True)
module.forward = _mm_forward module.forward = _mm_forward

View File

@@ -1 +1 @@
VERSION = "0.13.3" VERSION = "0.13.4"