File size: 12,871 Bytes
8d31176
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
from __future__ import annotations

from bisect import bisect_right
from pathlib import Path
from typing import Iterator

import torch
from torch import Tensor
from torch.utils.data import Dataset, Sampler


def _load_shard(path: Path) -> dict[str, object]:
    try:
        value = torch.load(path, map_location="cpu", weights_only=True, mmap=True)
    except TypeError:
        # Kept for compatibility with PyTorch builds predating mmap support.
        value = torch.load(path, map_location="cpu", weights_only=True)
    if not isinstance(value, dict):
        raise ValueError(f"{path}: expected a dictionary of tensors")
    return value


class LatentTextShardDataset(Dataset[dict[str, Tensor]]):
    """Lazy reader for tensor shards, with one latent resolution per shard.

    Required keys: ``latents`` [N,C,H,W], ``text_embeddings`` [N,L,D], and
    ``text_mask`` [N,L]. Files are loaded with PyTorch mmap when available.
    """

    def __init__(self, data_dir: str | Path, max_text_tokens: int = 96):
        self.paths = sorted(Path(data_dir).glob("*.pt"))
        if not self.paths:
            raise FileNotFoundError(f"no .pt shard files found in {data_dir}")
        self.max_text_tokens = max_text_tokens
        self.metadata: list[dict[str, int]] = []
        self.cumulative_ends: list[int] = []
        self._cache: dict[int, dict[str, object]] = {}
        self._cache_limit = 4
        total = 0
        latent_channels: int | None = None
        text_dim: int | None = None
        repa_dim: int | None = None
        has_teacher_features: bool | None = None
        for shard_index, path in enumerate(self.paths):
            shard = _load_shard(path)
            latents = shard.get("latents")
            text = shard.get("text_embeddings")
            mask = shard.get("text_mask")
            if not all(isinstance(x, Tensor) for x in (latents, text, mask)):
                raise ValueError(f"{path}: required tensor keys are latents, text_embeddings, text_mask")
            assert isinstance(latents, Tensor) and isinstance(text, Tensor) and isinstance(mask, Tensor)
            if latents.ndim != 4 or text.ndim != 3 or mask.ndim != 2:
                raise ValueError(f"{path}: expected [N,C,H,W], [N,L,D], and [N,L] tensors")
            count, channels, height, width = latents.shape
            if count < 1 or text.shape[0] != count or mask.shape != text.shape[:2]:
                raise ValueError(f"{path}: batch/sample counts do not match")
            if text.shape[1] > max_text_tokens:
                raise ValueError(f"{path}: {text.shape[1]} text tokens exceeds limit {max_text_tokens}")
            if latent_channels is None:
                latent_channels = channels
                text_dim = text.shape[2]
            if channels != latent_channels or text.shape[2] != text_dim:
                raise ValueError(f"{path}: channel or text embedding dimensions differ from earlier shards")
            if not latents.is_floating_point() or not text.is_floating_point():
                raise ValueError(f"{path}: latents and text_embeddings must be floating point")
            teacher = shard.get("teacher_features")
            shard_has_teacher = isinstance(teacher, Tensor)
            if has_teacher_features is None:
                has_teacher_features = shard_has_teacher
            elif shard_has_teacher != has_teacher_features:
                raise ValueError("either every shard or no shard must contain teacher_features")
            if shard_has_teacher:
                assert isinstance(teacher, Tensor)
                if teacher.ndim != 3 or teacher.shape[0] != count or teacher.shape[1] != height * width:
                    raise ValueError(f"{path}: teacher_features must have shape [N,H*W,feature_dim]")
                if not teacher.is_floating_point():
                    raise ValueError(f"{path}: teacher_features must be floating point values")
                if repa_dim is None:
                    repa_dim = teacher.shape[2]
                elif teacher.shape[2] != repa_dim:
                    raise ValueError("teacher feature dimensions differ across shards")
            self.metadata.append(
                {"start": total, "count": count, "height": height, "width": width, "shard": shard_index}
            )
            total += count
            self.cumulative_ends.append(total)
        self.total_samples = total
        self.latent_channels = int(latent_channels or 0)
        self.text_dim = int(text_dim or 0)
        self.has_teacher_features = bool(has_teacher_features)
        self.repa_dim = int(repa_dim or 0)

    def __len__(self) -> int:
        return self.total_samples

    def _locate(self, index: int) -> tuple[int, int]:
        shard_index = bisect_right(self.cumulative_ends, index)
        previous_end = self.cumulative_ends[shard_index - 1] if shard_index else 0
        return shard_index, index - previous_end

    def _get_shard(self, shard_index: int) -> dict[str, object]:
        if shard_index not in self._cache:
            self._cache[shard_index] = _load_shard(self.paths[shard_index])
            while len(self._cache) > self._cache_limit:
                del self._cache[next(iter(self._cache))]
        return self._cache[shard_index]

    def __getitem__(self, index: int) -> dict[str, Tensor]:
        shard_index, row = self._locate(index)
        shard = self._get_shard(shard_index)
        result = {
            "latents": shard["latents"][row].float(),  # type: ignore[index,union-attr]
            "text_embeddings": shard["text_embeddings"][row].float(),  # type: ignore[index,union-attr]
            "text_mask": shard["text_mask"][row].bool(),  # type: ignore[index,union-attr]
        }
        if self.has_teacher_features:
            result["teacher_features"] = shard["teacher_features"][row].float()  # type: ignore[index,union-attr]
        return result

    def bucket_shards(self) -> dict[tuple[int, int], list[dict[str, int]]]:
        buckets: dict[tuple[int, int], list[dict[str, int]]] = {}
        for meta in self.metadata:
            buckets.setdefault((meta["height"], meta["width"]), []).append(meta)
        return buckets


class ResolutionBucketBatchSampler(Sampler[list[int]]):
    """Shuffle batches while ensuring every batch has a consistent latent grid."""

    def __init__(
        self,
        dataset: LatentTextShardDataset,
        batch_size: int,
        *,
        seed: int = 1234,
        drop_last: bool = True,
        rank: int = 0,
        world_size: int = 1,
    ):
        if batch_size < 1:
            raise ValueError("batch_size must be positive")
        self.dataset = dataset
        self.batch_size = batch_size
        self.seed = seed
        self.drop_last = drop_last
        self.rank = rank
        self.world_size = world_size
        self.epoch = 0

    def set_epoch(self, epoch: int) -> None:
        self.epoch = epoch

    def _make_batches(self) -> list[list[int]]:
        generator = torch.Generator().manual_seed(self.seed + self.epoch)
        buckets = self.dataset.bucket_shards()
        bucket_items = list(buckets.items())
        if bucket_items:
            order = torch.randperm(len(bucket_items), generator=generator).tolist()
            bucket_items = [bucket_items[i] for i in order]
        batches: list[list[int]] = []
        for _, shards in bucket_items:
            pending: list[int] = []
            shard_order = torch.randperm(len(shards), generator=generator).tolist()
            for shard_position in shard_order:
                meta = shards[shard_position]
                rows = torch.randperm(meta["count"], generator=generator).tolist()
                pending.extend(meta["start"] + row for row in rows)
                while len(pending) >= self.batch_size:
                    batches.append(pending[: self.batch_size])
                    pending = pending[self.batch_size :]
            if pending and not self.drop_last:
                batches.append(pending)
        if batches:
            order = torch.randperm(len(batches), generator=generator).tolist()
            batches = [batches[i] for i in order]
        if self.world_size > 1:
            usable = len(batches) - (len(batches) % self.world_size)
            batches = batches[:usable][self.rank : usable : self.world_size]
        return batches

    def __iter__(self) -> Iterator[list[int]]:
        yield from self._make_batches()

    def __len__(self) -> int:
        count = 0
        for (height, width), shards in self.dataset.bucket_shards().items():
            samples = sum(meta["count"] for meta in shards)
            count += samples // self.batch_size if self.drop_last else (samples + self.batch_size - 1) // self.batch_size
        if self.world_size > 1:
            count -= count % self.world_size
            count //= self.world_size
        return count


def collate_latent_text(samples: list[dict[str, Tensor]]) -> dict[str, Tensor]:
    if not samples:
        raise ValueError("cannot collate an empty batch")
    latent_shapes = {tuple(sample["latents"].shape) for sample in samples}
    if len(latent_shapes) != 1:
        raise ValueError("all samples in a batch must have the same latent shape")
    max_tokens = max(sample["text_embeddings"].shape[0] for sample in samples)
    text_dim = samples[0]["text_embeddings"].shape[-1]
    batch = len(samples)
    embeddings = torch.zeros(batch, max_tokens, text_dim, dtype=torch.float32)
    masks = torch.zeros(batch, max_tokens, dtype=torch.bool)
    for row, sample in enumerate(samples):
        length = sample["text_embeddings"].shape[0]
        embeddings[row, :length] = sample["text_embeddings"]
        masks[row, :length] = sample["text_mask"]
    result = {
        "latents": torch.stack([sample["latents"] for sample in samples]),
        "text_embeddings": embeddings,
        "text_mask": masks,
    }
    teacher_presence = ["teacher_features" in sample for sample in samples]
    if any(teacher_presence) and not all(teacher_presence):
        raise ValueError("all samples in a batch must either have teacher features or omit them")
    if all(teacher_presence):
        teacher_shapes = {tuple(sample["teacher_features"].shape) for sample in samples}
        if len(teacher_shapes) != 1:
            raise ValueError("teacher feature shapes must match within a batch")
        result["teacher_features"] = torch.stack([sample["teacher_features"] for sample in samples])
    return result


def apply_classifier_free_dropout(
    text_embeddings: Tensor,
    text_mask: Tensor,
    empty_embeddings: Tensor,
    empty_mask: Tensor | None,
    probability: float,
) -> tuple[Tensor, Tensor]:
    """Replace a random fraction of conditions with pre-encoded empty prompts."""
    if not 0 <= probability <= 1:
        raise ValueError("probability must be between 0 and 1")
    if probability == 0:
        return text_embeddings, text_mask
    batch, current_length, dim = text_embeddings.shape
    if empty_embeddings.ndim == 2:
        empty_embeddings = empty_embeddings.unsqueeze(0)
    if empty_embeddings.ndim != 3 or empty_embeddings.shape[-1] != dim:
        raise ValueError("empty_embeddings must have shape [L,D] or [1,L,D]")
    if empty_embeddings.shape[0] not in (1, batch):
        raise ValueError("empty_embeddings batch dimension must be 1 or match the training batch")
    if empty_mask is None:
        empty_mask = torch.ones(empty_embeddings.shape[:2], device=empty_embeddings.device, dtype=torch.bool)
    elif empty_mask.ndim == 1:
        empty_mask = empty_mask.unsqueeze(0)
    if empty_mask.shape != empty_embeddings.shape[:2]:
        raise ValueError("empty_mask must match the first two empty_embeddings dimensions")
    max_length = max(current_length, empty_embeddings.shape[1])
    if empty_embeddings.shape[0] == 1:
        empty_embeddings = empty_embeddings.expand(batch, -1, -1)
    if empty_mask.shape[0] == 1:
        empty_mask = empty_mask.expand(batch, -1)
    expanded_text = text_embeddings.new_zeros(batch, max_length, dim)
    expanded_mask = text_mask.new_zeros(batch, max_length)
    expanded_text[:, :current_length] = text_embeddings
    expanded_mask[:, :current_length] = text_mask
    empty_embeddings = empty_embeddings.to(device=text_embeddings.device, dtype=text_embeddings.dtype)
    empty_mask = empty_mask.to(device=text_mask.device, dtype=text_mask.dtype)
    selected = torch.rand(batch, device=text_embeddings.device) < probability
    expanded_text[:, : empty_embeddings.shape[1]] = torch.where(
        selected[:, None, None], empty_embeddings, expanded_text[:, : empty_embeddings.shape[1]]
    )
    expanded_mask[:, : empty_mask.shape[1]] = torch.where(
        selected[:, None], empty_mask, expanded_mask[:, : empty_mask.shape[1]]
    )
    return expanded_text, expanded_mask