Spaces:
Sleeping
Sleeping
| import torch | |
| import tensorflow as tf | |
| from torchvision import datasets, transforms | |
| from torch.utils.data import DataLoader | |
| from pathlib import Path | |
| CLASSES = ["buildings", "forest", "glacier", "mountain", "sea", "street"] | |
| IMAGE_SIZE = (150, 150) | |
| # PYTORCH DATA LOADER | |
| def get_pytorch_data(data_dir="data", batch_size=64): | |
| data_dir = Path(data_dir) | |
| train_path = data_dir / "seg_train" | |
| test_path = data_dir / "seg_test" | |
| train_transform = transforms.Compose([ | |
| transforms.Resize(IMAGE_SIZE), | |
| transforms.RandomHorizontalFlip(), | |
| transforms.RandomVerticalFlip(p=0.1), | |
| transforms.RandomRotation(15), | |
| transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2), | |
| transforms.ToTensor(), | |
| transforms.Normalize( | |
| mean=[0.485, 0.456, 0.406], | |
| std=[0.229, 0.224, 0.225] | |
| ) | |
| ]) | |
| test_transform = transforms.Compose([ | |
| transforms.Resize(IMAGE_SIZE), | |
| transforms.ToTensor(), | |
| transforms.Normalize( | |
| mean=[0.485, 0.456, 0.406], | |
| std=[0.229, 0.224, 0.225] | |
| ) | |
| ]) | |
| full_train = datasets.ImageFolder( | |
| str(train_path), | |
| transform=train_transform | |
| ) | |
| test_dataset = datasets.ImageFolder( | |
| str(test_path), | |
| transform=test_transform | |
| ) | |
| print("Class mapping:", full_train.class_to_idx) | |
| # Split train / validation | |
| val_size = int(0.2 * len(full_train)) | |
| train_size = len(full_train) - val_size | |
| train_dataset, val_dataset = torch.utils.data.random_split( | |
| full_train, | |
| [train_size, val_size], | |
| generator=torch.Generator().manual_seed(42) | |
| ) | |
| train_loader = DataLoader( | |
| train_dataset, | |
| batch_size=batch_size, | |
| shuffle=True, | |
| num_workers=2 | |
| ) | |
| val_loader = DataLoader( | |
| val_dataset, | |
| batch_size=batch_size, | |
| shuffle=False, | |
| num_workers=2 | |
| ) | |
| test_loader = DataLoader( | |
| test_dataset, | |
| batch_size=batch_size, | |
| shuffle=False, | |
| num_workers=2 | |
| ) | |
| return train_loader, val_loader, test_loader | |
| # TENSORFLOW DATA PIPELINE | |
| def get_tensorflow_data(data_dir="data", batch_size=64): | |
| data_dir = Path(data_dir) | |
| train_dir = data_dir / "seg_train" | |
| test_dir = data_dir / "seg_test" | |
| train_ds = tf.keras.utils.image_dataset_from_directory( | |
| str(train_dir), | |
| image_size=IMAGE_SIZE, | |
| batch_size=batch_size, | |
| shuffle=True, | |
| seed=42, | |
| validation_split=0.2, | |
| subset="training", | |
| label_mode="int", | |
| class_names=CLASSES | |
| ) | |
| val_ds = tf.keras.utils.image_dataset_from_directory( | |
| str(train_dir), | |
| image_size=IMAGE_SIZE, | |
| batch_size=batch_size, | |
| shuffle=True, | |
| seed=42, | |
| validation_split=0.2, | |
| subset="validation", | |
| label_mode="int", | |
| class_names=CLASSES | |
| ) | |
| test_ds = tf.keras.utils.image_dataset_from_directory( | |
| str(test_dir), | |
| image_size=IMAGE_SIZE, | |
| batch_size=batch_size, | |
| shuffle=False, | |
| label_mode="int", | |
| class_names=CLASSES | |
| ) | |
| # AUGMENTATION | |
| augmentation = tf.keras.Sequential([ | |
| tf.keras.layers.RandomFlip("horizontal"), | |
| tf.keras.layers.RandomRotation(0.1), | |
| tf.keras.layers.RandomZoom(0.1), | |
| tf.keras.layers.RandomContrast(0.2), | |
| ]) | |
| normalization = tf.keras.layers.Rescaling(1.0 / 255) | |
| train_ds = ( | |
| train_ds | |
| .map(lambda x, y: (augmentation(x, training=True), y), | |
| num_parallel_calls=tf.data.AUTOTUNE) | |
| .map(lambda x, y: (normalization(x), y), | |
| num_parallel_calls=tf.data.AUTOTUNE) | |
| .cache() | |
| .shuffle(1000) | |
| .prefetch(tf.data.AUTOTUNE) | |
| ) | |
| val_ds = ( | |
| val_ds | |
| .map(lambda x, y: (normalization(x), y), | |
| num_parallel_calls=tf.data.AUTOTUNE) | |
| .cache() | |
| .prefetch(tf.data.AUTOTUNE) | |
| ) | |
| test_ds = ( | |
| test_ds | |
| .map(lambda x, y: (normalization(x), y), | |
| num_parallel_calls=tf.data.AUTOTUNE) | |
| .cache() | |
| .prefetch(tf.data.AUTOTUNE) | |
| ) | |
| return train_ds, val_ds, test_ds |