Files
2025-01-13 11:18:03 +03:00

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