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
Some checks failed
Close Stale Issues and PRs / close-stale (push) Has been cancelled
This commit is contained in:
@@ -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
|
||||||
|
|||||||
@@ -1 +1 @@
|
|||||||
VERSION = "0.13.3"
|
VERSION = "0.13.4"
|
||||||
|
|||||||
Reference in New Issue
Block a user