changed cnn architecture, added website image upload

This commit is contained in:
2026-04-29 14:02:56 +02:00
parent fd1c111739
commit 262dbe9f98
18 changed files with 264 additions and 51 deletions
+2 -1
View File
@@ -1 +1,2 @@
bird_cnn/data/ bird_cnn/data/
bird_cnn/saved_models/
+16
View File
@@ -0,0 +1,16 @@
FROM python:3.11-slim
# Set working directory
WORKDIR /app
# Copy requirements first (better Docker layer caching)
COPY . .
# Install Python dependencies
RUN pip install --no-cache-dir -r requirements.txt
# Make sure the startup script is executable
RUN chmod +x start_server.sh
# Use the script as the container entrypoint
ENTRYPOINT ["./start_server.sh"]
Binary file not shown.

After

Width:  |  Height:  |  Size: 218 KiB

Binary file not shown.
Binary file not shown.
+62 -11
View File
@@ -3,26 +3,77 @@ import os
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
import numpy as np
from tqdm import tqdm from tqdm import tqdm
class SeparableConvolution(nn.Module):
def __init__(self, c_in, c_out, kernel_size):
super().__init__()
self.depthwise = nn.Conv2d(c_in, c_in, kernel_size, groups=c_in, padding=kernel_size//2)
self.bn1 = nn.BatchNorm2d(c_in)
self.pointwise = nn.Conv2d(c_in, c_out, kernel_size=1)
self.bn2 = nn.BatchNorm2d(c_out)
def forward(self, x):
x = self.depthwise(x)
x = self.bn1(x)
x = F.relu(x)
x = self.pointwise(x)
x = self.bn2(x)
x = F.relu(x)
return x
class Bird_CNN(nn.Module): class Bird_CNN(nn.Module):
def __init__(self, c_in, c_hidden, c_out, kernel_size, img_width, img_height): def __init__(self, c_in, c_hidden, c_out):
super().__init__() super().__init__()
self.model = nn.Sequential(
nn.Conv2d(c_in, c_hidden, kernel_size, padding=kernel_size//2),
nn.ReLU(),
nn.Conv2d(c_hidden, c_hidden, kernel_size, padding=kernel_size//2), self.conv_init = nn.Sequential(
nn.Conv2d(c_in, c_hidden, kernel_size=3, padding=1),
nn.BatchNorm2d(c_hidden),
nn.ReLU(), nn.ReLU(),
nn.Conv2d(c_hidden, c_hidden, 3, stride=2, padding=1)
nn.Flatten(),
nn.Linear(c_hidden * img_height * img_width, c_out)
) )
# 1x1 conv branch
self.branch1 = SeparableConvolution(c_in=c_hidden, c_out=64, kernel_size=1)
# 1x1 -> 3x3 conv branch
self.branch2 = SeparableConvolution(c_in=c_hidden, c_out=128, kernel_size=3)
# 1x1 -> 5x5 conv branch
self.branch3 = SeparableConvolution(c_in=c_hidden, c_out=32, kernel_size=5)
# 3x3 max pooling -> 1x1 conv branch
self.branch4 = nn.Sequential(
nn.MaxPool2d(kernel_size=3, stride=1, padding=1),
nn.Conv2d(c_hidden, 32, kernel_size=1),
nn.ReLU()
)
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
self.flatten = nn.Flatten()
self.linear = nn.Linear(256, c_out)
self.dropout = nn.Dropout(0.3)
def forward(self, x): def forward(self, x):
return self.model(x) x = self.conv_init(x)
b1 = self.branch1(x)
b2 = self.branch2(x)
b3 = self.branch3(x)
b4 = self.branch4(x)
x = torch.cat([b1, b2, b3, b4], dim=1)
x = F.relu(x)
x = self.avgpool(x)
x = torch.flatten(x, 1)
x = self.dropout(x)
return self.linear(x)
def trainCNN(model, optimizer, loss_module, train_data_loader, validation_data_loader, device, num_epochs, SAVE_PATH, save=False): def trainCNN(model, optimizer, loss_module, train_data_loader, validation_data_loader, device, num_epochs, SAVE_PATH, save=False):
@@ -79,7 +130,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, "bird_cnn") save_path = os.path.join(save_dir, f"bird_cnn{epoch+1}")
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 -3
View File
@@ -34,13 +34,13 @@ val_size = len(dataset) - train_size
train_dataset, val_dataset = random_split(dataset, [train_size, val_size]) 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=26, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False) val_loader = DataLoader(val_dataset, batch_size=26, 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)
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]) model = Bird_CNN(c_in=3, c_hidden=15, c_out=7)
model.to(device) 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()
+7
View File
@@ -0,0 +1,7 @@
fastapi==0.136.1
networkx==3.6.1
numpy==2.3.4
torch==2.11.0+cu126
torchvision==0.26.0+cu126
tqdm==4.67.3
uvicorn==0.46.0
+13
View File
@@ -4,6 +4,7 @@ import os
import torch import torch
from torchvision import transforms from torchvision import transforms
from fastapi import FastAPI, File, UploadFile from fastapi import FastAPI, File, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from PIL import Image from PIL import Image
import io import io
import torch.nn.functional as F import torch.nn.functional as F
@@ -43,6 +44,18 @@ model.eval()
app = FastAPI() app = FastAPI()
origins = [
"http://localhost:5173",
]
app.add_middleware(
CORSMiddleware,
allow_origins=origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.post("/predict") @app.post("/predict")
async def predict(file: UploadFile = File(...)): async def predict(file: UploadFile = File(...)):
image_bytes = await file.read() image_bytes = await file.read()
+1
View File
@@ -0,0 +1 @@
python -m uvicorn server:app --reload
+60 -26
View File
@@ -9,7 +9,8 @@
"version": "0.0.0", "version": "0.0.0",
"dependencies": { "dependencies": {
"react": "^19.2.5", "react": "^19.2.5",
"react-dom": "^19.2.5" "react-dom": "^19.2.5",
"react-router-dom": "^7.14.2"
}, },
"devDependencies": { "devDependencies": {
"@eslint/js": "^10.0.1", "@eslint/js": "^10.0.1",
@@ -264,31 +265,6 @@
"node": ">=6.9.0" "node": ">=6.9.0"
} }
}, },
"node_modules/@emnapi/core": {
"version": "1.10.0",
"resolved": "https://registry.npmjs.org/@emnapi/core/-/core-1.10.0.tgz",
"integrity": "sha512-yq6OkJ4p82CAfPl0u9mQebQHKPJkY7WrIuk205cTYnYe+k2Z8YBh11FrbRG/H6ihirqcacOgl2BIO8oyMQLeXw==",
"dev": true,
"license": "MIT",
"optional": true,
"peer": true,
"dependencies": {
"@emnapi/wasi-threads": "1.2.1",
"tslib": "^2.4.0"
}
},
"node_modules/@emnapi/runtime": {
"version": "1.10.0",
"resolved": "https://registry.npmjs.org/@emnapi/runtime/-/runtime-1.10.0.tgz",
"integrity": "sha512-ewvYlk86xUoGI0zQRNq/mC+16R1QeDlKQy21Ki3oSYXNgLb45GV1P6A0M+/s6nyCuNDqe5VpaY84BzXGwVbwFA==",
"dev": true,
"license": "MIT",
"optional": true,
"peer": true,
"dependencies": {
"tslib": "^2.4.0"
}
},
"node_modules/@emnapi/wasi-threads": { "node_modules/@emnapi/wasi-threads": {
"version": "1.2.1", "version": "1.2.1",
"resolved": "https://registry.npmjs.org/@emnapi/wasi-threads/-/wasi-threads-1.2.1.tgz", "resolved": "https://registry.npmjs.org/@emnapi/wasi-threads/-/wasi-threads-1.2.1.tgz",
@@ -1056,6 +1032,19 @@
"dev": true, "dev": true,
"license": "MIT" "license": "MIT"
}, },
"node_modules/cookie": {
"version": "1.1.1",
"resolved": "https://registry.npmjs.org/cookie/-/cookie-1.1.1.tgz",
"integrity": "sha512-ei8Aos7ja0weRpFzJnEA9UHJ/7XQmqglbRwnf2ATjcB9Wq874VKH9kfjjirM6UhU2/E5fFYadylyhFldcqSidQ==",
"license": "MIT",
"engines": {
"node": ">=18"
},
"funding": {
"type": "opencollective",
"url": "https://opencollective.com/express"
}
},
"node_modules/cross-spawn": { "node_modules/cross-spawn": {
"version": "7.0.6", "version": "7.0.6",
"resolved": "https://registry.npmjs.org/cross-spawn/-/cross-spawn-7.0.6.tgz", "resolved": "https://registry.npmjs.org/cross-spawn/-/cross-spawn-7.0.6.tgz",
@@ -2110,6 +2099,7 @@
"resolved": "https://registry.npmjs.org/react-dom/-/react-dom-19.2.5.tgz", "resolved": "https://registry.npmjs.org/react-dom/-/react-dom-19.2.5.tgz",
"integrity": "sha512-J5bAZz+DXMMwW/wV3xzKke59Af6CHY7G4uYLN1OvBcKEsWOs4pQExj86BBKamxl/Ik5bx9whOrvBlSDfWzgSag==", "integrity": "sha512-J5bAZz+DXMMwW/wV3xzKke59Af6CHY7G4uYLN1OvBcKEsWOs4pQExj86BBKamxl/Ik5bx9whOrvBlSDfWzgSag==",
"license": "MIT", "license": "MIT",
"peer": true,
"dependencies": { "dependencies": {
"scheduler": "^0.27.0" "scheduler": "^0.27.0"
}, },
@@ -2117,6 +2107,44 @@
"react": "^19.2.5" "react": "^19.2.5"
} }
}, },
"node_modules/react-router": {
"version": "7.14.2",
"resolved": "https://registry.npmjs.org/react-router/-/react-router-7.14.2.tgz",
"integrity": "sha512-yCqNne6I8IB6rVCH7XUvlBK7/QKyqypBFGv+8dj4QBFJiiRX+FG7/nkdAvGElyvVZ/HQP5N19wzteuTARXi5Gw==",
"license": "MIT",
"dependencies": {
"cookie": "^1.0.1",
"set-cookie-parser": "^2.6.0"
},
"engines": {
"node": ">=20.0.0"
},
"peerDependencies": {
"react": ">=18",
"react-dom": ">=18"
},
"peerDependenciesMeta": {
"react-dom": {
"optional": true
}
}
},
"node_modules/react-router-dom": {
"version": "7.14.2",
"resolved": "https://registry.npmjs.org/react-router-dom/-/react-router-dom-7.14.2.tgz",
"integrity": "sha512-YZcM5ES8jJSM+KrJ9BdvHHqlnGTg5tH3sC5ChFRj4inosKctdyzBDhOyyHdGk597q2OT6NTrCA1OvB/YDwfekQ==",
"license": "MIT",
"dependencies": {
"react-router": "7.14.2"
},
"engines": {
"node": ">=20.0.0"
},
"peerDependencies": {
"react": ">=18",
"react-dom": ">=18"
}
},
"node_modules/rolldown": { "node_modules/rolldown": {
"version": "1.0.0-rc.17", "version": "1.0.0-rc.17",
"resolved": "https://registry.npmjs.org/rolldown/-/rolldown-1.0.0-rc.17.tgz", "resolved": "https://registry.npmjs.org/rolldown/-/rolldown-1.0.0-rc.17.tgz",
@@ -2174,6 +2202,12 @@
"semver": "bin/semver.js" "semver": "bin/semver.js"
} }
}, },
"node_modules/set-cookie-parser": {
"version": "2.7.2",
"resolved": "https://registry.npmjs.org/set-cookie-parser/-/set-cookie-parser-2.7.2.tgz",
"integrity": "sha512-oeM1lpU/UvhTxw+g3cIfxXHyJRc/uidd3yK1P242gzHds0udQBYzs3y8j4gCCW+ZJ7ad0yctld8RYO+bdurlvw==",
"license": "MIT"
},
"node_modules/shebang-command": { "node_modules/shebang-command": {
"version": "2.0.0", "version": "2.0.0",
"resolved": "https://registry.npmjs.org/shebang-command/-/shebang-command-2.0.0.tgz", "resolved": "https://registry.npmjs.org/shebang-command/-/shebang-command-2.0.0.tgz",
+2 -1
View File
@@ -11,7 +11,8 @@
}, },
"dependencies": { "dependencies": {
"react": "^19.2.5", "react": "^19.2.5",
"react-dom": "^19.2.5" "react-dom": "^19.2.5",
"react-router-dom": "^7.14.2"
}, },
"devDependencies": { "devDependencies": {
"@eslint/js": "^10.0.1", "@eslint/js": "^10.0.1",
+7
View File
@@ -1,15 +1,22 @@
import { useState } from 'react' import { useState } from 'react'
import { BrowserRouter as Router, Routes, Route } from "react-router-dom";
import reactLogo from './assets/react.svg' import reactLogo from './assets/react.svg'
import viteLogo from './assets/vite.svg' import viteLogo from './assets/vite.svg'
import heroImg from './assets/hero.png' import heroImg from './assets/hero.png'
import './App.css' import './App.css'
import Menu from './Menu' import Menu from './Menu'
import Homepage from './Homepage';
import Bird_CNN from './Bird_CNN';
function App() { function App() {
return ( return (
<> <>
<Menu /> <Menu />
<Routes>
<Route path="/" element={<Homepage />} />
<Route path="/bird_cnn" element={<Bird_CNN />} />
</Routes>
</> </>
); );
} }
+5
View File
@@ -0,0 +1,5 @@
.content-block {
display: flex;
align-items: baseline;
gap: 20px;
}
+68
View File
@@ -0,0 +1,68 @@
import { useState } from 'react'
import './Bird_CNN.css'
function Bird_CNN() {
const [file, setFile] = useState(null);
const [birdClass, setBirdClass] = useState(null);
const [confidence, setConfidence] = useState(null);
const [error, setError] = useState(false);
const handleImage = (e) => {
setFile(e.target.files[0]);
};
const sendImage = async (e) => {
if (!file) return;
const formData = new FormData();
formData.append("file", file);
try {
const response = await fetch("http://127.0.0.1:8000/predict", {
method: "POST",
body: formData,
});
if (!response.ok) {
setError(true);
}
else {
setError(false);
}
const result = await response.json();
setBirdClass(result["class"]);
const confidence = result["confidence"];
setConfidence(`${Math.round(confidence * 100)}%`);
}
catch (error) {
console.error("Upload failed:", error);
}
};
return (
<>
<h1>bird-cnn</h1>
<div>
<input type="file" accept="image/jpeg" onChange={handleImage} />
<button onClick={sendImage}>Upload</button>
<div className='response-block'>
{!error && <div className='content-block class'>
<h3>Bird Species: </h3>
<p id='bird-class-text'>{birdClass}</p>
</div>}
{!error && <div className='content-block confidence'>
<h4>Model Confidence: </h4>
<p id='bird-confidence-text'>{confidence}</p>
</div>}
{error && <h4>An Error has uccured. Please try again later.</h4>}
</div>
</div>
</>
);
}
export default Bird_CNN;
+9
View File
@@ -0,0 +1,9 @@
function Homepage() {
return (
<>
<p>Homepage</p>
</>
);
}
export default Homepage;
+3 -6
View File
@@ -1,15 +1,12 @@
import { Link } from 'react-router-dom';
import './Menu.css' import './Menu.css'
function Menu() { function Menu() {
return ( return (
<> <>
<nav id="navbar-main"> <nav id="navbar-main">
<li className="menu-item"> <Link to="/">Homepage</Link>
<a className="link-text">Test</a> <Link to="/bird_cnn">Birds</Link>
</li>
<li className="menu-item">
<a className="link-text">Test2</a>
</li>
</nav> </nav>
</> </>
); );
+6 -3
View File
@@ -2,9 +2,12 @@ import { StrictMode } from 'react'
import { createRoot } from 'react-dom/client' import { createRoot } from 'react-dom/client'
import './index.css' import './index.css'
import App from './App.jsx' import App from './App.jsx'
import { BrowserRouter } from 'react-router-dom'
createRoot(document.getElementById('root')).render( createRoot(document.getElementById('root')).render(
<StrictMode> <StrictMode>
<App /> <BrowserRouter>
</StrictMode>, <App />
) </BrowserRouter>
</StrictMode>
);