Download src/data/dataset.py from thanhhuyvan/Gaze-LIPE: direct link, hf CLI and curl.
- Browser
- Download file 4.07 kB
-
https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/src/data/dataset.py
- Command line
-
hf download hf://thanhhuyvan/Gaze-LIPE/src/data/dataset.py
-
curl -L -o dataset.py https://huggingface.co/thanhhuyvan/Gaze-LIPE/resolve/main/src/data/dataset.py
4.07 kB
| import torch | |
| from torch.utils.data import Dataset, DataLoader | |
| import h5py | |
| import numpy as np | |
| import os | |
| import cv2 | |
| # Disable HDF5 file locking for Windows compatibility | |
| os.environ["HDF5_USE_FILE_LOCKING"] = "FALSE" | |
| class GazeDataset(Dataset): | |
| def __init__(self, h5_files, transform=None, domain_id=0): | |
| """ | |
| h5_files: List of paths to .h5 files | |
| domain_id: 0 for Source (MPII), 1 for Target (Gaze360) | |
| """ | |
| self.h5_files = h5_files | |
| self.transform = transform | |
| self.domain_id = domain_id | |
| # We need to map global index to (file_idx, local_idx) | |
| self.indices = [] | |
| self.file_handles = {} | |
| total_samples = 0 | |
| for i, f_path in enumerate(h5_files): | |
| # Open file once to get count | |
| with h5py.File(f_path, 'r') as f: | |
| num_samples = f['left_patches'].shape[0] | |
| for j in range(num_samples): | |
| self.indices.append((i, j)) | |
| total_samples += num_samples | |
| print(f"Loaded {len(h5_files)} files, total samples: {total_samples} [Domain: {domain_id}]") | |
| def __len__(self): | |
| return len(self.indices) | |
| def __getitem__(self, idx): | |
| file_idx, local_idx = self.indices[idx] | |
| f_path = self.h5_files[file_idx] | |
| if f_path not in self.file_handles: | |
| self.file_handles[f_path] = h5py.File(f_path, 'r') | |
| f = self.file_handles[f_path] | |
| # RESTORE: Fixed Left Eye Only for peak accuracy (as of June 6th) | |
| patch = f['left_patches'][local_idx] | |
| if 'left_gaze' in f: | |
| gaze = f['left_gaze'][local_idx] | |
| else: | |
| gaze = f['gaze'][local_idx] | |
| landmarks = f['landmarks'][local_idx] | |
| # Augmentations enabled for ID 5 Gaze360 experiment | |
| if self.transform: | |
| from src.utils.hardening import apply_phase1_hardening | |
| patch, landmarks = apply_phase1_hardening(patch, landmarks) | |
| # Convert to tensors | |
| patch = torch.from_numpy(patch).float() / 255.0 | |
| landmarks = torch.from_numpy(landmarks).float().view(-1) # Flatten (956,) | |
| gaze = torch.from_numpy(gaze).float() # (2,) | |
| domain = torch.tensor(self.domain_id, dtype=torch.long) | |
| # Load teacher logits if available | |
| if 'teacher_pitch_logits' in f: | |
| pitch_logits = torch.from_numpy(f['teacher_pitch_logits'][local_idx]).float() | |
| yaw_logits = torch.from_numpy(f['teacher_yaw_logits'][local_idx]).float() | |
| else: | |
| # Consistent return shape for datasets without teacher labels | |
| pitch_logits = torch.zeros(90) | |
| yaw_logits = torch.zeros(90) | |
| return patch, landmarks, gaze, pitch_logits, yaw_logits, domain | |
| def get_dataloader(h5_dir, batch_size=32, shuffle=True, num_workers=0, version_suffix=""): | |
| """ | |
| version_suffix: e.g. "_v16" to load pXX_v16.h5 files. | |
| If empty, loads base pXX.h5 files (excluding those with _vXX). | |
| """ | |
| all_files = os.listdir(h5_dir) | |
| if version_suffix: | |
| files = [os.path.join(h5_dir, f) for f in all_files if f.endswith(f"{version_suffix}.h5")] | |
| else: | |
| # Load only files that DON'T have a _vXX suffix | |
| files = [os.path.join(h5_dir, f) for f in all_files if f.endswith('.h5') and '_v' not in f] | |
| if not files: | |
| print(f"Warning: No files found with suffix '{version_suffix}' in {h5_dir}") | |
| return None | |
| dataset = GazeDataset(files) | |
| return DataLoader(dataset, batch_size=batch_size, shuffle=shuffle, num_workers=num_workers) | |
| if __name__ == '__main__': | |
| # Test with the partially processed p00.h5 | |
| h5_dir = 'data/processed' | |
| loader = get_dataloader(h5_dir, batch_size=4) | |
| if loader: | |
| for p, l, g in loader: | |
| print(f"Patch batch: {p.shape}") | |
| print(f"Landmark batch: {l.shape}") | |
| print(f"Gaze batch: {g.shape}") | |
| break | |
| else: | |
| print("No processed files found.") | |