#Train a model from scratch in PyTorch and run evaluation
import tensorflow as tf
import sys
import io
import os
import subprocess
import math
import numpy as np
import pandas as pd
import glob
from tqdm import tqdm
from collections import OrderedDict
import random
import torch
import torch.nn as nn
from PIL import Image

#Train the model for an epoch
def train_model(trainset,trainlabels,model,optimizer,criterion,**kwargs):
    trainlen = trainset.shape[0]
    nbatches = math.ceil(trainlen/kwargs['batch_size'])
    total_loss = 0
    total_backs = 0
    with tqdm(total=nbatches,disable=(kwargs['verbose']<2)) as pbar:
        model = model.train()
        for b in range(nbatches):
            #Obtain batch
            X = trainset[b*kwargs['batch_size']:min(trainlen,(b+1)*kwargs['batch_size'])].clone().float().to(kwargs['device'])
            Y = trainlabels[b*kwargs['batch_size']:min(trainlen,(b+1)*kwargs['batch_size'])].clone().long().to(kwargs['device'])
            #Propagate
            posteriors = model(X)
            #Backpropagate
            loss = criterion(posteriors,Y)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            #Track loss
            if total_backs == 100:
                total_loss = total_loss*0.99+loss.detach().cpu().numpy()
            else:
                total_loss += loss.detach().cpu().numpy()
                total_backs += 1
            pbar.set_description(f'Training epoch. Loss {total_loss/(total_backs+1):.2f}')
            pbar.update()
    return total_loss/(total_backs+1)

class SimpleCNN(nn.Module):
    def __init__(self, **kwargs):
        super(SimpleCNN, self).__init__()
        #Arguments
        self.xsize = kwargs['xsize']
        self.ysize = kwargs['ysize']
        self.num_blocks = kwargs['num_blocks']
        self.channels = kwargs['channels']
        self.input_channels = (kwargs['input_channels'] if 'input_channels' in kwargs else 1)
        self.reduce_size = (kwargs['reduce_size'] if 'reduce_size' in kwargs else False)
        self.dropout = kwargs['dropout']
        self.embedding_size = kwargs['embedding_size']
        self.vocab = kwargs['vocab']
        self.num_classes = len(self.vocab)
        self.mean = kwargs['mean']
        self.std = kwargs['std']

        #Gaussian normalise the input
        self.inputnorm = InputNorm(self.mean,self.std)
        tmp_xsize = self.xsize
        tmp_ysize = self.ysize
        #Convolutional blocks
        for i in range(1,self.num_blocks+1):
            if self.reduce_size or i ==1:
                setattr(self,'convblock'+str(i),ConvBlock((self.channels if i>1 else self.input_channels), self.channels, kernel=(3, 3), stride=(2, 2), padding=(1, 1), groups=1, dropout=self.dropout, residual=False))
                tmp_xsize = int(tmp_xsize/2)
                tmp_ysize = int(tmp_ysize/2)
            setattr(self,'convblock'+str(i)+'residual',ConvBlock(self.channels, self.channels, kernel=(3, 3), stride=(1, 1), padding=(1, 1), groups=int(self.channels/2), dropout=self.dropout, residual=True))
            
        #Flatten the output
        self.flatten = nn.Flatten()
        #Reduce to embedding layer
        self.linear = nn.Linear(int(tmp_xsize*tmp_ysize*self.channels), self.embedding_size, bias=False)
        #Batch normalise
        self.batchnorm = nn.BatchNorm1d(self.embedding_size, momentum=0.9)
        #L2 normalise
        self.l2norm = L2Norm()
        #Classification layer and softmax
        self.output = nn.Linear(self.embedding_size, self.num_classes)
        self.softmax = nn.LogSoftmax(dim=1)

    def forward(self, x):
        out = self.inputnorm(x)
        for i in range(1,self.num_blocks+1):
            if self.reduce_size or i==1:
                conv = getattr(self,'convblock'+str(i))
                out = conv(out)
            conv = getattr(self,'convblock'+str(i)+'residual')
            out = conv(out)
        out = self.flatten(out)
        out = self.linear(out)
        out = self.batchnorm(out)
        out = self.l2norm(out)
        out = self.output(out)
        out = self.softmax(out)
        return out

#Performs gaussian normalisation of an input with mean and standard deviation
class InputNorm(nn.Module):
    def __init__(self, mean, std):
        super(InputNorm, self).__init__()
        self.mean = mean
        self.std = std
    def forward(self,x):
        out = torch.mul(torch.add(x,-self.mean),1/self.std)
        return out

#Residual convolutional block
class ConvBlock(nn.Module):
    def __init__(self, in_c, out_c, kernel=(1, 1), stride=(1, 1), padding=(0, 0), groups=1, dropout=0.2, residual=False):
        super(ConvBlock, self).__init__()
        #2D convolution
        self.conv = nn.Conv2d(in_c, out_channels=out_c, kernel_size=kernel, groups=groups, stride=stride, padding=padding, bias=False)
        #Batch normalisation
        self.bn = nn.BatchNorm2d(out_c, momentum=0.9)
        #Activation
        self.prelu = nn.PReLU(out_c)
        #Dropout
        self.dropout = nn.Dropout3d(p=dropout)
        self.residual = residual
    def forward(self, x):
        out = self.conv(x)
        out = self.bn(out)
        out = self.prelu(out)
        out = self.dropout(out)
        #Residual connection
        if self.residual:
            out = out + x
        return out

#Do L2 normalisation of embedding vectors
class L2Norm(nn.Module):
    def __init__(self, axis=1):
        super(L2Norm, self).__init__()
        self.axis = axis
    def forward(self,x):
        norm = torch.norm(x, 2, self.axis, True)
        output = torch.div(x, norm)
        return output

#Initialise all random numbers for reproducibility
def init_random(**kwargs):
    random.seed(kwargs['seed'])
    torch.manual_seed(kwargs['seed'])
    torch.cuda.manual_seed(kwargs['seed'])
    torch.backends.cudnn.deterministic = True
    
#Resize the set to the desired size
def resize_image(img,**kwargs):
    img = Image.fromarray(img,mode='RGB')
    img = img.resize((kwargs['ysize'],kwargs['xsize']))
    img = np.array(img)
    return img

#Augment by doing rotations and flips (7 per image plus original)
def augment_set(dataset,labels):
    outset = np.zeros((dataset.shape[0]*8,dataset.shape[1],dataset.shape[2],dataset.shape[3]),dtype=np.uint8)
    outset[0:dataset.shape[0]] = dataset
    for j in range(dataset.shape[0]):
        img = Image.fromarray(dataset[j],mode='RGB')
        mimg = img.transpose(Image.FLIP_LEFT_RIGHT)
        outset[dataset.shape[0]+j] = np.array(mimg)
        mimg = img.transpose(Image.ROTATE_90)
        outset[2*dataset.shape[0]+j] = np.array(mimg)
        mimg = img.transpose(Image.ROTATE_90).transpose(Image.FLIP_TOP_BOTTOM)
        outset[3*dataset.shape[0]+j] = np.array(mimg)
        mimg = img.transpose(Image.ROTATE_180)
        outset[4*dataset.shape[0]+j] = np.array(mimg)
        mimg = img.transpose(Image.ROTATE_180).transpose(Image.FLIP_LEFT_RIGHT)
        outset[5*dataset.shape[0]+j] = np.array(mimg)
        mimg = img.transpose(Image.ROTATE_270)
        outset[6*dataset.shape[0]+j] = np.array(mimg)
        mimg = img.transpose(Image.ROTATE_270).transpose(Image.FLIP_TOP_BOTTOM)
        outset[7*dataset.shape[0]+j] = np.array(mimg)
    outlabs = np.tile(labels,8)
    return outset,outlabs

#Read the images and preprocess
def read_tfrecords(filepaths,train=False,**kwargs):
    images = list()
    classes = list()
    ids = list()
    for path in filepaths:
        for record in tf.compat.v1.io.tf_record_iterator(path):
            example = tf.train.Example()
            example.ParseFromString(record)

            img = example.features.feature['image'].bytes_list.value[0]
            img = Image.open(io.BytesIO(img))
            img = np.asarray(img)
            if img.shape[1]!=kwargs['ysize'] or img.shape[2]!=kwargs['xsize']:
                img = resize_image(img,**kwargs)
            images.append(img)
            if 'target' in example.features.feature:
                label = example.features.feature['target'].int64_list.value[0]
                classes.append(label)
            iid = example.features.feature['image_name'].bytes_list.value[0].decode('ascii')
            ids.append(iid)
    images = np.array(images)
    classes = np.array(classes)
    images, classes = augment_set(images,classes)
    if train:
        idx = [i for i in range(images.shape[0])]
        random.shuffle(idx)
        images = images[idx]
        classes = classes[idx]
    images = torch.from_numpy(np.transpose(images,(0,3,1,2)))
    classes = torch.from_numpy(classes)
    return images, classes, ids

#Get posteriors for a test set
def evaluate_model(testset,model,**kwargs):
    testlen = testset.shape[0]
    predictions = np.zeros((testlen,len(kwargs['vocab'])))
    nbatches = math.ceil(testlen/kwargs['batch_size'])
    with torch.no_grad():
        model = model.eval()
        with tqdm(total=nbatches,disable=(kwargs['verbose']<2)) as pbar:
            for b in range(nbatches):
                #Obtain batch
                X = testset[b*kwargs['batch_size']:min(testlen,(b+1)*kwargs['batch_size'])].clone().float().to(kwargs['device'])
                #Propagate
                posteriors = model(X)
                predictions[b*kwargs['batch_size']:min(testlen,(b+1)*kwargs['batch_size']),:] = posteriors.detach().cpu().numpy()
                pbar.set_description('Testing')
                pbar.update()
    return predictions

#Arguments
args = {
    'cv_percentage': 0.1,
    'xsize': 64,
    'ysize': 64,
    'num_blocks': 4,
    'channels': 48,
    'input_channels': 3,
    'dropout': 0.0,
    'embedding_size': 256,
    'epochs': 5,
    'batch_size': 256,
    'learning_rate': 0.001,
    'seed': 0,
    'device': ('cuda:0' if torch.cuda.is_available() else 'cpu'),
    'verbose': 1,
}

#Initialise RNGs
init_random(**args)

print('Loading data...')
train_files = glob.glob("../input/cassava-leaf-disease-classification/train_tfrecords/*.tfrec")
val_files = np.sort(train_files[math.ceil(len(train_files)*(1-args['cv_percentage'])):])
train_files = np.sort(train_files[:math.ceil(len(train_files)*(1-args['cv_percentage']))])
test_files = glob.glob("../input/cassava-leaf-disease-classification/test_tfrecords/*.tfrec")

train_data, train_targets, _ = read_tfrecords(train_files,True,**args)
val_data, val_targets, _ = read_tfrecords(val_files,False,**args)
test_data, _, test_ids = read_tfrecords(test_files,False,**args)

#Mapping of outputs
args['vocab'] = OrderedDict({t:i for i,t in enumerate(np.unique(list(train_targets)))})

#Mean and standard deviation for input normalisation
args['mean'] = torch.mean(train_data.float())
args['std'] = torch.std(train_data.float())

#Build model, optimizer and weighted criterion
model = SimpleCNN(**args).to(args['device'])
optimizer = torch.optim.Adam(model.parameters(),lr=args['learning_rate'])
priors = torch.Tensor([len(np.where(train_targets.numpy()==t)[0])/train_targets.shape[0] for t in np.unique(train_targets)])
criterion = nn.NLLLoss(weight = 1 / (priors / torch.max(priors)),reduction='mean').to(args['device'])

print('Training...')
targets = val_targets[0:int(val_targets.shape[0]/8)].numpy()
best_acc = 0.0
for ep in range(1,args['epochs']+1):
    #Train an epoch
    loss = train_model(train_data,train_targets,model,optimizer,criterion,**args)
    #Get the posteriors for the validation set
    val_preds = evaluate_model(val_data,model,**args)
    #Compute accuracy and F1
    val_preds = np.mean([val_preds[i*int(val_preds.shape[0]/8):(i+1)*int(val_preds.shape[0]/8)] for i in range(8)],axis=0)
    acc = 100*len(np.where((np.argmax(val_preds,axis=1)-targets)==0)[0])/targets.shape[0]
    print('Epoch {0:d}, training loss: {1:.2f}, validation accuracy: {2:.2f}%'.format(ep,loss,acc))
    if acc >= best_acc:
        #Get the posteriors for the test set
        test_preds = evaluate_model(test_data,model,**args)
        
#Combine the posteriors for each of the 8 augmentations
predictions = np.argmax(np.mean([test_preds[i*int(test_preds.shape[0]/8):(i+1)*int(test_preds.shape[0]/8)] for i in range(8)],axis=0),axis=1)

#Write output
df_out = pd.DataFrame({'image_id': test_ids, 'label': predictions.astype(int)})
df_out.to_csv('/kaggle/working/submission.csv',index=False)