musk12 commited on
Commit
05fba58
·
verified ·
1 Parent(s): 0956343

Upload 4 files

Browse files
Files changed (4) hide show
  1. app.py +98 -0
  2. requirements.txt +6 -0
  3. unet.py +174 -0
  4. unet_car_final.pth +3 -0
app.py ADDED
@@ -0,0 +1,98 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import torch
3
+ import torchvision.transforms as T
4
+ import torchvision.transforms.functional as TF
5
+ import numpy as np
6
+ from PIL import Image
7
+ from flask import Flask, render_template, request
8
+
9
+ app = Flask(__name__)
10
+
11
+ device = "cuda" if torch.cuda.is_available() else "cpu"
12
+
13
+ # Load model (assuming UNet is defined in unet.py)
14
+ def load_model():
15
+ from unet import UNet
16
+ model = UNet().to(device)
17
+ model.load_state_dict(torch.load("unet_car_final.pth", map_location=device))
18
+ model.eval()
19
+ return model
20
+
21
+ model = load_model()
22
+
23
+ # Image transforms
24
+ img_transform = T.Compose([
25
+ T.Resize((256, 256)),
26
+ T.ToTensor(),
27
+ T.Normalize(mean=[0.485, 0.456, 0.406],
28
+ std=[0.229, 0.224, 0.225])
29
+ ])
30
+
31
+ UPLOAD_FOLDER = "static"
32
+ os.makedirs(UPLOAD_FOLDER, exist_ok=True)
33
+
34
+ @app.route("/", methods=["GET", "POST"])
35
+ def index():
36
+ orig = "static/input.jpg" if os.path.exists(os.path.join(UPLOAD_FOLDER, "input.jpg")) else None
37
+ mask = "static/mask.png" if os.path.exists(os.path.join(UPLOAD_FOLDER, "mask.png")) else None
38
+ overlay = "static/overlay.png" if os.path.exists(os.path.join(UPLOAD_FOLDER, "overlay.png")) else None
39
+
40
+ if request.method == "POST":
41
+ # Handle image upload
42
+ if "image" in request.files:
43
+ file = request.files["image"]
44
+ if file.filename == "":
45
+ return "No file selected", 400
46
+
47
+ # Save uploaded image
48
+ img_path = os.path.join(UPLOAD_FOLDER, "input.jpg")
49
+ file.save(img_path)
50
+ orig = "static/input.jpg"
51
+ # Clear previous results
52
+ mask = None
53
+ overlay = None
54
+ for path in [os.path.join(UPLOAD_FOLDER, "mask.png"), os.path.join(UPLOAD_FOLDER, "overlay.png")]:
55
+ if os.path.exists(path):
56
+ os.remove(path)
57
+
58
+ # Handle segmentation
59
+ if "segment" in request.form:
60
+ img_path = os.path.join(UPLOAD_FOLDER, "input.jpg")
61
+ if not os.path.exists(img_path):
62
+ return "No image available for segmentation", 400
63
+
64
+ image = Image.open(img_path).convert("RGB")
65
+ input_tensor = img_transform(image).unsqueeze(0).to(device)
66
+
67
+ # Predict
68
+ with torch.no_grad():
69
+ output = model(input_tensor)
70
+ pred_mask = torch.sigmoid(output)
71
+ pred_mask = (pred_mask > 0.5).float()
72
+
73
+ # Resize mask back to original image size
74
+ mask_resized = TF.resize(
75
+ TF.to_pil_image(pred_mask.squeeze().cpu()),
76
+ size=image.size[::-1],
77
+ interpolation=Image.NEAREST
78
+ )
79
+
80
+ # Save mask
81
+ mask_path = os.path.join(UPLOAD_FOLDER, "mask.png")
82
+ mask_resized.save(mask_path)
83
+
84
+ # Create overlay
85
+ mask_np = np.array(mask_resized)
86
+ overlay = np.array(image).copy()
87
+ overlay[mask_np > 128] = [255, 0, 0]
88
+ overlay_img = Image.fromarray(overlay)
89
+ overlay_path = os.path.join(UPLOAD_FOLDER, "overlay.png")
90
+ overlay_img.save(overlay_path)
91
+
92
+ mask = "static/mask.png"
93
+ overlay = "static/overlay.png"
94
+
95
+ return render_template("index.html", orig=orig, mask=mask, overlay=overlay)
96
+
97
+ if __name__ == "__main__":
98
+ app.run(debug=True)
requirements.txt ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ Flask==3.0.3
2
+ torch==2.3.1
3
+ torchvision==0.18.1
4
+ pillow==10.4.0
5
+ numpy==1.26.4
6
+ gunicorn==22.0.0
unet.py ADDED
@@ -0,0 +1,174 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from torch import nn
3
+
4
+ def center_crop(encoder_feature, target_tensor):
5
+ _, _, H, W = target_tensor.shape
6
+ enc_H, enc_W = encoder_feature.shape[2], encoder_feature.shape[3]
7
+ crop_H = (enc_H - H) // 2
8
+ crop_W = (enc_W - W) // 2
9
+ return encoder_feature[:, :, crop_H:crop_H+H, crop_W:crop_W+W]
10
+
11
+ class UNet(nn.Module):
12
+
13
+ def __init__(self):
14
+ super().__init__()
15
+
16
+ self.max_pool_2x2 = nn.MaxPool2d(kernel_size=2, stride=2)
17
+
18
+ # Encoder
19
+ self.encoder1 = nn.Sequential(
20
+ nn.Conv2d(in_channels=3, out_channels=64, kernel_size=3, padding=1),
21
+ nn.BatchNorm2d(64),
22
+ nn.ReLU(inplace=True),
23
+
24
+ nn.Conv2d(in_channels=64, out_channels=64, kernel_size=3, padding=1),
25
+ nn.BatchNorm2d(64),
26
+ nn.ReLU(inplace=True),
27
+ )
28
+
29
+ self.encoder2 = nn.Sequential(
30
+ nn.Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=1),
31
+ nn.BatchNorm2d(128),
32
+ nn.ReLU(inplace=True),
33
+
34
+ nn.Conv2d(in_channels=128, out_channels=128, kernel_size=3, padding=1),
35
+ nn.BatchNorm2d(128),
36
+ nn.ReLU(inplace=True),
37
+ )
38
+
39
+ self.encoder3 = nn.Sequential(
40
+ nn.Conv2d(in_channels=128, out_channels=256, kernel_size=3, padding=1),
41
+ nn.BatchNorm2d(256),
42
+ nn.ReLU(inplace=True),
43
+
44
+ nn.Conv2d(in_channels=256, out_channels=256, kernel_size=3, padding=1),
45
+ nn.BatchNorm2d(256),
46
+ nn.ReLU(inplace=True),
47
+ )
48
+
49
+ self.encoder4 = nn.Sequential(
50
+ nn.Conv2d(in_channels=256, out_channels=512, kernel_size=3, padding=1),
51
+ nn.BatchNorm2d(512),
52
+ nn.ReLU(inplace=True),
53
+
54
+ nn.Conv2d(in_channels=512, out_channels=512, kernel_size=3, padding=1),
55
+ nn.BatchNorm2d(512),
56
+ nn.ReLU(inplace=True),
57
+ )
58
+
59
+ self.encoder5 = nn.Sequential(
60
+ nn.Conv2d(in_channels=512, out_channels=1024, kernel_size=3, padding=1),
61
+ nn.BatchNorm2d(1024),
62
+ nn.ReLU(inplace=True),
63
+
64
+ nn.Conv2d(in_channels=1024, out_channels=1024, kernel_size=3, padding=1),
65
+ nn.BatchNorm2d(1024),
66
+ nn.ReLU(inplace=True),
67
+ )
68
+
69
+
70
+ # Decoder
71
+ self.upconv_1 = nn.Sequential(
72
+ nn.ConvTranspose2d(in_channels=1024, out_channels=512, kernel_size=2, stride=2),
73
+ nn.BatchNorm2d(512),
74
+ nn.ReLU(inplace=True)
75
+ )
76
+
77
+ self.decoder1 = nn.Sequential(
78
+ nn.Conv2d(in_channels=1024, out_channels=512, kernel_size=3, padding=1),
79
+ nn.BatchNorm2d(512),
80
+ nn.ReLU(inplace=True),
81
+
82
+ nn.Conv2d(in_channels=512, out_channels=256, kernel_size=3, padding=1),
83
+ nn.BatchNorm2d(256),
84
+ nn.ReLU(inplace=True),
85
+ )
86
+
87
+
88
+ self.upconv_2 = nn.Sequential(
89
+ nn.ConvTranspose2d(in_channels=256, out_channels=256, kernel_size=2, stride=2),
90
+ nn.BatchNorm2d(256),
91
+ nn.ReLU(inplace=True)
92
+ )
93
+
94
+ self.decoder2 = nn.Sequential(
95
+ nn.Conv2d(in_channels=512, out_channels=256, kernel_size=3, padding=1),
96
+ nn.BatchNorm2d(256),
97
+ nn.ReLU(inplace=True),
98
+
99
+ nn.Conv2d(in_channels=256, out_channels=128, kernel_size=3, padding=1),
100
+ nn.BatchNorm2d(128),
101
+ nn.ReLU(inplace=True),
102
+ )
103
+
104
+
105
+ self.upconv_3 = nn.Sequential(
106
+ nn.ConvTranspose2d(in_channels=128, out_channels=128, kernel_size=2, stride=2),
107
+ nn.BatchNorm2d(128),
108
+ nn.ReLU(inplace=True)
109
+ )
110
+
111
+ self.decoder3 = nn.Sequential(
112
+ nn.Conv2d(in_channels=256, out_channels=128, kernel_size=3, padding=1),
113
+ nn.BatchNorm2d(128),
114
+ nn.ReLU(inplace=True),
115
+
116
+ nn.Conv2d(in_channels=128, out_channels=64, kernel_size=3, padding=1),
117
+ nn.BatchNorm2d(64),
118
+ nn.ReLU(inplace=True),
119
+ )
120
+
121
+
122
+ self.upconv_4 = nn.Sequential(
123
+ nn.ConvTranspose2d(in_channels=64, out_channels=64, kernel_size=2, stride=2),
124
+ nn.BatchNorm2d(64),
125
+ nn.ReLU(inplace=True)
126
+ )
127
+
128
+ self.decoder4 = nn.Sequential(
129
+ nn.Conv2d(in_channels=128, out_channels=64, kernel_size=3, padding=1),
130
+ nn.BatchNorm2d(64),
131
+ nn.ReLU(inplace=True),
132
+
133
+ nn.Conv2d(in_channels=64, out_channels=64, kernel_size=3, padding=1),
134
+ nn.BatchNorm2d(64),
135
+ nn.ReLU(inplace=True),
136
+ )
137
+
138
+ self.final_conv = nn.Conv2d(64, 1, kernel_size=1)
139
+
140
+
141
+ def forward(self, image):
142
+
143
+ # Encoder
144
+ e1 = self.encoder1(image)
145
+ e2 = self.encoder2(self.max_pool_2x2(e1))
146
+ e3 = self.encoder3(self.max_pool_2x2(e2))
147
+ e4 = self.encoder4(self.max_pool_2x2(e3))
148
+ e5 = self.encoder5(self.max_pool_2x2(e4))
149
+
150
+ # Decoder
151
+ d1 = self.upconv_1(e5)
152
+ c1 = center_crop(e4, d1)
153
+ d1 = torch.cat([d1, c1], dim=1)
154
+ d1 = self.decoder1(d1)
155
+
156
+ d2 = self.upconv_2(d1)
157
+ c2 = center_crop(e3, d2)
158
+ d2 = torch.cat([d2, c2], dim=1)
159
+ d2 = self.decoder2(d2)
160
+
161
+ d3 = self.upconv_3(d2)
162
+ c3 = center_crop(e2, d3)
163
+ d3 = torch.cat([d3, c3], dim=1)
164
+ d3 = self.decoder3(d3)
165
+
166
+ d4 = self.upconv_4(d3)
167
+ c4 = center_crop(e1, d4)
168
+ d4 = torch.cat([d4, c4], dim=1)
169
+ d4 = self.decoder4(d4)
170
+
171
+ # Final output
172
+ out = self.final_conv(d4)
173
+
174
+ return out
unet_car_final.pth ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:014e1c1e1e7660b7697bf0c73abd65a6bc295ba102fdd64e86d3b4b10e73297a
3
+ size 116710042