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()