-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
78 lines (64 loc) · 2.5 KB
/
Copy pathmain.py
File metadata and controls
78 lines (64 loc) · 2.5 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
import torch
import os
from model import CustomCNN, VGG16FeatureExtractor, VGG16FineTuned
from data import get_data_loaders
from train import train_model, visualize_features, plot_training_history, evaluate_model
import pandas as pd
from datetime import datetime
def main():
# Set device
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")
# Create output directory
output_dir = "results"
os.makedirs(output_dir, exist_ok=True)
# Load data
data_dir = "Flowers/train" # Update this path to your dataset location
train_loader, val_loader, class_names = get_data_loaders(data_dir)
# Initialize models
models = {
'CustomCNN': CustomCNN(num_classes=len(class_names)).to(device),
'VGG16FeatureExtractor': VGG16FeatureExtractor(num_classes=len(class_names)).to(device),
'VGG16FineTuned': VGG16FineTuned(num_classes=len(class_names)).to(device)
}
# Train and evaluate each model
results = []
for model_name, model in models.items():
print(f"\nTraining {model_name}...")
# Train model
history, training_time = train_model(
model=model,
train_loader=train_loader,
val_loader=val_loader,
num_epochs=10,
learning_rate=0.001,
device=device
)
# Plot training history
plot_training_history(
history,
save_path=os.path.join(output_dir, f"{model_name}_training_history.png")
)
# Visualize features
visualize_features(
model=model,
data_loader=val_loader,
device=device,
save_path=os.path.join(output_dir, f"{model_name}_feature_maps.png")
)
# Evaluate model
metrics = evaluate_model(model, val_loader, device)
metrics['model'] = model_name
metrics['training_time'] = training_time
results.append(metrics)
# Save model
torch.save(model.state_dict(), os.path.join(output_dir, f"{model_name}.pth"))
# Create comparison table
results_df = pd.DataFrame(results)
results_df = results_df[['model', 'accuracy', 'precision', 'recall', 'f1', 'training_time']]
results_df.to_csv(os.path.join(output_dir, 'model_comparison.csv'), index=False)
# Print results
print("\nModel Comparison:")
print(results_df.to_string(index=False))
if __name__ == "__main__":
main()