49 lines
1.4 KiB
Python
49 lines
1.4 KiB
Python
import torch
|
|
import torch.optim as optim
|
|
import torch.nn as nn
|
|
from datasets import get_data_loaders
|
|
from model import SimpleCNN
|
|
from utils import save_model, calculate_accuracy
|
|
|
|
# Настройки
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
config = {
|
|
"batch_size": 128,
|
|
"epochs": 10,
|
|
"learning_rate": 0.001,
|
|
"data_dir": "./data",
|
|
"logs_dir": "./logs",
|
|
}
|
|
|
|
# Загрузка данных
|
|
train_loader, test_loader = get_data_loaders(config["data_dir"], config["batch_size"])
|
|
|
|
# Модель, критерий, оптимизатор
|
|
model = SimpleCNN().to(device)
|
|
criterion = nn.CrossEntropyLoss()
|
|
optimizer = optim.Adam(model.parameters(), lr=config["learning_rate"])
|
|
|
|
# Тренировка модели
|
|
for epoch in range(config["epochs"]):
|
|
model.train()
|
|
running_loss = 0.0
|
|
for inputs, labels in train_loader:
|
|
inputs, labels = inputs.to(device), labels.to(device)
|
|
|
|
optimizer.zero_grad()
|
|
outputs = model(inputs)
|
|
loss = criterion(outputs, labels)
|
|
loss.backward()
|
|
optimizer.step()
|
|
|
|
running_loss += loss.item()
|
|
|
|
print(f"Epoch {epoch+1}/{config['epochs']}, Loss: {running_loss/len(train_loader)}")
|
|
|
|
# Тестирование
|
|
accuracy = calculate_accuracy(model, test_loader, device)
|
|
print(f"Test Accuracy: {accuracy:.2f}%")
|
|
|
|
# Сохранение модели
|
|
save_model(model, "./model.pth")
|