{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "0",
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.text import read_jsonl, write_jsonl\n",
    "import pandas as pd"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1",
   "metadata": {},
   "outputs": [],
   "source": [
    "whosampled_all = read_jsonl(\"/home/sara/task_data/whosampled_processed.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "2",
   "metadata": {},
   "outputs": [],
   "source": [
    "whosampled_metadata = read_jsonl(\"/home/sara/who_sampled_victor/who_sampled/data/all_metadata_sara.jsonl\")\n",
    "whosampled_metadata_df = pd.DataFrame(whosampled_metadata)\n",
    "whosampled_metadata_df['date'] = whosampled_metadata_df['date'].fillna(0).astype(int)\n",
    "metadata_dict = whosampled_metadata_df.set_index('youtube').to_dict('index')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3",
   "metadata": {},
   "outputs": [],
   "source": [
    "whosampled_metadata_df.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "4",
   "metadata": {},
   "outputs": [],
   "source": [
    "for mashup in whosampled_all:\n",
    "    full_metadata = metadata_dict[mashup['output_id']]\n",
    "    artists = None\n",
    "    producers = None\n",
    "    if isinstance(full_metadata['artist'], str):\n",
    "        artists = full_metadata['artist']\n",
    "    elif isinstance(full_metadata['artist'], list):\n",
    "        artists = full_metadata['artist'].join(\",\")\n",
    "\n",
    "    if isinstance(full_metadata['producer'], str):\n",
    "        producers = full_metadata['producer']\n",
    "    elif isinstance(full_metadata['producer'], list):\n",
    "        producers = full_metadata['producer'].join(\",\")\n",
    "\n",
    "    if artists is not None and producers is not None:\n",
    "        artists = list(set((artists + \",\" + producers).split(\",\")))\n",
    "    elif artists is not None:\n",
    "        artists = list(set(artists.split(\",\")))\n",
    "    elif producers is not None:\n",
    "        artists = list(set(producers.split(\",\")))\n",
    "\n",
    "    print(artists)\n",
    "    mashup['artists'] = artists\n",
    "    mashup['song_name'] = full_metadata['song']\n",
    "    mashup['release_date'] = full_metadata['date']\n",
    "    mashup['album_name'] = full_metadata['release']\n",
    "    mashup['duration'] = full_metadata['duration']"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5",
   "metadata": {},
   "outputs": [],
   "source": [
    "whosampled_all[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "6",
   "metadata": {},
   "outputs": [],
   "source": [
    "full_df = pd.DataFrame(whosampled_all)\n",
    "full_df.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "7",
   "metadata": {},
   "outputs": [],
   "source": [
    "only_valid_mashups = full_df[full_df['is_valid_mashup'] == True]\n",
    "print(len(only_valid_mashups))\n",
    "only_valid_mashups.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "8",
   "metadata": {},
   "outputs": [],
   "source": [
    "only_valid_mashups = only_valid_mashups[only_valid_mashups['data_source'] != \"whosampled_cover\"]\n",
    "print(len(only_valid_mashups))\n",
    "only_valid_mashups.head()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "9",
   "metadata": {},
   "outputs": [],
   "source": [
    "import numpy as np\n",
    "\n",
    "def filter_sources(row):\n",
    "    \"\"\"\n",
    "    Remove sources with votes < (mean - 1 std dev) if row has > 2 sources.\n",
    "    Always keep at least 2 sources per row.\n",
    "    \"\"\"\n",
    "    source_ids = row['source_ids']\n",
    "    source_votes = row['source_votes']\n",
    "    \n",
    "    # Handle None/NaN values\n",
    "    if source_ids is None or source_votes is None or not isinstance(source_ids, list) or not isinstance(source_votes, list):\n",
    "        return pd.Series({'source_ids': source_ids, 'source_votes': source_votes})\n",
    "    \n",
    "    # If row has <= 2 sources, keep unaltered (always keep at least 2)\n",
    "    if len(source_ids) <= 2:\n",
    "        return pd.Series({'source_ids': source_ids, 'source_votes': source_votes})\n",
    "    \n",
    "    # Calculate mean and standard deviation\n",
    "    mean_votes = np.mean(source_votes)\n",
    "    std_votes = np.std(source_votes)\n",
    "    threshold = mean_votes - std_votes\n",
    "    \n",
    "    # Remove sources below (mean - 1 std dev)\n",
    "    filtered = [(sid, vote) for sid, vote in zip(source_ids, source_votes) \n",
    "                if vote >= threshold]\n",
    "    \n",
    "    # Always keep at least 2 sources - if filtering would leave < 2, keep original\n",
    "    if len(filtered) < 2:\n",
    "        return pd.Series({'source_ids': source_ids, 'source_votes': source_votes})\n",
    "    \n",
    "    filtered_ids, filtered_votes = zip(*filtered)\n",
    "    return pd.Series({'source_ids': list(filtered_ids), 'source_votes': list(filtered_votes)})\n",
    "\n",
    "# Apply the filter\n",
    "filtered_columns = only_valid_mashups.apply(filter_sources, axis=1)\n",
    "\n",
    "# Create new dataframe with filtered data\n",
    "filtered_mashups = only_valid_mashups.copy()\n",
    "filtered_mashups['source_ids'] = filtered_columns['source_ids']\n",
    "filtered_mashups['source_votes'] = filtered_columns['source_votes']\n",
    "\n",
    "# Print statistics\n",
    "original_total_sources = only_valid_mashups.apply(lambda row: len(row['source_ids']) if isinstance(row['source_ids'], list) else 0, axis=1).sum()\n",
    "filtered_total_sources = filtered_mashups.apply(lambda row: len(row['source_ids']) if isinstance(row['source_ids'], list) else 0, axis=1).sum()\n",
    "rows_with_no_sources = filtered_mashups['source_ids'].isna().sum()\n",
    "\n",
    "print(f\"Original total sources: {original_total_sources}\")\n",
    "print(f\"Filtered total sources: {filtered_total_sources}\")\n",
    "print(f\"Sources removed: {original_total_sources - filtered_total_sources}\")\n",
    "print(f\"Rows with all sources removed (now None): {rows_with_no_sources}\")\n",
    "print(f\"\\nOriginal dataframe size: {len(only_valid_mashups)}\")\n",
    "print(f\"Filtered dataframe size: {len(filtered_mashups)}\")\n",
    "\n",
    "# Drop source_ids and source_votes columns and reset index\n",
    "filtered_mashups = filtered_mashups.drop(columns=['is_valid_sample', 'is_valid_mashup']).reset_index(drop=True)\n",
    "print(f\"\\nFinal dataframe shape: {filtered_mashups.shape}\")\n",
    "filtered_mashups.head()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "10",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Convert to list of dicts and save as JSONL\n",
    "filtered_mashups_list = filtered_mashups.to_dict('records')\n",
    "print(f\"Total records to save: {len(filtered_mashups_list)}\")\n",
    "\n",
    "# Save to JSONL\n",
    "output_path = \"/home/sara/task_data/filtered_whosampled_mashups.jsonl\"\n",
    "write_jsonl(filtered_mashups_list, output_path)\n",
    "print(f\"Saved to: {output_path}\")\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "11",
   "metadata": {},
   "outputs": [],
   "source": [
    "all_other_mashups = read_jsonl(\"/home/sara/task_data/cleaned_mashup_data_wout_ws.jsonl\")\n",
    "print(len(all_other_mashups))\n",
    "all_other_mashups_df = pd.DataFrame(all_other_mashups)\n",
    "all_other_mashups_df.head()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "12",
   "metadata": {},
   "outputs": [],
   "source": [
    "all_other_mashups_df.to_csv(\"/home/sara/task_data/cleaned_mashup_data_wout_ws.csv\", index=False)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "13",
   "metadata": {},
   "outputs": [],
   "source": [
    "only_valid_samples = full_df[full_df['is_valid_sample'] == True].copy()\n",
    "print(len(only_valid_samples))\n",
    "only_valid_samples.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "14",
   "metadata": {},
   "outputs": [],
   "source": [
    "only_valid_samples['source_id'] = only_valid_samples['source_ids'].str[0]\n",
    "only_valid_samples['source_vote'] = only_valid_samples['source_votes'].str[0]\n",
    "\n",
    "only_valid_samples.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "15",
   "metadata": {},
   "outputs": [],
   "source": [
    "cleaned_samples = only_valid_samples.drop(columns=['is_valid_sample', 'is_valid_mashup', 'source_ids', 'source_votes']).reset_index(drop=True)\n",
    "cleaned_samples.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "16",
   "metadata": {},
   "outputs": [],
   "source": [
    "# Format output_id and source_id as YouTube URLs\n",
    "cleaned_samples['output_url'] = 'https://www.youtube.com/watch?v=' + cleaned_samples['output_id']\n",
    "cleaned_samples['source_url'] = 'https://www.youtube.com/watch?v=' + cleaned_samples['source_id']\n",
    "\n",
    "cleaned_samples[['output_id', 'output_url', 'source_id', 'source_url']].head()\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "17",
   "metadata": {},
   "outputs": [],
   "source": [
    "write_jsonl(cleaned_samples.to_dict('records'), \"/home/sara/task_data/final_whosampled_samples_11_12.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "18",
   "metadata": {},
   "outputs": [],
   "source": [
    "x = read_jsonl(\"/home/sara/task_data/cleaned_mashup_data_wout_ws.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "19",
   "metadata": {},
   "outputs": [],
   "source": [
    "sources = set()\n",
    "\n",
    "discogs = []\n",
    "\n",
    "for i in x:\n",
    "    sources.add(i['data_source'])\n",
    "    if i['data_source'] == 'discogs':\n",
    "        discogs.append(i)\n",
    "\n",
    "print(sources)\n",
    "print(len(discogs))\n",
    "\n",
    "write_jsonl(discogs, \"/home/sara/task_data/discogs_mashups_cleaned.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "20",
   "metadata": {},
   "outputs": [],
   "source": [
    "discogs[0]"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "21",
   "metadata": {},
   "outputs": [],
   "source": [
    "filtered_mashups.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "22",
   "metadata": {},
   "outputs": [],
   "source": [
    "filtered_mashups['output_url'] = 'https://www.youtube.com/watch?v=' + filtered_mashups['output_id']\n",
    "filtered_mashups['source_urls'] = filtered_mashups['source_ids'].apply(lambda ids: ['https://www.youtube.com/watch?v=' + id for id in ids] if isinstance(ids, list) else None)\n",
    "filtered_mashups.head()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "23",
   "metadata": {},
   "outputs": [],
   "source": [
    "write_jsonl(filtered_mashups.to_dict('records'), \"/home/sara/task_data/final_whosampled_mashups_11_12.jsonl\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "24",
   "metadata": {},
   "outputs": [],
   "source": [
    "x = read_jsonl(\"/home/sara/task_data/final_whosampled_samples_11_12.jsonl\")\n",
    "len(x)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "25",
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_clean",
   "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.10.15"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
