import tensorflow as tf
import numpy as np


class MorphoNet:
  def __init__(self, low_memory=False, kind='mn170', weights_folder=None):
    self.low_memory = low_memory
    self.kind = kind
    self.weights_folder = weights_folder
    self.members_sequence = ['inception_v4', 'efficientnetb2', 'densenet169', 'densenet201']
    self.members = []
    self.ensemble = None
    self.meta_standardize = {
      'mean': {
        'mn170': np.array([0.63030227885852785352227556359139, 0.68862444766805175344615008725668, 
                          0.66542678031938407023915260651847, 0.62889294177184085210541297783493, 
                          0.36969772078785573254933183307003, 0.31137555309005104930974994204007, 
                          0.33457321970229814134256685065338, 0.37110705833947049692156383571273]),
        'mn175': np.array([0.70050570887301311095995970390504, 0.75598683535962729607149412913714, 
                          0.73423295518560482975090053514577, 0.70544850342324827430218192603206, 
                          0.29949429097061641691368549800245, 0.24401316510100320988918554121483, 
                          0.26576704506784609716163458870142, 0.29455149681745618206463177557453])
      },
      'std': {
        'mn170': np.array([0.44140710557436596550573426611663, 0.42648277771076659181659351816052, 
                          0.43942331274475704416815347030933, 0.45133516923923167052379312735866, 
                          0.44140710541812488987289953001891, 0.42648277750523055917852843776927, 
                          0.43942331222158304004921092200675, 0.45133516890355390716038641585328]),
        'mn175': np.array([0.41159684724160133795223259767226, 0.37865357383476050401327483996283,
                          0.40167239366955120871693907247391, 0.42007152146889570332177754607983,
                          0.41159684715779715213912481885927, 0.37865357311726360878267882981163,
                          0.40167239323576364729007082132739, 0.42007152134568670476255647372454])
      }
    }


  def set_weights_folder(self, folder):
    self.weights_folder = folder


  def set_kind(self, kind):
    self.kind = kind


  def load(self):
    if self.weights_folder is None:
      raise ValueError('weights folder not defined.')
  
    if self.low_memory:
      raise NotImplementedError('low_memory feature not implemented, all members will be loaded in memory.')
    
    self.members = [] # TODO: check if members is loaded
    for member in self.members_sequence:
      model = tf.keras.models.load_model(f'{self.weights_folder}/{member}')
      self.members.append(model)
    
    self.ensemble = tf.keras.models.load_model(f'{self.weights_folder}/ensemble')

  
  def load_meta(self, path=None):
    if path is None:
      self.ensemble = tf.keras.models.load_model(f'{self.weights_folder}/ensemble')
    else:
      self.ensemble = tf.keras.models.load_model(path)


  def predict(self, X):
    if isinstance(X, list):
      X = np.array(X)

    if len(self.members) == 0:
      raise ValueError('Load models first calling load() method.')
    
    meta_input = [member.predict(X) for member in self.members]
    meta_input = np.dstack(tuple(meta_input))
    meta_input = np.reshape(meta_input, (meta_input.shape[0], meta_input.shape[1] * meta_input.shape[2]))
    
    meta_input -= self.meta_standardize['mean'][self.kind]
    meta_input /= self.meta_standardize['std'][self.kind]

    y_hat = self.ensemble.predict(meta_input)

    return y_hat