# -- Code Cell --
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader
from torchvision import models, transforms
from PIL import Image
from torchvision.transforms import functional as TF

# -- Code Cell --
import numpy as np
import os

# -- Code Cell --
class TrainDataset(Dataset):
    def __init__(self,path):
        self.img_tf = transforms.Compose([
            transforms.Resize((128, 128)),
            transforms.ToTensor(),
        ])

        self.mask_tf = transforms.Compose([
            transforms.Resize((128, 128)),
        ])
        self.files = os.listdir(path)
    def __len__(self):
        return len(self.files)
    def __getitem__(self, index):
        img = Image.open(f"./train/images/img_{index+1:03d}.png")
        mask = Image.open(f"./train/masks/mask_{index+1:03d}.png")
        img = self.img_tf(img)
        mask = self.mask_tf(mask)
        mask = np.array(mask, dtype=np.int64)
        mask_bin = (mask > 1).astype(np.float32)
        mask_bin = torch.tensor(mask_bin).unsqueeze(0)  
        return img, mask_bin

# -- Code Cell --
train_ds = TrainDataset("./train/images")
train = DataLoader(train_ds, batch_size=16, num_workers=0, shuffle=True)

# -- Code Cell --
class DoubleConv(nn.Module):
    def __init__(self, in_ch, out_ch):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
            nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
        )

    def forward(self, x):
        return self.conv(x)


class UNet(nn.Module):
    def __init__(self, in_channels=3, out_channels=1):
        super().__init__()
       
        self.enc1 = DoubleConv(in_channels, 64)
        self.pool1 = nn.MaxPool2d(2)

        self.enc2 = DoubleConv(64, 128)
        self.pool2 = nn.MaxPool2d(2)

        self.enc3 = DoubleConv(128, 256)
        self.pool3 = nn.MaxPool2d(2)

        self.enc4 = DoubleConv(256, 512)
        self.pool4 = nn.MaxPool2d(2)

        self.bottleneck = DoubleConv(512, 1024)
        
        self.up4 = nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2)
        self.dec4 = DoubleConv(1024, 512)

        self.up3 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)
        self.dec3 = DoubleConv(512, 256)

        self.up2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2)
        self.dec2 = DoubleConv(256, 128)

        self.up1 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)
        self.dec1 = DoubleConv(128, 64)

        self.out_conv = nn.Conv2d(64, out_channels, kernel_size=1)

    def forward(self, x):
        
        c1 = self.enc1(x)
        p1 = self.pool1(c1)

        c2 = self.enc2(p1)
        p2 = self.pool2(c2)

        c3 = self.enc3(p2)
        p3 = self.pool3(c3)
        c4 = self.enc4(p3)
        p4 = self.pool4(c4)
        bn = self.bottleneck(p4)
        u4 = self.up4(bn)
        x4 = torch.cat([u4, c4], dim=1)
        d4 = self.dec4(x4)

        u3 = self.up3(d4)
        x3 = torch.cat([u3, c3], dim=1)
        d3 = self.dec3(x3)

        u2 = self.up2(d3)
        x2 = torch.cat([u2, c2], dim=1)
        d2 = self.dec2(x2)

        u1 = self.up1(d2)
        x1 = torch.cat([u1, c1], dim=1)
        d1 = self.dec1(x1)

        logits = self.out_conv(d1)
        return logits

# -- Code Cell --
device = torch.device("cuda")

# -- Code Cell --
model = UNet()
optimizer = optim.Adam(model.parameters(), lr = 1e-3)
criterion = nn.BCEWithLogitsLoss()
model = model.to(device)

# -- Code Cell --
scaler = torch.amp.GradScaler()

for epoch in range(30):
    for img, masks in train:
        img = img.to(device)
        masks = masks.to(device)

        optimizer.zero_grad()

        with torch.amp.autocast(device_type='cuda'):
            output = model(img)
            loss = criterion(output, masks)

        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()

    print(f"Epoch{epoch+1}---Loss{loss.item():.4f}")

# -- Code Cell --
class TestDataset(Dataset):
    def __init__(self,path):
        self.img_tf = transforms.Compose([
            transforms.Resize((128,128)),
            transforms.ToTensor(),
        ])

        # self.mask_tf = transforms.Compose([
        #     transforms.Resize((128, 128)),
        # ])
        self.files = os.listdir(path)
    def __len__(self):
        return len(self.files)
    def __getitem__(self, index):
        img = Image.open(f"./test/images/img_{index+1:03d}.png")
        # mask = Image.open(f"./train/masks/mask_{index+1:03d}.png")
        img = self.img_tf(img)
        # mask = self.mask_tf(mask)
        # mask = np.array(mask, dtype=np.int64)
        # mask_bin = (mask > 1).astype(np.float32)
        # mask_bin = torch.tensor(mask_bin).unsqueeze(0)  
        return img

# -- Code Cell --
test_ds =TestDataset("./test/images")
test= DataLoader(test_ds, batch_size=16, num_workers=0)

# -- Code Cell --
from torchvision.transforms.functional import to_pil_image

# -- Code Cell --
model.eval()
count = 1

for images in test:
    images = images.to(device)

    with torch.no_grad():
        preds = model(images)

    for j in range(preds.shape[0]):
        mask = (preds[j] > 0.3).float() * 255   
        mask = mask.squeeze(0).cpu().byte()     

        pil_img = to_pil_image(mask)
        pil_img.save(f"C:/Users/David/Desktop/ProblemeML/vanatorul de glitchuri/preds/img_{count:03d}.png")
        count += 1

# -- Code Cell --
