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