Gaze-LIPE / src /data /dataset.py
thanhhuyvan's picture
Publish KD reproducibility investigation
178f61f
Raw
History Blame Contribute Delete
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.")