davidwardan commited on
Commit
f0196c3
·
verified ·
1 Parent(s): 0d61e2c

Upload folder using huggingface_hub

Browse files
.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: Da Detection
3
- emoji: 💻
4
- colorFrom: blue
5
- colorTo: yellow
6
  sdk: gradio
7
- sdk_version: 5.6.0
8
- app_file: app.py
9
- pinned: false
10
  ---
 
 
11
 
12
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ ![image](cover.jpg)
9
 
10
+ ![alt text](https://img.shields.io/badge/Status-Under%20Development-red)
11
+ ![alt text](https://img.shields.io/badge/Version-0.1.0-blue)
12
+ ![alt text](https://img.shields.io/badge/License-MIT-green)
13
+ ![alt text](https://img.shields.io/badge/Institution-UN%20Habitat-blue)
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