ColorGradingAE / trainer.py
Suchinthana Wijesundara
Init commit
fcd8868
Raw
History Blame Contribute Delete
4.59 kB
import os
import cv2
import torch
import random
import numpy as np
from torch import nn
from torch.utils.data import Dataset,DataLoader
import tqdm
IMAGE_DIR="data/images"
DEVICE="cuda" if torch.cuda.is_available() else "cpu"
# -------------------------
# Feature extraction
# -------------------------
HUE_BINS=16
def extract_features(path):
img=cv2.imread(path)
img=cv2.cvtColor(
img,
cv2.COLOR_BGR2HSV
)
h=img[:,:,0]
s=img[:,:,1]/255
v=img[:,:,2]/255
features=[]
zones=[
("shadow",v<0.3),
("mid", (v>=0.3)&(v<0.7)),
("highlight",v>=0.7)
]
for _,mask in zones:
total=np.sum(mask)+1
for i in range(HUE_BINS):
hue_start=i*180/HUE_BINS
hue_end=(i+1)*180/HUE_BINS
hm=(
(h>=hue_start)
&
(h<hue_end)
&
mask
)
pixels=np.sum(hm)
if pixels>0:
features.extend([
pixels/total,
np.mean(s[hm]),
np.mean(v[hm])
])
else:
features.extend([
0,
0,
0
])
# global features
features.extend([
np.mean(s),
np.mean(v),
np.std(v)
])
return np.array(
features,
dtype=np.float32
)
# -------------------------
# Dataset
# -------------------------
class ColorDataset(Dataset):
def __init__(self):
self.files=[
os.path.join(
IMAGE_DIR,x
)
for x in os.listdir(IMAGE_DIR)
]
def __len__(self):
return len(self.files)
def __getitem__(self,i):
return torch.tensor(
extract_features(
self.files[i]
)
)
# -------------------------
# Masking
# -------------------------
def random_mask(x,ratio=0.3):
mask=torch.rand_like(x)<ratio
y=x.clone()
y[mask]=0
return y,mask
# -------------------------
# Model
# -------------------------
INPUT_SIZE=147
class Encoder(nn.Module):
def __init__(self):
super().__init__()
self.net=nn.Sequential(
nn.Linear(INPUT_SIZE,256),
nn.ReLU(),
nn.Linear(256,256),
nn.ReLU(),
nn.Linear(256,128)
)
def forward(self,x):
return self.net(x)
class Decoder(nn.Module):
def __init__(self):
super().__init__()
self.net=nn.Sequential(
nn.Linear(128,256),
nn.ReLU(),
nn.Linear(256,INPUT_SIZE)
)
def forward(self,x):
return self.net(x)
# -------------------------
# Contrastive Loss
# -------------------------
def contrastive_loss(
a,
b,
temperature=0.1):
a=nn.functional.normalize(a,dim=1)
b=nn.functional.normalize(b,dim=1)
logits=a@b.T / temperature
labels=torch.arange(
len(a)
).to(DEVICE)
return nn.functional.cross_entropy(
logits,
labels
)
# -------------------------
# Training
# -------------------------
def train():
dataset=ColorDataset()
loader=DataLoader(
dataset,
batch_size=2,
shuffle=True
)
encoder=Encoder().to(DEVICE)
decoder=Decoder().to(DEVICE)
optimizer=torch.optim.Adam(
list(encoder.parameters())
+
list(decoder.parameters()),
lr=1e-3
)
for epoch in tqdm.tqdm(range(15)):
total=0
for x in loader:
x=x.to(DEVICE)
masked,_=random_mask(x)
z1=encoder(masked)
z2=encoder(x)
recon=decoder(z1)
loss_rec=nn.functional.mse_loss(
recon,
x
)
loss_con=contrastive_loss(
z1,
z2
)
loss=(
0.7*loss_con
+
0.3*loss_rec
)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total+=loss.item()
print(
epoch,
total/len(loader)
)
torch.save(
encoder.state_dict(),
"encoder.pt"
)
torch.save(
decoder.state_dict(),
"decoder.pt"
)
if __name__=="__main__":
train()