6
import torchvision.transforms as T7
from torchvision.models.segmentation import deeplabv3_resnet5015
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")20
IMAGE_LIST = os.listdir(IMAGE_DIR)22
# --- Image Transformations ---23
transform_img = T.Compose(26
T.Resize((HEIGHT, WIDTH)),28
T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),32
transform_ann = T.Compose(35
T.Resize((HEIGHT, WIDTH), interpolation=T.InterpolationMode.NEAREST),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")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)51
ann_map = np.zeros(img.shape[:2], dtype=np.float32)52
if vessel is not None:53
ann_map[vessel == 1] = 154
if filled is not None:55
ann_map[filled == 1] = 257
img = transform_img(img)58
ann_map = transform_ann(ann_map)64
images = torch.zeros((BATCH_SIZE, 3, HEIGHT, WIDTH))65
annotations = torch.zeros((BATCH_SIZE, HEIGHT, WIDTH))67
for i in range(BATCH_SIZE):68
images[i], annotations[i] = read_random_image()69
return images, annotations72
# --- Model and Optimizer Setup ---73
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")75
net = deeplabv3_resnet50(pretrained=True)76
net.classifier[4] = torch.nn.Conv2d(256, 3, kernel_size=1)79
optimizer = torch.optim.Adam(net.parameters(), lr=LEARNING_RATE)80
criterion = torch.nn.CrossEntropyLoss()82
# --- Training Loop ---83
for step in range(10_000):84
images, annotations = load_batch()86
images = images.to(device)87
annotations = annotations.to(device)90
outputs = net(images)["out"]92
loss = criterion(outputs, annotations.long())96
prediction = torch.argmax(outputs[0], dim=0).cpu().detach().numpy()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]")101
torch.save(net.state_dict(), f"[object Object],step,[object Object].torch")102
print(f"Model saved: [object Object],step,[object Object].torch")