working server
This commit is contained in:
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -16,6 +16,7 @@ class Bird_CNN(nn.Module):
|
|||||||
|
|
||||||
nn.Conv2d(c_hidden, c_hidden, kernel_size, padding=kernel_size//2),
|
nn.Conv2d(c_hidden, c_hidden, kernel_size, padding=kernel_size//2),
|
||||||
nn.ReLU(),
|
nn.ReLU(),
|
||||||
|
|
||||||
nn.Flatten(),
|
nn.Flatten(),
|
||||||
nn.Linear(c_hidden * img_height * img_width, c_out)
|
nn.Linear(c_hidden * img_height * img_width, c_out)
|
||||||
)
|
)
|
||||||
|
|||||||
+5
-5
@@ -10,6 +10,7 @@ from enum import Enum
|
|||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
SAVE_PATH = "./saved_models"
|
SAVE_PATH = "./saved_models"
|
||||||
|
#IMAGE_SIZE = (1141, 850)
|
||||||
IMAGE_SIZE = (300, 300)
|
IMAGE_SIZE = (300, 300)
|
||||||
|
|
||||||
class bird_species(Enum):
|
class bird_species(Enum):
|
||||||
@@ -28,11 +29,10 @@ transform = transforms.Compose([
|
|||||||
|
|
||||||
dataset = datasets.ImageFolder("data/train", transform=transform)
|
dataset = datasets.ImageFolder("data/train", transform=transform)
|
||||||
|
|
||||||
train_size = int(0.3 * len(dataset))
|
train_size = int(0.8 * len(dataset))
|
||||||
val_size = int(len(dataset) * 0.3)
|
val_size = len(dataset) - train_size
|
||||||
throw_away = len(dataset) - val_size - train_size
|
|
||||||
|
|
||||||
train_dataset, val_dataset, _ = random_split(dataset, [train_size, val_size, throw_away])
|
train_dataset, val_dataset = random_split(dataset, [train_size, val_size])
|
||||||
|
|
||||||
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
|
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
|
||||||
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)
|
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)
|
||||||
@@ -48,7 +48,7 @@ loss_module = nn.CrossEntropyLoss()
|
|||||||
trainCNN(model, optimizer, loss_module, train_loader, val_loader, device, 50, SAVE_PATH=SAVE_PATH, save=True)
|
trainCNN(model, optimizer, loss_module, train_loader, val_loader, device, 50, SAVE_PATH=SAVE_PATH, save=True)
|
||||||
exit()
|
exit()
|
||||||
|
|
||||||
image_path = "test.jpg"
|
image_path = "test1.jpg"
|
||||||
|
|
||||||
image = Image.open(image_path).convert("RGB")
|
image = Image.open(image_path).convert("RGB")
|
||||||
image = transform(image)
|
image = transform(image)
|
||||||
|
|||||||
Binary file not shown.
@@ -0,0 +1,61 @@
|
|||||||
|
from enum import Enum
|
||||||
|
import os
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torchvision import transforms
|
||||||
|
from fastapi import FastAPI, File, UploadFile
|
||||||
|
from PIL import Image
|
||||||
|
import io
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from bird_cnn import Bird_CNN
|
||||||
|
|
||||||
|
SAVE_PATH = "./saved_models"
|
||||||
|
#IMAGE_SIZE = (1141, 850)
|
||||||
|
IMAGE_SIZE = (300, 300)
|
||||||
|
|
||||||
|
transform = transforms.Compose([
|
||||||
|
transforms.Resize(IMAGE_SIZE),
|
||||||
|
transforms.ToTensor()
|
||||||
|
])
|
||||||
|
|
||||||
|
class bird_species(Enum):
|
||||||
|
Common_Kingfisher = 0
|
||||||
|
CommonMyna = 1
|
||||||
|
House_Crow = 2
|
||||||
|
Indian_Peacock = 3
|
||||||
|
Indian_Pitta = 4
|
||||||
|
Ruddy_Shelduck = 5
|
||||||
|
Sarus_Crane = 6
|
||||||
|
|
||||||
|
transform = transforms.Compose([
|
||||||
|
transforms.Resize(IMAGE_SIZE),
|
||||||
|
transforms.ToTensor()
|
||||||
|
])
|
||||||
|
|
||||||
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
|
||||||
|
model = Bird_CNN(c_in=3, c_hidden=15, c_out=7, kernel_size=3, img_width=IMAGE_SIZE[0], img_height=IMAGE_SIZE[1])
|
||||||
|
full_path = os.path.join(SAVE_PATH, "bird_cnn", "bird_cnn")
|
||||||
|
model.load_state_dict(torch.load(full_path, weights_only=False))
|
||||||
|
model.to(device)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
app = FastAPI()
|
||||||
|
|
||||||
|
@app.post("/predict")
|
||||||
|
async def predict(file: UploadFile = File(...)):
|
||||||
|
image_bytes = await file.read()
|
||||||
|
|
||||||
|
image = Image.open(io.BytesIO(image_bytes)).convert("RGB")
|
||||||
|
image = transform(image).unsqueeze(0).to(device)
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
pred = model(image)
|
||||||
|
probs = F.softmax(pred, dim=1)
|
||||||
|
confidence, cls = torch.max(probs, dim=1)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"class": bird_species(cls.item()).name,
|
||||||
|
"confidence": confidence.item()
|
||||||
|
}
|
||||||
Binary file not shown.
|
After Width: | Height: | Size: 57 KiB |
@@ -0,0 +1,22 @@
|
|||||||
|
from PIL import Image
|
||||||
|
import os
|
||||||
|
|
||||||
|
folder = "data/train"
|
||||||
|
|
||||||
|
widths = []
|
||||||
|
heights = []
|
||||||
|
|
||||||
|
for root, _, files in os.walk(folder):
|
||||||
|
for file in files:
|
||||||
|
if file.endswith((".jpg", ".png", ".jpeg")):
|
||||||
|
path = os.path.join(root, file)
|
||||||
|
img = Image.open(path)
|
||||||
|
w, h = img.size
|
||||||
|
widths.append(w)
|
||||||
|
heights.append(h)
|
||||||
|
|
||||||
|
avg_w = sum(widths) / len(widths)
|
||||||
|
avg_h = sum(heights) / len(heights)
|
||||||
|
|
||||||
|
print("Average width:", avg_w)
|
||||||
|
print("Average height:", avg_h)
|
||||||
Reference in New Issue
Block a user