SolPix / solpix /data.py
j0no12's picture
Publish SolPix step 210000 research checkpoint
8d31176 verified
Raw History Blame Contribute Delete
12.9 kB
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