danielle2035's picture
Add file
d32e728
Raw
History Blame Contribute Delete
4.15 kB
import argparse
import torch
import tensorflow as tf
from torch.utils.data import DataLoader, random_split
from torchvision import datasets, transforms
from models.cnn_pytorch import CNN
from models.cnn_tensorflow import build_model
from utils.prep import CLASSES
def get_data_loaders(data_path, batch_size=32):
transform = transforms.Compose([
transforms.Resize((150, 150)),
transforms.ToTensor()
])
train_data = datasets.ImageFolder(
f"{data_path}/seg_train/seg_train",
transform=transform
)
test_data = datasets.ImageFolder(
f"{data_path}/seg_test/seg_test",
transform=transform
)
val_size = int(0.2 * len(train_data))
train_size = len(train_data) - val_size
train_data, val_data = random_split(train_data, [train_size, val_size])
train_loader = DataLoader(train_data, batch_size=batch_size, shuffle=True)
val_loader = DataLoader(val_data, batch_size=batch_size, shuffle=False)
test_loader = DataLoader(test_data, batch_size=batch_size, shuffle=False)
return train_loader, val_loader, test_loader
def train_pytorch(model, train_loader, val_loader, epochs, device):
model.to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-4)
criterion = torch.nn.CrossEntropyLoss()
best_loss = float("inf")
for epoch in range(epochs):
# TRAIN
model.train()
total_loss, correct, total = 0, 0, 0
for x, y in train_loader:
x, y = x.to(device), y.to(device)
optimizer.zero_grad()
outputs = model(x)
loss = criterion(outputs, y)
loss.backward()
optimizer.step()
total_loss += loss.item() * y.size(0)
preds = outputs.argmax(1)
correct += (preds == y).sum().item()
total += y.size(0)
train_acc = 100 * correct / total
train_loss = total_loss / total
# VALIDATION
model.eval()
val_loss, val_correct, val_total = 0, 0, 0
with torch.no_grad():
for x, y in val_loader:
x, y = x.to(device), y.to(device)
outputs = model(x)
loss = criterion(outputs, y)
val_loss += loss.item() * y.size(0)
preds = outputs.argmax(1)
val_correct += (preds == y).sum().item()
val_total += y.size(0)
val_acc = 100 * val_correct / val_total
val_loss = val_loss / val_total
print(f"Epoch {epoch+1}/{epochs} | "
f"Train Loss {train_loss:.4f} Acc {train_acc:.2f}% | "
f"Val Loss {val_loss:.4f} Acc {val_acc:.2f}%")
# save best model
if val_loss < best_loss:
best_loss = val_loss
torch.save(model.state_dict(), "pytorch_model.pth")
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--model", required=True, choices=["pytorch", "tensorflow"])
parser.add_argument("--epochs", type=int, default=25)
parser.add_argument("--data", type=str, required=True)
args = parser.parse_args()
device = "cuda" if torch.cuda.is_available() else "cpu"
# PYTORCH
if args.model == "pytorch":
train_loader, val_loader, test_loader = get_data_loaders(args.data)
model = CNN(num_classes=len(CLASSES))
train_pytorch(model, train_loader, val_loader, args.epochs, device)
# TENSORFLOW
else:
train_ds = tf.keras.preprocessing.image_dataset_from_directory(
f"{args.data}/seg_train/seg_train",
image_size=(150, 150),
batch_size=32
)
val_ds = tf.keras.preprocessing.image_dataset_from_directory(
f"{args.data}/seg_test/seg_test",
image_size=(150, 150),
batch_size=32
)
model = build_model(num_classes=len(CLASSES))
model.fit(
train_ds,
validation_data=val_ds,
epochs=args.epochs
)
model.save("tensorflow_model.keras")
if __name__ == "__main__":
main()