1
import os
2

3
import cv2
4
import numpy as np
5
import torch
6
import torchvision.transforms as T
7
from torchvision.models.segmentation import deeplabv3_resnet50
8

9
# --- Constants ---
10
LEARNING_RATE = 1e-5
11
WIDTH = 900
12
HEIGHT = 900
13
BATCH_SIZE = 3
14

15
TRAIN_FOLDER = "LabPics/Simple/Train/"
16
IMAGE_DIR = os.path.join(TRAIN_FOLDER, "Image")
17
FILLED_DIR = os.path.join(TRAIN_FOLDER, "Semantic/16_Filled")
18
VESSEL_DIR = os.path.join(TRAIN_FOLDER, "Semantic/1_Vessel")
19

20
IMAGE_LIST = os.listdir(IMAGE_DIR)
21

22
# --- Image Transformations ---
23
transform_img = T.Compose(
24
[
25
T.ToPILImage(),
26
T.Resize((HEIGHT, WIDTH)),
27
T.ToTensor(),
28
T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),
29
]
30
)
31

32
transform_ann = T.Compose(
33
[
34
T.ToPILImage(),
35
T.Resize((HEIGHT, WIDTH), interpolation=T.InterpolationMode.NEAREST),
36
T.ToTensor(),
37
]
38
)
39

40

41
# --- Load Random Image + Annotations ---
42
def read_random_image():
43
idx = np.random.randint(len(IMAGE_LIST))
44
image_path = os.path.join(IMAGE_DIR, IMAGE_LIST[idx])
45
base_name = IMAGE_LIST[idx].replace("jpg", "png")
46

47
img = cv2.imread(image_path)[:, :, :3]
48
filled = cv2.imread(os.path.join(FILLED_DIR, base_name), 0)
49
vessel = cv2.imread(os.path.join(VESSEL_DIR, base_name), 0)
50

51
ann_map = np.zeros(img.shape[:2], dtype=np.float32)
52
if vessel is not None:
53
ann_map[vessel == 1] = 1
54
if filled is not None:
55
ann_map[filled == 1] = 2
56

57
img = transform_img(img)
58
ann_map = transform_ann(ann_map)
59
return img, ann_map
60

61

62
# --- Load Batch ---
63
def load_batch():
64
images = torch.zeros((BATCH_SIZE, 3, HEIGHT, WIDTH))
65
annotations = torch.zeros((BATCH_SIZE, HEIGHT, WIDTH))
66

67
for i in range(BATCH_SIZE):
68
images[i], annotations[i] = read_random_image()
69
return images, annotations
70

71

72
# --- Model and Optimizer Setup ---
73
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
74

75
net = deeplabv3_resnet50(pretrained=True)
76
net.classifier[4] = torch.nn.Conv2d(256, 3, kernel_size=1)
77
net.to(device)
78

79
optimizer = torch.optim.Adam(net.parameters(), lr=LEARNING_RATE)
80
criterion = torch.nn.CrossEntropyLoss()
81

82
# --- Training Loop ---
83
for step in range(10_000):
84
images, annotations = load_batch()
85

86
images = images.to(device)
87
annotations = annotations.to(device)
88

89
net.zero_grad()
90
outputs = net(images)["out"]
91

92
loss = criterion(outputs, annotations.long())
93
loss.backward()
94
optimizer.step()
95

96
prediction = torch.argmax(outputs[0], dim=0).cpu().detach().numpy()
97

98
print(f"[object Object],step,[object Object]) Loss = [object Object],loss,[object Object],item,[object Object],[object Object],[object Object],[object Object],[object Object]")
99

100
if step % 1000 == 0:
101
torch.save(net.state_dict(), f"[object Object],step,[object Object].torch")
102
print(f"Model saved: [object Object],step,[object Object].torch")

0

WPM •0 •0

100%

ACC •0 •0

0s

TIME •0