-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathgenerate.py
More file actions
42 lines (38 loc) · 1.53 KB
/
Copy pathgenerate.py
File metadata and controls
42 lines (38 loc) · 1.53 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
from core import simpleNN
import torch
import torch.nn.functional as F
import numpy as np
import matplotlib.pyplot as plt
import torchvision.transforms as transforms
from torchvision import datasets
def dataset_loader():
transform = transforms.Compose([
transforms.ToTensor(),transforms.Normalize((0.5,), (0.5,))
])
train_loader = torch.utils.data.DataLoader(datasets.MNIST('./data', train=True, download=True,
transform=transform), batch_size=64, shuffle=True)
data_iterator = iter(train_loader)
images, _ = next(data_iterator) # İlk batch'in resimlerini al
return images
def generate(model, device, images):
model.eval()
with torch.no_grad():
outputs = model(images.to(device))
_, predicted = torch.max(outputs, 1)
labels = predicted.cpu().numpy()
images = images.cpu().numpy()
plt.figure(figsize=(10, 5))
for i in range(10):
plt.subplot(2, 5, i + 1)
plt.imshow(images[i].reshape(28, 28), cmap='gray')
plt.title(f'Tahmin: {labels[i]}')
plt.axis('off')
plt.show()
model_path='mnist_model.pth'
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = simpleNN().to(device)
try:
model.load_state_dict(torch.load(model_path, map_location=device))
except FileNotFoundError:
print(f"Model dosyası '{model_path}' bulunamadı. Lütfen önce modeli eğitin ve kaydedin.")
generate(model,device,images=dataset_loader())