[Diffusion][LTX-2] Allocate AdaLN outputs from one contiguous slab (#34508)

This commit is contained in:
Xiaoyu Zhang
2026-08-12 16:28:27 +08:00
committed by GitHub
parent daae3acb36
commit 1f008dc226
2 changed files with 27 additions and 3 deletions
@@ -165,10 +165,12 @@ def ltx2_ada_values9(
batch, seq, _ = timestep.shape
rows = int(batch * seq)
outs = tuple(
torch.empty((batch, seq, hidden), device=timestep.device, dtype=timestep.dtype)
for _ in range(9)
# Each returned output is a disjoint, contiguous view, so one allocation
# avoids nine allocator round trips per transformer block.
output_storage = torch.empty(
(9, batch, seq, hidden), device=timestep.device, dtype=timestep.dtype
)
outs = tuple(output_storage.unbind(dim=0))
_ltx2_ada_values9_kernel[(rows,)](
timestep,
scale_shift_table,
@@ -54,6 +54,28 @@ def test_ltx2_ada_values9(
assert len(actual) == 9
for actual_value, expected_value in zip(actual, expected):
assert actual_value.is_contiguous()
torch.testing.assert_close(actual_value, expected_value, atol=0, rtol=0)
@torch.no_grad()
def test_ltx2_ada_values9_torch_compile_fullgraph() -> None:
hidden = 4096
scale_shift_table = torch.randn(
9, hidden, device=DEVICE, dtype=torch.bfloat16
).contiguous()
timestep = torch.randn(
1, 1, 9 * hidden, device=DEVICE, dtype=torch.bfloat16
).contiguous()
actual = torch.compile(ltx2_ada_values9, fullgraph=True)(
scale_shift_table, timestep
)
expected = _reference(scale_shift_table, timestep)
assert len(actual) == 9
for actual_value, expected_value in zip(actual, expected):
assert actual_value.is_contiguous()
torch.testing.assert_close(actual_value, expected_value, atol=0, rtol=0)