Automated PR - 2026-04-23

This commit is contained in:
github-actions[bot]
2026-04-23 12:43:54 +00:00
parent a2c3f24078
commit b604d3fab3
49 changed files with 2664 additions and 568 deletions
@@ -1,3 +1,4 @@
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
@@ -155,40 +156,77 @@ class PrecomputedDataset(Dataset):
return source_paths
def _discover_samples(self) -> dict[str, list[Path]]:
"""Discover all valid sample files across all data sources."""
# Use first data source as the reference to discover samples
"""Discover all valid sample files across all data sources.
Uses a fast two-pass approach: first globs all sources in parallel to build
full-path sets in memory, then checks expected paths via set membership.
This avoids O(N * num_sources) stat calls on networked filesystems while
correctly handling path remapping (e.g. latent_X.pt -> condition_X.pt).
"""
if not self.data_sources:
raise ValueError("No data sources configured")
data_key = "latents" if "latents" in self.data_sources else next(iter(self.data_sources.keys()))
data_path = self.source_paths[data_key]
data_files = list(data_path.glob("**/*.pt"))
# Pass 1: Glob all sources in parallel, build full-path sets
def _glob_source(dir_name: str) -> tuple[list[Path], set[str]]:
source_path = self.source_paths[dir_name]
paths = list(source_path.glob("**/*.pt"))
path_set = {str(p) for p in paths}
return paths, path_set
with ThreadPoolExecutor(max_workers=len(self.data_sources)) as executor:
glob_results = dict(
zip(
self.data_sources.keys(),
executor.map(_glob_source, self.data_sources.keys()),
strict=True,
)
)
# Get primary source files (cached from glob, no second scan)
data_files, _ = glob_results[data_key]
if not data_files:
raise ValueError(f"No data files found in {data_path}")
data_files.sort()
# Initialize sample files dict
sample_files = {output_key: [] for output_key in self.data_sources.values()}
# Log source sizes
for dir_name, (paths, _) in glob_results.items():
logger.debug(f"Source {dir_name}: {len(paths)} files")
# Build path sets for non-primary sources
other_path_sets = {
dir_name: path_set for dir_name, (_, path_set) in glob_results.items() if dir_name != data_key
}
# Pass 2: For each primary file, check if expected paths exist in other sources' sets
sample_files: dict[str, list[Path]] = {output_key: [] for output_key in self.data_sources.values()}
valid_count = 0
# For each data file, find corresponding files in other sources
for data_file in data_files:
rel_path = data_file.relative_to(data_path)
# Check if corresponding files exist in ALL sources
if self._all_source_files_exist(data_file, rel_path):
# Check all other sources via set lookup (O(1) per source, no stat calls)
all_exist = True
for dir_name, path_set in other_path_sets.items():
expected = self._get_expected_file_path(dir_name, data_file, rel_path)
if str(expected) not in path_set:
logger.debug(f"Skipping {data_file.name}: no matching {dir_name} file at {expected}")
all_exist = False
break
if all_exist:
self._fill_sample_data_files(data_file, rel_path, sample_files)
valid_count += 1
skipped = len(data_files) - valid_count
if skipped > 0:
logger.info(f"Fast index: {valid_count} valid samples from {len(data_files)} total ({skipped} skipped)")
else:
logger.debug(f"Fast index: {valid_count} valid samples from {len(data_files)} total")
return sample_files
def _all_source_files_exist(self, data_file: Path, rel_path: Path) -> bool:
"""Check if corresponding files exist in all data sources."""
for dir_name in self.data_sources:
expected_path = self._get_expected_file_path(dir_name, data_file, rel_path)
if not expected_path.exists():
logger.warning(
f"No matching {dir_name} file found for: {data_file.name} (expected in: {expected_path})"
)
return False
return True
def _get_expected_file_path(self, dir_name: str, data_file: Path, rel_path: Path) -> Path:
"""Get the expected file path for a given data source."""
source_path = self.source_paths[dir_name]
@@ -207,11 +245,14 @@ class PrecomputedDataset(Dataset):
def _validate_setup(self) -> None:
"""Validate that the dataset setup is correct."""
if not self.sample_files:
raise ValueError("No valid samples found - all data sources must have matching files")
sample_counts = {key: len(files) for key, files in self.sample_files.items()}
if not sample_counts or all(count == 0 for count in sample_counts.values()):
raise ValueError(
f"No valid samples found in {self.data_root} - all configured data sources "
f"({list(self.data_sources)}) must have matching files (per-source counts: {sample_counts})"
)
# Verify all output keys have the same number of samples
sample_counts = {key: len(files) for key, files in self.sample_files.items()}
if len(set(sample_counts.values())) > 1:
raise ValueError(f"Mismatched sample counts across sources: {sample_counts}")