import os
import time
import argparse
import sys
import datetime

from modelutils import *
from model_library import DVSCIFAR10NET_DOWNSIZED

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.data as data
from torch import amp
from torch.utils.tensorboard import SummaryWriter
import numpy as np

from spikingjelly.activation_based import neuron, functional, surrogate, layer
from spikingjelly.datasets import cifar10_dvs, split_to_train_test_set, transforms

class DVStransform:
    def __init__(self, transform):
        self.transform = transform

    def __call__(self, img):
        img = torch.from_numpy(img).float()
        shape = [img.shape[0], img.shape[1]]
        img = img.flatten(0, 1)
        img = self.transform(img)
        shape.extend(img.shape[1:])
        return img.view(shape)

def spikeRateLoss(true_rate, false_rate, model_out, label):
    modified_label = torch.zeros_like(label).to(DEV)
    modified_label[label < 0.5] = false_rate
    modified_label[label > 0.5] = true_rate
    return F.mse_loss(model_out, modified_label)

def main():
    
    parser = argparse.ArgumentParser(description='CIFAR10-DVS Training')
    parser.add_argument('-T', default=100, type=int, help='simulating time-steps')
    parser.add_argument('-device', default='cuda:0', help='device')
    parser.add_argument('-b', default=64, type=int, help='batch size')
    parser.add_argument('-epochs', default=100, type=int, metavar='N',
                        help='number of total epochs to run')
    parser.add_argument('-j', default=4, type=int, metavar='N',
                        help='number of data loading workers (default: 4)')
    parser.add_argument('-data-dir', type=str, help='root dir of MNIST dataset')
    parser.add_argument('-out-dir', type=str, default='./logs', help='root dir for saving logs and checkpoint')
    parser.add_argument('-resume', type=str, help='resume from the checkpoint path')
    parser.add_argument('-amp', action='store_true', help='automatic mixed precision training')
    parser.add_argument('-opt', type=str, choices=['sgd', 'adam'], default='adam', help='use which optimizer. SGD or Adam')
    parser.add_argument('-momentum', default=0.9, type=float, help='momentum for SGD')
    parser.add_argument('-lr', default=1e-3, type=float, help='learning rate')
    parser.add_argument('-tau', default=2.0, type=float, help='parameter tau of LIF neuron')

    args = parser.parse_args()
    print(args)

    net = DVSCIFAR10NET_DOWNSIZED(channels=128, tau=args.tau)
    net.to(args.device)
    functional.set_step_mode(net, 's')
    print(net)
    
    T_MAX = 64

    # 初始化数据加载器
    data_path = "./datas/CIFAR10DVS"
    data_transform = DVStransform(transforms.Resize(size=(64, 64), antialias=True))
    full_dataset = cifar10_dvs.CIFAR10DVS(
        root=data_path,
        data_type='frame',
        frames_number=args.T,
        split_by='number',
        transform=data_transform,
        duration=10000,
    )
        
    np.random.seed(TORCH_SEED)
    train_dataset, test_dataset = split_to_train_test_set(
        train_ratio=0.9,
        origin_dataset=full_dataset,
        num_classes=10,
        random_split=False,      # no shuffle within each class before splitting
    )
    
    train_data_loader = torch.utils.data.DataLoader(train_dataset, batch_size=args.b, shuffle=True,
                                                num_workers=args.j, pin_memory=True, sampler=None, 
                                                persistent_workers=True, prefetch_factor=2)
    test_data_loader = torch.utils.data.DataLoader(test_dataset, batch_size=args.b, shuffle=True,
                                                num_workers=args.j, pin_memory=True, 
                                                persistent_workers=True, prefetch_factor=2)

    
    scaler = None
    if args.amp:
        scaler = amp.GradScaler()

    start_epoch = 0
    max_test_acc = -1


    optimizer = None
    if args.opt == 'sgd':
        optimizer = torch.optim.SGD(net.parameters(), lr=args.lr, momentum=args.momentum)
    elif args.opt == 'adam':
        optimizer = torch.optim.Adam(net.parameters(), lr=args.lr, weight_decay=1e-5)
    else:
        raise NotImplementedError(args.opt)

    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=T_MAX)

    if args.resume:
        checkpoint = torch.load(args.resume, map_location='cpu')
        net.load_state_dict(checkpoint['net'])
        optimizer.load_state_dict(checkpoint['optimizer'])
        scheduler.load_state_dict(checkpoint['scheduler'])
        start_epoch = checkpoint['epoch'] + 1
        max_test_acc = checkpoint['max_test_acc']
    
    out_dir = os.path.join(args.out_dir, f'{type(net).__name__}_T{args.T}_b{args.b}_{args.opt}_lr{args.lr}')

    if args.amp:
        out_dir += '_amp'
    
    out_dir += '_dvs'

    if not os.path.exists(out_dir):
        os.makedirs(out_dir)
        print(f'Mkdir {out_dir}.')

    with open(os.path.join(out_dir, 'args.txt'), 'w', encoding='utf-8') as args_txt:
        args_txt.write(str(args))

    writer = SummaryWriter(out_dir, purge_step=start_epoch)
    with open(os.path.join(out_dir, 'args.txt'), 'w', encoding='utf-8') as args_txt:
        args_txt.write(str(args))
        args_txt.write('\n')
        args_txt.write(' '.join(sys.argv))

    for epoch in range(start_epoch, args.epochs):
        start_time = time.time()
        net.train()
        train_loss = 0
        train_acc = 0
        train_samples = 0
        idx = 0
        time1 = 0
        time2 = 0
        time3 = 0
        time4 = 0
        for frame, label in train_data_loader:
            print("\r" + str(idx) + "/" + str(len(train_data_loader)) + " " + str(frame.shape), time1, time2, time3, time4, end='')
            idx += 1
            optimizer.zero_grad()
            # for param in net.parameters():
            #     param.grad = None
            frame = frame.to(args.device, non_blocking=True)
            frame = frame.transpose(0, 1)  # [N, T, C, H, W] -> [T, N, C, H, W]
            label = label.to(args.device, non_blocking=True)
            label_onehot = F.one_hot(label, 10).float()

            time1 = time.time_ns()
            with amp.autocast(device_type=DEV, dtype=torch.bfloat16):
                out_fr = 0.
                for t in range(args.T):
                    out_fr += net(frame[t])
                out_fr /= args.T
                loss = spikeRateLoss(true_rate=0.82, false_rate=0.02, model_out = out_fr, label = label_onehot)
            time1 = time.time_ns() - time1
            
            time2 = time.time_ns()
            scaler.scale(loss).backward()
            time2 = time.time_ns() - time2

            time3 = time.time_ns()
            scaler.step(optimizer)
            time3 = time.time_ns() - time3

            time4 = time.time_ns()
            scaler.update()

            train_samples += label.numel()
            train_loss += loss.item() * label.numel()
            train_acc += (out_fr.argmax(1) == label).float().sum().item()

            functional.reset_net(net)
            time4 = time.time_ns() - time4
        
        print()
        scheduler.step()

        train_time = time.time()
        train_speed = train_samples / (train_time - start_time)
        train_loss /= train_samples
        train_acc /= train_samples

        writer.add_scalar('train_loss', train_loss, epoch)
        writer.add_scalar('train_acc', train_acc, epoch)

        net.eval()
        test_loss = 0
        test_acc = 0
        test_samples = 0
        with torch.no_grad():
            for frame, label in test_data_loader:
                frame = frame.to(args.device)
                frame = frame.transpose(0, 1)  # [N, T, C, H, W] -> [T, N, C, H, W]
                label = label.to(args.device)
                label_onehot = F.one_hot(label, 10).float()
                out_fr = 0.
                for t in range(args.T):
                    out_fr += net(frame[t])
                out_fr /= args.T
                loss = spikeRateLoss(true_rate=0.82, false_rate=0.02, model_out = out_fr, label = label_onehot)
                
                test_samples += label.numel()
                test_loss += loss.item() * label.numel()
                test_acc += (out_fr.argmax(1) == label).float().sum().item()
                functional.reset_net(net)
        test_time = time.time()
        test_speed = test_samples / (test_time - train_time)
        test_loss /= test_samples
        test_acc /= test_samples
        writer.add_scalar('test_loss', test_loss, epoch)
        writer.add_scalar('test_acc', test_acc, epoch)

        save_max = False
        if test_acc > max_test_acc:
            max_test_acc = test_acc
            save_max = True

        checkpoint = {
            'net': net.state_dict(),
            'optimizer': optimizer.state_dict(),
            'scheduler': scheduler.state_dict(),
            'epoch': epoch,
            'max_test_acc': max_test_acc
        }

        if save_max:
            torch.save(checkpoint, os.path.join(out_dir, 'checkpoint_max.pth'))

        torch.save(checkpoint, os.path.join(out_dir, 'checkpoint_latest.pth'))

        #print_neuron_parameters(net.conv_fc)

        print(args)
        print(out_dir)
        print(f'epoch ={epoch}, train_loss ={train_loss: .4f}, train_acc ={train_acc: .4f}, test_loss ={test_loss: .4f}, test_acc ={test_acc: .4f}, max_test_acc ={max_test_acc: .4f}')
        print(f'train speed ={train_speed: .4f} images/s, test speed ={test_speed: .4f} images/s')
        print(f'escape time = {(datetime.datetime.now() + datetime.timedelta(seconds=(time.time() - start_time) * (args.epochs - epoch))).strftime("%Y-%m-%d %H:%M:%S")}\n')


if __name__ == '__main__':
    main()