{
 "cells": [
  {
   "cell_type": "code",
   "execution_count": 2,
   "metadata": {},
   "outputs": [],
   "source": [
    "import sys\n",
    "sys.path.append('/home/christian_c')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 3,
   "metadata": {},
   "outputs": [],
   "source": [
    "import os\n",
    "import time\n",
    "import torch\n",
    "import random\n",
    "import numpy as np\n",
    "from suno_utils.audio import Audio\n",
    "import torchaudio\n",
    "import IPython.display as ipd\n",
    "from tqdm import tqdm\n",
    "import matplotlib.pyplot as plt\n",
    "from suno_seal import SunoSeal"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 4,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams[\"figure.figsize\"] = (20,3)\n",
    "os.environ[\"CUDA_VISIBLE_DEVICES\"] = \"1\"\n",
    "os.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\""
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 5,
   "metadata": {},
   "outputs": [],
   "source": [
    "def load_sample_audio(path):\n",
    "    audio = Audio.from_file(path)\n",
    "    sample_rate = audio.sample_rate\n",
    "    wav = torch.from_numpy(audio.array_float)\n",
    "    return wav, sample_rate\n",
    "\n",
    "def plot_waveform_and_specgram(waveform, sample_rate, title):\n",
    "    waveform = waveform.squeeze().detach().cpu().numpy()\n",
    "\n",
    "    num_frames = waveform.shape[-1]\n",
    "    time_axis = torch.arange(0, num_frames) / sample_rate\n",
    "\n",
    "    figure, (ax1, ax2) = plt.subplots(1, 2)\n",
    "\n",
    "    ax1.plot(time_axis, waveform, linewidth=1)\n",
    "    ax1.grid(True)\n",
    "    ax2.specgram(waveform, Fs=sample_rate)\n",
    "\n",
    "    figure.suptitle(f\"{title} - Waveform and specgram\")\n",
    "    plt.show()\n",
    "\n",
    "def play_audio(waveform, sample_rate):\n",
    "    waveform = waveform.unsqueeze(0).detach().cpu().numpy()\n",
    "\n",
    "    num_channels, *_ = waveform.shape\n",
    "    if num_channels == 1:\n",
    "        ipd.display(ipd.Audio(waveform[0], rate=sample_rate))\n",
    "    elif num_channels == 2:\n",
    "        ipd.display(ipd.Audio((waveform[0], waveform[1]), rate=sample_rate))\n",
    "    else:\n",
    "        raise ValueError(\"Waveform with more than 2 channels are not supported.\")"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": 6,
   "metadata": {},
   "outputs": [],
   "source": [
    "class AudioDataset(torch.utils.data.Dataset):\n",
    "    def __init__(self, data_dir: str):\n",
    "        self.data_dir = data_dir\n",
    "        self.samps_list = os.listdir(data_dir)\n",
    "\n",
    "    def __len__(self):\n",
    "        return len(self.samps_list)\n",
    "    \n",
    "    def __getitem__(self, idx):\n",
    "        sample = self.samps_list[idx]\n",
    "        filepath = os.path.join(self.data_dir, sample)\n",
    "        wav, sr = load_sample_audio(filepath)\n",
    "        sample_length = len(wav) / sr\n",
    "        return wav, sr, sample, sample_length"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Evaluate the Generator. Speed, memory, mel spectrogram + audio indistinguishability"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "First just a few suno samples"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "samps_dir = '/home/christian_c/christian_c/samps/30b_v7'\n",
    "output_dir = '/home/christian_c/christian_c/wm_test'\n",
    "generator = SunoSeal.load_generator('/home/christian_c/christian_c/wm_models/checkpoint_generator_fixed_augs_no_speed_minibatch_epoch_125.pth', nbits=16)\n",
    "\n",
    "for sample in tqdm(os.listdir(samps_dir)[:5]):\n",
    "    filepath = os.path.join(samps_dir, sample)\n",
    "    audio, sr = load_sample_audio(filepath)\n",
    "    plot_waveform_and_specgram(audio, sr, title=f'Non-watermarked audio: {sample}')\n",
    "    play_audio(audio, sr)\n",
    "    unsqueezed_audio = audio.unsqueeze(0).unsqueeze(0)\n",
    "    watermark = generator.get_watermark(unsqueezed_audio, sample_rate=sr)\n",
    "    watermarked_audio = (watermark + unsqueezed_audio).squeeze()\n",
    "    plot_waveform_and_specgram(watermarked_audio, sr, title=f'Watermarked audio: {sample}')\n",
    "    play_audio(watermarked_audio, sr)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Also a few normal music high-quality samples for good measure"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "samps = ['01 J.S. Bach Suite No.1, S.1007, G major - I. Prelude.wav', '02 Dreams.wav', '02 Freddie Freeloader.wav', \"10 Blue Ridge Mountains.wav\", \"11 High Hopes.wav\"]\n",
    "reference_audio_dir = '/home/christian/audio/reference-audio-wav'\n",
    "hq_samps = [os.path.join(reference_audio_dir, file) for file in samps]\n",
    "for file in hq_samps:\n",
    "    print(os.path.exists(file))"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "for sample in hq_samps:\n",
    "    audio, sr = load_sample_audio(sample)\n",
    "    plot_waveform_and_specgram(audio, sr, title=f'Non-watermarked audio: {sample}')\n",
    "    play_audio(audio, sr)\n",
    "    unsqueezed_audio = audio.unsqueeze(0).unsqueeze(0)\n",
    "    print(sr)\n",
    "    watermark = generator.get_watermark(unsqueezed_audio, sample_rate=sr)\n",
    "    watermarked_audio = (watermark + unsqueezed_audio).squeeze()\n",
    "    watermarked_audio = Audio.from_array_float(watermarked_audio, sample_rate=sr[i].item())\n",
    "    watermarked_audio.write_wav(output_filename)\n",
    "    plot_waveform_and_specgram(watermarked_audio, sr, title=f'Watermarked audio: {sample}')\n",
    "    play_audio(watermarked_audio, sr)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Move to GPU for speed test"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "samps_dir = '/home/christian_c/christian_c/samps/v13'\n",
    "wm_dir = '/home/christian_c/christian_c/sunoseal_wm_test'\n",
    "\n",
    "sample_audio_dataset = AudioDataset(samps_dir)\n",
    "sample_audio_dataloader = torch.utils.data.DataLoader(sample_audio_dataset, batch_size=1, num_workers=4, shuffle=False)\n",
    "\n",
    "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
    "generator = AudioSeal.load_generator('/home/christian_c/christian_c/wm_models/checkpoint_generator_fixed_augs_no_speed_minibatch_epoch_50.pth', nbits=16)\n",
    "generator = generator.to(device)\n",
    "\n",
    "lengths = []\n",
    "watermark_times = []\n",
    "\n",
    "for batch_idx, (audio, sr, filename, sample_length) in enumerate(tqdm(sample_audio_dataloader)):\n",
    "    with torch.no_grad():\n",
    "        samp_length = round(sample_length.item(), 2)\n",
    "        lengths.append(samp_length)\n",
    "        start_time = time.time()\n",
    "        audio = audio.to(device)\n",
    "        output = generator(audio.unsqueeze(0), sample_rate=sr.item())\n",
    "        end_time = time.time()\n",
    "        watermark_times.append(end_time - start_time)\n",
    "        watermarked_audio = output.cpu().numpy().squeeze()\n",
    "        output_filename = os.path.join(wm_dir, f'watermarked_{filename[0]}')\n",
    "        watermarked_audio = Audio.from_array_float(watermarked_audio, sample_rate=sr.item())\n",
    "        watermarked_audio.write_wav(output_filename)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams[\"figure.figsize\"] = (6.4,4.8)\n",
    "plt.scatter(lengths, watermark_times)\n",
    "plt.xlabel('Sample length (s)')\n",
    "plt.ylabel('Watermark time (s)')\n",
    "plt.title('SunoSeal Watermark Time Scales Linearly w/ Sample Length')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Batching (for bulk watermarking if desired)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "samps_dir = '/home/christian_c/christian_c/samps/v13'\n",
    "wm_dir = '/home/christian_c/christian_c/wm_test'\n",
    "\n",
    "sample_audio_dataset = AudioDataset(samps_dir)\n",
    "sample_audio_dataloader = torch.utils.data.DataLoader(sample_audio_dataset, batch_size=16, num_workers=4, shuffle=False)\n",
    "\n",
    "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
    "generator = generator.to(device)\n",
    "\n",
    "for batch_idx, (batch, sr, filenames) in enumerate(tqdm(sample_audio_dataloader)):\n",
    "    with torch.no_grad():\n",
    "        batch = batch.to(device)\n",
    "        outputs = generator(batch, sample_rate=sr[0].item())\n",
    "\n",
    "        for i in range(batch.size(0)):\n",
    "            watermarked_audio = outputs[i].cpu().numpy().squeeze()\n",
    "            output_filename = os.path.join(wm_dir, f'watermarked_{filenames[i]}')\n",
    "            watermarked_audio = Audio.from_array_float(watermarked_audio, sample_rate=sr[i].item())\n",
    "            watermarked_audio.write_wav(output_filename)\n",
    "        \n",
    "        \n",
    "\n",
    "for sample in tqdm(os.listdir(samps_dir)):\n",
    "    filepath = os.path.join(samps_dir, sample)\n",
    "    audio, sr = load_sample_audio(filepath)\n",
    "    plot_waveform_and_specgram(audio, sr, title=f'Non-watermarked audio: {sample}')\n",
    "    play_audio(audio, sr)\n",
    "    unsqueezed_audio = audio.unsqueeze(0).unsqueeze(0)\n",
    "    watermark = generator.get_watermark(unsqueezed_audio, sample_rate=sr)\n",
    "    watermarked_audio = (watermark + unsqueezed_audio).squeeze()\n",
    "    plot_waveform_and_specgram(watermarked_audio, sr, title=f'Watermarked audio: {sample}')\n",
    "    play_audio(watermarked_audio, sr)\n",
    "    "
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### Evaluating the Detector. Speed (can't batch here), confusion matrix"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "samps_dir = '/home/christian_c/christian_c/samps/v13'\n",
    "\n",
    "sample_audio_dataset = AudioDataset(samps_dir)\n",
    "sample_audio_dataloader = torch.utils.data.DataLoader(sample_audio_dataset, batch_size=1, num_workers=4, shuffle=False)\n",
    "\n",
    "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
    "detector = AudioSeal.load_detector('/home/christian_c/christian_c/wm_models/checkpoint_detector_fixed_augs_no_speed_minibatch_epoch_50.pth', nbits=16)\n",
    "detector = detector.to(device)\n",
    "\n",
    "lengths = []\n",
    "non_watermarked_detect_times = []\n",
    "non_watermarked_results = []\n",
    "\n",
    "for batch_idx, (audio, sr, filename, sample_length) in enumerate(tqdm(sample_audio_dataloader)):\n",
    "    with torch.no_grad():\n",
    "        samp_length = round(sample_length.item(), 2)\n",
    "        lengths.append(samp_length)\n",
    "        start_time = time.time()\n",
    "        audio = audio.to(device)\n",
    "        output_prob, msg = detector(audio.unsqueeze(0), sample_rate=sr.item())\n",
    "        detected = (\n",
    "            torch.count_nonzero(torch.gt(output_prob[:, 1, :], 0.5)) / output_prob.shape[-1]\n",
    "        )\n",
    "        detect_prob = detected.cpu().item() \n",
    "        end_time = time.time()\n",
    "        non_watermarked_detect_times.append(end_time - start_time)\n",
    "        non_watermarked_results.append(detect_prob)\n",
    "        "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "wm_dir = '/home/christian_c/christian_c/sunoseal_wm_test'\n",
    "\n",
    "wm_audio_dataset = AudioDataset(wm_dir)\n",
    "wm_audio_dataloader = torch.utils.data.DataLoader(wm_audio_dataset, batch_size=1, num_workers=4, shuffle=False)\n",
    "\n",
    "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
    "detector = SunoSeal.load_detector('/home/christian_c/christian_c/wm_models/checkpoint_detector_fixed_augs_no_speed_minibatch_epoch_125.pth', nbits=16)\n",
    "detector = detector.to(device)\n",
    "\n",
    "lengths = []\n",
    "watermarked_detect_times = []\n",
    "watermarked_results = []\n",
    "\n",
    "for batch_idx, (audio, sr, filename, sample_length) in enumerate(tqdm(wm_audio_dataloader)):\n",
    "    with torch.no_grad():\n",
    "        samp_length = round(sample_length.item(), 2)\n",
    "        lengths.append(samp_length)\n",
    "        start_time = time.time()\n",
    "        audio = audio.to(device)\n",
    "        output_prob, msg = detector(audio.unsqueeze(0), sample_rate=sr.item())\n",
    "        detected = (\n",
    "            torch.count_nonzero(torch.gt(output_prob[:, 1, :], 0.5)) / output_prob.shape[-1]\n",
    "        )\n",
    "        detect_prob = detected.cpu().item() \n",
    "        end_time = time.time()\n",
    "        watermarked_detect_times.append(end_time - start_time)\n",
    "        watermarked_results.append(detect_prob)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "labels = ['Non-watermarked', 'Watermarked']\n",
    "detector = AudioSeal.load_detector(('audioseal_detector_16bits'))\n",
    "non_watermarked_results = []\n",
    "watermarked_results = []\n",
    "# evaluate non-watermarked\n",
    "for sample in tqdm(os.listdir(samps_dir)):\n",
    "    filepath = os.path.join(samps_dir, sample)\n",
    "    audio, sr = load_sample_audio(filepath)\n",
    "    non_watermarked_result = detector.detect_watermark(audio.unsqueeze(0).unsqueeze(0), sr)\n",
    "    non_watermarked_results.append(non_watermarked_result)\n",
    "# evaluate watermarked\n",
    "for watermarked_sample in tqdm(os.listdir(wm_dir)):\n",
    "    filepath = os.path.join(wm_dir, watermarked_sample)\n",
    "    audio, sr = load_sample_audio(filepath)\n",
    "    watermarked_result = detector.detect_watermark(watermarked_audio.unsqueeze(0).unsqueeze(0), sr)\n",
    "    watermarked_results.append(watermarked_result)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams[\"figure.figsize\"] = (6.4,4.8)\n",
    "plt.scatter(lengths, non_watermarked_detect_times)\n",
    "plt.xlabel('Sample length (s)')\n",
    "plt.ylabel('Detect time (s)')\n",
    "plt.title('Detect Time for Non-watermarked Samples')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams[\"figure.figsize\"] = (6.4,4.8)\n",
    "plt.scatter(lengths, watermarked_detect_times)\n",
    "plt.xlabel('Sample length (s)')\n",
    "plt.ylabel('Detect time (s)')\n",
    "plt.title('Detect Time for Watermarked Samples')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams[\"figure.figsize\"] = (6.4,4.8)\n",
    "plt.hist(non_watermarked_results)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.hist(watermarked_results)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from sklearn.metrics import ConfusionMatrixDisplay\n",
    "labels = ['Non-watermarked', 'Watermarked']\n",
    "non_watermarked_labels = np.array([round(prob) for prob in non_watermarked_results])\n",
    "watermarked_labels = np.array([round(prob) for prob in watermarked_results])\n",
    "\n",
    "true_labels = np.concatenate((np.zeros(len(non_watermarked_labels)), np.ones(len(watermarked_labels))))\n",
    "pred_labels = np.concatenate((non_watermarked_labels, watermarked_labels))\n",
    "\n",
    "ConfusionMatrixDisplay.from_predictions(true_labels, pred_labels, normalize='true', display_labels=labels)"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "### What if we only detect part of the sample? Is performance still as good?"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "class RandomCroppedAudioDataset(torch.utils.data.Dataset):\n",
    "    def __init__(self, data_dir: str, crop_length = float):\n",
    "        self.data_dir = data_dir\n",
    "        self.samps_list = os.listdir(data_dir)\n",
    "        self.crop_length = crop_length # in seconds\n",
    "\n",
    "    def __len__(self):\n",
    "        return len(self.samps_list)\n",
    "    \n",
    "    def __getitem__(self, idx):\n",
    "        sample = self.samps_list[idx]\n",
    "        filepath = os.path.join(self.data_dir, sample)\n",
    "        wav, sr = load_sample_audio(filepath)\n",
    "        actual_crop = self.crop_length * sr # convert from s\n",
    "        start_ix = random.randint(0, len(wav) - actual_crop - 1)\n",
    "        wav = wav[start_ix: start_ix + actual_crop]\n",
    "        sample_length = len(wav) / sr\n",
    "        wav = wav.unsqueeze(0)\n",
    "        return wav, sr, sample, sample_length"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "samps_dir = '/home/christian_c/christian_c/samps/v13'\n",
    "\n",
    "cropped_sample_audio_dataset = RandomCroppedAudioDataset(samps_dir, 60)\n",
    "cropped_sample_audio_dataloader = torch.utils.data.DataLoader(cropped_sample_audio_dataset, batch_size=1, num_workers=4, shuffle=False)\n",
    "\n",
    "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
    "detector = SunoSeal.load_detector('/home/christian_c/christian_c/wm_models/checkpoint_detector_fixed_augs_no_speed_minibatch_epoch_125.pth', nbits=16)\n",
    "detector = detector.to(device)\n",
    "\n",
    "lengths = []\n",
    "cropped_non_watermarked_detect_times = []\n",
    "cropped_non_watermarked_results = []\n",
    "\n",
    "for batch_idx, (audio, sr, filename, sample_length) in enumerate(tqdm(cropped_sample_audio_dataloader)):\n",
    "    with torch.no_grad():\n",
    "        samp_length = round(sample_length.item(), 2)\n",
    "        lengths.append(samp_length)\n",
    "        start_time = time.time()\n",
    "        audio = audio.to(device)\n",
    "        output_prob, msg = detector(audio, sample_rate=sr.item())\n",
    "        detected = (\n",
    "            torch.count_nonzero(torch.gt(output_prob[:, 1, :], 0.5)) / output_prob.shape[-1]\n",
    "        )\n",
    "        detect_prob = detected.cpu().item() \n",
    "        end_time = time.time()\n",
    "        cropped_non_watermarked_detect_times.append(end_time - start_time)\n",
    "        cropped_non_watermarked_results.append(detect_prob)\n",
    "        "
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "wm_dir = '/home/christian_c/christian_c/sunoseal_wm_test'\n",
    "\n",
    "cropped_wm_audio_dataset = RandomCroppedAudioDataset(wm_dir, 60)\n",
    "cropped_wm_audio_dataloader = torch.utils.data.DataLoader(cropped_wm_audio_dataset, batch_size=1, num_workers=4, shuffle=False)\n",
    "\n",
    "device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n",
    "detector = AudioSeal.load_detector('/home/christian_c/christian_c/wm_models/checkpoint_detector_fixed_augs_no_speed_minibatch_epoch_125.pth', nbits=16)\n",
    "detector = detector.to(device)\n",
    "\n",
    "lengths = []\n",
    "cropped_watermarked_detect_times = []\n",
    "cropped_watermarked_results = []\n",
    "\n",
    "for batch_idx, (audio, sr, filename, sample_length) in enumerate(tqdm(cropped_wm_audio_dataloader)):\n",
    "    with torch.no_grad():\n",
    "        samp_length = round(sample_length.item(), 2)\n",
    "        lengths.append(samp_length)\n",
    "        start_time = time.time()\n",
    "        audio = audio.to(device)\n",
    "        output_prob, msg = detector(audio, sample_rate=sr.item())\n",
    "        detected = (\n",
    "            torch.count_nonzero(torch.gt(output_prob[:, 1, :], 0.5)) / output_prob.shape[-1]\n",
    "        )\n",
    "        detect_prob = detected.cpu().item() \n",
    "        end_time = time.time()\n",
    "        cropped_watermarked_detect_times.append(end_time - start_time)\n",
    "        cropped_watermarked_results.append(detect_prob)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "labels = ['Non-watermarked', 'Watermarked']\n",
    "detector = AudioSeal.load_detector(('audioseal_detector_16bits'))\n",
    "non_watermarked_results = []\n",
    "watermarked_results = []\n",
    "# evaluate non-watermarked\n",
    "for sample in tqdm(os.listdir(samps_dir)):\n",
    "    filepath = os.path.join(samps_dir, sample)\n",
    "    audio, sr = load_sample_audio(filepath)\n",
    "    non_watermarked_result = detector.detect_watermark(audio.unsqueeze(0).unsqueeze(0), sr)\n",
    "    non_watermarked_results.append(non_watermarked_result)\n",
    "# evaluate watermarked\n",
    "for watermarked_sample in tqdm(os.listdir(wm_dir)):\n",
    "    filepath = os.path.join(wm_dir, watermarked_sample)\n",
    "    audio, sr = load_sample_audio(filepath)\n",
    "    watermarked_result = detector.detect_watermark(watermarked_audio.unsqueeze(0).unsqueeze(0), sr)\n",
    "    watermarked_results.append(watermarked_result)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams[\"figure.figsize\"] = (6.4,4.8)\n",
    "plt.scatter(lengths, cropped_non_watermarked_detect_times)\n",
    "plt.xlabel('Sample length (s)')\n",
    "plt.ylabel('Detect time (s)')\n",
    "plt.title('Detect Time for Non-watermarked Samples')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams[\"figure.figsize\"] = (6.4,4.8)\n",
    "plt.scatter(lengths, cropped_watermarked_detect_times)\n",
    "plt.xlabel('Sample length (s)')\n",
    "plt.ylabel('Detect time (s)')\n",
    "plt.title('Detect Time for Watermarked Samples')"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.rcParams[\"figure.figsize\"] = (6.4,4.8)\n",
    "plt.hist(cropped_non_watermarked_results)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "plt.hist(cropped_watermarked_results)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "from sklearn.metrics import ConfusionMatrixDisplay\n",
    "labels = ['Non-watermarked', 'Watermarked']\n",
    "cropped_non_watermarked_labels = np.array([round(prob) for prob in cropped_non_watermarked_results])\n",
    "cropped_watermarked_labels = np.array([round(prob) for prob in cropped_watermarked_results])\n",
    "\n",
    "true_labels = np.concatenate((np.zeros(len(cropped_non_watermarked_labels)), np.ones(len(cropped_watermarked_labels))))\n",
    "pred_labels = np.concatenate((cropped_non_watermarked_labels, cropped_watermarked_labels))\n",
    "\n",
    "ConfusionMatrixDisplay.from_predictions(true_labels, pred_labels, normalize='true', display_labels=labels)\n",
    "plt.title('60-second Context Performance')"
   ]
  },
  {
   "cell_type": "markdown",
   "metadata": {},
   "source": [
    "Golden Sunshine Demo"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "filepath = '/home/christian_c/golden_sunshine.mp3'\n",
    "audio, sr = load_sample_audio(filepath)\n",
    "plot_waveform_and_specgram(audio, sr, title='Non-watermarked Golden Sunshine')\n",
    "play_audio(audio, sr)\n",
    "unsqueezed_audio = audio.unsqueeze(0).unsqueeze(0)\n",
    "watermark = generator.get_watermark(unsqueezed_audio, sample_rate=sr)\n",
    "watermark_audio = watermark.squeeze()\n",
    "watermarked_audio = (watermark + unsqueezed_audio).squeeze()\n",
    "play_audio(watermark_audio, sr)\n",
    "plot_waveform_and_specgram(watermarked_audio, sr, title='Watermarked Golden Sunshine')\n",
    "play_audio(watermarked_audio, sr)"
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "metadata": {},
   "outputs": [],
   "source": [
    "detector = SunoSeal.load_detector('/home/christian_c/christian_c/wm_models/checkpoint_detector_fixed_augs_no_speed_minibatch_epoch_125.pth', nbits=16)\n",
    "output_prob, msg = detector(audio.unsqueeze(0), sample_rate=sr.item())\n",
    "detected = (\n",
    "    torch.count_nonzero(torch.gt(output_prob[:, 1, :], 0.5)) / output_prob.shape[-1]\n",
    "    )\n",
    "detect_prob = detected.cpu().item() "
   ]
  }
 ],
 "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.14"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 2
}
