Download solpix/data.py from solintellegence/SolPix: direct link, hf CLI and curl.
- Browser
- Download file 12.9 kB
-
https://huggingface.co/solintellegence/SolPix/resolve/main/solpix/data.py
- Command line
-
hf download hf://solintellegence/SolPix/solpix/data.py
-
curl -L -o data.py https://huggingface.co/solintellegence/SolPix/resolve/main/solpix/data.py
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 | |