Unisharp / unisharp /cli /mixed_sampler.py
Insta360-Research's picture
Upload 47 files
c7a88d2 verified
Raw
History Blame Contribute Delete
2.85 kB
from __future__ import annotations
import random
from typing import Any, Iterator
from torch.utils.data import Dataset, IterableDataset
class LazyDataLoaderIterator:
def __init__(self, dataloader: Any):
self.dataloader = dataloader
self.iterator: Iterator[Any] | None = None
def __next__(self) -> Any:
if self.iterator is None:
self.iterator = iter(self.dataloader)
return next(self.iterator)
class MixedDatasetSampler:
def __init__(
self,
datasets: dict[str, Dataset | IterableDataset],
weights: dict[str, float],
iterators: dict[str, Iterator[Any]],
seed: int | None = None,
):
self.datasets = datasets
self.weights = weights
self.iterators = iterators
self._rng = random.Random(seed)
if len(weights) == 0:
raise ValueError("weights is empty")
for name, w in weights.items():
if float(w) <= 0.0:
raise ValueError(f"Dataset weight must be > 0, got {name}={float(w)}")
if name not in datasets:
raise ValueError(f"Unknown dataset in weights: {name}")
if name not in iterators:
raise ValueError(f"Missing iterator for dataset: {name}")
total_weight = float(sum(float(v) for v in weights.values()))
self.probs = {name: float(w) / total_weight for name, w in weights.items()}
self.dataset_names = list(datasets.keys())
self.prob_list = [self.probs[name] for name in self.dataset_names]
def sample(self) -> tuple[str, Any]:
dataset_name = self.choose_dataset_name()
batch = self.next_batch(dataset_name)
return dataset_name, batch
def choose_dataset_name(self, allowed_dataset_names: list[str] | None = None) -> str:
if allowed_dataset_names is None:
names = self.dataset_names
probs = self.prob_list
else:
names = [name for name in self.dataset_names if name in set(allowed_dataset_names)]
if len(names) == 0:
raise ValueError("No allowed dataset names available for sampling.")
probs = [self.probs[name] for name in names]
return self._rng.choices(names, weights=probs, k=1)[0]
def next_batch(self, dataset_name: str) -> Any:
if dataset_name not in self.iterators:
raise ValueError(f"Unknown dataset iterator: {dataset_name}")
try:
batch = next(self.iterators[dataset_name])
except StopIteration as exc:
raise StopIteration(f"Dataset {dataset_name} exhausted") from exc
return batch
def get_sampling_stats(self) -> dict[str, float]:
return {
"probabilities": self.probs.copy(),
"sampling": self.weights.copy(),
}