{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 78,
   "metadata": {},
   "outputs": [],
   "source": [
    "# corruptions include\n",
    "# lowpass\n",
    "# highpass\n",
    "# stereo -> mono\n",
    "# highsehlf\n",
    "# lowshelf\n",
    "# white noise\n",
    "import os\n",
    "import torch\n",
    "import random\n",
    "import torchaudio\n",
    "import numpy as np\n",
    "import scipy.signal as signal\n",
    "from typing import Tuple\n",
    "\n",
    "def biquad(\n",
    "    gain_db: float,\n",
    "    cutoff_freq: float,\n",
    "    q_factor: float,\n",
    "    sample_rate: float,\n",
    "    filter_type: str,\n",
    ") -> Tuple[np.ndarray, np.ndarray]:\n",
    "    \"\"\"Use design parameters to generate coefficients for a specific filter type.\"\"\"\n",
    "    A = 10 ** (gain_db / 40.0)\n",
    "    w0 = 2.0 * np.pi * (cutoff_freq / sample_rate)\n",
    "    alpha = np.sin(w0) / (2.0 * q_factor)\n",
    "    cos_w0 = np.cos(w0)\n",
    "    sqrt_A = np.sqrt(A)\n",
    "\n",
    "    if filter_type == \"high_shelf\":\n",
    "        b0 = A * ((A + 1) + (A - 1) * cos_w0 + 2 * sqrt_A * alpha)\n",
    "        b1 = -2 * A * ((A - 1) + (A + 1) * cos_w0)\n",
    "        b2 = A * ((A + 1) + (A - 1) * cos_w0 - 2 * sqrt_A * alpha)\n",
    "        a0 = (A + 1) - (A - 1) * cos_w0 + 2 * sqrt_A * alpha\n",
    "        a1 = 2 * ((A - 1) - (A + 1) * cos_w0)\n",
    "        a2 = (A + 1) - (A - 1) * cos_w0 - 2 * sqrt_A * alpha\n",
    "    elif filter_type == \"low_shelf\":\n",
    "        b0 = A * ((A + 1) - (A - 1) * cos_w0 + 2 * sqrt_A * alpha)\n",
    "        b1 = 2 * A * ((A - 1) - (A + 1) * cos_w0)\n",
    "        b2 = A * ((A + 1) - (A - 1) * cos_w0 - 2 * sqrt_A * alpha)\n",
    "        a0 = (A + 1) + (A - 1) * cos_w0 + 2 * sqrt_A * alpha\n",
    "        a1 = -2 * ((A - 1) + (A + 1) * cos_w0)\n",
    "        a2 = (A + 1) + (A - 1) * cos_w0 - 2 * sqrt_A * alpha\n",
    "    elif filter_type == \"peaking\":\n",
    "        b0 = 1 + alpha * A\n",
    "        b1 = -2 * cos_w0\n",
    "        b2 = 1 - alpha * A\n",
    "        a0 = 1 + alpha / A\n",
    "        a1 = -2 * cos_w0\n",
    "        a2 = 1 - alpha / A\n",
    "\n",
    "    b = np.array([b0, b1, b2]) / a0\n",
    "    a = np.array([1.0, a1 / a0, a2 / a0])\n",
    "    return b, a\n",
    "\n",
    "\n",
    "def apply_stereo_to_mono(audio: torch.Tensor, sample_rate: float):\n",
    "    return audio.mean(dim=0, keepdims=True).repeat(2, 1)\n",
    "\n",
    "\n",
    "def apply_channel_imbalance(\n",
    "    audio: torch.Tensor, sample_rate: float, imbalance: float = 0.0\n",
    "):\n",
    "    if not -1 <= imbalance <= 1 or audio.shape[-2] != 2:\n",
    "        raise ValueError(\"Invalid input\")\n",
    "    out = audio.clone()\n",
    "    l_gain, r_gain = (1.0 - imbalance, 1.0) if imbalance > 0 else (1.0, 1.0 + imbalance)\n",
    "    out[0, :], out[1, :] = out[0, :] * l_gain, out[1, :] * r_gain\n",
    "    return out\n",
    "\n",
    "\n",
    "def apply_highpass(audio: torch.Tensor, sample_rate: float, cutoff_hz: float = 1000.0):\n",
    "    return torchaudio.functional.highpass_biquad(audio, sample_rate, cutoff_hz)\n",
    "\n",
    "\n",
    "def apply_lowpass(audio: torch.Tensor, sample_rate: float, cutoff_hz: float = 1000.0):\n",
    "    return torchaudio.functional.lowpass_biquad(audio, sample_rate, cutoff_hz)\n",
    "\n",
    "\n",
    "def apply_noise(\n",
    "    audio: torch.Tensor,\n",
    "    sample_rate: float,\n",
    "    gain_db: float = 0.0,\n",
    "    noise_type: str = \"white\",\n",
    "):\n",
    "    gain_lin = 10 ** (gain_db / 20.0)\n",
    "    noise = torch.randn_like(audio)\n",
    "\n",
    "    if noise_type == \"white\":\n",
    "        return audio + gain_lin * noise\n",
    "    elif noise_type == \"pink\":\n",
    "        b = torch.tensor([0.049922035, -0.095993537, 0.050612699, -0.004408786])\n",
    "        a = torch.tensor([1, -2.494956002, 2.017265875, -0.522189400])\n",
    "        noise = torchaudio.functional.filtfilt(noise, a, b)\n",
    "        noise /= noise.abs().max()\n",
    "        return audio + gain_lin * noise\n",
    "    else:\n",
    "        raise ValueError(f\"Invalid noise type: {noise_type}\")\n",
    "\n",
    "def apply_shelving_filter(\n",
    "    audio: torch.Tensor, \n",
    "    sample_rate: float, \n",
    "    gain_db: float, \n",
    "    cutoff_freq: float, \n",
    "    q_factor: float, \n",
    "    filter_type: str\n",
    "):\n",
    "    # convert x to numpy\n",
    "    audio = audio.numpy()\n",
    "    b, a = biquad(\n",
    "        gain_db,\n",
    "        cutoff_freq,\n",
    "        q_factor,\n",
    "        sample_rate,\n",
    "        filter_type,\n",
    "    )\n",
    "    x = signal.lfilter(b, a, audio).astype(np.float32)\n",
    "    return torch.from_numpy(x)\n",
    "\n",
    "\n",
    "# randomized corrputions\n",
    "\n",
    "def apply_random_noise(audio: torch.Tensor, sample_rate: float):\n",
    "    noise_type = random.choice([\"white\", \"pink\"])\n",
    "    if noise_type == \"white\":\n",
    "        noise_gain = random.uniform(-96, -48)\n",
    "    else:\n",
    "        noise_gain = random.uniform(-48, -12)\n",
    "    return apply_noise(audio, sample_rate, noise_gain, noise_type)\n",
    "\n",
    "def apply_random_stereo_to_mono(audio: torch.Tensor, sample_rate: float):\n",
    "    return apply_stereo_to_mono(audio, sample_rate)\n",
    "\n",
    "def apply_random_channel_imbalance(audio: torch.Tensor, sample_rate: float):\n",
    "    imbalance = random.uniform(-1.0, 1.0)\n",
    "    return apply_channel_imbalance(audio, sample_rate, imbalance)\n",
    "\n",
    "def apply_random_filter(audio: torch.Tensor, sample_rate: float):\n",
    "    filter_type = random.choice([\"highpass\", \"lowpass\", \"high_shelf\", \"low_shelf\"])\n",
    "    if filter_type == \"highpass\":\n",
    "        cutoff_freq = random.uniform(20, 4000)\n",
    "        return apply_highpass(audio, sample_rate, cutoff_freq)\n",
    "    elif filter_type == \"lowpass\":\n",
    "        cutoff_freq = random.uniform(1000, 16000)\n",
    "        return apply_lowpass(audio, sample_rate, cutoff_freq)\n",
    "    else:\n",
    "        gain_db = random.uniform(-12, 12)\n",
    "        if filter_type == \"high_shelf\":\n",
    "            cutoff_freq = random.uniform(6000, 20000)\n",
    "        else:\n",
    "            cutoff_freq = random.uniform(20, 2000)\n",
    "        q_factor = random.uniform(0.1, 10.0)\n",
    "        return apply_shelving_filter(audio, sample_rate, gain_db, cutoff_freq, q_factor, filter_type)\n",
    "\n",
    "def corrupt(waveform_tensor, sample_rate):\n",
    "    \"\"\"\n",
    "    Apply a random number of corruptions (at least one) to the input waveform.\n",
    "    \"\"\"\n",
    "    # List of available random corruption functions\n",
    "    corruption_fns = [\n",
    "        apply_random_noise,\n",
    "        apply_random_stereo_to_mono,\n",
    "        apply_random_channel_imbalance,\n",
    "        apply_random_filter,\n",
    "    ]\n",
    "    n_corr = random.randint(1, len(corruption_fns))  # at least one\n",
    "    selected = random.sample(corruption_fns, n_corr)\n",
    "    print(selected)\n",
    "    out = waveform_tensor.clone()\n",
    "    for fn in selected:\n",
    "        out = fn(out, sample_rate)\n",
    "    return torch.tanh(out)\n",
    "    "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "import IPython\n",
    "\n",
    "# load some audio and test this\n",
    "audio, sample_rate = torchaudio.load(\"/home/christian/audio/reference-audio-wav/02 Dreams.wav\")\n",
    "audio = audio[:, :10*sample_rate]\n",
    "\n",
    "out = corrupt(audio, sample_rate)\n",
    "#out = torch.tanh(apply_random_noise(audio, sample_rate))\n",
    "\n",
    "IPython.display.display(IPython.display.Audio(audio, rate=sample_rate, normalize=False))\n",
    "IPython.display.display(IPython.display.Audio(out, rate=sample_rate, normalize=False))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 83,
   "metadata": {},
   "outputs": [],
   "source": [
    "from suno_utils.audio import Audio\n"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "audio_obj = Audio.from_array_float(audio.numpy(), sample_rate)\n",
    "audio_obj.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 106,
   "metadata": {},
   "outputs": [],
   "source": [
    "audio_obj = audio_obj.convert(\n",
    "    sample_rate=48_000, byte_width=2, n_channels=2\n",
    ").normalize_volume(target_db=-32)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "audio_obj.play()"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": []
  }
 ],
 "metadata": {
  "kernelspec": {
   "display_name": "suno_env",
   "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.9"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
