
train_station_lstm <- function(x_train,y_train,sample.weight,model_file) {
  # Build the LSTM model
  if(length(dim(x_train))==2){shape_last=1}
  if(length(dim(x_train))==3){shape_last=tail(dim(x_train),1)}
  model <- keras_model_sequential() %>%
    layer_lstm(units = 100
               ,batch_input_shape = c(batch.size,lag,shape_last)
               ,return_sequences = TRUE,
               stateful=stateful.lab) %>%
    layer_lstm(units=50
               ,return_sequences=TRUE,
               stateful= stateful.lab) %>%
    layer_lstm(units=20
               ,return_sequences=TRUE,
               stateful= stateful.lab) %>%
    layer_dropout(rate = 0.1) 

    if(non.neg){
      model <- model %>% layer_dense(units = 1,activation = output.activation,
                           kernel_constraint = constraint_nonneg()
      )
    }else{
      model <- model %>% layer_dense(units = 1,activation = output.activation
      )
    }
  
  # Compile the model
  model %>% compile(
    optimizer = optimizer_adam(learning_rate=lr.),  
    loss = "mean_squared_error",
    weighted_metrics = list("mean_squared_error")
  )
  summary(model)
  
  if(stateful.lab){
    cat("stateful model.............\n")
    # Define the ModelCheckpoint callback
    mc <- callback_model_checkpoint(
      filepath = paste0("best_model",i,".h5"), 
      monitor = "val_loss", 
      mode = "min", 
      verbose = 2, 
      save_best_only = TRUE
    )
    
    # Define the LambdaCallback for resetting model states
    es <- callback_lambda(
      on_epoch_end = function(batch, logs) {
        model$reset_states()
      }
    )
    
    hist <- model %>% fit(x_train, y_train, epochs=Epochs, batch_size=batch.size, 
                          sample_weight=sample.weight,
                          verbose=2, 
                          shuffle=FALSE,
                          validation_split=0.1,
                          callbacks = list(mc, es)
    )
  }else{
    cat("unstateful model.............\n")
    # Train the model
    if(model.type=="single"){
      check.p.id=i
    }else{
      check.p.id=cluster.id
    }
    mc <- callback_model_checkpoint(
      filepath = paste0("best_model",check.p.id,".h5"), 
      monitor = "val_loss", 
      mode = "min", 
      verbose = 2, 
      save_best_only = TRUE
    )
    
    
    history <- model %>% fit(
      x=x_train,
      y=y_train,
      batch_size=batch.size,
      epochs = Epochs,
      callbacks=list(mc,callback_early_stopping(monitor='val_loss',verbose=1,
                                        patience = 10,restore_best_weights=TRUE)),
      sample_weight=sample.weight,
      verbose=2,
      validation_split=0.1,
      shuffle=shuffle.lab
    )
  }
  # Save the model to an HDF5 file
  save_model_hdf5(model, model_file)
}


## prepare data
prepare_data <- function(i,n.v,lag){
    qsim=simCOUT[,i]
    vsim=qsim
    vsim.lag=embed(vsim,lag)[,lag:1]
  return(vsim.lag) 
}


scale_data = function(train, test, feature_range = c(0, 1)) {
  if(is.character(feature_range)){
    if(feature_range=="std"){
    x=asub(train,lag,dim=2)
    if(is.null(dim(x))){
      scale.=sd(x,na.rm=TRUE)
      center.=mean(x,na.rm=TRUE)
      scale.=replicate(lag,scale.)
      center.=replicate(lag,center.)
    }else{
      scale.=apply(x,tail(1:length(dim(x)),-1),sd,na.rm=TRUE)
      center.=apply(x,tail(1:length(dim(x)),-1),mean,na.rm=TRUE)
      scale.=t(replicate(lag,scale.))
      center.=t(replicate(lag,center.))
    }

    scaled_train = apply_normalization(train,center.,scale.)
    scaled_test = apply_normalization(test,center.,scale.)
    return( list(scaled_train = scaled_train, scaled_test = scaled_test ,center=center., scale= scale.) )
    }
  }else{
    x = train
    fr_min = feature_range[1]
    fr_max = feature_range[2]
    
    min.=apply(x,tail(1:length(dim(x)),-1),min,na.rm=TRUE)
    max.=apply(x,tail(1:length(dim(x)),-1),max,na.rm=TRUE)
    
    scale.=max.-min.
    center.=min.
    
    std_train = apply_normalization(train,center.,scale.)
    std_test  = apply_normalization(test,center.,scale.)
    
    scaled_train = std_train *(fr_max -fr_min) + fr_min
    scaled_test = std_test *(fr_max -fr_min) + fr_min
    return( list(scaled_train = scaled_train, scaled_test = scaled_test , center=center., scale= scale.) )
    
  }
  
}


evaluate_model <- function(x,y,qsim,qobs,note){
  temp <- model %>%
    predict(x, batch_size = batch.size,verbose=FALSE) %>%
    .[, , 1]
  
  
    if (length(dim(qsim)) > 2) {
      # qsim has more than 2 dimensions (e.g., 3D array)
      qsim_ev <- qsim[, lag, 1]
    } else {
      # qsim has only 2 dimensions (matrix)
      qsim_ev <- qsim[, lag]
    }
    qlstm_ev=temp[, lag]
    qobs_ev=qobs
  
  
  if (target.lab == "res2") {
    qlstm_ev = (qlstm_ev - 1) * (qsim_ev+0.1) + qsim_ev
  }
  
  tt.plot=get(paste0("tt.",note))
  all_dates <- data.frame(DATE = seq(tt.plot[1],tail(tt.plot,1), by = "day"))
  pp.data <- data.frame(DATE=tt.plot,qobs=qobs_ev,qsim=qsim_ev,pp=qlstm_ev)
  pp.data <- merge(pp.data,all_dates,by="DATE",all=TRUE)  
  plot_station(pp.data,i,"lstm",note)
  res=calculate_metrics(qsim_ev, qlstm_ev, qobs_ev, "ls")
  return(res)
}


plot_station <- function(data,station.ind,prefix,note){
  png(paste0(prefix,"_station_timeserie",station.ind,note,n.v,".png"),width = 10, height = 4, units = "in",res=300)
  qobs_ev=data[,2]
  qsim_ev=data[,3]
  qpp_ev=data[,4]
  X=data[,1]
  plot(X,qobs_ev,type="l",lty=2,col="black",xlab="Date",ylab="Runoff (mm/day)")
  lines(X,qsim_ev,type="l",lty=1,col="#d95f02")
  lines(X,qpp_ev,type="l",lty=1,col="#1b9e77")
  abline(h=0, col="black")
  legend("topright", legend=c("Obs", "EHYPE",toupper(prefix)),
         col=c("black", "#d95f02","#1b9e77"), lty=c(2,1,1), cex=0.8)
  dev.off()
}