Spaces:
Sleeping
Sleeping
Upload folder using huggingface_hub
Browse files- .dockerignore +24 -0
- .gitignore +29 -0
- .gradio/certificate.pem +31 -0
- Dockerfile +21 -0
- LICENSE +21 -0
- README.md +40 -8
- cover.jpg +0 -0
- data/download_data.py +69 -0
- requirements.txt +12 -0
- src/data_preperation.py +57 -0
- src/main.py +125 -0
- src/resnet.py +273 -0
- src/utils.py +430 -0
- tests/test_reconstruct_image.py +152 -0
- tests/test_sliding_window.py +186 -0
- weights/download_weights.py +4 -0
.dockerignore
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
data/
|
| 2 |
+
tests/
|
| 3 |
+
*.jpg
|
| 4 |
+
|
| 5 |
+
# Ignore version control system directories
|
| 6 |
+
.git/
|
| 7 |
+
.gitignore
|
| 8 |
+
|
| 9 |
+
# Ignore Python cache and environment files
|
| 10 |
+
__pycache__/
|
| 11 |
+
*.pyc
|
| 12 |
+
*.pyo
|
| 13 |
+
*.pyd
|
| 14 |
+
.env
|
| 15 |
+
|
| 16 |
+
# Ignore IDE and editor files
|
| 17 |
+
.vscode/
|
| 18 |
+
*.swp
|
| 19 |
+
.idea/
|
| 20 |
+
|
| 21 |
+
# Ignore logs and temporary files
|
| 22 |
+
*.log
|
| 23 |
+
*.tmp
|
| 24 |
+
*.bak
|
.gitignore
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# IntelliJ project files
|
| 2 |
+
.idea
|
| 3 |
+
*.iml
|
| 4 |
+
out
|
| 5 |
+
gen
|
| 6 |
+
|
| 7 |
+
# Data files
|
| 8 |
+
/original
|
| 9 |
+
/raw_data
|
| 10 |
+
/cropped_data
|
| 11 |
+
/data/highres
|
| 12 |
+
|
| 13 |
+
#Python cache
|
| 14 |
+
*.pyc
|
| 15 |
+
__pycache__
|
| 16 |
+
|
| 17 |
+
#Jupyter checkpoint
|
| 18 |
+
.ipynb_checkpoints
|
| 19 |
+
|
| 20 |
+
# Model weights
|
| 21 |
+
*.pth
|
| 22 |
+
*.h5
|
| 23 |
+
*.pkl
|
| 24 |
+
|
| 25 |
+
# Logs
|
| 26 |
+
logs
|
| 27 |
+
|
| 28 |
+
# gradio flag
|
| 29 |
+
flagged
|
.gradio/certificate.pem
ADDED
|
@@ -0,0 +1,31 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
-----BEGIN CERTIFICATE-----
|
| 2 |
+
MIIFazCCA1OgAwIBAgIRAIIQz7DSQONZRGPgu2OCiwAwDQYJKoZIhvcNAQELBQAw
|
| 3 |
+
TzELMAkGA1UEBhMCVVMxKTAnBgNVBAoTIEludGVybmV0IFNlY3VyaXR5IFJlc2Vh
|
| 4 |
+
cmNoIEdyb3VwMRUwEwYDVQQDEwxJU1JHIFJvb3QgWDEwHhcNMTUwNjA0MTEwNDM4
|
| 5 |
+
WhcNMzUwNjA0MTEwNDM4WjBPMQswCQYDVQQGEwJVUzEpMCcGA1UEChMgSW50ZXJu
|
| 6 |
+
ZXQgU2VjdXJpdHkgUmVzZWFyY2ggR3JvdXAxFTATBgNVBAMTDElTUkcgUm9vdCBY
|
| 7 |
+
MTCCAiIwDQYJKoZIhvcNAQEBBQADggIPADCCAgoCggIBAK3oJHP0FDfzm54rVygc
|
| 8 |
+
h77ct984kIxuPOZXoHj3dcKi/vVqbvYATyjb3miGbESTtrFj/RQSa78f0uoxmyF+
|
| 9 |
+
0TM8ukj13Xnfs7j/EvEhmkvBioZxaUpmZmyPfjxwv60pIgbz5MDmgK7iS4+3mX6U
|
| 10 |
+
A5/TR5d8mUgjU+g4rk8Kb4Mu0UlXjIB0ttov0DiNewNwIRt18jA8+o+u3dpjq+sW
|
| 11 |
+
T8KOEUt+zwvo/7V3LvSye0rgTBIlDHCNAymg4VMk7BPZ7hm/ELNKjD+Jo2FR3qyH
|
| 12 |
+
B5T0Y3HsLuJvW5iB4YlcNHlsdu87kGJ55tukmi8mxdAQ4Q7e2RCOFvu396j3x+UC
|
| 13 |
+
B5iPNgiV5+I3lg02dZ77DnKxHZu8A/lJBdiB3QW0KtZB6awBdpUKD9jf1b0SHzUv
|
| 14 |
+
KBds0pjBqAlkd25HN7rOrFleaJ1/ctaJxQZBKT5ZPt0m9STJEadao0xAH0ahmbWn
|
| 15 |
+
OlFuhjuefXKnEgV4We0+UXgVCwOPjdAvBbI+e0ocS3MFEvzG6uBQE3xDk3SzynTn
|
| 16 |
+
jh8BCNAw1FtxNrQHusEwMFxIt4I7mKZ9YIqioymCzLq9gwQbooMDQaHWBfEbwrbw
|
| 17 |
+
qHyGO0aoSCqI3Haadr8faqU9GY/rOPNk3sgrDQoo//fb4hVC1CLQJ13hef4Y53CI
|
| 18 |
+
rU7m2Ys6xt0nUW7/vGT1M0NPAgMBAAGjQjBAMA4GA1UdDwEB/wQEAwIBBjAPBgNV
|
| 19 |
+
HRMBAf8EBTADAQH/MB0GA1UdDgQWBBR5tFnme7bl5AFzgAiIyBpY9umbbjANBgkq
|
| 20 |
+
hkiG9w0BAQsFAAOCAgEAVR9YqbyyqFDQDLHYGmkgJykIrGF1XIpu+ILlaS/V9lZL
|
| 21 |
+
ubhzEFnTIZd+50xx+7LSYK05qAvqFyFWhfFQDlnrzuBZ6brJFe+GnY+EgPbk6ZGQ
|
| 22 |
+
3BebYhtF8GaV0nxvwuo77x/Py9auJ/GpsMiu/X1+mvoiBOv/2X/qkSsisRcOj/KK
|
| 23 |
+
NFtY2PwByVS5uCbMiogziUwthDyC3+6WVwW6LLv3xLfHTjuCvjHIInNzktHCgKQ5
|
| 24 |
+
ORAzI4JMPJ+GslWYHb4phowim57iaztXOoJwTdwJx4nLCgdNbOhdjsnvzqvHu7Ur
|
| 25 |
+
TkXWStAmzOVyyghqpZXjFaH3pO3JLF+l+/+sKAIuvtd7u+Nxe5AW0wdeRlN8NwdC
|
| 26 |
+
jNPElpzVmbUq4JUagEiuTDkHzsxHpFKVK7q4+63SM1N95R1NbdWhscdCb+ZAJzVc
|
| 27 |
+
oyi3B43njTOQ5yOf+1CceWxG1bQVs5ZufpsMljq4Ui0/1lvh+wjChP4kqKOJ2qxq
|
| 28 |
+
4RgqsahDYVvTH9w7jXbyLeiNdd8XM2w9U/t7y0Ff/9yi0GE44Za4rF2LN9d11TPA
|
| 29 |
+
mRGunUHBcnWEvgJBQl9nJEiU0Zsnvgc/ubhPgXRR4Xq37Z0j4r7g1SgEEzwxA57d
|
| 30 |
+
emyPxgcYxn/eR44/KJ4EBs+lVDR3veyJm+kXQ99b21/+jh5Xos1AnX5iItreGCc=
|
| 31 |
+
-----END CERTIFICATE-----
|
Dockerfile
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Use an official PyTorch image with CUDA support or Python base image
|
| 2 |
+
FROM python:3.10-slim
|
| 3 |
+
|
| 4 |
+
# Set the working directory to /app
|
| 5 |
+
WORKDIR /app
|
| 6 |
+
|
| 7 |
+
# Copy the contents of your project to /app in the container
|
| 8 |
+
COPY . .
|
| 9 |
+
|
| 10 |
+
# Set the PYTHONPATH to include the src directory
|
| 11 |
+
ENV PYTHONPATH=/app
|
| 12 |
+
|
| 13 |
+
# Install Python dependencies
|
| 14 |
+
COPY requirements.txt /app/requirements.txt
|
| 15 |
+
RUN pip install --no-cache-dir -r /app/requirements.txt
|
| 16 |
+
|
| 17 |
+
# Expose the port Gradio will run on (default: 7860)
|
| 18 |
+
EXPOSE 7860
|
| 19 |
+
|
| 20 |
+
# Command to run the application
|
| 21 |
+
CMD ["python", "./src/main.py"]
|
LICENSE
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
MIT License
|
| 2 |
+
|
| 3 |
+
Copyright (c) 2024 Auxiliary.ai
|
| 4 |
+
|
| 5 |
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 6 |
+
of this software and associated documentation files (the "Software"), to deal
|
| 7 |
+
in the Software without restriction, including without limitation the rights
|
| 8 |
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 9 |
+
copies of the Software, and to permit persons to whom the Software is
|
| 10 |
+
furnished to do so, subject to the following conditions:
|
| 11 |
+
|
| 12 |
+
The above copyright notice and this permission notice shall be included in all
|
| 13 |
+
copies or substantial portions of the Software.
|
| 14 |
+
|
| 15 |
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 16 |
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 17 |
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 18 |
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 19 |
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 20 |
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
| 21 |
+
SOFTWARE.
|
README.md
CHANGED
|
@@ -1,12 +1,44 @@
|
|
| 1 |
---
|
| 2 |
-
title:
|
| 3 |
-
|
| 4 |
-
colorFrom: blue
|
| 5 |
-
colorTo: yellow
|
| 6 |
sdk: gradio
|
| 7 |
-
sdk_version: 5.
|
| 8 |
-
app_file: app.py
|
| 9 |
-
pinned: false
|
| 10 |
---
|
|
|
|
|
|
|
| 11 |
|
| 12 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
+
title: da_detection
|
| 3 |
+
app_file: src/main.py
|
|
|
|
|
|
|
| 4 |
sdk: gradio
|
| 5 |
+
sdk_version: 5.3.0
|
|
|
|
|
|
|
| 6 |
---
|
| 7 |
+
# Beirut Disadvantaged Areas Mapping
|
| 8 |
+

|
| 9 |
|
| 10 |
+

|
| 11 |
+

|
| 12 |
+

|
| 13 |
+

|
| 14 |
+
|
| 15 |
+
## Description
|
| 16 |
+
This project aims to map the disadvantaged areas in Beirut, Lebanon. The deployed model will be able to detect urban disadvantaged areas through HR and LR satellite imagery. The project is still under development.
|
| 17 |
+
|
| 18 |
+
## Installation
|
| 19 |
+
First, clone the repository using the following command:
|
| 20 |
+
```bash
|
| 21 |
+
git clone https://github.com/auxiliary-ai/DA_detector.git
|
| 22 |
+
```
|
| 23 |
+
Then, create and activate your conda environment using the following commands (refer to the miniconda documentation https://docs.anaconda.com/free/miniconda/miniconda-install/):
|
| 24 |
+
```bash
|
| 25 |
+
conda create -n DA_env python=3.10
|
| 26 |
+
conda activate DA_env
|
| 27 |
+
```
|
| 28 |
+
Finally, install the required packages using the following command:
|
| 29 |
+
```bash
|
| 30 |
+
cd DA_detector
|
| 31 |
+
pip install -r requirements.txt
|
| 32 |
+
```
|
| 33 |
+
Remember to replace the path to the repository with the correct path on your machine.
|
| 34 |
+
|
| 35 |
+
## Usage
|
| 36 |
+
To run the application, use the following command:
|
| 37 |
+
```bash
|
| 38 |
+
python main.py
|
| 39 |
+
```
|
| 40 |
+
|
| 41 |
+
## TODO
|
| 42 |
+
|
| 43 |
+
- [ ] Train model on data with different resolutions.
|
| 44 |
+
- [ ] Add a low resolution model.
|
cover.jpg
ADDED
|
data/download_data.py
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import msal
|
| 2 |
+
import requests
|
| 3 |
+
|
| 4 |
+
# TODO: Need to purchase sharepoint license to access files. For now use below link:
|
| 5 |
+
#link: https://onedrive.live.com/?id=388473CC408C5CF1%21sa863b43cc1ce49f0a092e21d2d0514e7&cid=388473CC408C5CF1
|
| 6 |
+
|
| 7 |
+
'''This script demonstrates how to authenticate with Microsoft Graph API using the MSAL library and list files in a
|
| 8 |
+
specified OneDrive folder. The script uses the client credentials flow to obtain an access token for the Microsoft
|
| 9 |
+
Graph API. The access token is then used to make a request to list files in a specified folder. This can be used to
|
| 10 |
+
download files from OneDrive or perform other operations on files and folders.'''
|
| 11 |
+
|
| 12 |
+
# Microsoft App credentials
|
| 13 |
+
client_id = '671ad178-15a3-4c2d-ae73-3e0efa00b91f'
|
| 14 |
+
client_secret = 'GfY8Q~c-IwupL_Q5qefWRKS2.F.LXVk5GSSc8chA'
|
| 15 |
+
tenant_id = 'e168e317-9725-4890-a65d-05a8925848ca'
|
| 16 |
+
|
| 17 |
+
# Authority URL for authentication
|
| 18 |
+
authority = f"https://login.microsoftonline.com/{tenant_id}"
|
| 19 |
+
|
| 20 |
+
# Initialize MSAL ConfidentialClientApplication
|
| 21 |
+
app = msal.ConfidentialClientApplication(
|
| 22 |
+
client_id,
|
| 23 |
+
authority=authority,
|
| 24 |
+
client_credential=client_secret,
|
| 25 |
+
)
|
| 26 |
+
|
| 27 |
+
# Acquire token for Microsoft Graph API
|
| 28 |
+
token_response = app.acquire_token_for_client(scopes=["https://graph.microsoft.com/.default"])
|
| 29 |
+
|
| 30 |
+
# Check if token is obtained
|
| 31 |
+
if 'access_token' in token_response:
|
| 32 |
+
access_token = token_response['access_token']
|
| 33 |
+
print("Access token obtained successfully.")
|
| 34 |
+
else:
|
| 35 |
+
print("Failed to obtain access token.")
|
| 36 |
+
print(token_response)
|
| 37 |
+
exit() # Exit the script if the token isn't obtained
|
| 38 |
+
|
| 39 |
+
# Function to list files in a specified OneDrive folder
|
| 40 |
+
def list_files(folder_id, drive_id, access_token):
|
| 41 |
+
url = f"https://graph.microsoft.com/v1.0/drives/{drive_id}/items/{folder_id}/children"
|
| 42 |
+
headers = {
|
| 43 |
+
'Authorization': f'Bearer {access_token}'
|
| 44 |
+
}
|
| 45 |
+
response = requests.get(url, headers=headers)
|
| 46 |
+
|
| 47 |
+
# Check if the request was successful
|
| 48 |
+
if response.status_code == 200:
|
| 49 |
+
files = response.json()
|
| 50 |
+
print(files) # Inspect the structure of the response
|
| 51 |
+
return files
|
| 52 |
+
else:
|
| 53 |
+
print(f"Error: {response.status_code}")
|
| 54 |
+
print(response.text)
|
| 55 |
+
return None
|
| 56 |
+
|
| 57 |
+
# Example usage:
|
| 58 |
+
# From your shared link: https://onedrive.live.com/?id=388473CC408C5CF1%21sa863b43cc1ce49f0a092e21d2d0514e7&cid=388473CC408C5CF1
|
| 59 |
+
drive_id = '388473CC408C5CF1' # Extracted from the URL as the drive ID
|
| 60 |
+
folder_id = 'sa863b43cc1ce49f0a092e21d2d0514e7' # Extracted from the URL as the item ID
|
| 61 |
+
|
| 62 |
+
# List files in the specified folder
|
| 63 |
+
files = list_files(folder_id, drive_id, access_token)
|
| 64 |
+
|
| 65 |
+
if files and 'value' in files:
|
| 66 |
+
for file in files['value']:
|
| 67 |
+
print(f"File Name: {file['name']} | File ID: {file['id']}")
|
| 68 |
+
else:
|
| 69 |
+
print("No files found or 'value' key is missing in the response.")
|
requirements.txt
ADDED
|
@@ -0,0 +1,12 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
numpy~=2.1.0
|
| 2 |
+
pillow~=10.4.0
|
| 3 |
+
torch
|
| 4 |
+
torchvision
|
| 5 |
+
scikit-learn~=1.5.1
|
| 6 |
+
matplotlib~=3.8.4
|
| 7 |
+
tqdm~=4.66.4
|
| 8 |
+
requests~=2.32.3
|
| 9 |
+
msal
|
| 10 |
+
parameterized
|
| 11 |
+
gradio
|
| 12 |
+
imagehash
|
src/data_preperation.py
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from src import utils
|
| 2 |
+
import numpy as np
|
| 3 |
+
import os
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
def main(in_dir: str, out_dir: str, size: tuple, stride: tuple, classes: list):
|
| 7 |
+
# Read all files in the in_dir
|
| 8 |
+
files = os.listdir(in_dir)
|
| 9 |
+
|
| 10 |
+
# Create the out_dir if it does not exist
|
| 11 |
+
if not os.path.exists(out_dir):
|
| 12 |
+
os.makedirs(out_dir)
|
| 13 |
+
|
| 14 |
+
dataset = []
|
| 15 |
+
# Loop through all files in the in_dir
|
| 16 |
+
for file in files:
|
| 17 |
+
# apply the sliding window method to each file
|
| 18 |
+
imgs = os.listdir(in_dir + file)
|
| 19 |
+
for img in imgs:
|
| 20 |
+
path = in_dir + file + "/" + img
|
| 21 |
+
data, _ = utils.sliding_window(path, window_size=size, stride=stride)
|
| 22 |
+
|
| 23 |
+
for x in data:
|
| 24 |
+
dataset.append((np.array(x[0]), classes.index(file)))
|
| 25 |
+
|
| 26 |
+
# Split the dataset into training, validation, and test sets
|
| 27 |
+
train, val, test = utils.split_data(dataset, 0.7, 0.1, 0.2)
|
| 28 |
+
utils.plot_distribution(
|
| 29 |
+
train, val, test, classes, title="Data Distribution Before Balancing"
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
# Balance the dataset
|
| 33 |
+
train, excess_data = utils.balance_data(train)
|
| 34 |
+
test += excess_data
|
| 35 |
+
utils.plot_distribution(
|
| 36 |
+
train, val, test, classes, title="Data Distribution After Balancing"
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
# Shuffle the data
|
| 40 |
+
train = utils.shuffle_data(train)
|
| 41 |
+
val = utils.shuffle_data(val)
|
| 42 |
+
test = utils.shuffle_data(test)
|
| 43 |
+
|
| 44 |
+
# Pickle the data
|
| 45 |
+
utils.save_to_pickle(train, out_dir + "/train.pkl")
|
| 46 |
+
utils.save_to_pickle(val, out_dir + "/val.pkl")
|
| 47 |
+
utils.save_to_pickle(test, out_dir + "/test.pkl")
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
if __name__ == "__main__":
|
| 51 |
+
main(
|
| 52 |
+
in_dir="raw_data/high_res/",
|
| 53 |
+
out_dir="data/high_res/",
|
| 54 |
+
size=(224, 224),
|
| 55 |
+
stride=(224, 224),
|
| 56 |
+
classes=["formal", "informal"],
|
| 57 |
+
)
|
src/main.py
ADDED
|
@@ -0,0 +1,125 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
from torchvision import transforms, models
|
| 4 |
+
from torchvision.transforms import InterpolationMode
|
| 5 |
+
from torch.utils.data import Dataset, DataLoader
|
| 6 |
+
import gradio as gr
|
| 7 |
+
from src import utils
|
| 8 |
+
from PIL import Image
|
| 9 |
+
import numpy as np
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
def main(window_size: int = 224):
|
| 13 |
+
# Set the device
|
| 14 |
+
np.random.seed(0)
|
| 15 |
+
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 16 |
+
print("Running on {}".format(device))
|
| 17 |
+
|
| 18 |
+
def predict(inp, stride_size, threshold):
|
| 19 |
+
|
| 20 |
+
# Apply sliding window on the image
|
| 21 |
+
windows, pad_size = utils.sliding_window(
|
| 22 |
+
inp, (window_size, window_size), (stride_size, stride_size)
|
| 23 |
+
)
|
| 24 |
+
|
| 25 |
+
# Define the transformations to be applied to the images
|
| 26 |
+
transform = transforms.Compose(
|
| 27 |
+
[
|
| 28 |
+
transforms.Resize([256], interpolation=InterpolationMode.BICUBIC),
|
| 29 |
+
transforms.CenterCrop([224]),
|
| 30 |
+
transforms.ToTensor(), # Converts the image to [0.0, 1.0] range
|
| 31 |
+
transforms.Normalize(
|
| 32 |
+
mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]
|
| 33 |
+
),
|
| 34 |
+
]
|
| 35 |
+
)
|
| 36 |
+
|
| 37 |
+
# Custom dataset class
|
| 38 |
+
class CustomDataset(Dataset):
|
| 39 |
+
def __init__(self, data, transform=None):
|
| 40 |
+
self.data = data
|
| 41 |
+
self.transform = transform
|
| 42 |
+
|
| 43 |
+
def __len__(self):
|
| 44 |
+
return len(self.data)
|
| 45 |
+
|
| 46 |
+
def __getitem__(self, idx):
|
| 47 |
+
image, position = self.data[idx]
|
| 48 |
+
if self.transform:
|
| 49 |
+
image = self.transform(image)
|
| 50 |
+
return image, position
|
| 51 |
+
|
| 52 |
+
# Create custom datasets with transformations
|
| 53 |
+
dataset = CustomDataset(windows, transform=transform)
|
| 54 |
+
dataloader = DataLoader(dataset, batch_size=128, shuffle=False)
|
| 55 |
+
|
| 56 |
+
# Load a pretrained ResNet model
|
| 57 |
+
model = models.resnet50().to(device)
|
| 58 |
+
|
| 59 |
+
# Modify the output layer directly to match binary classification
|
| 60 |
+
model.fc = nn.Sequential(
|
| 61 |
+
nn.Linear(model.fc.in_features, 128),
|
| 62 |
+
nn.ReLU(inplace=True),
|
| 63 |
+
nn.Linear(128, 1),
|
| 64 |
+
).to(device)
|
| 65 |
+
|
| 66 |
+
# load model weights (Ensure map_location points to the correct device)
|
| 67 |
+
model.load_state_dict(
|
| 68 |
+
torch.load("./weights/resnet.pth", map_location=device, weights_only=True)
|
| 69 |
+
)
|
| 70 |
+
model = model.to(device)
|
| 71 |
+
|
| 72 |
+
model.eval()
|
| 73 |
+
|
| 74 |
+
heatmap = []
|
| 75 |
+
for data in dataloader:
|
| 76 |
+
images, positions = data
|
| 77 |
+
images = images.to(device)
|
| 78 |
+
with torch.no_grad():
|
| 79 |
+
outputs = model(images)
|
| 80 |
+
for i in range(len(outputs)):
|
| 81 |
+
probability = torch.sigmoid(outputs[i]).item()
|
| 82 |
+
img = Image.new(
|
| 83 |
+
"RGB",
|
| 84 |
+
(window_size, window_size),
|
| 85 |
+
color=(int(255 * probability), 0, 0),
|
| 86 |
+
)
|
| 87 |
+
heatmap.append((img, (positions[0][i], positions[1][i])))
|
| 88 |
+
|
| 89 |
+
# overlay the heatmap on the original image
|
| 90 |
+
result, _ = utils.reconstruct_image(
|
| 91 |
+
heatmap, pad_size, (window_size, window_size)
|
| 92 |
+
)
|
| 93 |
+
prob_threshold = round(threshold * 255, 0)
|
| 94 |
+
result = result.point(
|
| 95 |
+
lambda p: p * 0 if p < prob_threshold else p
|
| 96 |
+
) # set any pixel value less than
|
| 97 |
+
# threshold to 0
|
| 98 |
+
|
| 99 |
+
reconstructed_image, _ = utils.reconstruct_image(
|
| 100 |
+
windows, pad_size, (window_size, window_size)
|
| 101 |
+
)
|
| 102 |
+
result = Image.blend(reconstructed_image, result, alpha=0.5)
|
| 103 |
+
|
| 104 |
+
# slice the image to remove the padding
|
| 105 |
+
# result = result.crop((0, 0, inp.width, inp.height))
|
| 106 |
+
|
| 107 |
+
return result
|
| 108 |
+
|
| 109 |
+
# Gradio Interface with stride size input slider
|
| 110 |
+
demo = gr.Interface(
|
| 111 |
+
fn=predict,
|
| 112 |
+
inputs=[
|
| 113 |
+
gr.Image(type="pil"),
|
| 114 |
+
gr.Slider(minimum=32, maximum=224, step=32, label="Stride Size"),
|
| 115 |
+
gr.Slider(minimum=0.5, maximum=1.0, step=0.1, label="Threshold"),
|
| 116 |
+
],
|
| 117 |
+
outputs=gr.Image(type="pil"),
|
| 118 |
+
)
|
| 119 |
+
|
| 120 |
+
# launch the interface in a new tab
|
| 121 |
+
demo.launch(inbrowser=True, share=True, debug=True)
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
if __name__ == "__main__":
|
| 125 |
+
main()
|
src/resnet.py
ADDED
|
@@ -0,0 +1,273 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from torchvision import models
|
| 3 |
+
import torchvision.transforms as transforms
|
| 4 |
+
from torchvision.transforms import InterpolationMode
|
| 5 |
+
from torch.utils.data import Dataset, DataLoader
|
| 6 |
+
import torch.nn as nn
|
| 7 |
+
from sklearn.metrics import confusion_matrix
|
| 8 |
+
import tqdm
|
| 9 |
+
import matplotlib.pyplot as plt
|
| 10 |
+
import numpy as np
|
| 11 |
+
from src import utils
|
| 12 |
+
|
| 13 |
+
plt.style.use("ggplot")
|
| 14 |
+
plt.rcParams.update({"font.size": 14})
|
| 15 |
+
plt.rcParams.update({"figure.autolayout": True})
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def main(batch_size=64, epochs=50, classes=("formal", "informal"), train: bool = True):
|
| 19 |
+
train_transform = transforms.Compose(
|
| 20 |
+
[
|
| 21 |
+
transforms.ToPILImage(),
|
| 22 |
+
transforms.RandomHorizontalFlip(p=0.5),
|
| 23 |
+
transforms.RandomVerticalFlip(p=0.5),
|
| 24 |
+
transforms.ColorJitter(
|
| 25 |
+
brightness=0.3, contrast=0.3, saturation=0.3, hue=0.1
|
| 26 |
+
),
|
| 27 |
+
transforms.RandomRotation(
|
| 28 |
+
degrees=30, interpolation=InterpolationMode.BICUBIC
|
| 29 |
+
),
|
| 30 |
+
transforms.RandomResizedCrop(
|
| 31 |
+
size=224, scale=(0.8, 1.0), interpolation=InterpolationMode.BICUBIC
|
| 32 |
+
),
|
| 33 |
+
transforms.ToTensor(),
|
| 34 |
+
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
|
| 35 |
+
]
|
| 36 |
+
)
|
| 37 |
+
|
| 38 |
+
transform = transforms.Compose(
|
| 39 |
+
[
|
| 40 |
+
transforms.ToPILImage(),
|
| 41 |
+
transforms.Resize([256], interpolation=InterpolationMode.BICUBIC),
|
| 42 |
+
transforms.CenterCrop([224]),
|
| 43 |
+
transforms.ToTensor(), # Converts the image to [0.0, 1.0] range
|
| 44 |
+
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
|
| 45 |
+
]
|
| 46 |
+
)
|
| 47 |
+
|
| 48 |
+
# Custom dataset class
|
| 49 |
+
class CustomDataset(Dataset):
|
| 50 |
+
def __init__(self, data, transform=None):
|
| 51 |
+
self.data = data
|
| 52 |
+
self.transform = transform
|
| 53 |
+
|
| 54 |
+
def __len__(self):
|
| 55 |
+
return len(self.data)
|
| 56 |
+
|
| 57 |
+
def __getitem__(self, idx):
|
| 58 |
+
image, label = self.data[idx]
|
| 59 |
+
if self.transform:
|
| 60 |
+
image = self.transform(image)
|
| 61 |
+
return image, label
|
| 62 |
+
|
| 63 |
+
# Load the data
|
| 64 |
+
train_data = utils.load_from_pickle("./data/high_res/train.pkl")
|
| 65 |
+
val_data = utils.load_from_pickle("./data/high_res/val.pkl")
|
| 66 |
+
test_data = utils.load_from_pickle("./data/high_res/test.pkl")
|
| 67 |
+
|
| 68 |
+
# process the data to keep only three channels
|
| 69 |
+
train_data = [(image[:, :, :3], label) for image, label in train_data]
|
| 70 |
+
val_data = [(image[:, :, :3], label) for image, label in val_data]
|
| 71 |
+
test_data = [(image[:, :, :3], label) for image, label in test_data]
|
| 72 |
+
|
| 73 |
+
# Create custom datasets with transformations
|
| 74 |
+
trainset = CustomDataset(train_data, transform=train_transform)
|
| 75 |
+
valset = CustomDataset(val_data, transform=transform)
|
| 76 |
+
testset = CustomDataset(test_data, transform=transform)
|
| 77 |
+
|
| 78 |
+
# Create DataLoaders
|
| 79 |
+
trainloader = DataLoader(trainset, batch_size=batch_size, shuffle=True)
|
| 80 |
+
valloader = DataLoader(valset, batch_size=batch_size, shuffle=False)
|
| 81 |
+
testloader = DataLoader(testset, batch_size=batch_size, shuffle=False)
|
| 82 |
+
|
| 83 |
+
# Set the device
|
| 84 |
+
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
| 85 |
+
print("Running on {}".format(device))
|
| 86 |
+
|
| 87 |
+
# Load a pretrained ResNet model
|
| 88 |
+
model = models.resnet50(weights="ResNet50_Weights.IMAGENET1K_V1").to(
|
| 89 |
+
device
|
| 90 |
+
) # You can use resnet18, resnet50, etc.
|
| 91 |
+
|
| 92 |
+
# Modify the output layer directly to match binary classification
|
| 93 |
+
model.fc = nn.Sequential(
|
| 94 |
+
nn.Linear(model.fc.in_features, 128),
|
| 95 |
+
nn.ReLU(inplace=True),
|
| 96 |
+
nn.Linear(128, 1), # 1 output unit for binary classification
|
| 97 |
+
).to(
|
| 98 |
+
device
|
| 99 |
+
) # Make sure the head is also on the correct device
|
| 100 |
+
|
| 101 |
+
# Print the model architecture
|
| 102 |
+
print(model)
|
| 103 |
+
|
| 104 |
+
if train:
|
| 105 |
+
# Define loss function and optimizer
|
| 106 |
+
criterion = nn.BCEWithLogitsLoss() # Binary classification
|
| 107 |
+
optimizer = torch.optim.SGD(model.parameters(), lr=0.001, momentum=0.9)
|
| 108 |
+
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.25)
|
| 109 |
+
|
| 110 |
+
# Train the model
|
| 111 |
+
optimal_accuracy = 0
|
| 112 |
+
patience = 10
|
| 113 |
+
|
| 114 |
+
for epoch in tqdm.tqdm(range(epochs)):
|
| 115 |
+
model.train()
|
| 116 |
+
running_loss = 0.0
|
| 117 |
+
|
| 118 |
+
for i, data in enumerate(trainloader, 0):
|
| 119 |
+
# Move inputs and labels to the correct device
|
| 120 |
+
inputs, labels = data
|
| 121 |
+
inputs, labels = (
|
| 122 |
+
inputs.to(device),
|
| 123 |
+
labels.to(device).float(),
|
| 124 |
+
) # Ensure labels are float for BCEWithLogitsLoss
|
| 125 |
+
|
| 126 |
+
# Zero the parameter gradients
|
| 127 |
+
optimizer.zero_grad()
|
| 128 |
+
|
| 129 |
+
# Forward + backward + optimize
|
| 130 |
+
outputs = model(inputs)
|
| 131 |
+
loss = criterion(
|
| 132 |
+
outputs, labels.unsqueeze(1)
|
| 133 |
+
) # Match output shape (N, 1)
|
| 134 |
+
loss.backward()
|
| 135 |
+
optimizer.step()
|
| 136 |
+
|
| 137 |
+
running_loss += loss.item()
|
| 138 |
+
if i % 2000 == 1999:
|
| 139 |
+
print(f"[{epoch + 1}, {i + 1}] loss: {running_loss / 2000:.3f}")
|
| 140 |
+
running_loss = 0.0
|
| 141 |
+
|
| 142 |
+
scheduler.step()
|
| 143 |
+
|
| 144 |
+
# Validate the model
|
| 145 |
+
model.eval()
|
| 146 |
+
correct = 0
|
| 147 |
+
total = 0
|
| 148 |
+
|
| 149 |
+
with torch.no_grad():
|
| 150 |
+
for data in valloader:
|
| 151 |
+
images, labels = data
|
| 152 |
+
images, labels = images.to(device), labels.to(device).float()
|
| 153 |
+
outputs = model(images)
|
| 154 |
+
predicted = (
|
| 155 |
+
torch.sigmoid(outputs) > 0.5
|
| 156 |
+
).float() # Apply sigmoid and threshold at 0.5
|
| 157 |
+
total += labels.size(0)
|
| 158 |
+
correct += (predicted == labels.unsqueeze(1)).sum().item()
|
| 159 |
+
|
| 160 |
+
accuracy = 100 * correct / total
|
| 161 |
+
print(f"Validation accuracy: {accuracy:.2f}%")
|
| 162 |
+
|
| 163 |
+
# early stopping
|
| 164 |
+
if accuracy > optimal_accuracy:
|
| 165 |
+
optimal_accuracy = accuracy
|
| 166 |
+
optimal_model = model.state_dict()
|
| 167 |
+
patience = 10
|
| 168 |
+
else:
|
| 169 |
+
patience -= 1
|
| 170 |
+
|
| 171 |
+
if patience == 0:
|
| 172 |
+
print("Early stopping")
|
| 173 |
+
break
|
| 174 |
+
|
| 175 |
+
print("Finished Training")
|
| 176 |
+
|
| 177 |
+
# Save the model
|
| 178 |
+
torch.save(optimal_model, "./weights/ResNet50.pth")
|
| 179 |
+
print("Model saved")
|
| 180 |
+
|
| 181 |
+
# Load the model and move it to the correct device
|
| 182 |
+
model.load_state_dict(
|
| 183 |
+
torch.load("./weights/ResNet50.pth", map_location=device, weights_only=True)
|
| 184 |
+
)
|
| 185 |
+
model.to(device)
|
| 186 |
+
|
| 187 |
+
# Test the model
|
| 188 |
+
model.eval()
|
| 189 |
+
correct = 0
|
| 190 |
+
total = 0
|
| 191 |
+
y_pred = []
|
| 192 |
+
y_true = []
|
| 193 |
+
|
| 194 |
+
with torch.no_grad():
|
| 195 |
+
for data in tqdm.tqdm(testloader):
|
| 196 |
+
images, labels = data
|
| 197 |
+
images, labels = images.to(device), labels.to(device).float()
|
| 198 |
+
outputs = model(images)
|
| 199 |
+
predicted = (
|
| 200 |
+
torch.sigmoid(outputs) > 0.5
|
| 201 |
+
).float() # Apply sigmoid and threshold at 0.5
|
| 202 |
+
|
| 203 |
+
# store predictions for CM
|
| 204 |
+
y_pred.extend(predicted.cpu().numpy())
|
| 205 |
+
y_true.extend(labels.cpu().numpy())
|
| 206 |
+
|
| 207 |
+
total += labels.size(0)
|
| 208 |
+
correct += (predicted == labels.unsqueeze(1)).sum().item()
|
| 209 |
+
|
| 210 |
+
accuracy = 100 * correct / total
|
| 211 |
+
print(f"Test accuracy: {accuracy:.2f}%")
|
| 212 |
+
|
| 213 |
+
print("Finished Testing")
|
| 214 |
+
|
| 215 |
+
# Confusion matrix
|
| 216 |
+
cm = confusion_matrix(y_true, y_pred, normalize="true")
|
| 217 |
+
plt.figure(figsize=(8, 8))
|
| 218 |
+
plt.imshow(cm, interpolation="nearest", cmap=plt.cm.Blues)
|
| 219 |
+
plt.title("Confusion Matrix")
|
| 220 |
+
plt.colorbar()
|
| 221 |
+
tick_marks = np.arange(len(classes))
|
| 222 |
+
plt.xticks(tick_marks, classes, rotation=45)
|
| 223 |
+
plt.yticks(tick_marks, classes)
|
| 224 |
+
plt.xlabel("Predicted")
|
| 225 |
+
plt.ylabel("True")
|
| 226 |
+
plt.savefig("confusion_matrix.pdf", format="pdf")
|
| 227 |
+
plt.show()
|
| 228 |
+
|
| 229 |
+
# Show test images with predicted label and actual label
|
| 230 |
+
def imshow(img, title=None):
|
| 231 |
+
"""This function plots a tensor"""
|
| 232 |
+
img = img / 2 + 0.5 # unnormalize
|
| 233 |
+
npimg = img.cpu().numpy() # convert to numpy for display
|
| 234 |
+
plt.imshow(np.transpose(npimg, (1, 2, 0))) # reshape to (H, W, C)
|
| 235 |
+
if title is not None:
|
| 236 |
+
plt.title(title)
|
| 237 |
+
plt.show()
|
| 238 |
+
|
| 239 |
+
def show_predictions(model, dataloader, device, classes):
|
| 240 |
+
"""Function to show images with predicted and actual labels."""
|
| 241 |
+
model.eval()
|
| 242 |
+
|
| 243 |
+
with torch.no_grad():
|
| 244 |
+
for i, data in enumerate(dataloader):
|
| 245 |
+
images, labels = data
|
| 246 |
+
images, labels = images.to(device), labels.to(device).float()
|
| 247 |
+
|
| 248 |
+
# Forward pass to get predictions
|
| 249 |
+
outputs = model(images)
|
| 250 |
+
outputs = outputs.squeeze(dim=1) # Ensure shape is [batch_size]
|
| 251 |
+
predicted = (torch.sigmoid(outputs) > 0.5).float()
|
| 252 |
+
|
| 253 |
+
# Plot each image with its predicted and actual labels
|
| 254 |
+
for j in range(images.size(0)):
|
| 255 |
+
imshow(images[j].cpu()) # Unnormalize and plot image
|
| 256 |
+
|
| 257 |
+
# Convert predictions and labels to text (formal/informal)
|
| 258 |
+
pred_label = classes[int(predicted[j].item())]
|
| 259 |
+
actual_label = classes[int(labels[j].item())]
|
| 260 |
+
|
| 261 |
+
# Display the predicted and actual labels
|
| 262 |
+
print(f"Predicted: {pred_label}, Actual: {actual_label}")
|
| 263 |
+
|
| 264 |
+
# Optionally, stop after displaying N images
|
| 265 |
+
if i * len(images) + j >= 20: # Show 5 images, adjust as needed
|
| 266 |
+
return
|
| 267 |
+
|
| 268 |
+
# Call the function to display images along with predictions
|
| 269 |
+
# show_predictions(model, testloader, device, classes)
|
| 270 |
+
|
| 271 |
+
|
| 272 |
+
if __name__ == "__main__":
|
| 273 |
+
main()
|
src/utils.py
ADDED
|
@@ -0,0 +1,430 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
from PIL import Image, ImageOps
|
| 3 |
+
import os
|
| 4 |
+
from tqdm import tqdm
|
| 5 |
+
import matplotlib.pyplot as plt
|
| 6 |
+
from sklearn.model_selection import train_test_split
|
| 7 |
+
import pickle
|
| 8 |
+
from collections import Counter
|
| 9 |
+
import random
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class utils:
|
| 13 |
+
|
| 14 |
+
def __init__(self):
|
| 15 |
+
pass
|
| 16 |
+
|
| 17 |
+
@staticmethod
|
| 18 |
+
def calculate_padding(
|
| 19 |
+
image_height: int, image_width: int, window_size: tuple, stride: tuple
|
| 20 |
+
):
|
| 21 |
+
"""
|
| 22 |
+
Calculate padding needed for height and width
|
| 23 |
+
|
| 24 |
+
:param image_height:
|
| 25 |
+
:param image_width:
|
| 26 |
+
:param window_size:
|
| 27 |
+
:param stride:
|
| 28 |
+
:return: pad_height, pad_width
|
| 29 |
+
"""
|
| 30 |
+
|
| 31 |
+
# Calculate padding needed for height and width
|
| 32 |
+
if (image_height - window_size[0]) // stride[0] != (
|
| 33 |
+
image_height - window_size[0]
|
| 34 |
+
) / stride[0]:
|
| 35 |
+
if window_size[0] == stride[0]:
|
| 36 |
+
pad_height = window_size[0] - (image_height % stride[0])
|
| 37 |
+
else:
|
| 38 |
+
pad_height = image_height - (
|
| 39 |
+
((image_height - window_size[0]) // stride[0]) * stride[0]
|
| 40 |
+
+ window_size[0]
|
| 41 |
+
)
|
| 42 |
+
else:
|
| 43 |
+
pad_height = 0
|
| 44 |
+
if (image_width - window_size[1]) // stride[1] != (
|
| 45 |
+
image_width - window_size[1]
|
| 46 |
+
) / stride[1]:
|
| 47 |
+
if window_size[1] == stride[1]:
|
| 48 |
+
pad_width = window_size[1] - (image_width % stride[1])
|
| 49 |
+
else:
|
| 50 |
+
pad_width = image_width - (
|
| 51 |
+
((image_width - window_size[1]) // stride[1]) * stride[1]
|
| 52 |
+
+ window_size[1]
|
| 53 |
+
)
|
| 54 |
+
else:
|
| 55 |
+
pad_width = 0
|
| 56 |
+
|
| 57 |
+
return pad_height, pad_width
|
| 58 |
+
|
| 59 |
+
@staticmethod
|
| 60 |
+
def sliding_window(
|
| 61 |
+
image_dir: str, window_size: tuple, stride: tuple, padding_value: int = 0
|
| 62 |
+
):
|
| 63 |
+
"""
|
| 64 |
+
Slide a window across the image and extract patches.
|
| 65 |
+
:param image_dir: Path to the image file.
|
| 66 |
+
:param window_size: Tuple (height, width) of the window.
|
| 67 |
+
:param stride: Tuple (height, width) of the stride.
|
| 68 |
+
:param padding_value: Value to use for padding.
|
| 69 |
+
|
| 70 |
+
:return: List of tuples (window, (y, x)) and the size of the padded image.
|
| 71 |
+
"""
|
| 72 |
+
# Open the image
|
| 73 |
+
if isinstance(image_dir, str):
|
| 74 |
+
image = Image.open(image_dir)
|
| 75 |
+
else:
|
| 76 |
+
image = image_dir
|
| 77 |
+
# Convert image to a numpy array
|
| 78 |
+
image_np = np.array(image)
|
| 79 |
+
|
| 80 |
+
# Get image dimensions
|
| 81 |
+
image_height, image_width = image_np.shape[:2]
|
| 82 |
+
|
| 83 |
+
# if window_size[0] > image_height or window_size[1] > image_width:
|
| 84 |
+
# raise ValueError("Window size should be smaller than the image size.")
|
| 85 |
+
# if stride[0] > window_size[0] or stride[1] > window_size[1]:
|
| 86 |
+
# raise ValueError("Stride should be smaller than the window size.")
|
| 87 |
+
|
| 88 |
+
# Calculate padding needed for height and width
|
| 89 |
+
pad_height, pad_width = utils.calculate_padding(
|
| 90 |
+
image_height, image_width, window_size, stride
|
| 91 |
+
)
|
| 92 |
+
|
| 93 |
+
# Pad the image using Pillow (symmetric padding)
|
| 94 |
+
padded_image = ImageOps.expand(
|
| 95 |
+
image, border=(0, 0, pad_width, pad_height), fill=padding_value
|
| 96 |
+
)
|
| 97 |
+
|
| 98 |
+
# Convert padded image back to a numpy array
|
| 99 |
+
padded_image_np = np.array(padded_image)
|
| 100 |
+
|
| 101 |
+
# Get padded image dimensions
|
| 102 |
+
padded_height, padded_width = padded_image_np.shape[:2]
|
| 103 |
+
|
| 104 |
+
# List to store cropped windows and their positions
|
| 105 |
+
windows = []
|
| 106 |
+
|
| 107 |
+
# Slide the window across the padded image
|
| 108 |
+
for y in range(0, padded_height - window_size[0] + 1, stride[0]):
|
| 109 |
+
for x in range(0, padded_width - window_size[1] + 1, stride[1]):
|
| 110 |
+
# Crop the window from the image
|
| 111 |
+
window = padded_image_np[y : y + window_size[0], x : x + window_size[1]]
|
| 112 |
+
# Append the window and its position to the list
|
| 113 |
+
windows.append((Image.fromarray(window), (y, x)))
|
| 114 |
+
|
| 115 |
+
return windows, padded_image.size # Return windows and padded image size
|
| 116 |
+
|
| 117 |
+
@staticmethod
|
| 118 |
+
def save_windows(windows: list, out_dir: str):
|
| 119 |
+
"""
|
| 120 |
+
Save the cropped windows to the output directory.
|
| 121 |
+
|
| 122 |
+
Args:
|
| 123 |
+
windows: List of tuples (window, (y, x)) from the sliding_window function.
|
| 124 |
+
out_dir: Path to the output directory.
|
| 125 |
+
"""
|
| 126 |
+
# Create the output directory if it doesn't exist
|
| 127 |
+
if not os.path.exists(out_dir):
|
| 128 |
+
os.makedirs(out_dir)
|
| 129 |
+
|
| 130 |
+
# Save each window to the output directory
|
| 131 |
+
for i, (window, _) in enumerate(windows):
|
| 132 |
+
window.save(os.path.join(out_dir, f"window_{i}.png"))
|
| 133 |
+
|
| 134 |
+
@staticmethod
|
| 135 |
+
def reconstruct_image(windows: list, padded_size: tuple, window_size: tuple):
|
| 136 |
+
"""
|
| 137 |
+
Reconstruct the original image from the sliding window patches. Strides are averaged.
|
| 138 |
+
|
| 139 |
+
:param windows: List of tuples (window, (y, x)) from the sliding_window function.
|
| 140 |
+
:param padded_size: Tuple (height, width) of the padded image.
|
| 141 |
+
:param window_size: Tuple (height, width) of the window.
|
| 142 |
+
|
| 143 |
+
:return: Reconstructed image and count map.
|
| 144 |
+
"""
|
| 145 |
+
# Initialize an empty numpy array for the reconstructed image
|
| 146 |
+
num_channels = 4 if windows[0][0].mode == "RGBA" else 3
|
| 147 |
+
reconstructed_image = np.zeros(
|
| 148 |
+
(padded_size[1], padded_size[0], num_channels), dtype=np.float32
|
| 149 |
+
)
|
| 150 |
+
|
| 151 |
+
# Initialize an array to count the number of overlapping windows
|
| 152 |
+
count_map = np.zeros(
|
| 153 |
+
(padded_size[1], padded_size[0], num_channels), dtype=np.float32
|
| 154 |
+
)
|
| 155 |
+
|
| 156 |
+
# Place each window back into the reconstructed image
|
| 157 |
+
for window, (y, x) in windows:
|
| 158 |
+
window_np = np.array(window, dtype=np.float32)
|
| 159 |
+
|
| 160 |
+
# Add the window to the corresponding position in the reconstructed image
|
| 161 |
+
reconstructed_image[
|
| 162 |
+
y : y + window_size[1], x : x + window_size[0], :
|
| 163 |
+
] += window_np
|
| 164 |
+
|
| 165 |
+
# Increment the count map to handle overlaps
|
| 166 |
+
count_map[y : y + window_size[1], x : x + window_size[0], :] += 1
|
| 167 |
+
|
| 168 |
+
# Divide by the count map to average overlapping areas
|
| 169 |
+
reconstructed_image = np.divide(
|
| 170 |
+
reconstructed_image, count_map, where=count_map != 0
|
| 171 |
+
)
|
| 172 |
+
|
| 173 |
+
# Clip and convert to 8-bit image
|
| 174 |
+
reconstructed_image = np.clip(reconstructed_image, 0, 255).astype(np.uint8)
|
| 175 |
+
mode = "RGBA" if num_channels == 4 else "RGB"
|
| 176 |
+
reconstructed_image = Image.fromarray(reconstructed_image, mode=mode)
|
| 177 |
+
|
| 178 |
+
return reconstructed_image, count_map
|
| 179 |
+
|
| 180 |
+
@staticmethod
|
| 181 |
+
def zip_images(directory: str, label: int):
|
| 182 |
+
"""
|
| 183 |
+
This function reads images from a directory and returns them as a list of tuples (image, label).
|
| 184 |
+
:param directory: The directory containing images
|
| 185 |
+
:param label: The label to assign to all images in the directory
|
| 186 |
+
:return: List of tuples where each tuple is (image, label)
|
| 187 |
+
"""
|
| 188 |
+
data = []
|
| 189 |
+
|
| 190 |
+
range_dir = os.listdir(directory)
|
| 191 |
+
for file_name in tqdm(range_dir):
|
| 192 |
+
img_path = os.path.join(directory, file_name)
|
| 193 |
+
img = Image.open(img_path).convert("RGB") # Open image and convert to RGB
|
| 194 |
+
img_array = np.array(img) # Convert image to a numpy array
|
| 195 |
+
data.append((img_array, label)) # Append tuple (image, label) to the list
|
| 196 |
+
|
| 197 |
+
return data
|
| 198 |
+
|
| 199 |
+
@staticmethod
|
| 200 |
+
def unzip_images(data: list, base_directory: str):
|
| 201 |
+
"""
|
| 202 |
+
Unzips a list of tuples (image, label) and saves the images into directories named after their labels.
|
| 203 |
+
:param data: List of tuples (image, label)
|
| 204 |
+
:param base_directory: Base directory where the images will be saved
|
| 205 |
+
"""
|
| 206 |
+
for i, (img_array, label) in enumerate(data):
|
| 207 |
+
label_directory = os.path.join(base_directory, str(label))
|
| 208 |
+
|
| 209 |
+
# Create the label directory if it doesn't exist
|
| 210 |
+
if not os.path.exists(label_directory):
|
| 211 |
+
os.makedirs(label_directory)
|
| 212 |
+
|
| 213 |
+
# Define the image path (e.g., "label_directory/image_0.png")
|
| 214 |
+
img_path = os.path.join(label_directory, f"image_{i}.png")
|
| 215 |
+
|
| 216 |
+
# Convert the numpy array back to an image and save it
|
| 217 |
+
img = Image.fromarray(img_array)
|
| 218 |
+
img.save(img_path)
|
| 219 |
+
|
| 220 |
+
@staticmethod
|
| 221 |
+
def plot_image(image: np.array, title: str = None):
|
| 222 |
+
"""
|
| 223 |
+
This function plots an image using matplotlib.
|
| 224 |
+
:param image: Numpy array representing the image.
|
| 225 |
+
:param title: Title of the plot.
|
| 226 |
+
:return: None
|
| 227 |
+
"""
|
| 228 |
+
plt.imshow(image)
|
| 229 |
+
plt.axis("off")
|
| 230 |
+
if title:
|
| 231 |
+
plt.title(title)
|
| 232 |
+
plt.show()
|
| 233 |
+
|
| 234 |
+
@staticmethod
|
| 235 |
+
def plot_images(images: list, titles: list):
|
| 236 |
+
"""
|
| 237 |
+
This function plots multiple images side by side.
|
| 238 |
+
:param images: List of numpy arrays representing the images.
|
| 239 |
+
:param titles: List of titles for each image.
|
| 240 |
+
:return: None
|
| 241 |
+
"""
|
| 242 |
+
fig, axes = plt.subplots(1, len(images), figsize=(20, 20))
|
| 243 |
+
for i, (image, title) in enumerate(zip(images, titles)):
|
| 244 |
+
axes[i].imshow(image)
|
| 245 |
+
axes[i].axis("off")
|
| 246 |
+
axes[i].set_title(title)
|
| 247 |
+
plt.show()
|
| 248 |
+
|
| 249 |
+
@staticmethod
|
| 250 |
+
def split_data(
|
| 251 |
+
data,
|
| 252 |
+
train_size: float,
|
| 253 |
+
val_size: float,
|
| 254 |
+
test_size: float,
|
| 255 |
+
random_state: int = None,
|
| 256 |
+
):
|
| 257 |
+
"""
|
| 258 |
+
This function splits the dataset into training, validation, and test sets.
|
| 259 |
+
|
| 260 |
+
:param data: List of tuples where each tuple is (image, label)
|
| 261 |
+
:param train_size: Proportion of the dataset to include in the training set
|
| 262 |
+
:param val_size: Proportion of the dataset to include in the validation set
|
| 263 |
+
:param test_size: Proportion of the dataset to include in the test set
|
| 264 |
+
:param random_state: Controls the shuffling applied to the data before applying the split
|
| 265 |
+
:return: Tuple of (train_data, val_data, test_data) where each is a list of (image, label)
|
| 266 |
+
"""
|
| 267 |
+
# Ensure the split sizes add up to 1
|
| 268 |
+
assert np.isclose(
|
| 269 |
+
train_size + val_size + test_size, 1.0
|
| 270 |
+
), "Split sizes must add up to 1"
|
| 271 |
+
|
| 272 |
+
# First split: Train + (Val + Test)
|
| 273 |
+
train_data, temp_data = train_test_split(
|
| 274 |
+
data, train_size=train_size, random_state=random_state
|
| 275 |
+
)
|
| 276 |
+
|
| 277 |
+
# Second split: Val + Test
|
| 278 |
+
val_ratio = val_size / (
|
| 279 |
+
val_size + test_size
|
| 280 |
+
) # Adjust val_size to be relative to the size of temp_data
|
| 281 |
+
val_data, test_data = train_test_split(
|
| 282 |
+
temp_data, train_size=val_ratio, random_state=random_state
|
| 283 |
+
)
|
| 284 |
+
|
| 285 |
+
return train_data, val_data, test_data
|
| 286 |
+
|
| 287 |
+
@staticmethod
|
| 288 |
+
def save_to_pickle(data: list, file_path: str):
|
| 289 |
+
"""
|
| 290 |
+
This function saves data to a pickle file.
|
| 291 |
+
:param data: Data to save.
|
| 292 |
+
:param file_path: Path to save the pickle file.
|
| 293 |
+
:return: None
|
| 294 |
+
"""
|
| 295 |
+
with open(file_path, "wb") as f:
|
| 296 |
+
pickle.dump(data, f)
|
| 297 |
+
|
| 298 |
+
@staticmethod
|
| 299 |
+
def load_from_pickle(file_path: str):
|
| 300 |
+
"""
|
| 301 |
+
This function loads data from a pickle file.
|
| 302 |
+
:param file_path: Path to the pickle file.
|
| 303 |
+
:return: Loaded data.
|
| 304 |
+
"""
|
| 305 |
+
with open(file_path, "rb") as f:
|
| 306 |
+
return pickle.load(f)
|
| 307 |
+
|
| 308 |
+
@staticmethod
|
| 309 |
+
def plot_distribution(
|
| 310 |
+
train_data,
|
| 311 |
+
val_data,
|
| 312 |
+
test_data,
|
| 313 |
+
class_names,
|
| 314 |
+
title: str = "Split Data Distribution",
|
| 315 |
+
):
|
| 316 |
+
"""
|
| 317 |
+
Plots the distribution of data after splitting into training, validation, and test sets.
|
| 318 |
+
:param train_data: List of tuples (image, label) for the training set
|
| 319 |
+
:param val_data: List of tuples (image, label) for the validation set
|
| 320 |
+
:param test_data: List of tuples (image, label) for the test set
|
| 321 |
+
:param class_names: List of class names corresponding to the label
|
| 322 |
+
:param title: Title of the plot
|
| 323 |
+
:return: None
|
| 324 |
+
"""
|
| 325 |
+
# Count the number of samples for each label in each dataset
|
| 326 |
+
train_labels = [label for _, label in train_data]
|
| 327 |
+
val_labels = [label for _, label in val_data]
|
| 328 |
+
test_labels = [label for _, label in test_data]
|
| 329 |
+
|
| 330 |
+
train_counter = Counter(train_labels)
|
| 331 |
+
val_counter = Counter(val_labels)
|
| 332 |
+
test_counter = Counter(test_labels)
|
| 333 |
+
|
| 334 |
+
# Prepare the data for plotting
|
| 335 |
+
labels = sorted(set(train_labels + val_labels + test_labels))
|
| 336 |
+
train_counts = [train_counter[label] for label in labels]
|
| 337 |
+
val_counts = [val_counter[label] for label in labels]
|
| 338 |
+
test_counts = [test_counter[label] for label in labels]
|
| 339 |
+
|
| 340 |
+
# Plot the distribution
|
| 341 |
+
x = range(len(labels))
|
| 342 |
+
width = 0.25 # Width of the bars
|
| 343 |
+
|
| 344 |
+
plt.figure(figsize=(10, 6))
|
| 345 |
+
|
| 346 |
+
plt.bar(
|
| 347 |
+
x, train_counts, width=width, label="Train", color="blue", align="center"
|
| 348 |
+
)
|
| 349 |
+
plt.bar(
|
| 350 |
+
[p + width for p in x],
|
| 351 |
+
val_counts,
|
| 352 |
+
width=width,
|
| 353 |
+
label="Validation",
|
| 354 |
+
color="orange",
|
| 355 |
+
align="center",
|
| 356 |
+
)
|
| 357 |
+
plt.bar(
|
| 358 |
+
[p + width * 2 for p in x],
|
| 359 |
+
test_counts,
|
| 360 |
+
width=width,
|
| 361 |
+
label="Test",
|
| 362 |
+
color="green",
|
| 363 |
+
align="center",
|
| 364 |
+
)
|
| 365 |
+
|
| 366 |
+
plt.xlabel("Classes")
|
| 367 |
+
plt.ylabel("Number of Samples")
|
| 368 |
+
plt.title(title)
|
| 369 |
+
plt.xticks([p + width for p in x], [class_names[label] for label in labels])
|
| 370 |
+
plt.legend()
|
| 371 |
+
plt.show()
|
| 372 |
+
|
| 373 |
+
@staticmethod
|
| 374 |
+
def balance_data(train_data: list, random_state=None):
|
| 375 |
+
"""
|
| 376 |
+
This function balances the data by randomly selecting samples from the majority class to match the number of
|
| 377 |
+
samples in the minority class.
|
| 378 |
+
:param train_data: List of tuples (image, label)
|
| 379 |
+
:param random_state: int, random seed for reproducibility
|
| 380 |
+
:return: List of tuples (image, label)
|
| 381 |
+
with balanced classes and a list of tuples (image, label) with the excess data
|
| 382 |
+
"""
|
| 383 |
+
|
| 384 |
+
# get all labels from the train data
|
| 385 |
+
labels = [sample[1] for sample in train_data]
|
| 386 |
+
|
| 387 |
+
# count the number of samples for each label
|
| 388 |
+
counter = Counter(labels)
|
| 389 |
+
|
| 390 |
+
# find the label with the fewest samples
|
| 391 |
+
min_samples = min(counter.values())
|
| 392 |
+
|
| 393 |
+
# create a list to store the balanced data
|
| 394 |
+
balanced_data = []
|
| 395 |
+
|
| 396 |
+
# create a list to store the excess data
|
| 397 |
+
excess_data = []
|
| 398 |
+
|
| 399 |
+
# get how many classes we have
|
| 400 |
+
num_classes = len(set(labels))
|
| 401 |
+
|
| 402 |
+
# for each class
|
| 403 |
+
for i in range(num_classes):
|
| 404 |
+
# get all the samples for that class
|
| 405 |
+
samples = [sample for sample in train_data if sample[1] == i]
|
| 406 |
+
|
| 407 |
+
# shuffle the samples
|
| 408 |
+
random.seed(random_state)
|
| 409 |
+
|
| 410 |
+
random.shuffle(samples)
|
| 411 |
+
|
| 412 |
+
# add the first min_samples samples to the balanced data
|
| 413 |
+
balanced_data += samples[:min_samples]
|
| 414 |
+
|
| 415 |
+
# add the remaining samples to the excess data
|
| 416 |
+
excess_data += samples[min_samples:]
|
| 417 |
+
|
| 418 |
+
return balanced_data, excess_data
|
| 419 |
+
|
| 420 |
+
@staticmethod
|
| 421 |
+
def shuffle_data(data: list, random_state=None):
|
| 422 |
+
"""
|
| 423 |
+
This function shuffles the data.
|
| 424 |
+
:param data: List of tuples (image, label)
|
| 425 |
+
:param random_state: int, random seed for reproducibility
|
| 426 |
+
:return: List of tuples (image, label) with shuffled data
|
| 427 |
+
"""
|
| 428 |
+
random.seed(random_state)
|
| 429 |
+
random.shuffle(data)
|
| 430 |
+
return data
|
tests/test_reconstruct_image.py
ADDED
|
@@ -0,0 +1,152 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from src.utils import utils
|
| 2 |
+
import unittest
|
| 3 |
+
from PIL import Image, ImageOps, ImageDraw
|
| 4 |
+
import imagehash
|
| 5 |
+
import numpy as np
|
| 6 |
+
import os
|
| 7 |
+
import random
|
| 8 |
+
from parameterized import parameterized
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
class TestReconstructImageFunction(unittest.TestCase):
|
| 12 |
+
"""
|
| 13 |
+
Unit tests for the reconstruct image function from the utils module.
|
| 14 |
+
The test compares original images with reconstructed images
|
| 15 |
+
Stride can't be larger than image size or else there'll be loss of information
|
| 16 |
+
"""
|
| 17 |
+
|
| 18 |
+
def setUp(self):
|
| 19 |
+
"""Prepare the test images and their paths before each test."""
|
| 20 |
+
self.prepare_test_images()
|
| 21 |
+
|
| 22 |
+
def tearDown(self):
|
| 23 |
+
"""Clean up the test images after each test."""
|
| 24 |
+
self.clean_up_images()
|
| 25 |
+
|
| 26 |
+
@staticmethod
|
| 27 |
+
def create_random_test_image(width, height):
|
| 28 |
+
"""
|
| 29 |
+
Create a complex image with various shapes and colors, randomly generated.
|
| 30 |
+
|
| 31 |
+
Args:
|
| 32 |
+
width: defining the width of the image.
|
| 33 |
+
height: defining the height of the image.
|
| 34 |
+
|
| 35 |
+
Returns:
|
| 36 |
+
A PIL Image object.
|
| 37 |
+
"""
|
| 38 |
+
size = (width, height)
|
| 39 |
+
# Create a white image
|
| 40 |
+
image = Image.new('RGB', size, color='white')
|
| 41 |
+
draw = ImageDraw.Draw(image)
|
| 42 |
+
|
| 43 |
+
# Randomly generate shape parameters
|
| 44 |
+
num_shapes = random.randint(5, 10)
|
| 45 |
+
shape_colors = [random.choice(['red', 'blue', 'green', 'yellow', 'purple', 'cyan']) for _ in range(num_shapes)]
|
| 46 |
+
# shape_coords = [(random.randint(0, size[0]), random.randint(0, size[1]), random.randint(0, size[0]),
|
| 47 |
+
# random.randint(0, size[1])) for _ in range(num_shapes)]
|
| 48 |
+
shape_coords = []
|
| 49 |
+
shape_types = [random.choice(['rectangle', 'ellipse', 'line', 'polygon']) for _ in range(num_shapes)]
|
| 50 |
+
# Generate valid coordinates for shapes
|
| 51 |
+
for _ in range(num_shapes):
|
| 52 |
+
x0, y0 = random.randint(0, size[0]), random.randint(0, size[1])
|
| 53 |
+
x1, y1 = random.randint(x0, size[0]), random.randint(y0, size[1])
|
| 54 |
+
shape_coords.append((x0, y0, x1, y1))
|
| 55 |
+
# Draw random shapes
|
| 56 |
+
for i in range(num_shapes):
|
| 57 |
+
if shape_types[i] == 'rectangle':
|
| 58 |
+
draw.rectangle(shape_coords[i], outline=shape_colors[i],
|
| 59 |
+
fill=random.choice(['red', 'blue', 'green', 'yellow', 'purple', 'cyan']))
|
| 60 |
+
elif shape_types[i] == 'ellipse':
|
| 61 |
+
draw.ellipse(shape_coords[i], outline=shape_colors[i],
|
| 62 |
+
fill=random.choice(['red', 'blue', 'green', 'yellow', 'purple', 'cyan']))
|
| 63 |
+
elif shape_types[i] == 'line':
|
| 64 |
+
draw.line([shape_coords[i][0], shape_coords[i][1], shape_coords[i][2], shape_coords[i][3]],
|
| 65 |
+
fill=shape_colors[i], width=random.randint(1, 5))
|
| 66 |
+
elif shape_types[i] == 'polygon':
|
| 67 |
+
draw.polygon([shape_coords[i][0], shape_coords[i][1], shape_coords[i][2], shape_coords[i][3]],
|
| 68 |
+
outline=shape_colors[i],
|
| 69 |
+
fill=random.choice(['red', 'blue', 'green', 'yellow', 'purple', 'cyan']))
|
| 70 |
+
|
| 71 |
+
# Add random text
|
| 72 |
+
draw.text((random.randint(0, size[0]), random.randint(0, size[1])),
|
| 73 |
+
random.choice(['Sample Image', 'Random Art', 'Generated Artwork']),
|
| 74 |
+
fill=random.choice(['black', 'white']))
|
| 75 |
+
|
| 76 |
+
return image
|
| 77 |
+
|
| 78 |
+
def prepare_test_images(self):
|
| 79 |
+
"""
|
| 80 |
+
Create and save test images of various sizes for testing.
|
| 81 |
+
"""
|
| 82 |
+
self.test_images = {
|
| 83 |
+
'image_1x1': self.create_random_test_image(1, 1),
|
| 84 |
+
'image_10x10': self.create_random_test_image(10, 10),
|
| 85 |
+
'image_11x11': self.create_random_test_image(11, 11),
|
| 86 |
+
'image_7x5': self.create_random_test_image(7, 5),
|
| 87 |
+
'image_15x15': self.create_random_test_image(15, 15),
|
| 88 |
+
'image_512x512': self.create_random_test_image(512, 512),
|
| 89 |
+
'image_1024x1024': self.create_random_test_image(1024, 1024),
|
| 90 |
+
'image_2048x2048': self.create_random_test_image(2048, 2048),
|
| 91 |
+
}
|
| 92 |
+
|
| 93 |
+
self.image_dirs = {}
|
| 94 |
+
for name, img in self.test_images.items():
|
| 95 |
+
img_dir = f'{name}.png'
|
| 96 |
+
img.save(img_dir)
|
| 97 |
+
self.image_dirs[name] = img_dir
|
| 98 |
+
|
| 99 |
+
def clean_up_images(self):
|
| 100 |
+
"""
|
| 101 |
+
Remove the image files created for testing.
|
| 102 |
+
"""
|
| 103 |
+
for img_dir in self.image_dirs.values():
|
| 104 |
+
if os.path.exists(img_dir):
|
| 105 |
+
os.remove(img_dir)
|
| 106 |
+
|
| 107 |
+
@staticmethod
|
| 108 |
+
def compare_images(original_image, reconstructed_image):
|
| 109 |
+
# reconstructed_image = Image.open(reconstructed_image_dir)
|
| 110 |
+
hash0 = imagehash.average_hash(original_image)
|
| 111 |
+
hash1 = imagehash.average_hash(reconstructed_image)
|
| 112 |
+
hashDiff = hash0 - hash1 # Finds the distance between the hashes of images
|
| 113 |
+
if abs(hashDiff) < 10e-10: # added a tolerance of 10e-10
|
| 114 |
+
return "Identical"
|
| 115 |
+
else:
|
| 116 |
+
return "Non Identical"
|
| 117 |
+
|
| 118 |
+
@parameterized.expand([
|
| 119 |
+
# Parameters: (image_name, window_size, stride)
|
| 120 |
+
('image_11x11', (3, 3), (3, 3)),
|
| 121 |
+
('image_1x1', (1, 1), (1, 1)),
|
| 122 |
+
('image_10x10', (3, 3), (3, 3)),
|
| 123 |
+
('image_7x5', (3, 3), (2, 2)),
|
| 124 |
+
('image_15x15', (5, 5), (5, 5)),
|
| 125 |
+
('image_10x10', (10, 10), (1, 1)), # Large window
|
| 126 |
+
('image_15x15', (7, 7), (5, 5)), # Window larger but within bounds
|
| 127 |
+
('image_10x10', (3, 3), (3, 3)), # Stride equal to window size
|
| 128 |
+
('image_10x10', (3, 3), (2, 2)), # Stride smaller than window size
|
| 129 |
+
('image_7x5', (2, 2), (1, 1)), # Small stride
|
| 130 |
+
('image_15x15', (3, 3), (3, 3)), # Regular case
|
| 131 |
+
('image_512x512', (50, 50), (50, 50)),
|
| 132 |
+
('image_1024x1024', (100, 100), (100, 100)),
|
| 133 |
+
('image_2048x2048', (200, 200), (200, 200)),
|
| 134 |
+
])
|
| 135 |
+
def test_images(self, image_name, window_size, stride):
|
| 136 |
+
|
| 137 |
+
image_dir = self.image_dirs[image_name]
|
| 138 |
+
windows, padded_size = utils.sliding_window(image_dir, window_size, stride)
|
| 139 |
+
reconstructed_image, count_map = utils.reconstruct_image(windows, padded_size, window_size)
|
| 140 |
+
|
| 141 |
+
#Add padding to original image
|
| 142 |
+
original_image = Image.open(image_dir)
|
| 143 |
+
image_np = np.array(original_image)
|
| 144 |
+
image_height, image_width = image_np.shape[:2]
|
| 145 |
+
pad_height, pad_width = utils.calculate_padding(image_height, image_width, window_size, stride)
|
| 146 |
+
original_image = ImageOps.expand(original_image, border=(0, 0, pad_width, pad_height), fill=0)
|
| 147 |
+
comparison = self.compare_images(original_image, reconstructed_image)
|
| 148 |
+
self.assertEqual(comparison, "Identical")
|
| 149 |
+
|
| 150 |
+
|
| 151 |
+
if __name__ == '__main__':
|
| 152 |
+
unittest.main()
|
tests/test_sliding_window.py
ADDED
|
@@ -0,0 +1,186 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from src.utils import utils
|
| 2 |
+
import unittest
|
| 3 |
+
from PIL import Image
|
| 4 |
+
import numpy as np
|
| 5 |
+
import os
|
| 6 |
+
from parameterized import parameterized
|
| 7 |
+
|
| 8 |
+
|
| 9 |
+
# TODO: remove test cases where stride>window & window>image
|
| 10 |
+
class TestSlidingWindowFunction(unittest.TestCase):
|
| 11 |
+
"""
|
| 12 |
+
Unit tests for the sliding window function from the utils module.
|
| 13 |
+
The test does not dummy-proof the function:
|
| 14 |
+
windows_size <= 0, strides <=0, image_size = 0, invalid image format, invalid image path, etc.
|
| 15 |
+
were not taken into consideration in this test.
|
| 16 |
+
Window size can't be larger than image size
|
| 17 |
+
Stride can't be larger than window size
|
| 18 |
+
"""
|
| 19 |
+
|
| 20 |
+
def setUp(self):
|
| 21 |
+
"""Prepare the test images and their paths before each test."""
|
| 22 |
+
self.prepare_test_images()
|
| 23 |
+
|
| 24 |
+
def tearDown(self):
|
| 25 |
+
"""Clean up the test images after each test."""
|
| 26 |
+
self.clean_up_images()
|
| 27 |
+
|
| 28 |
+
@staticmethod
|
| 29 |
+
def create_test_image(width, height):
|
| 30 |
+
"""
|
| 31 |
+
Helper function to create a test image with pixel values increasing row-wise.
|
| 32 |
+
|
| 33 |
+
Args:
|
| 34 |
+
width (int): The width of the image.
|
| 35 |
+
height (int): The height of the image.
|
| 36 |
+
|
| 37 |
+
Returns:
|
| 38 |
+
Image: A PIL Image object.
|
| 39 |
+
"""
|
| 40 |
+
return Image.fromarray(np.arange(width * height).reshape(height, width).astype(np.uint8))
|
| 41 |
+
|
| 42 |
+
def prepare_test_images(self):
|
| 43 |
+
"""
|
| 44 |
+
Create and save test images of various sizes for testing.
|
| 45 |
+
"""
|
| 46 |
+
self.test_images = {
|
| 47 |
+
'image_1x1': self.create_test_image(1, 1),
|
| 48 |
+
'image_10x10': self.create_test_image(10, 10),
|
| 49 |
+
'image_11x11': self.create_test_image(11, 11),
|
| 50 |
+
'image_7x5': self.create_test_image(7, 5),
|
| 51 |
+
'image_15x15': self.create_test_image(15, 15),
|
| 52 |
+
'image_512x512': self.create_test_image(512, 512),
|
| 53 |
+
'image_1024x1024': self.create_test_image(1024, 1024),
|
| 54 |
+
'image_2048x2048': self.create_test_image(2048, 2048),
|
| 55 |
+
}
|
| 56 |
+
|
| 57 |
+
self.image_dirs = {}
|
| 58 |
+
for name, img in self.test_images.items():
|
| 59 |
+
img_dir = f'{name}.png'
|
| 60 |
+
img.save(img_dir)
|
| 61 |
+
self.image_dirs[name] = img_dir
|
| 62 |
+
|
| 63 |
+
def clean_up_images(self):
|
| 64 |
+
"""
|
| 65 |
+
Remove the image files created for testing.
|
| 66 |
+
"""
|
| 67 |
+
for img_dir in self.image_dirs.values():
|
| 68 |
+
if os.path.exists(img_dir):
|
| 69 |
+
os.remove(img_dir)
|
| 70 |
+
|
| 71 |
+
@staticmethod
|
| 72 |
+
def calculate_expected_number_of_windows(padded_size: tuple, window_size: tuple, stride: tuple):
|
| 73 |
+
"""
|
| 74 |
+
Calculate the expected number of windows that can be extracted from a padded image.
|
| 75 |
+
|
| 76 |
+
Args:
|
| 77 |
+
padded_size (tuple): The size of the padded image (height, width).
|
| 78 |
+
window_size (tuple): The size of the sliding window (height, width).
|
| 79 |
+
stride (tuple): The stride of the sliding window (height, width).
|
| 80 |
+
|
| 81 |
+
Returns:
|
| 82 |
+
int: The number of windows that can be extracted.
|
| 83 |
+
"""
|
| 84 |
+
number_of_positions_height = (padded_size[0] - window_size[0]) // stride[0] + 1
|
| 85 |
+
number_of_positions_width = (padded_size[1] - window_size[1]) // stride[1] + 1
|
| 86 |
+
return number_of_positions_height * number_of_positions_width
|
| 87 |
+
|
| 88 |
+
@staticmethod
|
| 89 |
+
def calculate_expected_padding(image_dir: str, window_size: tuple, stride: tuple):
|
| 90 |
+
"""
|
| 91 |
+
Calculate the expected dimensions of the padded image based on the stride.
|
| 92 |
+
|
| 93 |
+
Args:
|
| 94 |
+
image_dir (str): The path to the image file.
|
| 95 |
+
window_size (tuple): The size of the sliding window (height, width).
|
| 96 |
+
stride (tuple): The stride of the sliding window (height, width).
|
| 97 |
+
Returns:
|
| 98 |
+
tuple: The expected padded size (height, width).
|
| 99 |
+
"""
|
| 100 |
+
with (Image.open(image_dir) as image):
|
| 101 |
+
image_np = np.array(image)
|
| 102 |
+
image_height, image_width = image_np.shape[:2]
|
| 103 |
+
expected_padded_height, expected_padded_width = utils.calculate_padding(image_height, image_width,
|
| 104 |
+
window_size, stride)
|
| 105 |
+
|
| 106 |
+
return image_height + expected_padded_height, image_width + expected_padded_width
|
| 107 |
+
|
| 108 |
+
@parameterized.expand([
|
| 109 |
+
# Parameters: (image_name, window_size, stride)
|
| 110 |
+
('image_11x11', (3, 3), (3, 3)),
|
| 111 |
+
('image_1x1', (1, 1), (1, 1)),
|
| 112 |
+
('image_10x10', (3, 3), (3, 3)),
|
| 113 |
+
('image_7x5', (3, 3), (2, 2)),
|
| 114 |
+
('image_15x15', (5, 5), (5, 5)),
|
| 115 |
+
('image_10x10', (1, 1), (3, 3)), # Small window
|
| 116 |
+
('image_10x10', (10, 10), (1, 1)), # Large window
|
| 117 |
+
('image_7x5', (8, 8), (2, 2)), # Window larger than image
|
| 118 |
+
('image_15x15', (7, 7), (5, 5)), # Window larger but within bounds
|
| 119 |
+
('image_10x10', (3, 3), (3, 3)), # Stride equal to window size
|
| 120 |
+
('image_10x10', (3, 3), (4, 4)), # Stride larger than window size
|
| 121 |
+
('image_10x10', (3, 3), (2, 2)), # Stride smaller than window size
|
| 122 |
+
('image_7x5', (2, 2), (1, 1)), # Small stride
|
| 123 |
+
('image_7x5', (1, 1), (2, 2)), # Small window on larger image
|
| 124 |
+
('image_10x10', (1, 1), (5, 5)), # Large stride, small window
|
| 125 |
+
('image_15x15', (3, 3), (3, 3)), # Regular case
|
| 126 |
+
('image_10x10', (3, 3), (10, 10)), # Stride equals image size
|
| 127 |
+
('image_512x512', (50, 50), (50, 50)),
|
| 128 |
+
('image_1024x1024', (100, 100), (100, 100)),
|
| 129 |
+
('image_2048x2048', (200, 200), (200, 200)),
|
| 130 |
+
])
|
| 131 |
+
def test_padded_size(self, image_name, window_size, stride):
|
| 132 |
+
"""
|
| 133 |
+
Test that the image is correctly padded to the expected size.
|
| 134 |
+
|
| 135 |
+
Args:
|
| 136 |
+
image_name (str): The name of the test image.
|
| 137 |
+
window_size (tuple): The size of the sliding window (height, width).
|
| 138 |
+
stride (tuple): The stride of the sliding window (height, width).
|
| 139 |
+
"""
|
| 140 |
+
image_dir = self.image_dirs[image_name]
|
| 141 |
+
windows, padded_size = utils.sliding_window(image_dir, window_size, stride)
|
| 142 |
+
expected_padded_height, expected_padded_width = self.calculate_expected_padding(image_dir, window_size, stride)
|
| 143 |
+
self.assertEqual(padded_size, (expected_padded_width, expected_padded_height))
|
| 144 |
+
|
| 145 |
+
@parameterized.expand([
|
| 146 |
+
# Parameters: (image_name, window_size, stride)
|
| 147 |
+
('image_512x512', (500, 500), (50, 50)),
|
| 148 |
+
('image_1x1', (1, 1), (1, 1)),
|
| 149 |
+
('image_1x1', (2, 2), (1, 1)),
|
| 150 |
+
('image_10x10', (3, 3), (3, 3)),
|
| 151 |
+
('image_11x11', (3, 3), (3, 3)),
|
| 152 |
+
('image_7x5', (3, 3), (2, 2)),
|
| 153 |
+
('image_15x15', (5, 5), (5, 5)),
|
| 154 |
+
('image_10x10', (1, 1), (3, 3)), # Small window
|
| 155 |
+
('image_10x10', (10, 10), (1, 1)), # Large window
|
| 156 |
+
('image_7x5', (8, 8), (2, 2)), # Window larger than image
|
| 157 |
+
('image_15x15', (7, 7), (5, 5)), # Window larger but within bounds
|
| 158 |
+
('image_10x10', (3, 3), (3, 3)), # Stride equal to window size
|
| 159 |
+
('image_10x10', (3, 3), (4, 4)), # Stride larger than window size
|
| 160 |
+
('image_10x10', (3, 3), (2, 2)), # Stride smaller than window size
|
| 161 |
+
('image_7x5', (2, 2), (1, 1)), # Small stride
|
| 162 |
+
('image_7x5', (1, 1), (2, 2)), # Small window on larger image
|
| 163 |
+
('image_10x10', (1, 1), (5, 5)), # Large stride, small window
|
| 164 |
+
('image_15x15', (3, 3), (3, 3)), # Regular case
|
| 165 |
+
('image_10x10', (3, 3), (10, 10)), # Stride equals image size
|
| 166 |
+
('image_512x512', (50, 50), (50, 50)),
|
| 167 |
+
('image_1024x1024', (100, 100), (100, 100)),
|
| 168 |
+
('image_2048x2048', (200, 200), (200, 200)),
|
| 169 |
+
])
|
| 170 |
+
def test_number_of_windows(self, image_name, window_size, stride):
|
| 171 |
+
"""
|
| 172 |
+
Test that the number of windows generated matches the expected count.
|
| 173 |
+
|
| 174 |
+
Args:
|
| 175 |
+
image_name (str): The name of the test image.
|
| 176 |
+
window_size (tuple): The size of the sliding window (height, width).
|
| 177 |
+
stride (tuple): The stride of the sliding window (height, width).
|
| 178 |
+
"""
|
| 179 |
+
image_dir = self.image_dirs[image_name]
|
| 180 |
+
windows, padded_size = utils.sliding_window(image_dir, window_size, stride)
|
| 181 |
+
expected_number_of_windows = self.calculate_expected_number_of_windows(padded_size, window_size, stride)
|
| 182 |
+
self.assertEqual(len(windows), expected_number_of_windows)
|
| 183 |
+
|
| 184 |
+
|
| 185 |
+
if __name__ == '__main__':
|
| 186 |
+
unittest.main()
|
weights/download_weights.py
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
This script downloads the weights of the ResNet model from the following link:
|
| 3 |
+
"""
|
| 4 |
+
# TODO: Add the link to the weights
|