added review sql database

This commit is contained in:
2026-05-12 17:56:31 +02:00
parent 58af9f0f82
commit 941cfba332
8 changed files with 124 additions and 29 deletions
+2
View File
@@ -2,3 +2,5 @@ api/data/
api/saved_models/ api/saved_models/
api/testimages/ api/testimages/
api/__pycache__/ api/__pycache__/
**/*.env
+1 -1
View File
@@ -134,7 +134,7 @@ def trainCNN(model, optimizer, loss_module, train_data_loader, validation_data_l
save_dir = os.path.join(SAVE_PATH, "bird_cnn") save_dir = os.path.join(SAVE_PATH, "bird_cnn")
os.makedirs(save_dir, exist_ok=True) os.makedirs(save_dir, exist_ok=True)
save_path = os.path.join(save_dir, f"bird_cnn{epoch+1}") save_path = os.path.join(save_dir, f"bird_cnn")
torch.save(model.state_dict(), save_path) torch.save(model.state_dict(), save_path)
print(f"epoch: {epoch+1} | train accuracy: {int(train_acc * 1000) / 10}% | validation accuracy: {int(val_acc * 1000) / 10}%") print(f"epoch: {epoch+1} | train accuracy: {int(train_acc * 1000) / 10}% | validation accuracy: {int(val_acc * 1000) / 10}%")
+3
View File
@@ -7,3 +7,6 @@ python-multipart==0.0.27
opencv-python==4.13.0.92 opencv-python==4.13.0.92
pillow==12.2.0 pillow==12.2.0
pycocotools==2.0.11 pycocotools==2.0.11
pydantic==2.13.3
pydantic_core==2.46.3
+63 -1
View File
@@ -1,9 +1,12 @@
from enum import Enum from enum import Enum
import os import os
import sqlite3
from dotenv import load_dotenv
from pydantic import BaseModel
import torch import torch
from torchvision import transforms from torchvision import transforms
from fastapi import FastAPI, File, UploadFile from fastapi import Depends, FastAPI, File, HTTPException, Header, UploadFile
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from PIL import Image from PIL import Image
import io import io
@@ -73,3 +76,62 @@ async def predict(file: UploadFile = File(...)):
"class": bird_species(cls.item()).name, "class": bird_species(cls.item()).name,
"confidence": confidence.item() "confidence": confidence.item()
} }
load_dotenv("database/.env")
API_KEY = os.getenv("API_KEY")
def get_api_key(authorization: str = Header(None)):
if authorization != f"Bearer {API_KEY}":
raise HTTPException(status_code=401, detail="Unauthorized")
return authorization
class Review(BaseModel):
website: str
rating: int
text: str
date: str
@app.get("/review")
def get_reviews():
conn = sqlite3.connect("database/reviews.db")
cur = conn.cursor()
cur.execute("""
SELECT * FROM reviews;
""")
data = cur.fetchall()
print(data)
conn.commit()
conn.close()
return {"data": data}
@app.post("/review")
def post_review(review: Review):
conn = sqlite3.connect("database/reviews.db")
cur = conn.cursor()
cur.execute("""
INSERT INTO reviews (website, rating, text, created_at)
VALUES (?, ?, ?, ?)
""", (review.website, review.rating, review.text, review.date))
conn.commit()
conn.close()
return {"status": "ok"}
@app.delete("/review")
def post_review(auth=Depends(get_api_key)):
conn = sqlite3.connect("database/reviews.db")
cur = conn.cursor()
cur.execute("""
DELETE FROM reviews;
""")
conn.commit()
conn.close()
return {"status": "ok"}
+20 -9
View File
@@ -5,6 +5,7 @@ from torchvision import datasets, transforms
from torch.utils.data import DataLoader, random_split from torch.utils.data import DataLoader, random_split
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
from util import TransformedSubset, visualizeData
from bird_cnn import Bird_CNN, sample, trainCNN from bird_cnn import Bird_CNN, sample, trainCNN
from enum import Enum from enum import Enum
@@ -12,7 +13,7 @@ from enum import Enum
from PIL import Image from PIL import Image
SAVE_PATH = "./saved_models" SAVE_PATH = "./saved_models"
IMAGE_SIZE = 64 IMAGE_SIZE = 128
class bird_species(Enum): class bird_species(Enum):
Common_Kingfisher = 0 Common_Kingfisher = 0
@@ -23,24 +24,34 @@ class bird_species(Enum):
Ruddy_Shelduck = 5 Ruddy_Shelduck = 5
Sarus_Crane = 6 Sarus_Crane = 6
transform = transforms.Compose([ transform_augemnt = transforms.Compose([
transforms.RandomAffine(
degrees=35, # no rotation
translate=(0.2, 0.2) # shift up to 20% horizontally/vertically
),
transforms.RandomHorizontalFlip(p=0.5), transforms.RandomHorizontalFlip(p=0.5),
transforms.RandomRotation(35),
transforms.Resize(IMAGE_SIZE), transforms.Resize(IMAGE_SIZE),
transforms.CenterCrop(IMAGE_SIZE), transforms.CenterCrop(IMAGE_SIZE),
transforms.ToTensor() transforms.ToTensor()
]) ])
dataset = datasets.ImageFolder("data/CUB_200_2011/images", transform=transform) transform = transforms.Compose([
transforms.Resize(IMAGE_SIZE),
transforms.CenterCrop(IMAGE_SIZE),
transforms.ToTensor()
])
dataset = datasets.ImageFolder("data/CUB_200_2011/images")
train_size = int(0.8 * len(dataset)) train_size = int(0.8 * len(dataset))
val_size = len(dataset) - train_size val_size = len(dataset) - train_size
train_dataset, val_dataset = random_split(dataset, [train_size, val_size]) train_subset, val_subset = random_split(dataset, [train_size, val_size])
train_dataset = TransformedSubset(train_subset, transform_augemnt)
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True) val_dataset = TransformedSubset(val_subset, transform)
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=64, shuffle=False)
device = torch.device("cpu") if not torch.cuda.is_available() else torch.device("cuda:0") device = torch.device("cpu") if not torch.cuda.is_available() else torch.device("cuda:0")
print("Using device", device) print("Using device", device)
@@ -50,7 +61,7 @@ model.to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
loss_module = nn.CrossEntropyLoss() 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, 500, SAVE_PATH=SAVE_PATH, save=True)
exit() exit()
image_path = "test1.jpg" image_path = "test1.jpg"
+30 -15
View File
@@ -1,22 +1,37 @@
from PIL import Image from PIL import Image
import os import os
from matplotlib import pyplot as plt
import torch
folder = "data/train" class TransformedSubset(torch.utils.data.Dataset):
def __init__(self, subset, transform=None):
self.subset = subset
self.transform = transform
widths = [] def __getitem__(self, idx):
heights = [] x, y = self.subset[idx]
for root, _, files in os.walk(folder): if self.transform:
for file in files: x = self.transform(x)
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) return x, y
avg_h = sum(heights) / len(heights)
print("Average width:", avg_w) def __len__(self):
print("Average height:", avg_h) return len(self.subset)
def visualizeData(dataset):
images, labels = next(iter(dataset))
for i in range(4):
img = images[i]
# Convert tensor shape from [C,H,W] -> [H,W,C]
img = img.permute(1, 2, 0)
plt.figure(figsize=(3,3))
plt.imshow(img)
plt.title(f"Label: {labels[i].item()}")
plt.axis("off")
plt.show()
+2
View File
@@ -29,6 +29,8 @@ services:
- cnn_network - cnn_network
pull_policy: never pull_policy: never
container_name: cnn_api container_name: cnn_api
volumes:
- ./database:/api/database
cnn_website: cnn_website:
Binary file not shown.