File size: 5,597 Bytes
087921f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
import torch
import torch.nn as nn
import torch.nn.functional as F
from .base_network import BaseNetwork
import random
from .blocks import Conv2dBlock

def inst_id_to_label_id(inst_id: int) -> int:
    return int(inst_id) % 120

class DeeperZencoder(BaseNetwork):
    def __init__(self,cfg, input_nc = 3, output_nc = 512, ngf=32, n_downsampling=2, norm_layer=nn.InstanceNorm2d):
        super().__init__()
        self.cfg = cfg
        self.n_downsampling = 6
        self.non_spade_norm_layer = norm_layer
        self.output_nc = cfg["style_length"]
        input_nc = input_nc +1 #RGB + Valid

        self.gamma = nn.Linear(1, 1)
        self.beta = nn.Linear(1, 1)
        ### downsample
        input_nc = cfg['input_nc']
        ngf = cfg['ngf']
        output_nc = cfg['style_length']
        lab_nc = cfg['lab_dim'] + 1
        g_norm = cfg['G_norm_type']
        self.enc1 = Conv2dBlock(4, 16, kernel_size=3, stride=1, padding=1, norm=g_norm, activation='lrelu') # 16, 256, 256
        self.enc2 = Conv2dBlock(16, 32, kernel_size=3, stride=2, padding=1, norm=g_norm, activation='lrelu') # 32, 128, 128
        self.enc3 = Conv2dBlock(32, 64, kernel_size=3, stride=2, padding=1, norm=g_norm, activation='lrelu') # 64, 64, 64
        self.enc4 = Conv2dBlock(64, 128, kernel_size=3, stride=2, padding=1, norm=g_norm, activation='lrelu') # 128, 32, 32
        self.enc5 = Conv2dBlock(128, 256, kernel_size=3, stride=2, padding=1, norm=g_norm, activation='lrelu') # 256, 16, 16
        self.enc6 = Conv2dBlock(256, 512, kernel_size=3, stride=1, padding=1, dilation=1, norm=g_norm, activation='lrelu') # 512, 16, 16

        # Decoder layers
        self.dec5 = Conv2dBlock(512+128, 512, kernel_size=3, stride=1, padding=1, norm=g_norm,
                                   activation='lrelu')
        self.dec4 = Conv2dBlock(512+64, 256, kernel_size=3, stride=1, padding=1, norm=g_norm,
                                   activation='lrelu')
        self.dec3 = Conv2dBlock(256+32, 256, kernel_size=3, stride=1, padding=1, norm=g_norm,
                                   activation='lrelu')
        self.dec2 = Conv2dBlock(256+16, 256, kernel_size=3, stride=1, padding=1, norm=g_norm,
                                   activation='lrelu')
        self.dec1 = Conv2dBlock(256, output_nc, kernel_size=3, stride=1, padding=1, norm='none', activation='tanh')

    def forward(self, input, segmap, valids, instance_map=None, cached_codes=None):
        masked_input = input*valids
        masked_input = torch.cat((masked_input, valids), 1)
        codes = self._forward_layers(masked_input)

        original_input = torch.cat((input, torch.ones(valids.shape, device=valids.device, dtype=valids.dtype)), 1)
        unmasked_codes = self._forward_layers(original_input)

        if instance_map is None:
            styles_map = segmap
        else:
            styles_map = instance_map
        styles_map = F.interpolate(styles_map, size=codes.size()[2:], mode='nearest')
        valids = F.interpolate(valids, size=codes.size()[2:], mode='nearest')
        # instance-wise average pooling
        style_codes = codes.clone()
        for b in range(input.size()[0]):
            inst_list = torch.unique(styles_map[b]).to(torch.long)
            for i in inst_list:
                indices = (styles_map[b:b+1] == int(i)).nonzero() # n x 4 
                valid_ins = valids[indices[:,0] + b, :, indices[:,2], indices[:,3]]
                if valid_ins.sum() > 0:
                    output_ins = codes[indices[:,0] + b, :, indices[:,2], indices[:,3]]
                    mean_feat = output_ins.mean(dim=0).expand_as(output_ins)
                    valid_ratio = valid_ins.mean().unsqueeze(0)
                    mean_feat = mean_feat * torch.sigmoid(self.gamma(valid_ratio)) + torch.sigmoid(self.beta(valid_ratio))
                elif self.cfg["is_train"]:
                    unmasked_output_ins = unmasked_codes[indices[:,0] + b, :, indices[:,2], indices[:,3]]
                    mean_feat = unmasked_output_ins.mean(dim=0).expand_as(unmasked_output_ins)
                    valid_ratio = torch.ones((1,))
                    mean_feat = mean_feat * torch.sigmoid(self.gamma(valid_ratio)) + torch.sigmoid(self.beta(valid_ratio))
                else:
                    code_list = cached_codes[inst_id_to_label_id(i)]
                    random_code = random.choice(code_list)
                    mean_feat = random_code.expand_as(codes[indices[:,0] + b, :, indices[:,2], indices[:,3]])
                    valid_ratio = torch.ones((1,))
                    mean_feat = mean_feat * torch.sigmoid(self.gamma(valid_ratio)) + torch.sigmoid(self.beta(valid_ratio))
                style_codes[indices[:,0] + b, :, indices[:,2], indices[:,3]] = mean_feat
        return style_codes

    def _forward_layers(self, input):
        # Encoder
        e1 = self.enc1(input)
        e2 = self.enc2(e1)
        e3 = self.enc3(e2)
        e4 = self.enc4(e3)
        e5 = self.enc5(e4)
        x = self.enc6(e5)
        
        x = F.interpolate(x, scale_factor=2, mode='bilinear')
        x = torch.cat((x, e4), dim=1)
        x = self.dec5(x)
        
        x = F.interpolate(x, scale_factor=2, mode='bilinear')
        x = torch.cat((x, e3), dim=1)
        x = self.dec4(x)
        
        x = F.interpolate(x, scale_factor=2, mode='bilinear')
        x = torch.cat((x, e2), dim=1)
        x = self.dec3(x)
        
        x = F.interpolate(x, scale_factor=2, mode='bilinear')
        x = torch.cat((x, e1), dim=1)
        x = self.dec2(x)
        
        x = self.dec1(x)

        return x