import os
import cv2
from tqdm import tqdm
from py_sod_metrics import MAE, Emeasure, Fmeasure, Smeasure, WeightedFmeasure

method = 'CPNet'
# for _data_name in ['DUT-RGBD', 'NJU2K','NLPR','SIP']:
for _data_name in ['CAMO','CHAMELEON','COD10K','NC4K']:
# for _data_name in ['VT821', 'VT1000', 'VT5000']:
    print("eval-dataset: {}".format(_data_name))
    mask_root = './RGBT_dataset/test/{}/{}/'.format(_data_name,"GT") # change path
    # mask_root = './COD-TestDataset/{}/{}/'.format(_data_name, "GT")  # change path
    pred_root = './test_maps/CPNet/{}/'.format( _data_name) # change path
    mask_name_list = sorted(os.listdir(mask_root))
    FM = Fmeasure()
    WFM = WeightedFmeasure()
    SM = Smeasure()
    EM = Emeasure()
    M = MAE()
    for mask_name in tqdm(mask_name_list, total=len(mask_name_list)):
        mask_path = os.path.join(mask_root, mask_name)
        pred_path = os.path.join(pred_root, mask_name)
        mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)
        pred = cv2.imread(pred_path, cv2.IMREAD_GRAYSCALE)
        FM.step(pred=pred, gt=mask)
        WFM.step(pred=pred, gt=mask)
        SM.step(pred=pred, gt=mask)
        EM.step(pred=pred, gt=mask)
        M.step(pred=pred, gt=mask)

    fm = FM.get_results()["fm"]
    wfm = WFM.get_results()["wfm"]
    sm = SM.get_results()["sm"]
    em = EM.get_results()["em"]
    mae = M.get_results()["mae"]

    results = {
        "Smeasure": sm,
        "wFmeasure": wfm,
        "MAE": mae,
        "adpEm": em["adp"],
        "meanEm": em["curve"].mean(),
        "maxEm": em["curve"].max(),
        "adpFm": fm["adp"],
        "meanFm": fm["curve"].mean(),
        "maxFm": fm["curve"].max(),
    }

    print(results)
    file=open("./path_to_eval_results/eval_results.txt", "a")
    file.write(method+' '+_data_name+' '+str(results)+'\n')

print("Eval finished!")