source("LoadFunctions.R")
library(reticulate)
library(abind)
library(hydroGOF)
library(tensorflow)
env_path <-
  file.path(.libPaths()[1], "reticulate_env", "r-reticulate.env")
use_virtualenv(env_path)
library(keras)

args = commandArgs(trailingOnly = TRUE)

lag = as.integer(args[1])
setting = as.integer(args[2])
source("Settings.R")
model.type="single"


n.v = 1
load('Data.RData')

st = as.integer(args[3])
end = min(st + as.integer(args[4]) - 1, nrow(Qsta.info))

setwd(fold.out)

for (i in st:end) {
    print(i)
    sta.name = as.character(Qsta.info$SUBID[i])
    qobs = obsCOUT[, i]
    qsim = simCOUT[, i]
    data.x = prepare_data(i, n.v, lag)
    if(target.lab=="ori"){
      data.y=qobs[-(1:lag - 1)]
      data.y.sw=data.y
    }
    if(target.lab=="res2"){
      qres = (qobs - qsim)/(qsim+0.1)+1
      qres = qres[-(1:lag - 1)]
      data.y = qres
      data.y.sw=qobs[-(1:lag - 1)]
      tt.=tt[-(1:lag - 1)]
    }
   
    
    
    ## drop rows with NA
    na.idx = which(is.na(data.y))
    
    if (length(na.idx) != 0) {
      data.x = asub(data.x, -na.idx, dims = 1)
      data.y = data.y[-na.idx]
      data.y.sw=data.y.sw[-na.idx]
      tt.=tt.[-na.idx]
    }
    
    
    ## split into tr,ts by time series
    tr.idx = seq(1, ceiling(length(data.y) * 0.8))
    
    
    x.tr = asub(data.x, tr.idx, dims = 1)
    x.ts = asub(data.x, -tr.idx, dims = 1)
    
    y.tr = data.y[tr.idx]
    y.ts = data.y[-tr.idx]
    y.tr.sw=data.y.sw[tr.idx]
    y.ts.sw=data.y.sw[-tr.idx]
    tt.tr=tt.[tr.idx]
    tt.ts=tt.[-tr.idx]
    
    Scaled.x = scale_data(x.tr, x.ts, x.range)
    data.x.tr= Scaled.x$scaled_train
    data.x.ts= Scaled.x$scaled_test
    center.=Scaled.x$center
    scale.=Scaled.x$scale
    
    if(exists("y.range")){
    Scaled.y = scale_data(y.tr, y.ts, y.range)
    data.y.tr= Scaled.y$scaled_train
    data.y.ts= Scaled.y$scaled_test
    }else{
      data.y.tr=y.tr
      data.y.ts=y.ts
    }
    
    
    
    ## calculate sample_weight
    if (sample.weight) {
      x=y.tr.sw
      br = quantile(x,c(0.1,0.33,0.66,0.9))
      gp.n=table(br)
      weig.assign=c(gp.n*0.2,0.2)*length(x)
      br=unique(br)
      freq   = hist(x,
                    breaks = c(-Inf,br,Inf),
                    include.lowest = TRUE,
                    plot = FALSE)
      weig= weig.assign/ freq$counts
      weight.va = cut(x, c(-Inf,br,Inf), include.lowest = TRUE, weig)
      weight.va = as.numeric(levels(weight.va))[weight.va]
    }
    
    
    
    
    
    batch.size = 1
    cat(paste('\n lstm for station', i, "\n", sep = " "))
    model_file <- paste0("lstm_model_station_", i,"nv",n.v, ".h5")
    if (!file.exists(paste0("lstm_model_station_", i, "nv",n.v,".h5"))) {
    train_station_lstm(data.x.tr, data.y.tr,weight.va, model_file)
    }
    
    model <- load_model_hdf5(model_file)
    
    eva.tr = evaluate_model(data.x.tr, data.y.tr, x.tr,y.tr.sw,"tr")
    eva.ts = evaluate_model(data.x.ts, data.y.ts, x.ts,y.ts.sw,"ts")
    
    save(
      eva.tr,
      eva.ts,
      center.,
      scale.,
      file = paste("Eva_lstm_nv", n.v, "lag", lag, "_station", i, ".RData", sep =
                     "")
    )
}