from __future__ import print_function, division
import math
import gc

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.data

from .shared import conv_block, up_conv


class UNet_foreground(nn.Module):

    def __init__(self, in_ch=3, out_ch=1):
        super().__init__()

        n1 = 64
        filters = [n1, n1 * 2, n1 * 4, n1 * 8, n1 * 16]

        self.Maxpool1 = nn.MaxPool2d(kernel_size=2, stride=2)
        self.Maxpool2 = nn.MaxPool2d(kernel_size=2, stride=2)
        self.Maxpool3 = nn.MaxPool2d(kernel_size=2, stride=2)
        self.Maxpool4 = nn.MaxPool2d(kernel_size=2, stride=2)

        self.Conv1 = conv_block(in_ch, filters[0])
        self.Conv2 = conv_block(filters[0], filters[1])
        self.Conv3 = conv_block(filters[1], filters[2])
        self.Conv4 = conv_block(filters[2], filters[3])
        self.Conv5 = conv_block(filters[3], filters[4])

        self.Up5 = up_conv(filters[4], filters[3])
        self.Up_conv5 = conv_block(filters[4], filters[3])

        self.Up4 = up_conv(filters[3], filters[2])
        self.Up_conv4 = conv_block(filters[3], filters[2])

        self.Up3 = up_conv(filters[2], filters[1])
        self.Up_conv3 = conv_block(filters[2], filters[1])

        self.Up2 = up_conv(filters[1], filters[0])
        self.Up_conv2 = conv_block(filters[1], filters[0])

        self.Conv = nn.Conv2d(filters[0], out_ch,
                              kernel_size=1, stride=1, padding=0)

        self.scene_vector_conv = nn.Sequential(
            nn.MaxPool2d(kernel_size=2, stride=2),
            conv_block(filters[4], filters[2]),
            nn.MaxPool2d(kernel_size=2, stride=2),
            nn.Conv2d(filters[2], filters[2], kernel_size=4, stride=1, padding=0),
            nn.BatchNorm2d(filters[2]),
            nn.ReLU(inplace=True)
        )

        self.e5_reduce = conv_block(filters[4], filters[2])
        self.e4_reduce = conv_block(filters[3], filters[2])
        self.e3_reduce = conv_block(filters[2], filters[2])
        self.e2_reduce = conv_block(filters[1], filters[2])

        self.e5_enc = conv_block(filters[4], filters[4])
        self.e4_enc = conv_block(filters[3], filters[3])
        self.e3_enc = conv_block(filters[2], filters[2])
        self.e2_enc = conv_block(filters[1], filters[1])

        # self.active = torch.nn.Softmax(dim=1)

    def forward(self, x):
        e1 = self.Conv1(x)

        e2 = self.Maxpool1(e1)
        e2 = self.Conv2(e2)

        e3 = self.Maxpool2(e2)
        e3 = self.Conv3(e3)

        e4 = self.Maxpool3(e3)
        e4 = self.Conv4(e4)

        e5 = self.Maxpool4(e4)
        e5 = self.Conv5(e5)

        # foreground-aware
        scene = self.scene_vector_conv(e5)  # bs, 256, 1, 1
        e2r, e3r, e4r, e5r = self.e2_reduce(e2), self.e3_reduce(e3), self.e4_reduce(e4), self.e5_reduce(e5)  # bs, 256, .., ..
        e2e, e3e, e4e, e5e = self.e2_enc(e2), self.e3_enc(e3), self.e4_enc(e4), self.e5_enc(e5)
        r2, r3, r4, r5 = [(scene*i).sum(dim=1, keepdim=True).sigmoid() for i in [e2r, e3r, e4r, e5r]]
        z2, z3, z4, z5 = r2*e2e, r3*e3e, r4*e4e, r5*e5e

        d5 = self.Up5(z5)
        d5 = torch.cat((z4, d5), dim=1)

        d5 = self.Up_conv5(d5)

        d4 = self.Up4(d5)
        d4 = torch.cat((z3, d4), dim=1)
        d4 = self.Up_conv4(d4)

        d3 = self.Up3(d4)
        d3 = torch.cat((z2, d3), dim=1)
        d3 = self.Up_conv3(d3)

        d2 = self.Up2(d3)
        d2 = torch.cat((e1, d2), dim=1)
        d2 = self.Up_conv2(d2)

        out = self.Conv(d2)

        # d1 = self.active(out)

        return out
