##############################################
# Semi-Automated Plant Community Classification
# Workflow implemented in R
# Packages: terra, sp, sf, rgdal, raster, rsample, MLmetrics, randomForest
##############################################

##############################################
# 0. Preparing R environment
##############################################

# Load required packages, install if missing
packages <- c("terra", "sp", "sf", "rgdal", "raster", "rsample", "MLmetrics", "randomForest")

for (pkg in packages) {
  if (!requireNamespace(pkg, quietly = TRUE)) {
    install.packages(pkg)
  }
  library(pkg, character.only = TRUE)
}

# Set working directory (in which folder the files will be saved) and site name
site = "viscosa"
path = paste("I:/R_", site, sep = "")
setwd(path)
dataset = "multispectral"

################################################
# 1. Raster Preprocessing
################################################

# Collect all raster files and create a stack
rst_files = list.files(path = dataset, pattern = "\\.tif$", full.names = T)
r_stack = stack(rst_files)

################################################
# 2. Data Extraction
################################################

# Import ground truth data
samples = readOGR("VI_GNNS_joined_data.gpkg") # Replace with your geopackage path

# Extract raster values for each sample polygon
data.df=terra::extract(r_stack,samples,df=TRUE)

# Merge extracted values with plant community codes
ID=c(1:length(samples$Number.code))
temp.df=cbind(ID,samples$Number.code,samples$Plant.community.code)

# Add columns for codes
N_CODE=c(rep(NA,length(data.df$ID)))
COM_CODE=c(rep(NA,length(data.df$ID)))

# Populate codes
data.df=cbind(data.df,N_CODE,COM_CODE)
temp.df=data.frame(temp.df)
for (i in 1:length(temp.df[,1])) {
  l=temp.df[i,]
  a1=as.character(l[1])
  b1=as.character(l[2])
  b2=as.character(l[3])
  for (j in 1:length(data.df[,1])) {
    w=data.df[j,]
    a2=as.character(w$ID)
    data.df$N_CODE[j][a1==a2] = b1
    data.df$COM_CODE[j][a1==a2] = b2
  }
}

# Convert to factors
data.df$N_CODE = as.factor(data.df$N_CODE)
data.df$COM_CODE = as.factor(data.df$COM_CODE)

# Export extracted data for reproducibility (Optional)
write.csv(data.df[,-1], paste("20251110_", site,"_extracted_", dataset, "_data_field.csv", sep = ""))

################################################
# 3. Dataset Preparation
################################################
set.seed(123)

# Split into training (75%) and validation (25%) sets using stratified sampling based on levels in "COM_CODE" to maintain class balance
split = initial_split(data.df, prop = 0.75, strata = "COM_CODE")

train = training(split)
vali = testing(split)

train = na.omit(train)
vali = na.omit(vali)

################################################
# 4. Model Tuning
################################################

# Define parameter grid for Random Forest
mtry_values <- seq(2, ncol(train), by = 2) # Adjust step size for finer search
ntree_values <- c(100, 200, 500, 1000) # Increase if dataset is large (while it improves stability, it increases computation time)

results <- data.frame(mtry = integer(), ntree = integer(), accuracy = numeric())

for (m in mtry_values) {
  for (n in ntree_values) {
    rf_model <- randomForest(COM_CODE ~ ., data = train[,-c(1, ncol(train)-1)], mtry = m, ntree = n)
    preds <- predict(rf_model, newdata = vali)
    acc <- mean(preds == vali$COM_CODE)
    
    results <- rbind(results, data.frame(mtry = m, ntree = n, accuracy = acc))
  }
}

# View tuning results
print(results)

# Select best parameters
best <- results[which.max(results$accuracy), ]
print(best)
best_mtry <- best$mtry
best_ntree <- best$ntree

################################################
# 5. Model Training
################################################

# Train final RF model using best parameters
final_rf <- randomForest(COM_CODE ~ ., data = train[,-c(1, ncol(train)-1)], mtry = best_mtry, ntree = best_ntree)
final_rf

# Plot variable importance (Optional)
varImpPlot(final_rf)

# Save final model for reuse (Optional)
saveRDS(final_rf, file = paste("final_rf_model_",site,"_",dataset,".rds", sep = ""))

################################################
# 6. Prediction
################################################

# Covert raster stack to data frame format
rst_df = as.data.frame(r_stack,xy=TRUE)

# Apply RF model to raster stack for spatial classification
community.df = predict(final_rf,newdata=rst_df, progress='text')

# Convert predictions to raster
classified_raster = rasterFromXYZ(data.frame(rst_df[, c(1, 2)], as.numeric(community.df)))
crs(classified_raster) = crs(r_stack)

# Save classified raste
writeRaster(classified_raster, paste("2025_",site,"_",dataset,"_PlantCommunities.tif", sep = ""), overwrite = TRUE)

################################################
# 7. Accuracy Assessment
################################################

final_preds <- predict(final_rf, newdata = vali)

# Compute confusion matrix, accuracy, F1-score, precision, and recall
cm = ConfusionMatrix(y_pred = final_preds, y_true = vali$COM_CODE)
print(cm)

# Overall metrics
F1_s = F1_Score(vali$COM_CODE, final_preds)
print(paste("F1 Score:", F1_s))

acc = Accuracy(y_pred = final_preds, y_true = vali$COM_CODE)
print(paste("Accuracy:", acc))

# Per-class metrics
classes <- levels(as.factor(vali$COM_CODE))

for (cls in classes) {
  cat("\nClass:", cls, "\n")
  cat("Precision:", Precision(final_preds, vali$COM_CODE, positive = cls), "\n")
  cat("Recall:", Recall(final_preds, vali$COM_CODE, positive = cls), "\n")
  cat("F1 Score:", F1_Score(final_preds, vali$COM_CODE, positive = cls), "\n")
}