{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "f4a0eb89",
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import pandas as pd\n",
    "\n",
    "base_dir = \"/home/tony/Data/Preference/up_v2_d4\"\n",
    "npz_dir = \"/app2/suno/data/dpo/diff2_v2_d4\"\n",
    "pkl_filename = \"fully_merged_up_v2_d4.pkl\"\n",
    "\n",
    "base_dir = \"/home/tony/Data/Preference/up_diff2_v1\"\n",
    "npz_dir = \"/app/suno/data/dpo/diff2_v1\"\n",
    "pkl_filename = \"interesting_clips_upv2_u1_20250417_full.pkl\"\n",
    "\n",
    "df = pd.read_pickle(os.path.join(base_dir, pkl_filename))\n",
    "\n",
    "# filter to ensure source is \"web\"\n",
    "#df = df[df[\"source\"] == \"web\"]\n",
    "\n",
    "# filter to ensure user_n_clips >= 10\n",
    "df = df[df[\"user_n_clips\"] >= 10]\n",
    "print(len(df))\n"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "77fb83b8",
   "metadata": {},
   "source": [
    "# Filtering\n",
    "Apply some filters to the main set of metas to get training data pairs"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "07ed60de",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "from tqdm import tqdm\n",
    "\n",
    "neg_indices = np.arange(0, len(df), 2)\n",
    "\n",
    "metas = []\n",
    "\n",
    "# Pre-fetch all relevant columns as numpy arrays or lists to avoid repeated .iloc\n",
    "# This assumes df is not a view, but a DataFrame in memory\n",
    "metadata_col = df[\"metadata\"].values\n",
    "id_col = df[\"id\"].values\n",
    "play_count_col = df[\"play_count\"].values\n",
    "upvote_count_col = df[\"upvote_count\"].values\n",
    "dislike_count_col = df[\"dislike_count\"].values\n",
    "flag_count_col = df[\"flag_count\"].values\n",
    "user_n_clips_col = df[\"user_n_clips\"].values\n",
    "prompt_text_col = df[\"prompt_text\"].values\n",
    "source_col = df[\"source\"].values\n",
    "reaction_play_count_col = df[\"reaction_play_count\"].values\n",
    "\n",
    "for i in tqdm(neg_indices):\n",
    "    # Get both rows at once\n",
    "    neg_meta = metadata_col[i]\n",
    "    pos_meta = metadata_col[i+1]\n",
    "    neg_upsample_clip_id = neg_meta.get(\"upsample_clip_id\")\n",
    "    pos_upsample_clip_id = pos_meta.get(\"upsample_clip_id\")\n",
    "\n",
    "    if neg_upsample_clip_id != pos_upsample_clip_id:\n",
    "        continue\n",
    "\n",
    "    if neg_upsample_clip_id is None or pos_upsample_clip_id is None:\n",
    "        continue\n",
    "\n",
    "\n",
    "    neg_id = id_col[i]\n",
    "    pos_id = id_col[i+1]\n",
    "\n",
    "    neg_play_count = play_count_col[i]\n",
    "    pos_play_count = play_count_col[i+1]\n",
    "    neg_upvote_count = upvote_count_col[i]\n",
    "    pos_upvote_count = upvote_count_col[i+1]\n",
    "    neg_dislike_count = dislike_count_col[i]\n",
    "    pos_dislike_count = dislike_count_col[i+1]\n",
    "    neg_flag_count = flag_count_col[i]\n",
    "    pos_flag_count = flag_count_col[i+1]\n",
    "    user_n_clips = user_n_clips_col[i]\n",
    "    neg_reaction_play_count = reaction_play_count_col[i]\n",
    "    pos_reaction_play_count = reaction_play_count_col[i+1]\n",
    "    source = source_col[i]\n",
    "\n",
    "    try:\n",
    "        tags = neg_meta[\"tags\"]\n",
    "        text = prompt_text_col[i]\n",
    "    except Exception as e:\n",
    "        print(f\"Error loading metadata for {i}: {e}\")\n",
    "        break\n",
    "\n",
    "    pos_vae_latents_filename = f\"{pos_id}_vae.npz\"\n",
    "    neg_vae_latents_filename = f\"{neg_id}_vae.npz\"\n",
    "\n",
    "    if not os.path.exists(os.path.join(npz_dir, pos_vae_latents_filename)):\n",
    "        continue\n",
    "    if not os.path.exists(os.path.join(npz_dir, neg_vae_latents_filename)):\n",
    "        continue\n",
    "\n",
    "    metas.append({\n",
    "        \"upsample_clip_id\": str(neg_upsample_clip_id),\n",
    "        \"semantic_codes_filename\": f\"{neg_upsample_clip_id}.npz\",\n",
    "        \"pos_vae_latents_filename\": f\"{pos_id}_vae.npz\",\n",
    "        \"neg_vae_latents_filename\": f\"{neg_id}_vae.npz\",\n",
    "        \"pos_vae_latents_filepath\": os.path.join(npz_dir, pos_vae_latents_filename),\n",
    "        \"neg_vae_latents_filepath\": os.path.join(npz_dir, neg_vae_latents_filename),\n",
    "        \"neg_id\": str(neg_id),\n",
    "        \"pos_id\": str(pos_id),\n",
    "        \"tags\": str(tags),\n",
    "        \"text\": str(text),\n",
    "        \"neg_play_count\": int(neg_play_count),\n",
    "        \"pos_play_count\": int(pos_play_count),\n",
    "        \"user_n_clips\": int(user_n_clips),\n",
    "        \"neg_upvote_count\": int(neg_upvote_count),\n",
    "        \"pos_upvote_count\": int(pos_upvote_count),\n",
    "        \"neg_dislike_count\": int(neg_dislike_count),\n",
    "        \"pos_dislike_count\": int(pos_dislike_count),\n",
    "        \"neg_flag_count\": int(neg_flag_count),\n",
    "        \"pos_flag_count\": int(pos_flag_count),\n",
    "        \"neg_reaction_play_count\": int(neg_reaction_play_count),\n",
    "        \"pos_reaction_play_count\": int(pos_reaction_play_count),\n",
    "        \"source\": str(source),\n",
    "    })\n",
    "\n",
    "print(f\"Total metas: {len(metas)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1ac9a240",
   "metadata": {},
   "outputs": [],
   "source": [
    "# filter out where neg_play_count is greater than pos_play_count\n",
    "metas = [meta for meta in metas if meta[\"neg_play_count\"] < meta[\"pos_play_count\"]]\n",
    "print(f\"Total metas with neg_play_count < pos_play_count: {len(metas)}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d1eb1244",
   "metadata": {},
   "outputs": [],
   "source": [
    "# filter out where user_n_clips < 10\n",
    "metas = [meta for meta in metas if meta[\"user_n_clips\"] >= 10]\n",
    "print(f\"Total metas with user_n_clips >= 10: {len(metas)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7a002042",
   "metadata": {},
   "outputs": [],
   "source": [
    "# filter to ensure positive have upvote_count > 0\n",
    "metas = [meta for meta in metas if meta[\"pos_upvote_count\"] > 0]\n",
    "print(f\"Total metas with pos_upvote_count > 0: {len(metas)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7ce05c0c",
   "metadata": {},
   "outputs": [],
   "source": [
    "# filter so that positive does not have flag_count > 0\n",
    "metas = [meta for meta in metas if meta[\"pos_flag_count\"] == 0]\n",
    "print(f\"Total metas with pos_flag_count == 0: {len(metas)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2d52ab49",
   "metadata": {},
   "outputs": [],
   "source": [
    "# filter so positive does not have dislike_count > 0\n",
    "metas = [meta for meta in metas if meta[\"pos_dislike_count\"] == 0]\n",
    "print(f\"Total metas with pos_dislike_count == 0: {len(metas)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "189c58a9",
   "metadata": {},
   "outputs": [],
   "source": [
    "# filter so that positive and negative has reaction_play_count > 0\n",
    "metas = [meta for meta in metas if meta[\"neg_reaction_play_count\"] > 0 and meta[\"pos_reaction_play_count\"] > 0]\n",
    "print(f\"Total metas with neg_reaction_play_count > 0 and pos_reaction_play_count > 0: {len(metas)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "01835a8c",
   "metadata": {},
   "outputs": [],
   "source": [
    "# ensure source is \"web\"\n",
    "metas = [meta for meta in metas if meta[\"source\"] == \"web\"]\n",
    "print(f\"Total metas with source == 'web': {len(metas)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "d746bcb7",
   "metadata": {},
   "outputs": [],
   "source": [
    "metas = [meta for meta in metas if meta[\"cosine_similarity\"] < 0.18]\n",
    "print(f\"Total metas with cosine similarity < 0.18: {len(metas)}\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "cce01604",
   "metadata": {},
   "source": [
    "# Metas\n",
    "Use filtered data to create training and validation metas for reward model training"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "433a884f",
   "metadata": {},
   "outputs": [],
   "source": [
    "# t1 - only filter on play count \n",
    "# t2 - filter on play count and user_n_clips >= 100\n",
    "# t3 - more filters\n",
    "# t4 - filter where cosine similarity is < 0.18"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8da39f1f",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.text import write_jsonl\n",
    "\n",
    "exp_name = \"t4\"\n",
    "output_path = \"/home/christian/code/christian/metadata/reward_model\"\n",
    "\n",
    "# split into train and val\n",
    "# use 95% for train and 5% for val\n",
    "train_metas = metas[:int(len(metas) * 0.95)]\n",
    "val_metas = metas[int(len(metas) * 0.95):]\n",
    "\n",
    "print(f\"Train metas: {len(train_metas)}\")\n",
    "print(f\"Val metas: {len(val_metas)}\")\n",
    "\n",
    "tr_metas_filepath = os.path.join(output_path, f\"metas_tr_{exp_name}.jsonl\")\n",
    "val_metas_filepath = os.path.join(output_path, f\"metas_val_{exp_name}.jsonl\")\n",
    "\n",
    "write_jsonl(train_metas, tr_metas_filepath)\n",
    "write_jsonl(val_metas, val_metas_filepath)"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "217bdb04",
   "metadata": {},
   "source": [
    "# Analysis\n",
    "Here we want to load all VAE pairs and then measure the cosine similarity alone the time dim and then average. This should give us a notion of similiarty in the VAE space for each pair. We would expect pairs that are near in the space to sound the same and ones that are far to sound different. The hypothesis is that many of the pairs may be too close in the VAE space to learn something useful. "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "43dbb297",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "from tqdm import tqdm\n",
    "import os\n",
    "import numpy as np\n",
    "\n",
    "from joblib import Parallel, delayed\n",
    "\n",
    "def compute_distances(meta):\n",
    "    import torch\n",
    "    import numpy as np\n",
    "    import os\n",
    "\n",
    "    neg_vae_path = os.path.join(npz_dir, meta[\"neg_vae_latents_filename\"])\n",
    "    pos_vae_path = os.path.join(npz_dir, meta[\"pos_vae_latents_filename\"])\n",
    "    neg_vae = torch.from_numpy(np.load(neg_vae_path)[\"vae_latents\"])\n",
    "    pos_vae = torch.from_numpy(np.load(pos_vae_path)[\"vae_latents\"])\n",
    "\n",
    "    # compute cosine similarity between neg and pos\n",
    "    cosine_similarity = torch.nn.functional.cosine_similarity(neg_vae, pos_vae, dim=1)\n",
    "    cosine_similarity = cosine_similarity.mean().item()\n",
    "\n",
    "    # compute euclidean distance between neg and pos\n",
    "    euclidean_distance = torch.nn.functional.pairwise_distance(neg_vae, pos_vae, p=2)\n",
    "    euclidean_distance = euclidean_distance.mean().item()\n",
    "\n",
    "    return cosine_similarity, euclidean_distance\n",
    "\n",
    "results = Parallel(n_jobs=-1, prefer=\"processes\", verbose=10)(\n",
    "    delayed(compute_distances)(meta) for meta in metas\n",
    ")\n",
    "\n",
    "cosine_distances = [r[0] for r in results]\n",
    "euclidean_distances = [r[1] for r in results]\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "62393e1c",
   "metadata": {},
   "outputs": [],
   "source": [
    "result = compute_statistics(metas[0])"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0b0b2a0a",
   "metadata": {},
   "outputs": [],
   "source": [
    "import torch\n",
    "from tqdm import tqdm\n",
    "import os\n",
    "import numpy as np\n",
    "\n",
    "from joblib import Parallel, delayed\n",
    "from scipy.stats import kurtosis, skew\n",
    "from sklearn.linear_model import LogisticRegression\n",
    "from sklearn.preprocessing import StandardScaler\n",
    "from sklearn.model_selection import train_test_split\n",
    "from sklearn.metrics import classification_report\n",
    "\n",
    "def compute_statistics(meta):\n",
    "    feature_vectors = []\n",
    "\n",
    "    for filename in [meta[\"pos_vae_latents_filename\"], meta[\"neg_vae_latents_filename\"]]:\n",
    "        vae_path = os.path.join(npz_dir, filename)\n",
    "        vae = torch.from_numpy(np.load(vae_path)[\"vae_latents\"]).float()\n",
    "\n",
    "        # Basic statistics\n",
    "        mean = vae.mean()  # shape: (batch_size,)\n",
    "        std = vae.std()\n",
    "        min_ = vae.min()\n",
    "        max_ = vae.max()\n",
    "\n",
    "        # Higher-order stats: flatten to (batch_size, -1)\n",
    "        #vae_np = vae.numpy()\n",
    "        #kurt = torch.from_numpy(kurtosis(vae_np, axis=1, fisher=True)).float()\n",
    "        #skew_ = torch.from_numpy(skew(vae_np, axis=1)).float()\n",
    "\n",
    "        # Combine all stats into a single feature vector per sample\n",
    "        stats = [mean, std, min_, max_]\n",
    "        feature_vector = torch.stack(stats)  # shape: (num_features)\n",
    "        feature_vectors.append(feature_vector)\n",
    "\n",
    "    return feature_vectors\n",
    "\n",
    "results = Parallel(n_jobs=-1, prefer=\"processes\", verbose=10)(\n",
    "    delayed(compute_statistics)(meta) for meta in metas\n",
    ")\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9ff048f1",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Collect all feature vectors and labels from the new compute_statistics output\n",
    "all_features = []\n",
    "all_labels = []\n",
    "\n",
    "for meta, feats in zip(metas, results):\n",
    "    # feats is a list: [pos_feature_vector, neg_feature_vector]\n",
    "    # Each feature_vector is shape (batch_size, num_features)\n",
    "    # Assume meta[\"label\"] is 1 for pos, 0 for neg (or adjust as needed)\n",
    "    pos_label = 1\n",
    "    neg_label = 0\n",
    "\n",
    "    pos_feats = feats[0]\n",
    "    neg_feats = feats[1]\n",
    "\n",
    "    all_features.append(pos_feats)\n",
    "    all_labels.append(torch.full((1,), pos_label, dtype=torch.long))\n",
    "    all_features.append(neg_feats)\n",
    "    all_labels.append(torch.full((1,), neg_label, dtype=torch.long))\n",
    "\n",
    "# Concatenate all features and labels\n",
    "feature_vectors = torch.stack(all_features)\n",
    "labels = torch.cat(all_labels)\n",
    "\n",
    "print(feature_vectors.shape)\n",
    "print(labels.shape)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9fe52e95",
   "metadata": {},
   "outputs": [],
   "source": [
    "X_pairs = []\n",
    "y_pairs = []\n",
    "\n",
    "for (features_i, features_j) in results:\n",
    "    phi_i = features_i\n",
    "    phi_j = features_j\n",
    "\n",
    "    X_pairs.append(phi_i - phi_j)  # difference in features\n",
    "    y_pairs.append(1)              # i preferred over j\n",
    "\n",
    "    X_pairs.append(phi_j - phi_i)  # reverse pair\n",
    "    y_pairs.append(0)              # j preferred over i"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b1870ae2",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2d831c36",
   "metadata": {},
   "outputs": [],
   "source": [
    "# -------------------------\n",
    "# TRAIN LINEAR CLASSIFIER\n",
    "# -------------------------\n",
    "\n",
    "X = np.stack(X_pairs)\n",
    "y = np.stack(y_pairs)\n",
    "print(X.shape)\n",
    "print(y.shape)\n",
    "\n",
    "\n",
    "# Normalize features\n",
    "scaler = StandardScaler()\n",
    "X = scaler.fit_transform(X)\n",
    "\n",
    "# Train/test split\n",
    "X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "e889c787",
   "metadata": {},
   "outputs": [],
   "source": [
    "print(X)\n",
    "print(y)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "c5c833cc",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Train logistic regression\n",
    "clf = LogisticRegression(max_iter=1000)\n",
    "clf.fit(X, y)\n",
    "\n",
    "# Evaluate\n",
    "y_pred = clf.predict(X)\n",
    "print(\"Classification Report:\")\n",
    "print(classification_report(y, y_pred))   \n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "35570c5d",
   "metadata": {},
   "outputs": [],
   "source": []
  },
  {
   "cell_type": "markdown",
   "id": "e5d569fd",
   "metadata": {},
   "source": [
    "# Synthetic data "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "ad81f940",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n",
    "import numpy as np\n",
    "import os\n",
    "base_dir = \"/app2/suno/data/christian/outputs/v3-bootstrap-data-t4/\"\n",
    "\n",
    "# get all directories in base_dir\n",
    "dirs = os.listdir(base_dir)\n",
    "print(len(dirs))\n",
    "# get all files in each directory\n",
    "\n",
    "model_name = \"16n_25hz_v45_infill_shared_flow_resume_1_75m\"\n",
    "\n",
    "def process_dir(dirpath):\n",
    "    a_metadata_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_a__metadata.npz\")\n",
    "    a_upsampled_vae_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_a_upsampled_vae.npz\")\n",
    "    b_metadata_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_b__metadata.npz\")\n",
    "    b_upsampled_vae_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_b_upsampled_vae.npz\")\n",
    "\n",
    "    a_metadata = np.load(a_metadata_filepath, allow_pickle=True)\n",
    "    b_metadata = np.load(b_metadata_filepath, allow_pickle=True)\n",
    "\n",
    "    # put metadata into dict\n",
    "    a_metadata_dict = {}\n",
    "    b_metadata_dict = {}\n",
    "    for key in a_metadata.keys():\n",
    "        a_metadata_dict[key] = a_metadata[key].tolist()\n",
    "    for key in b_metadata.keys():\n",
    "        b_metadata_dict[key] = b_metadata[key].tolist()\n",
    "\n",
    "    a_upsampled_vae = np.load(a_upsampled_vae_filepath)\n",
    "    b_upsampled_vae = np.load(b_upsampled_vae_filepath)\n",
    "\n",
    "    # load the mp3s \n",
    "    #a_mp3_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_a.mp3\")\n",
    "    #b_mp3_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_b.mp3\")\n",
    "    #print(f\"a: ear {a_metadata_dict['ear_score']:0.1f}, shimmer {a_metadata_dict['shimmer_score']:0.1f}, stereo {a_metadata_dict['stereo_width']:0.1f}, hoot_cer {a_metadata_dict['hoot_cer']}\")\n",
    "    #a_mp3 = Audio.from_file(a_mp3_filepath, n_channels=2).play()\n",
    "\n",
    "    #print(f\"b: ear {b_metadata_dict['ear_score']:0.1f}, shimmer {b_metadata_dict['shimmer_score']:0.1f}, stereo {b_metadata_dict['stereo_width']:0.1f}, hoot_cer {b_metadata_dict['hoot_cer']}\")\n",
    "    #b_mp3 = Audio.from_file(b_mp3_filepath, n_channels=2).play()\n",
    "\n",
    "    return a_metadata_dict, a_upsampled_vae, b_metadata_dict, b_upsampled_vae"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "02ab1593",
   "metadata": {},
   "outputs": [],
   "source": [
    "# create metas in parallel with joblib\n",
    "from tqdm import tqdm\n",
    "from joblib import Parallel, delayed\n",
    "\n",
    "def process_dir_to_meta(dirpath):\n",
    "    try:    \n",
    "        a_metadata, a_upsampled_vae, b_metadata, b_upsampled_vae = process_dir(dirpath)\n",
    "    except Exception as e:\n",
    "        #print(f\"Error processing {dirpath}: {e}\")\n",
    "        return None\n",
    "\n",
    "    tags = a_metadata[\"tags\"]\n",
    "    text = a_metadata[\"text\"]\n",
    "    semantic_codes_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_semantic.npz\")\n",
    "\n",
    "    # figure out which is positive (better) and which is negative (worse) based on ear score\n",
    "    if a_metadata[\"ear_score\"] > b_metadata[\"ear_score\"]:\n",
    "        pos_metadata = a_metadata\n",
    "        neg_metadata = b_metadata\n",
    "        pos_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_a_upsampled_vae.npz\")\n",
    "        neg_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_b_upsampled_vae.npz\")\n",
    "    else:\n",
    "        pos_metadata = b_metadata\n",
    "        neg_metadata = a_metadata\n",
    "        pos_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_b_upsampled_vae.npz\")\n",
    "        neg_vae_latents_filepath = os.path.join(base_dir, dirpath, f\"{dirpath}_{model_name}_a_upsampled_vae.npz\")\n",
    "\n",
    "        # load the vae latents and make sure they are large enough\n",
    "        pos_vae_latents = np.load(pos_vae_latents_filepath)[\"vae_latents\"]\n",
    "        neg_vae_latents = np.load(neg_vae_latents_filepath)[\"vae_latents\"]\n",
    "        if pos_vae_latents.shape[0] < 1024 or neg_vae_latents.shape[0] < 1024:\n",
    "            return None\n",
    "\n",
    "    meta = {\n",
    "        \"id\": dirpath,\n",
    "        \"tags\": tags,\n",
    "        \"text\": str(text),\n",
    "        #\"pos_metadata\": pos_metadata,\n",
    "        #\"neg_metadata\": neg_metadata,\n",
    "        \"pos_vae_latents_filepath\": pos_vae_latents_filepath,\n",
    "        \"neg_vae_latents_filepath\": neg_vae_latents_filepath,\n",
    "        \"semantic_codes_filepath\": semantic_codes_filepath,\n",
    "    }\n",
    "    return meta\n",
    "\n",
    "# Use joblib to parallelize\n",
    "metas = []\n",
    "results = Parallel(n_jobs=-1)(\n",
    "    delayed(process_dir_to_meta)(dirpath) for dirpath in tqdm(dirs)\n",
    ")\n",
    "# Filter out any None results (from failed processing)\n",
    "metas = [meta for meta in results if meta is not None]\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "b249da65",
   "metadata": {},
   "outputs": [],
   "source": [
    "# split metas into train and val \n",
    "\n",
    "metas_tr = metas[:int(len(metas) * 0.98)]\n",
    "metas_val = metas[int(len(metas) * 0.98):]\n",
    "\n",
    "print(f\"len(metas_tr): {len(metas_tr)}\")\n",
    "print(f\"len(metas_val): {len(metas_val)}\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "028a3099",
   "metadata": {},
   "outputs": [],
   "source": [
    "# save the metas to a jsonl file\n",
    "from suno_utils.utils.text import write_jsonl\n",
    "\n",
    "exp_name = \"t6\"\n",
    "base_dir = \"/home/christian/code/christian/metadata/reward_model/\"\n",
    "\n",
    "write_jsonl(metas_tr, \"/home/christian/code/christian/metadata/reward_model/metas_tr_t6.jsonl\")\n",
    "write_jsonl(metas_val, \"/home/christian/code/christian/metadata/reward_model/metas_val_t6.jsonl\")"
   ]
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_diff",
   "language": "python",
   "name": "python3"
  },
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.12.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
