from suno_utils.utils.text import read_jsonl, write_jsonl

import polars as pl

# dataset sources
# discogs_subset, 2.9M
# genius, 2.1M
# imslp, 200k
# deezer, 700k
# podcast,
# sfx, 5M

# Define all dataset parquet filepaths in a dictionary for easier loading and management
parquet_filepaths = {
    "discogs_subset": "/app2/suno/data/christian/metadata/raw_discogs_subset_metas.parquet",
    # "genius": "/app2/suno/data/christian/metadata/raw_genius_metas.parquet",
    "genius": "/app2/suno/data/christian/metadata/genius_hq_metas_plus.parquet",
    "imslp": "/app2/suno/data/christian/metadata/raw_imslp_metas.parquet",
    "deezer": "/app2/suno/data/christian/metadata/raw_deezer_metas.parquet",
    # "sfx": "/app2/suno/data/christian/metadata/combined_v3_w_extreme_metas_v0_aligned.parquet"
}

# Load the parquet files into Polars DataFrames and keep them in a dictionary
dfs = {}
for key, path in parquet_filepaths.items():
    try:
        dfs[key] = pl.read_parquet(path)
        print(f"Loaded {key} ({dfs[key].shape[0]:,} rows)")
    except Exception as e:
        print(f"Error loading {key} from {path}: {e}")

# create a special map for genius that maps id to original id
genius_df = dfs.get("genius")
genius_id_map = {}
for row in genius_df.iter_rows(named=True):
    genius_id_map[row["id"]] = row["original_id"]

from suno_utils.utils.text import read_jsonl

version = "5"

# load some alignemnts info (h5 alignments)
genius_alignments_filepath = (
    f"/home/tony/Work/tony/hoot/tmp/genius_hq_alignments_h5_t480_v{version}.jsonl"
)
discogs_alignments_filepath = (
    f"/home/tony/Work/tony/hoot/tmp/discogs_hq_alignments_h5_t480_v{version}.jsonl"
)
deezer_alignments_filepath = (
    f"/home/tony/Work/tony/hoot/tmp/deezer_hq_alignments_h5_t480_v{version}.jsonl"
)

genius_alignments = read_jsonl(genius_alignments_filepath, progress=False)
print(len(genius_alignments))
discogs_alignments = read_jsonl(discogs_alignments_filepath, progress=False)
print(len(discogs_alignments))
deezer_alignments = read_jsonl(deezer_alignments_filepath, progress=False)
print(len(deezer_alignments))


def build_alignment_map(data):
    result = {}
    for k, v, cer in data:
        meta = v[0]
        lines, starts, ends = (
            meta["line_text"],
            meta["line_start_s"],
            meta["line_end_s"],
        )

        line_entries = []
        for text, start, end in zip(lines, starts, ends):
            if start is None or end is None:
                continue
            line_entries.append(
                (
                    start,
                    end,
                    text,
                )
            )

        result[k] = {
            "lines": line_entries,
            "cer": cer,
            "text": meta.get("text"),
            "start_s": meta.get("start_s"),
            "end_s": meta.get("end_s"),
            "vocal_start_s": meta.get("vocal_start_s"),
            "vocal_end_s": meta.get("vocal_end_s"),
        }
    return result


# usage
genius_alignments_map = build_alignment_map(genius_alignments)
discogs_alignments_map = build_alignment_map(discogs_alignments)
deezer_alignments_map = build_alignment_map(deezer_alignments)

# adjust the genius_alignments_map to have the original_id as the key
genius_alignments_map = {genius_id_map[k]: v for k, v in genius_alignments_map.items()}


print("Loading metas_v9_tr.parquet")
parquet_filepath = "/app2/suno/data/christian/metadata/metas_v9_tr.parquet"
df = pl.read_parquet(parquet_filepath)

# Filter the DataFrame `df` to only include rows where at least one of
# "text", "text_aligned", "stems", or "tags" is not null, AND "weight" is not null
df_filtered = df.filter(
    (
        pl.col("text").is_not_null()
        | pl.col("text_aligned").is_not_null()
        | pl.col("stems").is_not_null()
        | pl.col("tags").is_not_null()
    )
)
print(df_filtered.height)

# load audio quality scores
discogs_subset_ear_scores = pl.read_csv(
    "/home/christian/code/christian/metadata/organized/ear/discogs_subset_ear_scores.csv"
)
genius_ear_scores = pl.read_csv(
    "/home/christian/code/christian/metadata/organized/ear/genius_ear_scores.csv"
)
imslp_ear_scores = pl.read_csv(
    "/home/christian/code/christian/metadata/organized/ear/imslp_ear_scores.csv"
)

# convert the column "mean_score" to "ear_score" in all the dataframes
discogs_subset_ear_scores = discogs_subset_ear_scores.rename(
    {"mean_score": "ear_score"}
)
genius_ear_scores = genius_ear_scores.rename({"mean_score": "ear_score"})
# imslp_ear_scores = imslp_ear_scores.rename({"mean_score": "ear_score"})

# merge the df_filtered with the ear scores
df_filtered = df_filtered.join(discogs_subset_ear_scores, on="id", how="left")
df_filtered = df_filtered.join(genius_ear_scores, on="id", how="left")
# df_filtered = df_filtered.join(imslp_ear_scores, on="id", how="left")

# iterate over the rows in the df
from tqdm import tqdm

new_metas = []
missing_alignments = []
with_alignments = []

# i want to reset weights to 1 for all rows

# this is just going to repair the metas alignment issues
# i guess only use text when we have alignments?
for i, row in tqdm(
    enumerate(df_filtered.iter_rows(named=True)),
    total=df_filtered.height,
    desc="Processing rows",
):

    weight = row.get("weight", 1)
    if weight is None:
        weight = 1

    if weight < 1:
        continue

    # check if this has text_aligned
    has_text_aligned = row["text_aligned"] is not None
    have_genius_aligned = False
    have_discogs_aligned = False
    have_deezer_aligned = False
    have_podcast_aligned = True if "podcast" in row["id"] else False
    genius_cer = None
    discogs_cer = None
    deezer_cer = None
    # check if the id is in the alignments map
    if row["id"] in genius_alignments_map:
        # print(row["id"])
        # print(row)
        # print(genius_alignments_map[row["id"]])
        have_genius_aligned = True
        genius_cer = genius_alignments_map[row["id"]]["cer"]
        genius_text = genius_alignments_map[row["id"]]["text"]
        genius_alignments = genius_alignments_map[row["id"]]["lines"]
    elif row["id"] in discogs_alignments_map:
        # print(row["id"])
        # print(row)
        # print(discogs_alignments_map[row["id"]])
        have_discogs_aligned = True
        discogs_cer = discogs_alignments_map[row["id"]]["cer"]
        discogs_text = discogs_alignments_map[row["id"]]["text"]
        discogs_alignments = discogs_alignments_map[row["id"]]["lines"]
    elif row["id"] in deezer_alignments_map:
        # print(row["id"])
        # print(row)
        # print(deezer_alignments_map[row["id"]])
        have_deezer_aligned = True
        deezer_cer = deezer_alignments_map[row["id"]]["cer"]
        deezer_text = deezer_alignments_map[row["id"]]["text"]
        deezer_alignments = deezer_alignments_map[row["id"]]["lines"]

    # has alignments
    new_meta = row

    # set the weight to 1
    new_meta["weight"] = 1

    # only use text_aligned if it comes from genius, discogs, or deezer
    # otherwise we null the text, and text_aligned
    # also we check if the CER is lower than 0.8
    if have_genius_aligned:
        if genius_cer < 0.8:
            new_meta["text"] = genius_text
            new_meta["text_aligned"] = genius_alignments
        else:
            new_meta["text"] = ""
            new_meta["text_aligned"] = []
    elif have_discogs_aligned:
        if discogs_cer < 0.8:
            new_meta["text"] = discogs_text
            new_meta["text_aligned"] = discogs_alignments
        else:
            new_meta["text"] = ""
            new_meta["text_aligned"] = []
    elif have_deezer_aligned:
        if deezer_cer < 0.8:
            new_meta["text"] = deezer_text
            new_meta["text_aligned"] = deezer_alignments
        else:
            new_meta["text"] = ""
            new_meta["text_aligned"] = []
    elif have_podcast_aligned:
        pass  # keep the text as is
    else:
        new_meta["text"] = ""
        new_meta["text_aligned"] = []

    new_metas.append(new_meta)


print(len(missing_alignments))
print(len(with_alignments))

# count the number of rows with text, and text_aligned in df_filtered (as a DataFrame)
num_with_text = (df_filtered["text"] != "").sum()
# Use list comprehension to count the length of 'text_aligned' if it's a list or tuple
num_with_text_aligned = sum(
    isinstance(x, (list, tuple)) and len(x) > 0 for x in df_filtered["text_aligned"]
)
print(f"Number of rows with text: {num_with_text}")
print(f"Number of rows with text_aligned: {num_with_text_aligned}")

# count the number of rows with text, and text_aligned
# in new_metas
num_with_text = sum(1 for meta in new_metas if meta["text"] != "")
num_with_text_aligned = sum(1 for meta in new_metas if meta["text_aligned"] != [])
print(f"Number of rows with text: {num_with_text}")
print(f"Number of rows with text_aligned: {num_with_text_aligned}")

# Split new_metas into train and val sets with 1% in val, randomly distributed

import random

val_frac = 0.01
random.seed(42)
df_shuffled = new_metas[:]
random.shuffle(df_shuffled)
num_val = int(len(df_shuffled) * val_frac)
df_val = df_shuffled[:num_val]
df_train = df_shuffled[num_val:]

print("train:", len(df_train), "val:", len(df_val))

# save to new jsonl files
from suno_utils.utils.text import write_jsonl

train_filepath = "/app2/suno/data/diffusion/v1/metas_diff_v0_tr.jsonl"
val_filepath = "/app2/suno/data/diffusion/v1/metas_diff_v0_val.jsonl"
write_jsonl(df_train, train_filepath)
write_jsonl(df_val, val_filepath)
print("Done.")
