Spaces:
Sleeping
Sleeping
| 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() |