File size: 5,219 Bytes
9aa90e0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
# SPDX-FileCopyrightText: © 2025 Tenstorrent AI ULC

# SPDX-License-Identifier: Apache-2.0

"""Self-contained LTX-2 latent shape types and patchifier grid helpers."""

from __future__ import annotations

from typing import NamedTuple

import torch


class VideoPixelShape(NamedTuple):
    """Shape of a video in pixel space: (batch, frames, height, width, fps)."""

    batch: int
    frames: int
    height: int
    width: int
    fps: float


class VideoLatentShape(NamedTuple):
    """Shape of a video in VAE latent space: (batch, channels, frames, height, width)."""

    batch: int
    channels: int
    frames: int
    height: int
    width: int


class AudioLatentShape(NamedTuple):
    """Shape of audio in VAE latent space: (batch, channels, frames, mel_bins)."""

    batch: int
    channels: int
    frames: int
    mel_bins: int

    @staticmethod
    def from_duration(
        batch: int,
        duration: float,
        channels: int = 8,
        mel_bins: int = 16,
        sample_rate: int = 16000,
        hop_length: int = 160,
        audio_latent_downsample_factor: int = 4,
    ) -> AudioLatentShape:
        latents_per_second = float(sample_rate) / float(hop_length) / float(audio_latent_downsample_factor)
        return AudioLatentShape(
            batch=batch,
            channels=channels,
            frames=round(duration * latents_per_second),
            mel_bins=mel_bins,
        )

    @staticmethod
    def from_video_pixel_shape(
        shape: VideoPixelShape,
        channels: int = 8,
        mel_bins: int = 16,
        sample_rate: int = 16000,
        hop_length: int = 160,
        audio_latent_downsample_factor: int = 4,
    ) -> AudioLatentShape:
        return AudioLatentShape.from_duration(
            batch=shape.batch,
            duration=float(shape.frames) / float(shape.fps),
            channels=channels,
            mel_bins=mel_bins,
            sample_rate=sample_rate,
            hop_length=hop_length,
            audio_latent_downsample_factor=audio_latent_downsample_factor,
        )


def video_get_patch_grid_bounds(
    shape: VideoLatentShape,
    patch_size: tuple[int, int, int] = (1, 1, 1),
    device: torch.device | str = "cpu",
) -> torch.Tensor:
    """Per-patch [start, end) grid bounds for video latent tokens.

    Returns (batch, 3, num_patches, 2) where axis 1 is (frame, height, width)
    and axis 3 is [start, end).
    """
    grid_coords = torch.meshgrid(
        torch.arange(start=0, end=shape.frames, step=patch_size[0], device=device),
        torch.arange(start=0, end=shape.height, step=patch_size[1], device=device),
        torch.arange(start=0, end=shape.width, step=patch_size[2], device=device),
        indexing="ij",
    )
    patch_starts = torch.stack(grid_coords, dim=0)  # (3, F, H, W)
    patch_size_delta = torch.tensor(patch_size, device=patch_starts.device, dtype=patch_starts.dtype).view(3, 1, 1, 1)
    patch_ends = patch_starts + patch_size_delta
    latent_coords = torch.stack((patch_starts, patch_ends), dim=-1)  # (3, F, H, W, 2)

    F, H, W = latent_coords.shape[1], latent_coords.shape[2], latent_coords.shape[3]
    latent_coords = latent_coords.reshape(3, F * H * W, 2)
    latent_coords = latent_coords.unsqueeze(0).expand(shape.batch, -1, -1, -1)
    return latent_coords


def audio_get_patch_grid_bounds(
    shape: AudioLatentShape,
    sample_rate: int = 16000,
    hop_length: int = 160,
    audio_latent_downsample_factor: int = 4,
    is_causal: bool = True,
    shift: int = 0,
    device: torch.device | str = "cpu",
) -> torch.Tensor:
    """Per-patch temporal bounds (seconds) for audio latent tokens, shape (batch, 1, num_frames, 2)."""

    def _latent_to_seconds(start_idx: int, end_idx: int) -> torch.Tensor:
        frame = torch.arange(start_idx, end_idx, dtype=torch.float32, device=device)
        mel_frame = frame * audio_latent_downsample_factor
        if is_causal:
            mel_frame = (mel_frame + 1 - audio_latent_downsample_factor).clamp(min=0)
        return mel_frame * hop_length / sample_rate

    num_steps = shape.frames
    start_timings = _latent_to_seconds(shift, num_steps + shift)
    start_timings = start_timings.unsqueeze(0).expand(shape.batch, -1).unsqueeze(1)  # (B, 1, N)

    end_timings = _latent_to_seconds(shift + 1, num_steps + shift + 1)
    end_timings = end_timings.unsqueeze(0).expand(shape.batch, -1).unsqueeze(1)  # (B, 1, N)

    return torch.stack([start_timings, end_timings], dim=-1)  # (B, 1, N, 2)


def get_pixel_coords(
    latent_coords: torch.Tensor,
    scale_factors: tuple[int, int, int] = (8, 32, 32),
    causal_fix: bool = False,
) -> torch.Tensor:
    """Scale latent-space [start, end) coordinates to pixel space per axis.

    scale_factors is (temporal, height, width); causal_fix offsets the first
    temporal frame for causal encoding.
    """
    broadcast_shape = [1] * latent_coords.ndim
    broadcast_shape[1] = -1
    scale_tensor = torch.tensor(scale_factors, device=latent_coords.device).view(*broadcast_shape)

    pixel_coords = latent_coords * scale_tensor

    if causal_fix:
        pixel_coords[:, 0, ...] = (pixel_coords[:, 0, ...] + 1 - scale_factors[0]).clamp(min=0)

    return pixel_coords