from numpy.core.shape_base import block
from numpy.lib import stride_tricks
import torch
import numpy as np
from torch import einsum
from einops import rearrange 
import torch.nn as nn
import torch.nn.functional as F
import random

from torch.nn.modules.linear import Linear
from image_synthesis.utils.misc import instantiate_from_config
from image_synthesis.modeling.codecs.base_codec import BaseCodec
from image_synthesis.modeling.modules.vqgan_loss.vqperceptual import VQLPIPSWithDiscriminator
from image_synthesis.modeling.utils.misc import mask_with_top_k, logits_top_k, get_token_type
from image_synthesis.distributed.distributed import all_reduce


class LambdaWarmUpCosineScheduler:
    """
    note: use with a base_lr of 1.0
    """
    def __init__(self, warm_up_steps, lr_min, lr_max, lr_start, max_decay_steps, verbosity_interval=0):
        self.lr_warm_up_steps = warm_up_steps
        self.lr_start = lr_start
        self.lr_min = lr_min
        self.lr_max = lr_max
        self.lr_max_decay_steps = max_decay_steps
        self.last_lr = 0.
        self.verbosity_interval = verbosity_interval

    def schedule(self, n):
        if self.verbosity_interval > 0:
            if n % self.verbosity_interval == 0: print(f"current step: {n}, recent lr-multiplier: {self.last_lr}")
        if n < self.lr_warm_up_steps:
            lr = (self.lr_max - self.lr_start) / self.lr_warm_up_steps * n + self.lr_start
            self.last_lr = lr
            return lr
        else:
            t = (n - self.lr_warm_up_steps) / (self.lr_max_decay_steps - self.lr_warm_up_steps)
            t = min(t, 1.0)
            lr = self.lr_min + 0.5 * (self.lr_max - self.lr_min) * (
                    1 + np.cos(t * np.pi))
            self.last_lr = lr
            return lr

    def __call__(self, n):
        return self.schedule(n)

# class for quantization
class GumbelQuantizer(nn.Module):
    """
    see https://github.com/CompVis/taming-transformers/blob/04c8ad6c0fa4650e3d600faee793515c4aa2658c/taming/modules/vqvae/quantize.py
    ____________________________________________
    Discretization bottleneck part of the VQ-VAE.
    Inputs:
    - n_e : number of embeddings
    - e_dim : dimension of embedding
    - beta : commitment cost used in loss term, beta * ||z_e(x)-sg[e]||^2
    _____________________________________________
    """

    def __init__(self, num_hiddens, embedding_dim, n_embed, straight_through=True,
                 kl_weight=1e-9, temp_init=1.0, use_vqinterface=True,
                 remap=None, unknown_index="random"):
        super().__init__()
        self.embedding_dim = embedding_dim
        self.n_embed = n_embed
        self.n_e = n_embed

        self.straight_through = straight_through
        self.temperature = temp_init
        self.kl_weight = kl_weight

        self.proj = nn.Conv2d(num_hiddens, n_embed, 1)
        self.embed = nn.Embedding(n_embed, embedding_dim)

        self.use_vqinterface = use_vqinterface

        self.beta_scheduler = LambdaWarmUpCosineScheduler( warm_up_steps = 0,  lr_min = 3e-2, lr_max = 1e-9, lr_start = 0, max_decay_steps = 20000)
        self.temp_scheduler = LambdaWarmUpCosineScheduler( warm_up_steps = 0,  lr_min = 0.0625, lr_max = 1.0, lr_start = 0, max_decay_steps = 150000)

        self.remap = remap
        if self.remap is not None:
            self.register_buffer("used", torch.tensor(np.load(self.remap)))
            self.re_embed = self.used.shape[0]
            self.unknown_index = unknown_index # "random" or "extra" or integer
            if self.unknown_index == "extra":
                self.unknown_index = self.re_embed
                self.re_embed = self.re_embed+1
            print(f"Remapping {self.n_embed} indices to {self.re_embed} indices. "
                  f"Using {self.unknown_index} for unknown indices.")
        else:
            self.re_embed = n_embed

    def remap_to_used(self, inds):
        ishape = inds.shape
        assert len(ishape)>1
        inds = inds.reshape(ishape[0],-1)
        used = self.used.to(inds)
        match = (inds[:,:,None]==used[None,None,...]).long()
        new = match.argmax(-1)
        unknown = match.sum(2)<1
        if self.unknown_index == "random":
            new[unknown]=torch.randint(0,self.re_embed,size=new[unknown].shape).to(device=new.device)
        else:
            new[unknown] = self.unknown_index
        return new.reshape(ishape)

    def unmap_to_all(self, inds):
        ishape = inds.shape
        assert len(ishape)>1
        inds = inds.reshape(ishape[0],-1)
        used = self.used.to(inds)
        if self.re_embed > self.used.shape[0]: # extra token
            inds[inds>=self.used.shape[0]] = 0 # simply set to zero
        back=torch.gather(used[None,:][inds.shape[0]*[0],:], 1, inds)
        return back.reshape(ishape)


    def forward(self, z, temp=None, return_logits=False):
        # force hard = True when we are in eval mode, as we must quantize. actually, always true seems to work
        batch_size, _, height, width = z.shape
        hard = self.straight_through if self.training else True
        temp = self.temperature if temp is None else temp

        logits = self.proj(z)
        if self.remap is not None:
            # continue only with used logits
            full_zeros = torch.zeros_like(logits)
            logits = logits[:,self.used,...]

        soft_one_hot = F.gumbel_softmax(logits, tau=temp, dim=1, hard=hard)
        if self.remap is not None:
            # go back to all entries but unused set to zero
            full_zeros[:,self.used,...] = soft_one_hot
            soft_one_hot = full_zeros
        z_q = einsum('b n h w, n d -> b d h w', soft_one_hot, self.embed.weight)

        # + kl divergence to the prior loss
        qy = F.softmax(logits, dim=1)
        diff = self.kl_weight * torch.sum(qy * torch.log(qy * self.n_embed + 1e-10), dim=1).mean()

        ind = soft_one_hot.argmax(dim=1)

        # return z_q, loss, (perplexity, min_encodings, min_encoding_indices)
        # return z_q, loss, min_encoding_indices.view(batch_size, height, width)

        output = {
            'quantize': z_q,
            'quantize_loss': diff,
            'index': ind.view(batch_size, height, width)
        }

        return output

    def get_codebook_entry(self, indices, shape):
        #b, h, w, c = shape
        b, h, w = shape
        assert b*h*w == indices.shape[0]
        indices = rearrange(indices, '(b h w) -> b h w', b=b, h=h, w=w)
        if self.remap is not None:
            indices = self.unmap_to_all(indices)
        one_hot = F.one_hot(indices, num_classes=self.n_embed).permute(0, 3, 1, 2).float()
        z_q = einsum('b n h w, n d -> b d h w', one_hot, self.embed.weight)
        return z_q

    def update_temp_and_beta(self, global_step):
        self.temperature = self.temp_scheduler(global_step)
        self.kl_weight = self.beta_scheduler(global_step)
        

# blocks for encoder and decoder
def Normalize(in_channels):
    return torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)

def nonlinearity(x):
    # swish
    return x*torch.sigmoid(x)

class Upsample(nn.Module):
    def __init__(self, in_channels, with_conv, upsample_type='interpolate'):
        super().__init__()
        self.upsample_type = upsample_type
        self.with_conv = with_conv
        if self.upsample_type == 'conv':
            self.sample = nn.ConvTranspose2d(in_channels, in_channels, kernel_size=4, stride=2, padding=1),
        if self.with_conv:
            self.conv = torch.nn.Conv2d(in_channels,
                                        in_channels,
                                        kernel_size=3,
                                        stride=1,
                                        padding=1)

    def forward(self, x):
        if self.upsample_type == 'conv':
            x = self.sample(x)
        else:
            x = torch.nn.functional.interpolate(x, scale_factor=2.0, mode="nearest")
        if self.with_conv:
            x = self.conv(x)
        return x

class ResnetBlock(nn.Module):
    def __init__(self, *, in_channels, out_channels=None, conv_shortcut=False,
                 dropout=0.0, temb_channels=512):
        super().__init__()
        self.in_channels = in_channels
        out_channels = in_channels if out_channels is None else out_channels
        self.out_channels = out_channels
        self.use_conv_shortcut = conv_shortcut

        self.norm1 = Normalize(in_channels)
        self.conv1 = torch.nn.Conv2d(in_channels,
                                     out_channels,
                                     kernel_size=3,
                                     stride=1,
                                     padding=1)
        if temb_channels > 0:
            self.temb_proj = torch.nn.Linear(temb_channels,
                                             out_channels)
        self.norm2 = Normalize(out_channels)
        self.dropout = torch.nn.Dropout(dropout)
        self.conv2 = torch.nn.Conv2d(out_channels,
                                     out_channels,
                                     kernel_size=3,
                                     stride=1,
                                     padding=1)
        if self.in_channels != self.out_channels:
            if self.use_conv_shortcut:
                self.conv_shortcut = torch.nn.Conv2d(in_channels,
                                                     out_channels,
                                                     kernel_size=3,
                                                     stride=1,
                                                     padding=1)
            else:
                self.nin_shortcut = torch.nn.Conv2d(in_channels,
                                                    out_channels,
                                                    kernel_size=1,
                                                    stride=1,
                                                    padding=0)

    def forward(self, x, temb):
        h = x
        h = self.norm1(h)
        h = nonlinearity(h)
        h = self.conv1(h)

        if temb is not None:
            h = h + self.temb_proj(nonlinearity(temb))[:,:,None,None]

        h = self.norm2(h)
        h = nonlinearity(h)
        h = self.dropout(h)
        h = self.conv2(h)

        if self.in_channels != self.out_channels:
            if self.use_conv_shortcut:
                x = self.conv_shortcut(x)
            else:
                x = self.nin_shortcut(x)

        return x+h

class AttnBlock(nn.Module):
    def __init__(self, in_channels):
        super().__init__()
        self.in_channels = in_channels

        self.norm = Normalize(in_channels)
        self.q = torch.nn.Conv2d(in_channels,
                                 in_channels,
                                 kernel_size=1,
                                 stride=1,
                                 padding=0)
        self.k = torch.nn.Conv2d(in_channels,
                                 in_channels,
                                 kernel_size=1,
                                 stride=1,
                                 padding=0)
        self.v = torch.nn.Conv2d(in_channels,
                                 in_channels,
                                 kernel_size=1,
                                 stride=1,
                                 padding=0)
        self.proj_out = torch.nn.Conv2d(in_channels,
                                        in_channels,
                                        kernel_size=1,
                                        stride=1,
                                        padding=0)


    def forward(self, x):
        h_ = x
        h_ = self.norm(h_)
        q = self.q(h_)
        k = self.k(h_)
        v = self.v(h_)

        # compute attention
        b,c,h,w = q.shape
        q = q.reshape(b,c,h*w)
        q = q.permute(0,2,1)   # b,hw,c
        k = k.reshape(b,c,h*w) # b,c,hw
        w_ = torch.bmm(q,k)     # b,hw,hw    w[b,i,j]=sum_c q[b,i,c]k[b,c,j]
        w_ = w_ * (int(c)**(-0.5))
        w_ = torch.nn.functional.softmax(w_, dim=2)

        # attend to values
        v = v.reshape(b,c,h*w)
        w_ = w_.permute(0,2,1)   # b,hw,hw (first hw of k, second of q)
        h_ = torch.bmm(v,w_)     # b, c,hw (hw of q) h_[b,c,j] = sum_i v[b,c,i] w_[b,i,j]
        h_ = h_.reshape(b,c,h,w)

        h_ = self.proj_out(h_)

        return x+h_



# resblock only uses linear layer
class LinearResBlock(nn.Module):
    def __init__(self, in_channel, channel):
        super().__init__()

        self.layers = nn.Sequential(
            nn.ReLU(inplace=True),
            nn.Linear(in_channel, channel),
            nn.ReLU(inplace=True),
            nn.Linear(channel, in_channel),
        )

    def forward(self, x):
        out = self.layers(x)
        out = out + x

        return out

class PatchFullResEncoder(nn.Module):
    def __init__(self, *, in_ch, ch=128, z_channels, 
                ch_mult=[1, 2, 4], 
                n_res_block=2, 
                stride=8, 
                num_pos_emb=0,
                activate_output=True):
        super().__init__()
        in_dim = in_ch * stride * stride 
        self.stride = stride
        self.pre_layers = nn.Sequential(*[
            nn.Linear(in_dim, ch),
        ])


        res_layers = []
        for i in range(n_res_block):
            res_layers.append(LinearResBlock(ch, ch//2))
        #res_layers.append(nn.ReLU(inplace=True))
        self.res_layers = nn.Sequential(*res_layers)

        post_layer = [
            nn.ReLU(inplace=False),
            nn.Linear(ch, z_channels)
        ]
        if activate_output:
            post_layer.append(nn.ReLU(inplace=False))
        self.post_layer = nn.Sequential(*post_layer)

        if num_pos_emb != 0:
            self.pos_emb = nn.Embedding(num_pos_emb, ch)
        else:
            self.pos_emb = None


    def forward(self, x):
        """
        x: [B, 3, H, W]

        """
        in_size = [x.shape[-2], x.shape[-1]]
        out_size = [s//self.stride for s in in_size]

        x = torch.nn.functional.unfold(x, kernel_size=(self.stride, self.stride), stride=(self.stride, self.stride)) # B x 3*patch_size^2 x L
        x = x.permute(0, 2, 1).contiguous() # B x L x 3*patch_size

        x = self.pre_layers(x)
        # import pdb; pdb.set_trace()
        if self.pos_emb is not None:
            x = x + self.pos_emb.weight[0:x.shape[1]].unsqueeze(dim=0)

        x = self.res_layers(x)
        x = self.post_layer(x)

        x = x.permute(0, 2, 1).contiguous() # B x C x L
        # import pdb; pdb.set_trace()
        x = torch.nn.functional.fold(x, output_size=out_size, kernel_size=(1,1), stride=(1,1))

        return x


class PatchEncoder(nn.Module):
    def __init__(self, *, in_ch, ch=128, z_channels, 
                ch_mult=[1, 2, 4], 
                n_res_block=2, 
                stride=8, 
                num_pos_emb=0,
                activate_output=True):
        super().__init__()
        in_dim = in_ch * stride * stride 
        self.stride = stride
        self.pre_layers = nn.Sequential(*[
            nn.Linear(in_dim, ch//2),
            nn.ReLU(inplace=True),
            nn.Linear(ch//2, ch),
        ])

        mid_layers = []
        in_ch_ = ch
        for m in ch_mult:
            out_ch_ = ch * m
            mid_layers.append(nn.ReLU(inplace=True))
            mid_layers.append(nn.Linear(in_ch_, out_ch_))
            in_ch_ = out_ch_
        self.mid_layers = nn.Sequential(*mid_layers)

        res_layers = []
        for i in range(n_res_block):
            res_layers.append(LinearResBlock(out_ch_, out_ch_//4))
        #res_layers.append(nn.ReLU(inplace=True))
        self.res_layers = nn.Sequential(*res_layers)

        post_layer = [
            nn.ReLU(inplace=False),
            nn.Linear(out_ch_, z_channels)
        ]
        if activate_output:
            post_layer.append(nn.ReLU(inplace=False))
        self.post_layer = nn.Sequential(*post_layer)

        if num_pos_emb != 'none':
            self.pos_emb = nn.Embedding(num_pos_emb, ch)
        else:
            self.pos_emb = None


    def forward(self, x):
        """
        x: [B, 3, H, W]

        """
        in_size = [x.shape[-2], x.shape[-1]]
        out_size = [s//self.stride for s in in_size]

        x = torch.nn.functional.unfold(x, kernel_size=(self.stride, self.stride), stride=(self.stride, self.stride)) # B x 3*patch_size^2 x L
        x = x.permute(0, 2, 1).contiguous() # B x L x 3*patch_size

        x = self.pre_layers(x)
        # import pdb; pdb.set_trace()
        if self.pos_emb != 0 :
            x = x + self.pos_emb.weight[0:x.shape[1]].unsqueeze(dim=0)

        x = self.mid_layers(x)
        x = self.res_layers(x)
        x = self.post_layer(x)

        x = x.permute(0, 2, 1).contiguous() # B x C x L
        # import pdb; pdb.set_trace()
        x = torch.nn.functional.fold(x, output_size=out_size, kernel_size=(1,1), stride=(1,1))

        return x


class Decoder(nn.Module):
    def __init__(self, *, ch, out_ch, ch_mult=(1,2,4,8), scale_by_2=None, num_res_blocks,
                 attn_resolutions, dropout=0.0, resamp_with_conv=True,
                 resolution, z_channels, **ignorekwargs):
        super().__init__()
        
        if isinstance(resolution, int):
            resolution = [resolution, resolution] # H, W
        elif isinstance(resolution, (tuple, list)):
            resolution = list(resolution)
        else:
            raise ValueError('Unknown type of resolution:', resolution)
            
        attn_resolutions_ = []
        for ar in attn_resolutions:
            if isinstance(ar, (list, tuple)):
                attn_resolutions_.append(list(ar))
            else:
                attn_resolutions_.append([ar, ar])
        attn_resolutions = attn_resolutions_

        self.ch = ch
        self.temb_ch = 0
        self.num_resolutions = len(ch_mult)
        self.num_res_blocks = num_res_blocks
        self.resolution = resolution 
        self.requires_image = False

        # compute in_ch_mult, block_in and curr_res at lowest res
        in_ch_mult = (1,)+tuple(ch_mult)
        block_in = ch*ch_mult[self.num_resolutions-1]
        if scale_by_2 is None:
            curr_res = [r // 2**(self.num_resolutions-1) for r in self.resolution]
        else:
            scale_factor = sum([int(s) for s in scale_by_2])
            curr_res = [r // 2**scale_factor for r in self.resolution]

        self.z_shape = (1, z_channels, curr_res[0], curr_res[1])
        print("Working with z of shape {} = {} dimensions.".format(self.z_shape, np.prod(self.z_shape)))

        # z to block_in
        self.conv_in = torch.nn.Conv2d(z_channels,
                                       block_in,
                                       kernel_size=3,
                                       stride=1,
                                       padding=1)

        # middle
        self.mid = nn.Module()
        self.mid.block_1 = ResnetBlock(in_channels=block_in,
                                       out_channels=block_in,
                                       temb_channels=self.temb_ch,
                                       dropout=dropout)
        self.mid.attn_1 = AttnBlock(block_in)
        self.mid.block_2 = ResnetBlock(in_channels=block_in,
                                       out_channels=block_in,
                                       temb_channels=self.temb_ch,
                                       dropout=dropout)

        # upsampling
        self.up = nn.ModuleList()
        for i_level in reversed(range(self.num_resolutions)):
            block = nn.ModuleList()
            attn = nn.ModuleList()
            block_out = ch*ch_mult[i_level]
            for i_block in range(self.num_res_blocks+1):
                block.append(ResnetBlock(in_channels=block_in,
                                         out_channels=block_out,
                                         temb_channels=self.temb_ch,
                                         dropout=dropout))
                block_in = block_out
                if curr_res in attn_resolutions:
                    attn.append(AttnBlock(block_in))
            up = nn.Module()
            up.block = block
            up.attn = attn
            if scale_by_2 is None:
                if i_level != 0:
                    up.upsample = Upsample(block_in, resamp_with_conv)
                    curr_res = [r * 2 for r in curr_res]
            else:
                if scale_by_2[i_level]:
                    up.upsample = Upsample(block_in, resamp_with_conv)
                    curr_res = [r * 2 for r in curr_res]
            self.up.insert(0, up) # prepend to get consistent order

        # end
        self.norm_out = Normalize(block_in)
        self.conv_out = torch.nn.Conv2d(block_in,
                                        out_ch,
                                        kernel_size=3,
                                        stride=1,
                                        padding=1)

    def forward(self, z, **kwargs):
        #assert z.shape[1:] == self.z_shape[1:]
        self.last_z_shape = z.shape

        # timestep embedding
        temb = None

        # z to block_in
        h = self.conv_in(z)

        # middle
        h = self.mid.block_1(h, temb)
        h = self.mid.attn_1(h)
        h = self.mid.block_2(h, temb)

        # upsampling
        for i_level in reversed(range(self.num_resolutions)):
            for i_block in range(self.num_res_blocks+1):
                h = self.up[i_level].block[i_block](h, temb)
                if len(self.up[i_level].attn) > 0:
                    h = self.up[i_level].attn[i_block](h)
            # if i_level != 0:
            if getattr(self.up[i_level], 'upsample', None) is not None:
                h = self.up[i_level].upsample(h)

        h = self.norm_out(h)
        h = nonlinearity(h)
        h = self.conv_out(h)
        return h


class GumbelPatchVQGAN(BaseCodec):
    def __init__(self,
                 *,
                 encoder_config,
                 decoder_config,
                 lossconfig=None,
                 n_embed,
                 embed_dim,
                 ignore_keys=[],
                 data_info={'key': 'image'},
                 trainable=False,
                 ckpt_path=None,
                 token_shape=None
                 ):
        super().__init__()
        self.encoder = instantiate_from_config(encoder_config) # Encoder(**encoder_config)
        self.decoder = instantiate_from_config(decoder_config) # Decoder(**decoder_config)
        #self.quantize = VectorQuantizer(n_embed, embed_dim, beta=0.25)
        self.quantize = GumbelQuantizer(decoder_config['params']["z_channels"], embedding_dim = embed_dim, n_embed = n_embed, kl_weight = 1e-8, straight_through = False, temp_init = 1.0, remap=None)
        # import pdb; pdb.set_trace()
        self.quant_conv = torch.nn.Conv2d(encoder_config['params']["z_channels"], encoder_config['params']["z_channels"], 1)
        self.post_quant_conv = torch.nn.Conv2d(embed_dim, decoder_config['params']["z_channels"], 1)

        self.data_info = data_info
        self.cur_step = 1
    
        if lossconfig is not None and trainable:
            self.loss = instantiate_from_config(lossconfig)
        else:
            self.loss = None
        
        if ckpt_path is not None:
            self.init_from_ckpt(ckpt_path, ignore_keys=ignore_keys)
        
        self.trainable = trainable
        self._set_trainable()

        self.token_shape = token_shape

    def init_from_ckpt(self, path, ignore_keys=list()):
        sd = torch.load(path, map_location="cpu")
        if 'model' in sd:
            sd = sd['model']
        else:
            sd = sd["state_dict"]
        keys = list(sd.keys())
        for k in keys:
            for ik in ignore_keys:
                if k.startswith(ik):
                    print("VQGAN: Deleting key {} from state_dict.".format(k))
                    del sd[k]
        self.load_state_dict(sd, strict=False)
        print(f"VQGAN: Restored from {path}")

    @property
    def device(self):
        return self.quant_conv.weight.device

    def pre_process(self, data):
        data = data.to(self.device)
        data = data / 127.5 - 1.0
        return data

    def multi_pixels_with_mask(self, data, mask):
        if data.max() > 1:
            raise ValueError('The data need to be preprocessed!')
        mask = mask.to(self.device)
        data = data * mask
        data[~mask.repeat(1,3,1,1)] = -1.0
        return data

    def post_process(self, data):
        data = (data + 1.0) * 127.5
        data = torch.clamp(data, min=0.0, max=255.0)
        return data

    def get_number_of_tokens(self):
        return self.quantize.n_e    

    def get_tokens(self, data, mask=None, return_token_index=False, **kwargs):
        data = self.pre_process(data)
        x = self.encoder(data)
        x = self.quant_conv(x)
        idx = self.quantize(x)['index']
        if self.token_shape is None:
            self.token_shape = idx.shape[1:3]

        if self.decoder.requires_image:
            self.mask_im_tmp = self.multi_pixels_with_mask(data, mask)

        output = {}
        output['token'] = idx.view(idx.shape[0], -1)

        # import pdb; pdb.set_trace()
        if mask is not None: # mask should be B x 1 x H x W
            # downsampling
            # mask = F.interpolate(mask.float(), size=idx_mask.shape[-2:]).to(torch.bool)
            token_type = get_token_type(mask, self.token_shape) # B x 1 x H x W
            mask = token_type == 1
            output = {
                'target': idx.view(idx.shape[0], -1).clone(),
                'mask': mask.view(mask.shape[0], -1),
                'token': idx.view(idx.shape[0], -1),
                'token_type': token_type.view(token_type.shape[0], -1),
            }
        else:
            output = {
                'token': idx.view(idx.shape[0], -1)
                }

        # get token index
        # used for computing token frequency
        if return_token_index:
            token_index = output['token'] #.view(-1)
            output['token_index'] = token_index

        return output


    # def get_features(self, data, reshape_for_transformer=True):
    #     data = self.pre_process(data)
    #     x = self.encoder(data)
    #     x = self.quant_conv(x)
    #     quant = self.quantize(x)['quantize']
    #     # quant: B x C x H x W
    #     if reshape_for_transformer:
    #         quant = quant.view(*quant.shape[0:2], -1).permute(0, 2, 1) # B x HW x C
    #     output = {'feature': quant}
    #     raise NotImplementedError
    #     return output


    def decode(self, token):
        assert self.token_shape is not None
        # import pdb; pdb.set_trace()
        bhw = (token.shape[0], self.token_shape[0], self.token_shape[1])
        quant = self.quantize.get_codebook_entry(token.view(-1), shape=bhw)
        quant = self.post_quant_conv(quant)
        if self.decoder.requires_image:
            rec = self.decoder(quant, self.mask_im_tmp)
            self.mask_im_tmp = None
        else:
            rec = self.decoder(quant)
        rec = self.post_process(rec)
        return rec


    def get_rec_loss(self, input, rec):
        if input.max() > 1:
            input = self.pre_process(input)
        if rec.max() > 1:
            rec = self.pre_process(rec)

        rec_loss = F.mse_loss(rec, input)
        return rec_loss


    @torch.no_grad()
    def sample(self, batch):

        data = self.pre_process(batch[self.data_info['key']])
        x = self.encoder(data)
        x = self.quant_conv(x)
        quant = self.quantize(x)['quantize']
        quant = self.post_quant_conv(quant)
        if self.decoder.requires_image:
            mask_im = self.multi_pixels_with_mask(data, batch['mask'])
            rec = self.decoder(quant, mask_im)
        else:
            rec = self.decoder(quant)
        rec = self.post_process(rec)

        out = {'input': batch[self.data_info['key']], 'reconstruction': rec}
        if self.decoder.requires_image:
            out['mask_input'] = self.post_process(mask_im)
            out['mask'] = batch['mask'] * 255
            # import pdb; pdb.set_trace()
        return out

    def get_last_layer(self):
        if isinstance(self.decoder, Decoder):
            return self.decoder.conv_out.weight
        elif isinstance(self.decoder, PatchDecoder):
            return self.decoder.post_layers[-1].weight
    
    def parameters(self, recurse=True, name=None):
        if name is None or name == 'none':
            return super().parameters(recurse=recurse)
        else:
            if name == 'generator':
                params = list(self.encoder.parameters())+ \
                         list(self.decoder.parameters())+\
                         list(self.quantize.parameters())+\
                         list(self.quant_conv.parameters())+\
                         list(self.post_quant_conv.parameters())
            elif name == 'discriminator':
                params = self.loss.discriminator.parameters()
            else:
                raise ValueError("Unknown type of name {}".format(name))
            return params

    def forward(self, batch, name='none', return_loss=True, step=0, **kwargs):
        
        if name == 'generator':
            self.quantize.update_temp_and_beta(step)
            #print(self.quantize.kl_weight)
            #print(self.quantize.temperature)
            input = self.pre_process(batch[self.data_info['key']])
            x = self.encoder(input)
            x = self.quant_conv(x)
            quant_out = self.quantize(x)
            quant = quant_out['quantize']
            emb_loss = quant_out['quantize_loss']

            # recconstruction
            quant = self.post_quant_conv(quant)
            if self.decoder.requires_image:
                rec = self.decoder(quant, self.multi_pixels_with_mask(input, batch['mask']))
            else:
                rec = self.decoder(quant)
            # save some tensors for 
            self.input_tmp = input 
            self.rec_tmp = rec 

            if isinstance(self.loss, VQLPIPSWithDiscriminator):
                output = self.loss(codebook_loss=emb_loss,
                                inputs=input, 
                                reconstructions=rec, 
                                optimizer_name=name, 
                                global_step=step, 
                                last_layer=self.get_last_layer())
            else:
                raise NotImplementedError('{}'.format(type(self.loss)))

        elif name == 'discriminator':
            if isinstance(self.loss, VQLPIPSWithDiscriminator):
                output = self.loss(codebook_loss=None,
                                inputs=self.input_tmp, 
                                reconstructions=self.rec_tmp, 
                                optimizer_name=name, 
                                global_step=step, 
                                last_layer=self.get_last_layer())
            else:
                raise NotImplementedError('{}'.format(type(self.loss)))
        else:
            raise NotImplementedError('{}'.format(name))
        return output


if __name__ == '__main__':
    logits = torch.tensor([[0, 1, 2], [4,5,6]])
    mask = ~(logits > 2)

    print(mask)

