Spaces:
Sleeping
Sleeping
| 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() |