added review sql database
This commit is contained in:
@@ -2,3 +2,5 @@ api/data/
|
||||
api/saved_models/
|
||||
api/testimages/
|
||||
api/__pycache__/
|
||||
|
||||
**/*.env
|
||||
+1
-1
@@ -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")
|
||||
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)
|
||||
|
||||
print(f"epoch: {epoch+1} | train accuracy: {int(train_acc * 1000) / 10}% | validation accuracy: {int(val_acc * 1000) / 10}%")
|
||||
|
||||
@@ -7,3 +7,6 @@ python-multipart==0.0.27
|
||||
opencv-python==4.13.0.92
|
||||
pillow==12.2.0
|
||||
pycocotools==2.0.11
|
||||
|
||||
pydantic==2.13.3
|
||||
pydantic_core==2.46.3
|
||||
+63
-1
@@ -1,9 +1,12 @@
|
||||
from enum import Enum
|
||||
import os
|
||||
import sqlite3
|
||||
|
||||
from dotenv import load_dotenv
|
||||
from pydantic import BaseModel
|
||||
import torch
|
||||
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 PIL import Image
|
||||
import io
|
||||
@@ -73,3 +76,62 @@ async def predict(file: UploadFile = File(...)):
|
||||
"class": bird_species(cls.item()).name,
|
||||
"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
@@ -5,6 +5,7 @@ from torchvision import datasets, transforms
|
||||
from torch.utils.data import DataLoader, random_split
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
from util import TransformedSubset, visualizeData
|
||||
from bird_cnn import Bird_CNN, sample, trainCNN
|
||||
|
||||
from enum import Enum
|
||||
@@ -12,7 +13,7 @@ from enum import Enum
|
||||
from PIL import Image
|
||||
|
||||
SAVE_PATH = "./saved_models"
|
||||
IMAGE_SIZE = 64
|
||||
IMAGE_SIZE = 128
|
||||
|
||||
class bird_species(Enum):
|
||||
Common_Kingfisher = 0
|
||||
@@ -23,24 +24,34 @@ class bird_species(Enum):
|
||||
Ruddy_Shelduck = 5
|
||||
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.RandomRotation(35),
|
||||
transforms.Resize(IMAGE_SIZE),
|
||||
transforms.CenterCrop(IMAGE_SIZE),
|
||||
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))
|
||||
val_size = len(dataset) - train_size
|
||||
|
||||
train_dataset, val_dataset = random_split(dataset, [train_size, val_size])
|
||||
|
||||
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
|
||||
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)
|
||||
train_subset, val_subset = random_split(dataset, [train_size, val_size])
|
||||
train_dataset = TransformedSubset(train_subset, transform_augemnt)
|
||||
val_dataset = TransformedSubset(val_subset, transform)
|
||||
|
||||
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")
|
||||
print("Using device", device)
|
||||
@@ -50,7 +61,7 @@ model.to(device)
|
||||
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
|
||||
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()
|
||||
|
||||
image_path = "test1.jpg"
|
||||
|
||||
+30
-15
@@ -1,22 +1,37 @@
|
||||
from PIL import Image
|
||||
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 = []
|
||||
heights = []
|
||||
def __getitem__(self, idx):
|
||||
x, y = self.subset[idx]
|
||||
|
||||
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)
|
||||
if self.transform:
|
||||
x = self.transform(x)
|
||||
|
||||
avg_w = sum(widths) / len(widths)
|
||||
avg_h = sum(heights) / len(heights)
|
||||
return x, y
|
||||
|
||||
print("Average width:", avg_w)
|
||||
print("Average height:", avg_h)
|
||||
def __len__(self):
|
||||
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()
|
||||
@@ -29,6 +29,8 @@ services:
|
||||
- cnn_network
|
||||
pull_policy: never
|
||||
container_name: cnn_api
|
||||
volumes:
|
||||
- ./database:/api/database
|
||||
|
||||
|
||||
cnn_website:
|
||||
|
||||
Binary file not shown.
Reference in New Issue
Block a user