{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 1,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.s3 import read_from_s3\n",
    "from suno_utils.audio import Audio\n",
    "import numpy as np"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "base_path = \"s3://suno-data/christian/outputs/mountain-whisper-eclipse/VveQ3alEmXE/VveQ3alEmXE.mp3\"\n",
    "item_id = \"VveQ3alEmXE\" \n",
    "\n",
    "generated_mp3 = f\"{base_path}/{item_id}.mp3\"\n",
    "audio = Audio.from_s3(generated_mp3)\n",
    "audio.play()\n",
    "\n",
    "metadata = read_from_s3(f\"{base_path}/{item_id}_metadata.json\")\n",
    "print(metadata)\n",
    "\n",
    "gen_semantic = read_from_s3(f\"{base_path}/{item_id}_generated_semantic.npz\", read_f=np.load)[\"semantic_codes\"]\n",
    "print(gen_semantic.shape)\n",
    "\n",
    "original_semantic = read_from_s3(f\"{base_path}/original_semantic.npz\", read_f=np.load)[\"semantic_codes\"]\n",
    "print(original_semantic)\n",
    "\n",
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "import torch\n",
    "import math\n",
    "\n",
    "def append_zero(x):\n",
    "    return torch.cat([x, x.new_zeros([1])])\n",
    "\n",
    "def sigma_to_t(sigma):\n",
    "    return sigma.atan() / math.pi * 2\n",
    "\n",
    "def t_to_sigma(t):\n",
    "    return (t * math.pi / 2).tan()\n",
    "\n",
    "def get_sigmas_polyexponential(n, sigma_min, sigma_max, rho=1.0, device=\"cpu\"):\n",
    "    \"\"\"Constructs an polynomial in log sigma noise schedule.\"\"\"\n",
    "    ramp = torch.linspace(1, 0, n, device=device) ** rho\n",
    "    sigmas = torch.exp(ramp * (math.log(sigma_max) - math.log(sigma_min)) + math.log(sigma_min))\n",
    "    return append_zero(sigmas)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n",
    "steps = 3\n",
    "sigma_min = 0.5\n",
    "sigma_max = 100\n",
    "rho = 1.0\n",
    "\n",
    "sigmas = get_sigmas_polyexponential(steps, sigma_min, sigma_max, rho, device=\"cuda\")\n",
    "print(sigmas)\n",
    "\n",
    "t = sigma_to_t(sigmas)\n",
    "print(t)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "t = torch.tensor([0.999, 0.749, 0.499, 0.249])\n",
    "sigmas = t_to_sigma(t)\n",
    "print(sigmas)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os   \n",
    "\n",
    "def setup_datasets(data_dir: str, val_size: int = 100):\n",
    "    item_dirs = [\n",
    "        d for d in os.listdir(data_dir) if os.path.isdir(os.path.join(data_dir, d))\n",
    "    ]\n",
    "\n",
    "    # valid item dirs are those that contain both original and generated semantic files\n",
    "    item_ids = [\n",
    "        os.path.basename(d)\n",
    "        for d in item_dirs\n",
    "        if os.path.exists(os.path.join(data_dir, d, f\"{d}_original_semantic.npz\"))\n",
    "        and os.path.exists(os.path.join(data_dir, d, f\"{d}_generated_semantic.npz\"))\n",
    "    ]\n",
    "\n",
    "    # sort the item_ids by name\n",
    "    item_ids.sort()\n",
    "\n",
    "    # split into train and val\n",
    "    train_item_ids = item_ids[:-val_size]\n",
    "    val_item_ids = item_ids[-val_size:]\n",
    "\n",
    "    print(f\"Found {len(item_ids)} item ids in {data_dir}\")\n",
    "    print(f\"Splitting into {len(train_item_ids)} train and {len(val_item_ids)} val\")\n",
    "\n",
    "    return train_item_ids, val_item_ids\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# test loading everything into ram for now\n",
    "\n",
    "data_dir = \"/mnt/localdisk/christian\"\n",
    "\n",
    "train_item_ids, val_item_ids = setup_datasets(data_dir, val_size=100)\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "# load train set \n",
    "import torch\n",
    "from tqdm import tqdm\n",
    "\n",
    "examples = []\n",
    "for item_id in tqdm(train_item_ids):\n",
    "    # read from local disk\n",
    "    original_semantic = np.load(f\"{data_dir}/{item_id}/{item_id}_original_semantic.npz\")[\"semantic_codes\"]\n",
    "    generated_semantic = np.load(f\"{data_dir}/{item_id}/{item_id}_generated_semantic.npz\")[\"semantic_codes\"]\n",
    "    # convert to torch\n",
    "    original_semantic = torch.from_numpy(original_semantic).long()\n",
    "    generated_semantic = torch.from_numpy(generated_semantic).long()\n",
    "    examples.append((original_semantic, generated_semantic))\n",
    "\n",
    "\n",
    "print(len(examples))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.utils.s3 import list_s3_dir\n",
    "from suno_utils.utils.text import read_jsonl\n",
    "\n",
    "output_str = \"mountain-whisper-eclipse\"\n",
    "\n",
    "work_items = read_jsonl(\n",
    "    \"/home/christian/code/christian/metadata/discogs_subset_sampled_metas.jsonl\"\n",
    ")\n",
    "print(len(work_items))\n",
    "\n",
    "# first check for existing ids in the output path\n",
    "existing_ids = list_s3_dir(\n",
    "    f\"s3://suno-data/christian/outputs/{output_str}/\",\n",
    ")\n",
    "\n",
    "existing_ids = [os.path.dirname(f[0]).split(\"/\")[-1] for f in existing_ids]\n",
    "existing_ids = list(set(existing_ids))\n",
    "print(len(existing_ids))\n",
    "print(existing_ids[:3])\n",
    "\n",
    "# now remove these from the work items\n",
    "work_items = [f for f in work_items if f[\"id\"] not in existing_ids]\n",
    "print(len(work_items))\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "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": 2
}
