from sbi import analysis as analysis

# sbi
from sbi import utils as utils
from sbi.inference import NPE, simulate_for_sbi
from sbi.utils.user_input_checks import (
    check_sbi_inputs,
    process_prior,
    process_simulator,
)

import torch
import torch.nn as nn
import torch.nn.functional as F
from sbi.inference import SNPE, SNLE
from sbi.neural_nets import posterior_nn
from sbi import analysis as analysis
from sbi.utils import MultipleIndependent


from dustbi_simulator import *
from Functions import *

infos = load_kestrel("posteriors/sims.v5.NOM.yml.bk")
dicts = [infos['Functions'], infos['Splits'], infos['Priors'], infos['Correlations']]

simfilename = infos['Simbank_File'][0]
datfilename = infos['Data_File'][0]

df, dfdata = load_data(simfilename, datfilename)

param_names = infos['param_names']

params_to_fit = parameter_generation(param_names, dicts)
priors = prior_generator(param_names, dicts)

layout = build_layout(params_to_fit, dicts)
parameters_to_condition_on = infos['parameters_to_condition_on']

def add_distance(df_tensor):
    
    x1_obs = df_tensor['x1'] ; c_obs = df_tensor['c'] ; mB_obs = df_tensor['mB']
    
    beta = 3.1 ; alpha = 0.16 ; M0 = -19.3
    
    correction = alpha * x1_obs - beta * c_obs + M0 + mB_obs
        
    MURES = df_tensor['MU'] - correction
    
    return  MURES

output_distribution = preprocess_input_distribution(
    df, parameters_to_condition_on[:-1]+['x0', 'x0ERR', 'MU'])

MURES_sims = add_distance(output_distribution)

df['MURES'] = MURES_sims

output_distribution = preprocess_input_distribution(
    dfdata, parameters_to_condition_on[:-1]+['x0', 'x0ERR', 'MU'])

MURES_sims = add_distance(output_distribution)

dfdata['MURES'] = MURES_sims

sim_for_training = make_batched_simulator(layout, df,
                        param_names,parameters_to_condition_on,
                        dicts, dfdata, sub_batch=20, device='cpu', )
batched = True

import pickle
import io

#https://stackoverflow.com/questions/57081727/load-pickle-file-obtained-from-gpu-to-cpu
class CPU_Unpickler(pickle.Unpickler):
    def find_class(self, module, name):
        if module == 'torch.storage' and name == '_load_from_bytes':
            return lambda b: torch.load(io.BytesIO(b), map_location='cpu')
        else:
            return super().find_class(module, name)

with open("posteriors/posterior.v5.NOM.pt", "rb") as f:
    posterior = CPU_Unpickler(f).load()

x = preprocess_data(parameters_to_condition_on, dfdata)

true_params = priors.sample()

new_x = sim_for_training(true_params)

posterior_samples = posterior.sample((50000,), x=new_x)

