danielle2035's picture
Add file
d32e728
Raw
History Blame Contribute Delete
4.23 kB
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