| import torch |
| from torch.utils.data import Dataset |
| from PIL import Image |
| import requests |
| from io import BytesIO |
| import json |
| from pathlib import Path |
| import hashlib |
| import logging |
| from typing import Optional, Tuple, Dict, List |
| import time |
| import numpy as np |
| from concurrent.futures import ThreadPoolExecutor, as_completed |
| from tqdm import tqdm |
| from torchvision import transforms |
|
|
| |
| logging.basicConfig( |
| level=logging.INFO, |
| format='%(asctime)s - %(name)s - %(levelname)s - %(message)s', |
| handlers=[ |
| logging.StreamHandler(), |
| ] |
| ) |
| |
| logger = logging.getLogger(__name__) |
|
|
|
|
| class BaseDataset(Dataset): |
| """Base dataset class with common functionality""" |
|
|
| def __init__(self, split_file: str, transform=None, multi_task: bool = False): |
| """ |
| Args: |
| split_file: Path to JSON file with image metadata |
| transform: Torchvision transforms to apply |
| """ |
| |
| with open(split_file, 'r') as f: |
| all_data = json.load(f) |
| |
| |
| if multi_task: |
| self.data = [item for item in all_data if item.get('cluster', 0) != -1] |
| if len(self.data) < len(all_data): |
| logger.info(f"Filtered out {len(all_data) - len(self.data)} items with cluster=-1") |
| else: |
| self.data = all_data |
|
|
| self.transform = transform |
| self.multi_task = multi_task |
|
|
| |
| |
| |
| ALL_DECADES = ['1960s', '1970s', '1980s', '1990s', '2000s'] |
| decades_in_data = set(item['decade'] for item in self.data) |
| |
| |
| self.decades = ALL_DECADES |
| logger.info(f"Using fixed decade classes: {self.decades}") |
| logger.info(f"Decades actually in this split: {sorted(list(decades_in_data))}") |
| self.label_to_idx = {d: i for i, d in enumerate(self.decades)} |
| self.idx_to_label = {i: d for i, d in enumerate(self.decades)} |
| self.num_classes = len(self.decades) |
|
|
| if self.multi_task: |
| |
| |
| ALL_CLUSTERS = [0, 1, 2, 3, 4] |
| clusters_in_data = set() |
| for item in self.data: |
| cluster = item.get('cluster', 0) |
| |
| clusters_in_data.add(cluster) |
| |
| self.clusters = ALL_CLUSTERS |
| logger.info(f"Using fixed cluster classes: {self.clusters}") |
| logger.info(f"Clusters actually in this split: {sorted(list(clusters_in_data))}") |
| self.cluster_to_idx = {c: i for i, c in enumerate(self.clusters)} |
| self.idx_to_cluster = {i: c for i, c in enumerate(self.clusters)} |
| self.num_cluster_classes = len(self.clusters) |
| |
| |
| devices = set() |
| for item in self.data: |
| |
| classification = item.get('classification', 'unknown').lower() |
| if classification == 'phone': |
| devices.add('phone') |
| elif classification == 'calculator': |
| devices.add('calculator') |
| else: |
| devices.add('unknown') |
| |
| |
| if len(devices) > 2 and 'unknown' in devices: |
| devices.remove('unknown') |
| |
| self.devices = sorted(list(devices)) |
| self.device_to_idx = {d: i for i, d in enumerate(self.devices)} |
| self.idx_to_device = {i: d for i, d in enumerate(self.devices)} |
| self.num_device_classes = len(self.devices) |
| |
| logger.info(f"Multi-task mode: {self.num_classes} decades, {self.num_cluster_classes} clusters, {self.num_device_classes} device types") |
| logger.info(f"Clusters: {self.clusters}") |
| logger.info(f"Device types: {self.devices}") |
| else: |
| self.clusters = None |
| self.cluster_to_idx = None |
| self.idx_to_cluster = None |
| self.num_cluster_classes = 0 |
| self.devices = None |
| self.device_to_idx = None |
| self.idx_to_device = None |
| self.num_device_classes = 0 |
|
|
| logger.info(f"Loaded dataset from {split_file} with {len(self.data)} images") |
| logger.info(f"Decades in data: {self.decades}") |
| logger.info(f"Number of decade classes: {self.num_classes}") |
|
|
| def __len__(self) -> int: |
| return len(self.data) |
|
|
| def get_labels(self) -> List[int]: |
| """Get all labels for computing class weights""" |
| return [self.label_to_idx[item['decade']] for item in self.data] |
| |
| def get_cluster_labels(self) -> List[int]: |
| """Get all cluster labels for computing class weights (multi-task only)""" |
| if not self.multi_task: |
| raise ValueError("Cluster labels only available in multi-task mode") |
| |
| return [self.cluster_to_idx[item.get('cluster', 0)] for item in self.data] |
| |
| def get_device_labels(self) -> List[int]: |
| """Get all device labels for computing class weights (multi-task only)""" |
| if not self.multi_task: |
| raise ValueError("Device labels only available in multi-task mode") |
| labels = [] |
| for item in self.data: |
| classification = item.get('classification', 'unknown').lower() |
| if classification == 'phone': |
| device = 'phone' |
| elif classification == 'calculator': |
| device = 'calculator' |
| else: |
| device = 'unknown' if 'unknown' in self.devices else self.devices[0] |
| labels.append(self.device_to_idx[device]) |
| return labels |
|
|
| def get_metadata(self, idx: int) -> Dict: |
| """Get metadata for an item""" |
| item = self.data[idx] |
| metadata = { |
| 'id': item['id'], |
| 'product_id': item['product_id'], |
| 'name': item['name'], |
| 'decade': item['decade'], |
| 'url': item.get('url', ''), |
| 'classification': item.get('classification', 'unknown'), |
| 'makers': item.get('makers', 'unknown'), |
| 'country': item.get('country', 'unknown') |
| } |
| |
| if self.multi_task: |
| metadata['cluster'] = item.get('cluster', 0) |
| classification = item.get('classification', 'unknown').lower() |
| if classification == 'phone': |
| metadata['device'] = 'phone' |
| elif classification == 'calculator': |
| metadata['device'] = 'calculator' |
| else: |
| metadata['device'] = 'unknown' if 'unknown' in self.devices else self.devices[0] |
| |
| return metadata |
|
|
|
|
| class URLDataset(BaseDataset): |
| """Dataset that loads images from URLs with caching and error handling""" |
|
|
| def __init__( |
| self, |
| split_file: str, |
| transform=None, |
| cache_dir: Optional[str] = None, |
| max_retries: int = 3, |
| timeout: int = 10, |
| fallback_on_error: bool = True, |
| multi_task: bool = False |
| ): |
| """ |
| Args: |
| split_file: Path to JSON file with image metadata |
| transform: Torchvision transforms to apply |
| cache_dir: Directory to cache downloaded images |
| max_retries: Maximum download attempts per image |
| timeout: Download timeout in seconds |
| fallback_on_error: Use placeholder image on download failure |
| multi_task: Whether to use multi-task learning (decade + cluster) |
| """ |
| super().__init__(split_file, transform, multi_task) |
|
|
| self.max_retries = max_retries |
| self.timeout = timeout |
| self.fallback_on_error = fallback_on_error |
|
|
| |
| if cache_dir: |
| self.cache_dir = Path(cache_dir) |
| else: |
| |
| data_root = Path(split_file).parent.parent |
| self.cache_dir = data_root / 'cache' / 'images' |
|
|
| self.cache_dir.mkdir(parents=True, exist_ok=True) |
|
|
| |
| self.stats = { |
| 'cache_hits': 0, |
| 'downloads': 0, |
| 'failures': 0 |
| } |
|
|
| |
| self.failed_downloads = set() |
|
|
| logger.info(f"Cache directory: {self.cache_dir}") |
|
|
| def _get_cache_path(self, url: str) -> Path: |
| """Generate cache filename from URL""" |
| url_hash = hashlib.md5(url.encode()).hexdigest() |
| return self.cache_dir / f"{url_hash}.jpg" |
|
|
| def _download_image(self, url: str) -> Optional[Image.Image]: |
| """Download image from URL with retries and improved error handling""" |
|
|
| |
| headers = { |
| 'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36', |
| 'Accept': 'image/webp,image/apng,image/*,*/*;q=0.8', |
| 'Accept-Language': 'en-US,en;q=0.9', |
| 'Accept-Encoding': 'gzip, deflate, br', |
| 'DNT': '1', |
| 'Connection': 'keep-alive', |
| 'Upgrade-Insecure-Requests': '1', |
| } |
|
|
| for attempt in range(self.max_retries): |
| try: |
| |
| if attempt > 0: |
| delay = min(2 ** attempt, 10) |
| time.sleep(delay) |
| logger.debug(f"Retry {attempt + 1} for {url} after {delay}s delay") |
|
|
| |
| response = requests.get( |
| url, |
| headers=headers, |
| timeout=self.timeout, |
| stream=True, |
| allow_redirects=True, |
| verify=True |
| ) |
|
|
| |
| response.raise_for_status() |
|
|
| |
| content_type = response.headers.get('content-type', '').lower() |
| if not any(img_type in content_type for img_type in ['image/', 'application/octet-stream']): |
| raise ValueError(f"Invalid content type: {content_type}") |
|
|
| |
| content_length = response.headers.get('content-length') |
| if content_length and int(content_length) > 50 * 1024 * 1024: |
| raise ValueError(f"Image too large: {content_length} bytes") |
|
|
| |
| content = response.content |
|
|
| |
| if len(content) < 100: |
| raise ValueError(f"Content too small: {len(content)} bytes") |
|
|
| |
| image_signatures = [ |
| b'\xff\xd8\xff', |
| b'\x89PNG\r\n\x1a\n', |
| b'GIF87a', |
| b'GIF89a', |
| b'RIFF', |
| b'BM', |
| ] |
|
|
| if not any(content.startswith(sig) for sig in image_signatures): |
| logger.warning(f"Content doesn't appear to be a valid image: {url}") |
| |
|
|
| |
| try: |
| image = Image.open(BytesIO(content)).convert('RGB') |
| except Exception as img_error: |
| raise ValueError(f"Failed to decode image: {img_error}") |
|
|
| |
| if image.size[0] < 10 or image.size[1] < 10: |
| raise ValueError(f"Image too small: {image.size}") |
|
|
| |
| if image.size[0] * image.size[1] > 20000 * 20000: |
| logger.warning(f"Very large image: {image.size}, might resize") |
| |
|
|
| |
| self.stats['downloads'] += 1 |
| logger.debug(f"Successfully downloaded {url}: {image.size}") |
| return image |
|
|
| except requests.exceptions.HTTPError as e: |
| error_msg = f"HTTP error {response.status_code}" |
| if response.status_code == 403: |
| error_msg += " (Forbidden - website blocking requests)" |
| elif response.status_code == 404: |
| error_msg += " (Not Found - URL may be outdated)" |
| elif response.status_code == 429: |
| error_msg += " (Rate Limited - too many requests)" |
| |
| if attempt < self.max_retries - 1: |
| time.sleep(30) |
| elif response.status_code >= 500: |
| error_msg += " (Server Error - temporary issue)" |
|
|
| logger.debug(f"Attempt {attempt + 1}: {error_msg} for {url}") |
| last_error = error_msg |
|
|
| except requests.exceptions.Timeout: |
| error_msg = f"Timeout after {self.timeout}s" |
| logger.debug(f"Attempt {attempt + 1}: {error_msg} for {url}") |
| last_error = error_msg |
|
|
| except requests.exceptions.ConnectionError: |
| error_msg = "Connection error (network or DNS issue)" |
| logger.debug(f"Attempt {attempt + 1}: {error_msg} for {url}") |
| last_error = error_msg |
|
|
| except requests.exceptions.RequestException as e: |
| error_msg = f"Request error: {str(e)}" |
| logger.debug(f"Attempt {attempt + 1}: {error_msg} for {url}") |
| last_error = error_msg |
|
|
| except ValueError as e: |
| |
| error_msg = f"Image validation error: {str(e)}" |
| logger.debug(f"Attempt {attempt + 1}: {error_msg} for {url}") |
| last_error = error_msg |
|
|
| except Exception as e: |
| error_msg = f"Unexpected error: {str(e)}" |
| logger.debug(f"Attempt {attempt + 1}: {error_msg} for {url}") |
| last_error = error_msg |
|
|
| |
| if any(phrase in str(last_error).lower() for phrase in [ |
| 'not found', '404', 'invalid content type', 'too small', 'too large' |
| ]): |
| logger.debug(f"Not retrying {url} due to: {last_error}") |
| break |
|
|
| |
| logger.error(f"Failed to download {url} after {self.max_retries} attempts: {last_error}") |
| self.failed_downloads.add(url) |
| self.stats['failures'] += 1 |
| return None |
|
|
| def _load_image(self, item: Dict) -> Optional[Image.Image]: |
| """Load image with caching""" |
| url = item['url'] |
|
|
| |
| cache_path = self._get_cache_path(url) |
| if cache_path.exists(): |
| try: |
| image = Image.open(cache_path).convert('RGB') |
| self.stats['cache_hits'] += 1 |
| return image |
| except Exception as e: |
| logger.warning(f"Failed to load cached image {cache_path}: {e}") |
| cache_path.unlink() |
|
|
| |
| image = self._download_image(url) |
| if image: |
| |
| try: |
| image.save(cache_path, 'JPEG', quality=95) |
| except Exception as e: |
| logger.warning(f"Failed to cache image: {e}") |
|
|
| return image |
|
|
| def _get_placeholder_image(self, size: Tuple[int, int] = (224, 224)) -> Image.Image: |
| """Create a placeholder image for failed downloads""" |
| |
| placeholder = np.random.randint(100, 150, (*size, 3), dtype=np.uint8) |
| return Image.fromarray(placeholder) |
|
|
| def __getitem__(self, idx: int): |
| """ |
| Returns: |
| image: Transformed image tensor |
| labels: If multi_task: dict with 'decade' and 'cluster' keys |
| Otherwise: int decade label (0-4) |
| metadata: Dictionary with item metadata |
| """ |
| item = self.data[idx] |
|
|
| |
| image = self._load_image(item) |
|
|
| if image is None and self.fallback_on_error: |
| |
| image = self._get_placeholder_image() |
| logger.debug(f"Using placeholder for index {idx}, URL: {item['url']}") |
| elif image is None: |
| |
| raise ValueError(f"Failed to load image at index {idx}") |
|
|
| |
| if self.transform: |
| image = self.transform(image) |
| else: |
| |
| image = transforms.ToTensor()(image) |
|
|
| |
| if self.multi_task: |
| decade_label = self.label_to_idx[item['decade']] |
| cluster = item.get('cluster', 0) |
| |
| cluster_label = self.cluster_to_idx[cluster] |
| |
| |
| classification = item.get('classification', 'unknown').lower() |
| if classification == 'phone': |
| device = 'phone' |
| elif classification == 'calculator': |
| device = 'calculator' |
| else: |
| device = 'unknown' if 'unknown' in self.devices else self.devices[0] |
| device_label = self.device_to_idx[device] |
| |
| labels = { |
| 'decade': decade_label, |
| 'cluster': cluster_label, |
| 'device': device_label |
| } |
| else: |
| labels = self.label_to_idx[item['decade']] |
|
|
| |
| metadata = self.get_metadata(idx) |
|
|
| return image, labels, metadata |
|
|
| def get_statistics(self) -> Dict: |
| """Get dataset statistics""" |
| return { |
| **self.stats, |
| 'total_images': len(self.data), |
| 'failed_urls': len(self.failed_downloads), |
| 'cache_size_mb': sum(f.stat().st_size for f in self.cache_dir.glob('*.jpg')) / 1024 / 1024 |
| } |
|
|
|
|
| class CachedDataset(BaseDataset): |
| """Dataset for pre-downloaded images (faster than URLDataset)""" |
|
|
| def __init__( |
| self, |
| split_file: str, |
| images_dir: str, |
| transform=None, |
| verify_images: bool = True, |
| multi_task: bool = False |
| ): |
| """ |
| Args: |
| split_file: Path to JSON file with image metadata |
| images_dir: Directory containing downloaded images |
| transform: Torchvision transforms to apply |
| verify_images: Whether to verify all images exist on init |
| multi_task: Whether to use multi-task learning (decade + cluster) |
| """ |
| super().__init__(split_file, transform, multi_task) |
|
|
| self.images_dir = Path(images_dir) |
|
|
| if verify_images: |
| |
| self.valid_data = [] |
| missing_count = 0 |
|
|
| for item in self.data: |
| cache_path = self._get_cache_path(item['url']) |
| if cache_path.exists(): |
| self.valid_data.append(item) |
| else: |
| missing_count += 1 |
|
|
| if missing_count > 0: |
| logger.warning(f"Missing {missing_count} cached images out of {len(self.data)}") |
|
|
| self.data = self.valid_data |
| logger.info(f"Using {len(self.data)} cached images") |
|
|
| def _get_cache_path(self, url: str) -> Path: |
| """Generate cache filename from URL""" |
| url_hash = hashlib.md5(url.encode()).hexdigest() |
| return self.images_dir / f"{url_hash}.jpg" |
|
|
| def __getitem__(self, idx: int): |
| item = self.data[idx] |
|
|
| |
| image_path = self._get_cache_path(item['url']) |
| try: |
| image = Image.open(image_path).convert('RGB') |
| except Exception as e: |
| logger.error(f"Failed to load image {image_path}: {e}") |
| raise |
|
|
| |
| if self.transform: |
| image = self.transform(image) |
|
|
| |
| if self.multi_task: |
| decade_label = self.label_to_idx[item['decade']] |
| cluster = item.get('cluster', 0) |
| |
| cluster_label = self.cluster_to_idx[cluster] |
| |
| |
| classification = item.get('classification', 'unknown').lower() |
| if classification == 'phone': |
| device = 'phone' |
| elif classification == 'calculator': |
| device = 'calculator' |
| else: |
| device = 'unknown' if 'unknown' in self.devices else self.devices[0] |
| device_label = self.device_to_idx[device] |
| |
| labels = { |
| 'decade': decade_label, |
| 'cluster': cluster_label, |
| 'device': device_label |
| } |
| else: |
| labels = self.label_to_idx[item['decade']] |
|
|
| |
| metadata = self.get_metadata(idx) |
|
|
| return image, labels, metadata |
|
|
|
|
| def download_dataset_images( |
| split_files: List[str], |
| cache_dir: str, |
| num_workers: int = 8, |
| skip_existing: bool = True |
| ) -> Dict[str, int]: |
| """ |
| Pre-download all images for faster training |
| |
| Args: |
| split_files: List of split JSON files |
| cache_dir: Directory to save images |
| num_workers: Number of parallel download workers |
| skip_existing: Skip already downloaded images |
| |
| Returns: |
| Dictionary with download statistics |
| """ |
| cache_dir = Path(cache_dir) |
| cache_dir.mkdir(parents=True, exist_ok=True) |
|
|
| |
| all_urls = set() |
| url_to_metadata = {} |
|
|
| for split_file in split_files: |
| with open(split_file, 'r') as f: |
| data = json.load(f) |
| for item in data: |
| url = item['url'] |
| all_urls.add(url) |
| url_to_metadata[url] = item |
|
|
| logger.info(f"Found {len(all_urls)} unique URLs to download") |
|
|
| |
| if skip_existing: |
| urls_to_download = [] |
| for url in all_urls: |
| cache_path = cache_dir / f"{hashlib.md5(url.encode()).hexdigest()}.jpg" |
| if not cache_path.exists(): |
| urls_to_download.append(url) |
| logger.info(f"Skipping {len(all_urls) - len(urls_to_download)} existing images") |
| else: |
| urls_to_download = list(all_urls) |
|
|
| |
| def download_single(url): |
| cache_path = cache_dir / f"{hashlib.md5(url.encode()).hexdigest()}.jpg" |
|
|
| if cache_path.exists() and skip_existing: |
| return url, True, "cached" |
|
|
| try: |
| headers = { |
| 'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36', |
| 'Accept': 'image/webp,image/apng,image/*,*/*;q=0.8', |
| 'Accept-Language': 'en-US,en;q=0.9', |
| 'Accept-Encoding': 'gzip, deflate, br', |
| 'DNT': '1', |
| 'Connection': 'keep-alive', |
| 'Upgrade-Insecure-Requests': '1', |
| } |
| response = requests.get( |
| url, |
| headers=headers, |
| timeout=10 |
| ) |
| response.raise_for_status() |
| image = Image.open(BytesIO(response.content)).convert('RGB') |
|
|
| |
| if image.size[0] < 10 or image.size[1] < 10: |
| raise ValueError(f"Image too small: {image.size}") |
|
|
| image.save(cache_path, 'JPEG', quality=95) |
| return url, True, "downloaded" |
| except Exception as e: |
| return url, False, str(e) |
|
|
| |
| results = {"cached": 0, "downloaded": 0, "failed": 0} |
| failed_items = [] |
|
|
| with ThreadPoolExecutor(max_workers=num_workers) as executor: |
| futures = {executor.submit(download_single, url): url for url in urls_to_download} |
|
|
| with tqdm(total=len(urls_to_download), desc="Downloading images") as pbar: |
| for future in as_completed(futures): |
| url, success, status = future.result() |
| pbar.update(1) |
|
|
| if success: |
| if status == "cached": |
| results["cached"] += 1 |
| else: |
| results["downloaded"] += 1 |
| else: |
| results["failed"] += 1 |
| metadata = url_to_metadata.get(url, {}) |
| failed_items.append({ |
| 'url': url, |
| 'error': status, |
| 'name': metadata.get('name', 'unknown'), |
| 'decade': metadata.get('decade', 'unknown') |
| }) |
|
|
| |
| if failed_items: |
| failed_report_path = cache_dir.parent / 'download_failures.json' |
| with open(failed_report_path, 'w') as f: |
| json.dump(failed_items, f, indent=2) |
| logger.info(f"Saved failure report to {failed_report_path}") |
|
|
| |
| total_processed = results["cached"] + results["downloaded"] + results["failed"] |
| logger.info(f"\nDownload complete:") |
| logger.info(f" Total processed: {total_processed}") |
| logger.info(f" Already cached: {results['cached']}") |
| logger.info(f" Downloaded: {results['downloaded']}") |
| logger.info(f" Failed: {results['failed']}") |
|
|
| return results |
|
|
|
|
| def create_subset_dataset( |
| dataset: BaseDataset, |
| fraction: float = 0.1, |
| seed: int = 42 |
| ) -> BaseDataset: |
| """ |
| Create a subset of a dataset for quick testing |
| |
| Args: |
| dataset: Original dataset |
| fraction: Fraction of data to keep |
| seed: Random seed |
| |
| Returns: |
| Subset dataset |
| """ |
| np.random.seed(seed) |
|
|
| |
| class_indices = {i: [] for i in range(dataset.num_classes)} |
| for idx, item in enumerate(dataset.data): |
| label = dataset.label_to_idx[item['decade']] |
| class_indices[label].append(idx) |
|
|
| |
| subset_indices = [] |
| for label, indices in class_indices.items(): |
| n_samples = max(1, int(len(indices) * fraction)) |
| sampled = np.random.choice(indices, n_samples, replace=False) |
| subset_indices.extend(sampled) |
|
|
| |
| subset_data = [dataset.data[i] for i in subset_indices] |
|
|
| |
| subset_dataset = type(dataset).__new__(type(dataset)) |
| subset_dataset.__dict__.update(dataset.__dict__) |
| subset_dataset.data = subset_data |
|
|
| logger.info(f"Created subset with {len(subset_data)} samples ({fraction * 100:.1f}% of original)") |
|
|
| return subset_dataset |
|
|
|
|
| if __name__ == "__main__": |
| |
| from torchvision import transforms |
| from pathlib import Path |
| import sys |
|
|
| print("π§ͺ URL_DATASET.PY QUICK TEST") |
| print("=" * 40) |
|
|
| |
| current_file = Path(__file__) |
| project_root = current_file.parent.parent.parent |
|
|
| print(f"Project root: {project_root}") |
| print(f"Current file: {current_file}") |
|
|
| |
| data_dir = project_root / "data" |
| splits_dir = data_dir / "splits" |
| cache_dir = data_dir / "cache" / "images" |
|
|
| |
| train_split = splits_dir / "train.json" |
| val_split = splits_dir / "val.json" |
|
|
| print(f"\nChecking files:") |
| print(f" Data dir: {data_dir.exists()} - {data_dir}") |
| print(f" Splits dir: {splits_dir.exists()} - {splits_dir}") |
| print(f" Cache dir: {cache_dir.exists()} - {cache_dir}") |
| print(f" Train split: {train_split.exists()} - {train_split}") |
|
|
| if not train_split.exists(): |
| print(f"\nβ Train split file not found!") |
| print(f"Available files in splits directory:") |
| if splits_dir.exists(): |
| for file in splits_dir.iterdir(): |
| print(f" - {file.name}") |
| else: |
| print(f" Splits directory doesn't exist!") |
| sys.exit(1) |
|
|
| |
| if cache_dir.exists(): |
| cached_images = list(cache_dir.glob('*.jpg')) |
| print(f" Cached images: {len(cached_images)}") |
| else: |
| cached_images = [] |
| print(f" Cached images: 0 (cache dir doesn't exist)") |
|
|
| |
| transform = transforms.Compose([ |
| transforms.Resize((224, 224)), |
| transforms.ToTensor(), |
| transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) |
| ]) |
|
|
| try: |
| print(f"\n1. Testing BaseDataset...") |
| base_dataset = BaseDataset(str(train_split)) |
| print(f" β Loaded {len(base_dataset)} samples") |
| print(f" β Classes: {base_dataset.decades}") |
|
|
| |
| labels = base_dataset.get_labels() |
| from collections import Counter |
|
|
| label_counts = Counter(labels) |
| print(f" β Label distribution:") |
| for idx, count in label_counts.items(): |
| decade = base_dataset.idx_to_label[idx] |
| print(f" {decade}: {count} samples") |
|
|
| |
| if len(base_dataset) > 0: |
| metadata = base_dataset.get_metadata(0) |
| print(f" β First sample: {metadata['name'][:50]}... ({metadata['decade']})") |
|
|
| except Exception as e: |
| print(f"β BaseDataset test failed: {e}") |
| sys.exit(1) |
|
|
| try: |
| print(f"\n2. Testing URLDataset...") |
| url_dataset = URLDataset( |
| split_file=str(train_split), |
| transform=transform, |
| cache_dir=str(cache_dir), |
| fallback_on_error=True, |
| max_retries=2, |
| timeout=5 |
| ) |
|
|
| print(f" β URLDataset initialized with {len(url_dataset)} samples") |
| print(f" β Cache directory: {url_dataset.cache_dir}") |
|
|
| |
| subset = create_subset_dataset(url_dataset, fraction=0.001, seed=42) |
| print(f" β Created test subset with {len(subset)} samples") |
|
|
| |
| successful_loads = 0 |
| failed_loads = 0 |
|
|
| print(f" β Testing sample loading...") |
| for i in range(min(3, len(subset))): |
| try: |
| image, label, metadata = subset[i] |
| print(f" Sample {i}: shape={image.shape}, label={label} ({metadata['decade']})") |
| print(f" Name: {metadata['name'][:40]}...") |
| successful_loads += 1 |
|
|
| |
| assert image.shape == (3, 224, 224), f"Unexpected shape: {image.shape}" |
| assert 0 <= label < 5, f"Label out of range: {label}" |
|
|
| except Exception as e: |
| print(f" Sample {i} failed: {str(e)[:60]}...") |
| failed_loads += 1 |
|
|
| print(f" β Results: {successful_loads} successful, {failed_loads} failed") |
|
|
| |
| stats = url_dataset.get_statistics() |
| print(f" β Dataset statistics:") |
| for key, value in stats.items(): |
| if isinstance(value, float): |
| print(f" {key}: {value:.2f}") |
| else: |
| print(f" {key}: {value}") |
|
|
| except Exception as e: |
| print(f"β URLDataset test failed: {e}") |
| import traceback |
|
|
| print(f"Traceback: {traceback.format_exc()}") |
| sys.exit(1) |
|
|
| try: |
| print(f"\n3. Testing CachedDataset...") |
|
|
| if len(cached_images) > 0: |
| cached_dataset = CachedDataset( |
| split_file=str(train_split), |
| images_dir=str(cache_dir), |
| transform=transform, |
| verify_images=True |
| ) |
|
|
| print(f" β CachedDataset: {len(cached_dataset)} valid cached images") |
|
|
| if len(cached_dataset) > 0: |
| |
| image, label, metadata = cached_dataset[0] |
| print(f" β Cached sample: shape={image.shape}, label={label}") |
| print(f" Name: {metadata['name'][:40]}...") |
| else: |
| print(f" β οΈ No valid cached images found") |
| else: |
| print(f" β οΈ No cached images available, skipping CachedDataset test") |
|
|
| except Exception as e: |
| print(f"β CachedDataset test failed: {e}") |
| |
|
|
| try: |
| print(f"\n4. Testing create_subset_dataset...") |
|
|
| |
| for fraction in [0.1, 0.01]: |
| subset = create_subset_dataset(base_dataset, fraction=fraction, seed=42) |
| expected_size = max(5, int(len(base_dataset) * fraction)) |
| print(f" β Subset {fraction * 100}%: {len(subset)} samples (expected ~{expected_size})") |
|
|
| |
| subset_labels = subset.get_labels() |
| unique_classes = len(set(subset_labels)) |
| print(f" Classes represented: {unique_classes}/5") |
|
|
| except Exception as e: |
| print(f"β Subset creation test failed: {e}") |
|
|
| print(f"\n" + "=" * 40) |
| print(f"π URL_DATASET.PY TESTS COMPLETED!") |
| print(f"β
Core functionality is working") |
| print(f"") |
| print(f"Usage examples:") |
| print(f" # Basic dataset") |
| print(f" dataset = BaseDataset('{train_split}')") |
| print(f" ") |
| print(f" # URL dataset with caching") |
| print(f" dataset = URLDataset('{train_split}', cache_dir='{cache_dir}')") |
| print(f" ") |
| print(f" # Cached dataset (faster)") |
| print(f" dataset = CachedDataset('{train_split}', images_dir='{cache_dir}')") |
|
|
|
|