#include <Syslog.h>
#include <AttributeParser.h>

#include "HMM.hh"
#include "ConfigTable.hh"
#include "System.h"

#define OBS_PFX "cond"

/*********************************************************************************
 **
 ** class HMMState
 **
 **   The only really tricky thing here is that the vector of
 **   conditional probabilities is sparse, so we keep track of both
 **   the index and the probability
 **
 *********************************************************************************/
HMMState::HMMState(char *name,prob_t initprob)
  :statename(strdup(name))
{
  ssdbg(DBG_LOAD)("creating HMMState \"%s\" class",name);
  reset(initprob);
}

HMMState::~HMMState(){
  ssdbg(DBG_CLEAN)("destroying HMMState \"%s\" class",statename);
  free((void *)statename);
}

prob_t HMMState::prev(int i) {return(prevprob[i]);}
prob_t HMMState::maxprev(){
  int ii;
  prob_t mv;
  mv = 0;
  for(ii=1;ii<HMM_KEEP_STATES;ii++){
    if(mv < prevprob[ii]){
      mv = prevprob[ii];
    }
  }
  return(mv);
}
prob_t HMMState::cur() {return(curprob);}
char *HMMState::name() {return(statename);}
prob_t HMMState::trans(state_t tostate){
  if(tostate >= (state_t)transprob.size())
    return(0);
  return(transprob[tostate]);
}
prob_t HMMState::cond(obs_t obs){
  cpi_t i = findcp(obs);
  if(i == condprob.end()){
    return(0);
  }
  return(i->prob);
}

// Append zeros until the trans vector includes the given state, then
// change that transition to the given probability.
void HMMState::settrans(state_t tostate,prob_t tp){
  int i,addnum;
  addnum = tostate+1-transprob.size();
  if(addnum > 0){
    for(int i=0;i<addnum;i++){
      transprob.push_back((prob_t)0);
    }
  }
  transprob[tostate] = tp;
}
void HMMState::setcond(obs_t cob,prob_t cpr){
  cp_t cp = {cob,cpr};
  if(findcp(cob) == condprob.end())
    condprob.push_back(cp);
  else
    Syslog::write("HMM -- Warning! Ignoring duplicate initialization of conditional probability.");
}
bool HMMState::newcur(prob_t newprob) {
  int ii;
  for(ii=HMM_KEEP_STATES-1;ii>0;ii--){
    prevprob[ii] = prevprob[ii-1];
  }
  prevprob[1] = curprob; 
  curprob = newprob;
  return(prevprob[1] != newprob);
}
void HMMState::reset(prob_t pr) {
  int ii;
  curprob = pr;
  for(ii=1;ii<HMM_KEEP_STATES;ii++){
    prevprob[ii] = pr;
  }
};

HMMState::cpi_t HMMState::findcp(obs_t ob){
  for(cpi_t i = condprob.begin();i != condprob.end();++i){
    if(i->obs == ob){
      break;
    }
  }
  return(i);
}


/*********************************************************************************
 **
 ** class HMM
 **
 *********************************************************************************/

// structors :

HMM::HMM(char *name)
  :m_name(strdup(name)),m_curobs(NULL),m_prevobs(NULL),num_nostate_errs(0),
   change_not_logged(true){
  char logname[128];
  ssdbg(DBG_LOAD)("creating HMM \"%s\" class",m_name);
  sprintf(logname,"%s_%s",HMMLogBase,m_name);
  m_logger = new HMMLogger(logname,this);
}

HMM::~HMM() {
  HMMState *tmp;
  ssdbg(DBG_CLEAN)("destroying HMM \"%s\" class",m_name);
  free((void *)m_name);
  while(m_statelist.size()){
    tmp = m_statelist.back();
    m_statelist.pop_back();
    delete(tmp);
  }
  delete(m_logger);
}

prob_t HMM::stateprev(state_t st,int i){
  return(m_statelist[st]->prev(i));
}
prob_t HMM::maxprev(state_t st){
  return(m_statelist[st]->maxprev());
}
prob_t HMM::statecur(state_t st){
  return(m_statelist[st]->cur());
}

obs_t HMM::curobs(){
  return(m_curobs);
}

bool HMM::updatelog(){
  bool rval = change_not_logged;
  change_not_logged = false;
  return(rval);
}

void HMM::reset(){
  int i;
  for(i=0;i<m_statelist.size();++i){
    // Set cur state and prev state both back to initial values
    m_statelist[i]->reset(i==0);
  }
}

void HMM::update(obs_t obs){
  prob_t tp,pp,cp,psum,pcur;
  std::vector<prob_t> newprobs;
  state_t i,j;
  std::vector<HMMState*>::reverse_iterator ri;
  char str[2][256];
  bool stn = 0;

  change_not_logged |= (m_prevobs != m_curobs);
  m_prevobs = m_curobs;
  m_curobs = obs;

  //  Syslog::write("HMM -- Updating with observation %d",obs);
  // Next we calculate the marginal probabilities.
  psum = 0;
  for(i=0;i<m_statelist.size();++i){
    stn = 0;
    sprintf(str[stn],"    -- %s state: (",m_statelist[i]->name());
    pcur = 0;
    // Sum up p(s_i|s_j)*p_{t-1}(s_j)
    for(j=0;j<m_statelist.size();++j){
	tp = m_statelist[j]->trans(i);
	pp = m_statelist[j]->cur();
	pcur+= tp*pp;
	stn = !stn;
	sprintf(str[stn],"%s%f*%f",str[!stn],tp,pp);
	if(j < m_statelist.size()-1){
	  stn = !stn;
	  sprintf(str[stn],"%s + ",str[!stn]);
	}
    }
    // Multiply by p(s_i|C)
    cp = m_statelist[i]->cond(obs);
    pcur *= cp;
    psum += pcur;
    //    Syslog::write("%s)*%f = %f",str[stn],cp,pcur);
    newprobs.push_back(pcur);
  }
  if(psum == 0){
    if(num_nostate_errs++ < HMM_MAX_NOSTATE_ERR){
      Syslog::write("SmartSampler -- Cluster Warning! State probability vector zero!  Starting over from initial state!");
      if(num_nostate_errs == HMM_MAX_NOSTATE_ERR){
	Syslog::write("SmartSampler -- Suppressing future display of this warning message.");
      }
    }
    reset();
    // always log these
    change_not_logged = true;
  }else{
    //    Syslog::write("    -- Normalizing with sum %f",psum);
    for(ri=m_statelist.rbegin();ri!=m_statelist.rend();++ri){
      //      Syslog::write("    -- %s state: new prob = %f",(*ri)->name(),newprobs.back()/psum);
      change_not_logged |= (*ri)->newcur(newprobs.back()/psum);
      newprobs.pop_back();
    }
  }

  m_logger->callWrite();
}

void HMM::dumpstate(){
  char str[2][128],*stnm;
  bool stn = 0;
  std::vector<HMMState*>::iterator i;
  sprintf(str[stn],"SmartSampler -- State = ");
  for(i=m_statelist.begin();i!=m_statelist.end();++i){
    stnm = (*i)->name();
    if(strlen(str[stn]) > 128 - strlen(stnm) - 10){
      break;
    }
    stn=!stn;
    sprintf(str[stn],"%s, %s(%.04f)",str[!stn],stnm,(*i)->cur());
  }
  Syslog::write(str[stn]);
}

void HMM::loadConfig(char *transcfgname,char *condcfgname,std::vector<string> *statenames,int maxobs){
  int ii,jj;
  int numst = statenames->size();
  prob_t val;
  ConfigTableRow<double>* curRow;
  ConfigTable<double> cfgtrans("transition",statenames,0);
  ConfigTable<double> cfgcond("conditional",OBS_PFX,maxobs+1,0);
  string row;
  const char *cfgFileName;

  cfgFileName = System::configurationFile( transcfgname );
  System::copyToLogDir(cfgFileName);

#ifdef DEBUG_LOAD
  Syslog::write("HMM -- Reading transition table \"%s\".",cfgFileName);
#endif

  // pick up tables
  AttributeParser::reset();
  try{
      AttributeParser::parse(cfgFileName, &cfgtrans); 
  }catch(...){
    throw LoadError("Failed to parse transition table config file.");
    exit(0);
  }

  cfgFileName = System::configurationFile( condcfgname );
  System::copyToLogDir(cfgFileName);
#ifdef DEBUG_LOAD
  Syslog::write("HMM -- Reading conditional table \"%s\".",cfgFileName);
#endif
  AttributeParser::reset();
  try{
      AttributeParser::parse(cfgFileName, &cfgcond); 
  }catch(...){
    throw LoadError("Failed to parse conditional table config file.");
      exit(0);
  }

  // Next we pull all the info out of these tables and populate our
  // own data structure.  Note that these loops depend on the fact
  // that since there is a default, all elements of the ConfigTable
  // should be instantiated.
  for(ii=0;ii<numst;ii++){
    row = statenames->at(ii);
#ifdef DEBUG_LOAD
    Syslog::write("HMM -- Initializing %s state.",row.c_str());
#endif
    addstate((char *)row.c_str());
    curRow = (ConfigTableRow<double> *)cfgtrans.findItem(ii);
    for(jj=0;jj<numst;jj++){
      val = curRow->getval(jj);
      if(val){m_statelist[ii]->settrans(jj, val);};
    }
    curRow = (ConfigTableRow<double> *)cfgcond.findItem(ii);
    for(jj=0;jj<=maxobs;jj++){
      val = curRow->getval(jj);
      if(val){m_statelist[ii]->setcond(jj, val);}
    }
  }

#ifdef DEBUG_LOAD
  dumptrans();
#endif
}


// Private methods


// Note that the first state is initialized with probability 1, while
// all others start at zero.
void HMM::addstate(char *name){
  prob_t init = (m_statelist.size() == 0);
  HMMState *newstate = new HMMState(name, init);
  m_statelist.push_back(newstate);
  m_logger->addstate(name, init);
}

int HMM::findstate(char *name){
  for(int i=0;i<m_statelist.size();i++){
    if(!strcmp(name,m_statelist[i]->name())){
      return(i);
    }
  }
  return(-1);
}

void HMM::dumptrans(){
  int i;
  std::vector<HMMState*>::iterator j;
  char st[2][256];
  bool stn;

  Syslog::write("Transition table:");
  sprintf(st[0],"to\\from");
  stn = 0;
  for(j=m_statelist.begin();j!=m_statelist.end();j++){
    stn = !stn;
    sprintf(st[stn],"%s\t%s",st[!stn],(*j)->name());
  }
  Syslog::write(st[stn]);
  for(i=0;i<6;i++){
    sprintf(st[0],"%s",m_statelist[i]->name());
    stn = 0;
    for(j=m_statelist.begin();j!=m_statelist.end();j++){
      stn = !stn;
      sprintf(st[stn],"%s\t%f",st[!stn],(*j)->trans(i));
    }
    Syslog::write(st[stn]);
  }
}
