{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":70203,"databundleVersionId":8068726,"sourceType":"competition"},{"sourceId":8593024,"sourceType":"datasetVersion","datasetId":5140249},{"sourceId":8593456,"sourceType":"datasetVersion","datasetId":5140577},{"sourceId":8594194,"sourceType":"datasetVersion","datasetId":5141139},{"sourceId":8594854,"sourceType":"datasetVersion","datasetId":5141626},{"sourceId":8609139,"sourceType":"datasetVersion","datasetId":5151713},{"sourceId":8609854,"sourceType":"datasetVersion","datasetId":5152232},{"sourceId":8614441,"sourceType":"datasetVersion","datasetId":5155764},{"sourceId":8620686,"sourceType":"datasetVersion","datasetId":5160324},{"sourceId":8622637,"sourceType":"datasetVersion","datasetId":5161832},{"sourceId":8626967,"sourceType":"datasetVersion","datasetId":5164976},{"sourceId":8628204,"sourceType":"datasetVersion","datasetId":5165841}],"dockerImageVersionId":30715,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install /kaggle/input/einops-whl/einops-0.8.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:05:16.388646Z","iopub.execute_input":"2024-06-07T04:05:16.389053Z","iopub.status.idle":"2024-06-07T04:05:50.504680Z","shell.execute_reply.started":"2024-06-07T04:05:16.389020Z","shell.execute_reply":"2024-06-07T04:05:50.503195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nimport inspect\nfrom glob import glob\nimport pickle\nimport os\nfrom dataclasses import dataclass\nfrom datetime import datetime\n\nimport pandas as pd\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as pt_F\nimport torchaudio\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:05:54.342417Z","iopub.execute_input":"2024-06-07T04:05:54.342826Z","iopub.status.idle":"2024-06-07T04:05:58.779926Z","shell.execute_reply.started":"2024-06-07T04:05:54.342793Z","shell.execute_reply":"2024-06-07T04:05:58.778845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"start_dt = datetime.now()","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:05:58.781848Z","iopub.execute_input":"2024-06-07T04:05:58.782327Z","iopub.status.idle":"2024-06-07T04:05:58.787096Z","shell.execute_reply.started":"2024-06-07T04:05:58.782297Z","shell.execute_reply":"2024-06-07T04:05:58.786025Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SPECIES = [\n    'asbfly', 'ashdro1', 'ashpri1', 'ashwoo2', 'asikoe2', 'asiope1', 'aspfly1', 'aspswi1', 'barfly1', 'barswa',\n    'bcnher', 'bkcbul1', 'bkrfla1', 'bkskit1', 'bkwsti', 'bladro1', 'blaeag1', 'blakit1', 'blhori1', 'blnmon1',\n    'blrwar1', 'bncwoo3', 'brakit1', 'brasta1', 'brcful1', 'brfowl1', 'brnhao1', 'brnshr', 'brodro1', 'brwjac1',\n    'brwowl1', 'btbeat1', 'bwfshr1', 'categr', 'chbeat1', 'cohcuc1', 'comfla1', 'comgre', 'comior1', 'comkin1',\n    'commoo3', 'commyn', 'compea', 'comros', 'comsan', 'comtai1', 'copbar1', 'crbsun2', 'cregos1', 'crfbar1',\n    'crseag1', 'dafbab1', 'darter2', 'eaywag1', 'emedov2', 'eucdov', 'eurbla2', 'eurcoo', 'forwag1', 'gargan',\n    'gloibi', 'goflea1', 'graher1', 'grbeat1', 'grecou1', 'greegr', 'grefla1', 'grehor1', 'grejun2', 'grenig1',\n    'grewar3', 'grnsan', 'grnwar1', 'grtdro1', 'gryfra', 'grynig2', 'grywag', 'gybpri1', 'gyhcaf1', 'heswoo1',\n    'hoopoe', 'houcro1', 'houspa', 'inbrob1', 'indpit1', 'indrob1', 'indrol2', 'indtit1', 'ingori1', 'inpher1',\n    'insbab1', 'insowl1', 'integr', 'isbduc1', 'jerbus2', 'junbab2', 'junmyn1', 'junowl1', 'kenplo1', 'kerlau2',\n    'labcro1', 'laudov1', 'lblwar1', 'lesyel1', 'lewduc1', 'lirplo', 'litegr', 'litgre1', 'litspi1', 'litswi1',\n    'lobsun2', 'maghor2', 'malpar1', 'maltro1', 'malwoo1', 'marsan', 'mawthr1', 'moipig1', 'nilfly2', 'niwpig1',\n    'nutman', 'orihob2', 'oripip1', 'pabflo1', 'paisto1', 'piebus1', 'piekin1', 'placuc3', 'plaflo1', 'plapri1',\n    'plhpar1', 'pomgrp2', 'purher1', 'pursun3', 'pursun4', 'purswa3', 'putbab1', 'redspu1', 'rerswa1', 'revbul',\n    'rewbul', 'rewlap1', 'rocpig', 'rorpar', 'rossta2', 'rufbab3', 'ruftre2', 'rufwoo2', 'rutfly6', 'sbeowl1',\n    'scamin3', 'shikra1', 'smamin1', 'sohmyn1', 'spepic1', 'spodov', 'spoowl1', 'sqtbul1', 'stbkin1', 'sttwoo1',\n    'thbwar1', 'tibfly3', 'tilwar1', 'vefnut1', 'vehpar1', 'wbbfly1', 'wemhar1', 'whbbul2', 'whbsho3', 'whbtre1',\n    'whbwag1', 'whbwat1', 'whbwoo2', 'whcbar1', 'whiter2', 'whrmun', 'whtkin2', 'woosan', 'wynlau1', 'yebbab1',\n    'yebbul3', 'zitcis1'\n]","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:05:58.788693Z","iopub.execute_input":"2024-06-07T04:05:58.789041Z","iopub.status.idle":"2024-06-07T04:05:58.804498Z","shell.execute_reply.started":"2024-06-07T04:05:58.789013Z","shell.execute_reply":"2024-06-07T04:05:58.803438Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:05:58.806439Z","iopub.execute_input":"2024-06-07T04:05:58.806809Z","iopub.status.idle":"2024-06-07T04:05:58.822160Z","shell.execute_reply.started":"2024-06-07T04:05:58.806777Z","shell.execute_reply":"2024-06-07T04:05:58.820898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:05:58.823495Z","iopub.execute_input":"2024-06-07T04:05:58.823881Z","iopub.status.idle":"2024-06-07T04:05:58.839269Z","shell.execute_reply.started":"2024-06-07T04:05:58.823851Z","shell.execute_reply":"2024-06-07T04:05:58.838036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## load data files list","metadata":{}},{"cell_type":"code","source":"TEST_AUDIO_DIR = '/kaggle/input/birdclef-2024/test_soundscapes'\nUNLABELED_AUDIO_DIR = '/kaggle/input/birdclef-2024/unlabeled_soundscapes'","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:05:59.362216Z","iopub.execute_input":"2024-06-07T04:05:59.363251Z","iopub.status.idle":"2024-06-07T04:05:59.368003Z","shell.execute_reply.started":"2024-06-07T04:05:59.363211Z","shell.execute_reply":"2024-06-07T04:05:59.366775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_files = glob(os.path.join(TEST_AUDIO_DIR, '**/*.ogg'), recursive=True)\nif len(test_files) == 0:\n    print(f'loading from {UNLABELED_AUDIO_DIR}')\n    test_files = glob(os.path.join(UNLABELED_AUDIO_DIR, '**/*.ogg'), recursive=True)[:30]\n\n# test_files = test_files[:10]\ndf_test_files = pd.DataFrame(test_files, columns=['file_path'])","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:05:59.561261Z","iopub.execute_input":"2024-06-07T04:05:59.561657Z","iopub.status.idle":"2024-06-07T04:06:05.385328Z","shell.execute_reply.started":"2024-06-07T04:05:59.561628Z","shell.execute_reply":"2024-06-07T04:06:05.384285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test_files","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:06:05.387508Z","iopub.execute_input":"2024-06-07T04:06:05.388341Z","iopub.status.idle":"2024-06-07T04:06:05.409104Z","shell.execute_reply.started":"2024-06-07T04:06:05.388306Z","shell.execute_reply":"2024-06-07T04:06:05.407739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## spectrogram transform","metadata":{}},{"cell_type":"code","source":"slice_duration = 5\nsr = 32000\nfmin = 20\nfmax = 15000\nn_mels = 128\nn_fft = n_mels*8\nspec_time_bins = 512\nhop_length = int(sr*slice_duration / spec_time_bins)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:06:05.410579Z","iopub.execute_input":"2024-06-07T04:06:05.410921Z","iopub.status.idle":"2024-06-07T04:06:05.419561Z","shell.execute_reply.started":"2024-06-07T04:06:05.410891Z","shell.execute_reply":"2024-06-07T04:06:05.418418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sepc_trans = torchaudio.transforms.MelSpectrogram(\n    sample_rate=sr, hop_length=hop_length, n_fft=n_fft,\n    n_mels=n_mels, f_min=fmin, f_max=fmax, mel_scale='slaney', center=True, pad_mode='reflect'\n).to(device)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:06:05.421686Z","iopub.execute_input":"2024-06-07T04:06:05.422117Z","iopub.status.idle":"2024-06-07T04:06:05.542159Z","shell.execute_reply.started":"2024-06-07T04:06:05.422084Z","shell.execute_reply":"2024-06-07T04:06:05.540959Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model preparation","metadata":{}},{"cell_type":"markdown","source":"### load pretrained model files","metadata":{}},{"cell_type":"code","source":"model_state_dict = torch.load(\n    '/kaggle/input/birdclf2024-vq-spec-gpt-model-v2-7/vq_spec_gpt_v2_7.pt',\n    map_location=torch.device('cpu')\n)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:07:15.054084Z","iopub.execute_input":"2024-06-07T04:07:15.054509Z","iopub.status.idle":"2024-06-07T04:07:16.621165Z","shell.execute_reply.started":"2024-06-07T04:07:15.054475Z","shell.execute_reply":"2024-06-07T04:07:16.619859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### VQ Spectrogram Encoder","metadata":{}},{"cell_type":"code","source":"import typing as tp\nimport warnings\n\nfrom einops import rearrange, repeat\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\n\n\ndef world_size():\n    if torch.distributed.is_initialized():\n        return torch.distributed.get_world_size()\n    else:\n        return 1\n\n\ndef is_distributed():\n    return world_size() > 1\n\n\ndef all_reduce(tensor: torch.Tensor, op=torch.distributed.ReduceOp.SUM):\n    if is_distributed():\n        return torch.distributed.all_reduce(tensor, op)\n\n\ndef _is_complex_or_float(tensor):\n    return torch.is_floating_point(tensor) or torch.is_complex(tensor)\n\n\ndef _check_number_of_params(params: tp.List[torch.Tensor]):\n    # utility function to check that the number of params in all workers is the same,\n    # and thus avoid a deadlock with distributed all reduce.\n    if not is_distributed() or not params:\n        return\n    tensor = torch.tensor([len(params)], device=params[0].device, dtype=torch.long)\n    all_reduce(tensor)\n    if tensor.item() != len(params) * world_size():\n        # If not all the workers have the same number, for at least one of them,\n        # this inequality will be verified.\n        raise RuntimeError(f\"Mismatch in number of params: ours is {len(params)}, \"\n                           \"at least one worker has a different one.\")\n\n\ndef broadcast_tensors(tensors: tp.Iterable[torch.Tensor], src: int = 0):\n    \"\"\"Broadcast the tensors from the given parameters to all workers.\n    This can be used to ensure that all workers have the same model to start with.\n    \"\"\"\n    if not is_distributed():\n        return\n    tensors = [tensor for tensor in tensors if _is_complex_or_float(tensor)]\n    _check_number_of_params(tensors)\n    handles = []\n    for tensor in tensors:\n        handle = torch.distributed.broadcast(tensor.data, src=src, async_op=True)\n        handles.append(handle)\n    for handle in handles:\n        handle.wait()\n\n\ndef default(val: tp.Any, d: tp.Any) -> tp.Any:\n    return val if val is not None else d\n\n\ndef ema_inplace(moving_avg, new, decay: float):\n    moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay))\n\n\ndef laplace_smoothing(x, n_categories: int, epsilon: float = 1e-5):\n    return (x + epsilon) / (x.sum() + n_categories * epsilon)\n\n\ndef uniform_init(*shape: int):\n    t = torch.empty(shape)\n    nn.init.kaiming_uniform_(t)\n    return t\n\n\ndef sample_vectors(samples, num: int):\n    num_samples, device = samples.shape[0], samples.device\n\n    if num_samples >= num:\n        indices = torch.randperm(num_samples, device=device)[:num]\n    else:\n        indices = torch.randint(0, num_samples, (num,), device=device)\n\n    return samples[indices]\n\n\ndef kmeans(samples, num_clusters: int, num_iters: int = 10):\n    dim, dtype = samples.shape[-1], samples.dtype\n\n    means = sample_vectors(samples, num_clusters)\n\n    for _ in range(num_iters):\n        diffs = rearrange(samples, \"n d -> n () d\") - rearrange(\n            means, \"c d -> () c d\"\n        )\n        dists = -(diffs ** 2).sum(dim=-1)\n\n        buckets = dists.max(dim=-1).indices\n        bins = torch.bincount(buckets, minlength=num_clusters)\n        zero_mask = bins == 0\n        bins_min_clamped = bins.masked_fill(zero_mask, 1)\n\n        new_means = buckets.new_zeros(num_clusters, dim, dtype=dtype)\n        new_means.scatter_add_(0, repeat(buckets, \"n -> n d\", d=dim), samples)\n        new_means = new_means / bins_min_clamped[..., None]\n\n        means = torch.where(zero_mask[..., None], means, new_means)\n\n    return means, bins\n\n\nclass EuclideanCodebook(nn.Module):\n    \"\"\"Codebook with Euclidean distance.\n    Args:\n        dim (int): Dimension.\n        codebook_size (int): Codebook size.\n        kmeans_init (bool): Whether to use k-means to initialize the codebooks.\n            If set to true, run the k-means algorithm on the first training batch and use\n            the learned centroids as initialization.\n        kmeans_iters (int): Number of iterations used for k-means algorithm at initialization.\n        decay (float): Decay for exponential moving average over the codebooks.\n        epsilon (float): Epsilon value for numerical stability.\n        threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes\n            that have an exponential moving average cluster size less than the specified threshold with\n            randomly selected vector from the current batch.\n    \"\"\"\n    def __init__(\n        self,\n        dim: int,\n        codebook_size: int,\n        kmeans_init: int = False,\n        kmeans_iters: int = 10,\n        decay: float = 0.99,\n        epsilon: float = 1e-5,\n        threshold_ema_dead_code: int = 2,\n    ):\n        super().__init__()\n        self.decay = decay\n        init_fn: tp.Union[tp.Callable[..., torch.Tensor], tp.Any] = uniform_init if not kmeans_init else torch.zeros\n        embed = init_fn(codebook_size, dim)\n\n        self.codebook_size = codebook_size\n\n        self.kmeans_iters = kmeans_iters\n        self.epsilon = epsilon\n        self.threshold_ema_dead_code = threshold_ema_dead_code\n\n        self.register_buffer(\"inited\", torch.Tensor([not kmeans_init]))\n        self.register_buffer(\"cluster_size\", torch.zeros(codebook_size))\n        self.register_buffer(\"embed\", embed)\n        self.register_buffer(\"embed_avg\", embed.clone())\n\n    @torch.jit.ignore\n    def init_embed_(self, data):\n        if self.inited:\n            return\n\n        embed, cluster_size = kmeans(data, self.codebook_size, self.kmeans_iters)\n        self.embed.data.copy_(embed)\n        self.embed_avg.data.copy_(embed.clone())\n        self.cluster_size.data.copy_(cluster_size)\n        self.inited.data.copy_(torch.Tensor([True]))\n        # Make sure all buffers across workers are in sync after initialization\n        broadcast_tensors(self.buffers())\n\n    def replace_(self, samples, mask):\n        modified_codebook = torch.where(\n            mask[..., None], sample_vectors(samples, self.codebook_size), self.embed\n        )\n        self.embed.data.copy_(modified_codebook)\n\n    def expire_codes_(self, batch_samples):\n        if self.threshold_ema_dead_code == 0:\n            return\n\n        expired_codes = self.cluster_size < self.threshold_ema_dead_code\n        if not torch.any(expired_codes):\n            return\n\n        batch_samples = rearrange(batch_samples, \"... d -> (...) d\")\n        self.replace_(batch_samples, mask=expired_codes)\n        broadcast_tensors(self.buffers())\n\n    def preprocess(self, x):\n        x = rearrange(x, \"... d -> (...) d\")\n        return x\n\n    def quantize(self, x):\n        embed = self.embed.t()  # (codebook_size, emb_dim) => (emb_dim, codebook_size)\n        eps = 1e-20\n        dist = -(\n            x.pow(2).sum(1, keepdim=True)  # (N, emb_dim) => (N, 1)\n            - 2 * x @ (embed + eps - eps)  # (N, emb_dim) @ (emb_dim, codebook_size) => (N, codebook_size)\n            + embed.pow(2).sum(0, keepdim=True)  # (emb_dim, codebook_size) => (1, codebook_size)\n        )  # (N, 1) + (N, codebook_size) + (1, codebook_size) => (N, codebook_size)\n        # embed + eps - epsしとかないとx @ embed計算が途中でものすごく遅くなるよう？\n        embed_ind = dist.max(dim=-1).indices  # (N, codebook_size) => (N,)\n        return embed_ind\n\n    def postprocess_emb(self, embed_ind, shape):\n        return embed_ind.view(*shape[:-1])\n\n    def dequantize(self, embed_ind):\n        quantize = F.embedding(embed_ind, self.embed)\n        return quantize\n\n    def encode(self, x):\n        shape = x.shape\n        # pre-process\n        x = self.preprocess(x)\n        # quantize\n        embed_ind = self.quantize(x)\n        # post-process\n        embed_ind = self.postprocess_emb(embed_ind, shape)\n        return embed_ind\n\n    def decode(self, embed_ind):\n        quantize = self.dequantize(embed_ind)\n        return quantize\n\n    def forward(self, x):\n        shape, dtype = x.shape, x.dtype\n        x = self.preprocess(x)  # (A, B,... , emb_dim) -> (N, emb_dim)\n\n        self.init_embed_(x)\n\n        embed_ind = self.quantize(x)  # (N, emb_dim) -> (N,)\n        embed_onehot = F.one_hot(embed_ind, self.codebook_size).type(dtype)  # (N,) -> (N, codebook_size)\n        embed_ind = self.postprocess_emb(embed_ind, shape)  # (N,) -> (A, B,... )\n        quantize = self.dequantize(embed_ind)  # (A, B,... ) -> (A, B,..., emb_dim)\n\n        if self.training:\n            # We do the expiry of code at that point as buffers are in sync\n            # and all the workers will take the same decision.\n            self.expire_codes_(x)\n            ema_inplace(self.cluster_size, embed_onehot.sum(0), self.decay)\n            embed_sum = x.t() @ embed_onehot\n            ema_inplace(self.embed_avg, embed_sum.t(), self.decay)\n            cluster_size = (\n                laplace_smoothing(self.cluster_size, self.codebook_size, self.epsilon)\n                * self.cluster_size.sum()\n            )\n            embed_normalized = self.embed_avg / cluster_size.unsqueeze(1)\n            self.embed.data.copy_(embed_normalized)\n\n        return quantize, embed_ind\n\nclass VectorQuantization(nn.Module):\n    \"\"\"Vector quantization implementation.\n    Currently supports only euclidean distance.\n    Args:\n        dim (int): Input dimension\n        codebook_size (int): Codebook size\n        codebook_dim (int): Codebook dimension. If not defined, uses the specified dimension in dim.\n        decay (float): Decay for exponential moving average over the codebooks.\n        epsilon (float): Epsilon value for numerical stability.\n        kmeans_init (bool): Whether to use kmeans to initialize the codebooks.\n        kmeans_iters (int): Number of iterations used for kmeans initialization.\n        threshold_ema_dead_code (int): Threshold for dead code expiration. Replace any codes\n            that have an exponential moving average cluster size less than the specified threshold with\n            randomly selected vector from the current batch.\n        commitment_weight (float): Weight for commitment loss.\n    \"\"\"\n    def __init__(\n        self,\n        input_dim: int,\n        codebook_size: int,\n        codebook_dim: tp.Optional[int] = None,\n        decay: float = 0.99,\n        epsilon: float = 1e-5,\n        kmeans_init: bool = True,\n        kmeans_iters: int = 50,\n        threshold_ema_dead_code: int = 2,\n        commitment_weight: float = 1.,\n    ):\n        super().__init__()\n        _codebook_dim: int = default(codebook_dim, input_dim)\n\n        requires_projection = _codebook_dim != input_dim\n        self.project_in = (nn.Linear(input_dim, _codebook_dim) if requires_projection else nn.Identity())\n        self.project_out = (nn.Linear(_codebook_dim, input_dim) if requires_projection else nn.Identity())\n\n        self.epsilon = epsilon\n        self.commitment_weight = commitment_weight\n\n        self._codebook = EuclideanCodebook(dim=_codebook_dim, codebook_size=codebook_size,\n                                           kmeans_init=kmeans_init, kmeans_iters=kmeans_iters,\n                                           decay=decay, epsilon=epsilon,\n                                           threshold_ema_dead_code=threshold_ema_dead_code)\n        self.codebook_size = codebook_size\n\n    @property\n    def codebook(self):\n        return self._codebook.embed\n\n    def encode(self, x):\n        x = rearrange(x, \"b d n -> b n d\")\n        x = self.project_in(x)\n        embed_in = self._codebook.encode(x)\n        return embed_in\n\n    def decode(self, embed_ind):\n        quantize = self._codebook.decode(embed_ind)\n        quantize = self.project_out(quantize)\n        quantize = rearrange(quantize, \"b n d -> b d n\")\n        return quantize\n\n    def forward(self, x):\n        \"\"\"\n        x: (batch_size, seq_len, feat_dim)\n        \"\"\"\n        device = x.device\n        # x = rearrange(x, \"b d n -> b n d\")  # (batch_size, feat_dim, seq_len) -> (batch_size, seq_len, feat_dim)\n        x = self.project_in(x)  # (batch_size, seq_len, feat_dim) -> (batch_size, seq_len, codebook_dim)\n\n        quantize, embed_ind = self._codebook(x)  # quantize: (batch_size, seq_len, codebook_dim), (batch_size, seq_len)\n\n        if self.training:\n            quantize = x + (quantize - x).detach()\n\n        loss = torch.tensor([0.0], device=device, requires_grad=self.training)\n\n        if self.training:\n            # warnings.warn('When using RVQ in training model, first check '\n            #               'https://github.com/facebookresearch/encodec/issues/25 . '\n            #               'The bug wasn\\'t fixed here for reproducibility.')\n            if self.commitment_weight > 0:\n                commit_loss = F.mse_loss(quantize.detach(), x)\n                loss = loss + commit_loss * self.commitment_weight\n\n        quantize = self.project_out(quantize)  # (batch_size, seq_len, codebook_dim) -> (batch_size, seq_len, feat_dim)\n        # quantize = rearrange(quantize, \"b n d -> b d n\")  # (batch_size, seq_len, feat_dim) -> (batch_size, feat_dim, seq_len)\n        return quantize, embed_ind, loss","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:07:24.288211Z","iopub.execute_input":"2024-06-07T04:07:24.288640Z","iopub.status.idle":"2024-06-07T04:07:24.339873Z","shell.execute_reply.started":"2024-06-07T04:07:24.288606Z","shell.execute_reply":"2024-06-07T04:07:24.338548Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\n\nfrom torch.nn.utils import spectral_norm, weight_norm\n\n\nclass ConvLayerNorm(nn.LayerNorm):\n    \"\"\"\n    Convolution-friendly LayerNorm that moves channels to last dimensions\n    before running the normalization and moves them back to original position right after.\n    \"\"\"\n    def __init__(self, normalized_shape: tp.Union[int, tp.List[int], torch.Size], **kwargs):\n        super().__init__(normalized_shape, **kwargs)\n\n    def forward(self, x):\n        x = einops.rearrange(x, 'b ... t -> b t ...')\n        x = super().forward(x)\n        x = einops.rearrange(x, 'b t ... -> b ... t')\n\n\nCONV_NORMALIZATIONS = frozenset(['none', 'weight_norm', 'spectral_norm',\n                                 'time_layer_norm', 'layer_norm', 'time_group_norm'])\n\n\ndef apply_parametrization_norm(module: nn.Module, norm: str = 'none') -> nn.Module:\n    assert norm in CONV_NORMALIZATIONS\n    if norm == 'weight_norm':\n        return weight_norm(module)\n    elif norm == 'spectral_norm':\n        return spectral_norm(module)\n    else:\n        # We already check was in CONV_NORMALIZATION, so any other choice\n        # doesn't need reparametrization.\n        return module\n\n\ndef get_norm_module(module: nn.Module, causal: bool = False, norm: str = 'none', **norm_kwargs) -> nn.Module:\n    \"\"\"Return the proper normalization module. If causal is True, this will ensure the returned\n    module is causal, or return an error if the normalization doesn't support causal evaluation.\n    \"\"\"\n    assert norm in CONV_NORMALIZATIONS\n    if norm == 'layer_norm':\n        assert isinstance(module, nn.modules.conv._ConvNd)\n        return ConvLayerNorm(module.out_channels, **norm_kwargs)\n    elif norm == 'time_group_norm':\n        if causal:\n            raise ValueError(\"GroupNorm doesn't support causal evaluation.\")\n        assert isinstance(module, nn.modules.conv._ConvNd)\n        return nn.GroupNorm(1, module.out_channels, **norm_kwargs)\n    else:\n        return nn.Identity()\n\n\ndef get_extra_padding_for_conv1d(x: torch.Tensor, kernel_size: int, stride: int,\n                                 padding_total: int = 0) -> int:\n    \"\"\"See `pad_for_conv1d`.\n    \"\"\"\n    length = x.shape[-1]\n    n_frames = (length - kernel_size + padding_total) / stride + 1\n    ideal_length = (math.ceil(n_frames) - 1) * stride + (kernel_size - padding_total)\n    return ideal_length - length\n\n\ndef pad_for_conv1d(x: torch.Tensor, kernel_size: int, stride: int, padding_total: int = 0):\n    \"\"\"Pad for a convolution to make sure that the last window is full.\n    Extra padding is added at the end. This is required to ensure that we can rebuild\n    an output of the same length, as otherwise, even with padding, some time steps\n    might get removed.\n    For instance, with total padding = 4, kernel size = 4, stride = 2:\n        0 0 1 2 3 4 5 0 0   # (0s are padding)\n        1   2   3           # (output frames of a convolution, last 0 is never used)\n        0 0 1 2 3 4 5 0     # (output of tr. conv., but pos. 5 is going to get removed as padding)\n            1 2 3 4         # once you removed padding, we are missing one time step !\n    \"\"\"\n    extra_padding = get_extra_padding_for_conv1d(x, kernel_size, stride, padding_total)\n    return F.pad(x, (0, extra_padding))\n\n\ndef pad1d(x: torch.Tensor, paddings: tp.Tuple[int, int], mode: str = 'zero', value: float = 0.):\n    \"\"\"Tiny wrapper around F.pad, just to allow for reflect padding on small input.\n    If this is the case, we insert extra 0 padding to the right before the reflection happen.\n    \"\"\"\n    length = x.shape[-1]\n    padding_left, padding_right = paddings\n    assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)\n    if mode == 'reflect':\n        max_pad = max(padding_left, padding_right)\n        extra_pad = 0\n        if length <= max_pad:\n            extra_pad = max_pad - length + 1\n            x = F.pad(x, (0, extra_pad))\n        padded = F.pad(x, paddings, mode, value)\n        end = padded.shape[-1] - extra_pad\n        return padded[..., :end]\n    else:\n        return F.pad(x, paddings, mode, value)\n\n\ndef unpad1d(x: torch.Tensor, paddings: tp.Tuple[int, int]):\n    \"\"\"Remove padding from x, handling properly zero padding. Only for 1d!\"\"\"\n    padding_left, padding_right = paddings\n    assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right)\n    assert (padding_left + padding_right) <= x.shape[-1]\n    end = x.shape[-1] - padding_right\n    return x[..., padding_left: end]\n\n\nclass NormConv1d(nn.Module):\n    \"\"\"Wrapper around Conv1d and normalization applied to this conv\n    to provide a uniform interface across normalization approaches.\n    \"\"\"\n    def __init__(self, *args, causal: bool = False, norm: str = 'none',\n                 norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):\n        super().__init__()\n        self.conv = apply_parametrization_norm(nn.Conv1d(*args, **kwargs), norm)\n        self.norm = get_norm_module(self.conv, causal, norm, **norm_kwargs)\n        self.norm_type = norm\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.norm(x)\n        return x\n\n\nclass NormConv2d(nn.Module):\n    \"\"\"Wrapper around Conv2d and normalization applied to this conv\n    to provide a uniform interface across normalization approaches.\n    \"\"\"\n    def __init__(self, *args, norm: str = 'none',\n                 norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):\n        super().__init__()\n        self.conv = apply_parametrization_norm(nn.Conv2d(*args, **kwargs), norm)\n        self.norm = get_norm_module(self.conv, causal=False, norm=norm, **norm_kwargs)\n        self.norm_type = norm\n\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.norm(x)\n        return x\n\n\nclass NormConvTranspose1d(nn.Module):\n    \"\"\"Wrapper around ConvTranspose1d and normalization applied to this conv\n    to provide a uniform interface across normalization approaches.\n    \"\"\"\n    def __init__(self, *args, causal: bool = False, norm: str = 'none',\n                 norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):\n        super().__init__()\n        self.convtr = apply_parametrization_norm(nn.ConvTranspose1d(*args, **kwargs), norm)\n        self.norm = get_norm_module(self.convtr, causal, norm, **norm_kwargs)\n        self.norm_type = norm\n\n    def forward(self, x):\n        x = self.convtr(x)\n        x = self.norm(x)\n        return x\n\n\nclass NormConvTranspose2d(nn.Module):\n    \"\"\"Wrapper around ConvTranspose2d and normalization applied to this conv\n    to provide a uniform interface across normalization approaches.\n    \"\"\"\n    def __init__(self, *args, norm: str = 'none',\n                 norm_kwargs: tp.Dict[str, tp.Any] = {}, **kwargs):\n        super().__init__()\n        self.convtr = apply_parametrization_norm(nn.ConvTranspose2d(*args, **kwargs), norm)\n        self.norm = get_norm_module(self.convtr, causal=False, norm=norm, **norm_kwargs)\n\n    def forward(self, x):\n        x = self.convtr(x)\n        x = self.norm(x)\n        return x\n\n\nclass SConv1d(nn.Module):\n    \"\"\"Conv1d with some builtin handling of asymmetric or causal padding\n    and normalization.\n    \"\"\"\n    def __init__(self, in_channels: int, out_channels: int,\n                 kernel_size: int, stride: int = 1, dilation: int = 1,\n                 groups: int = 1, bias: bool = True, causal: bool = False,\n                 norm: str = 'none', norm_kwargs: tp.Dict[str, tp.Any] = {},\n                 pad_mode: str = 'reflect'):\n        super().__init__()\n        # warn user on unusual setup between dilation and stride\n        if stride > 1 and dilation > 1:\n            warnings.warn('SConv1d has been initialized with stride > 1 and dilation > 1'\n                          f' (kernel_size={kernel_size} stride={stride}, dilation={dilation}).')\n        self.conv = NormConv1d(in_channels, out_channels, kernel_size, stride,\n                               dilation=dilation, groups=groups, bias=bias, causal=causal,\n                               norm=norm, norm_kwargs=norm_kwargs)\n        self.causal = causal\n        self.pad_mode = pad_mode\n\n    def forward(self, x):\n        B, C, T = x.shape\n        kernel_size = self.conv.conv.kernel_size[0]\n        stride = self.conv.conv.stride[0]\n        dilation = self.conv.conv.dilation[0]\n        kernel_size = (kernel_size - 1) * dilation + 1  # effective kernel size with dilations\n        padding_total = kernel_size - stride\n        extra_padding = get_extra_padding_for_conv1d(x, kernel_size, stride, padding_total)\n        if self.causal:\n            # Left padding for causal\n            x = pad1d(x, (padding_total, extra_padding), mode=self.pad_mode)\n        else:\n            # Asymmetric padding required for odd strides\n            padding_right = padding_total // 2\n            padding_left = padding_total - padding_right\n            x = pad1d(x, (padding_left, padding_right + extra_padding), mode=self.pad_mode)\n        return self.conv(x)\n\n\nclass SConvTranspose1d(nn.Module):\n    \"\"\"ConvTranspose1d with some builtin handling of asymmetric or causal padding\n    and normalization.\n    \"\"\"\n    def __init__(self, in_channels: int, out_channels: int,\n                 kernel_size: int, stride: int = 1, causal: bool = False,\n                 norm: str = 'none', trim_right_ratio: float = 1.,\n                 norm_kwargs: tp.Dict[str, tp.Any] = {}):\n        super().__init__()\n        self.convtr = NormConvTranspose1d(in_channels, out_channels, kernel_size, stride,\n                                          causal=causal, norm=norm, norm_kwargs=norm_kwargs)\n        self.causal = causal\n        self.trim_right_ratio = trim_right_ratio\n        assert self.causal or self.trim_right_ratio == 1., \\\n            \"`trim_right_ratio` != 1.0 only makes sense for causal convolutions\"\n        assert self.trim_right_ratio >= 0. and self.trim_right_ratio <= 1.\n\n    def forward(self, x):\n        kernel_size = self.convtr.convtr.kernel_size[0]\n        stride = self.convtr.convtr.stride[0]\n        padding_total = kernel_size - stride\n\n        y = self.convtr(x)\n\n        # We will only trim fixed padding. Extra padding from `pad_for_conv1d` would be\n        # removed at the very end, when keeping only the right length for the output,\n        # as removing it here would require also passing the length at the matching layer\n        # in the encoder.\n        if self.causal:\n            # Trim the padding on the right according to the specified ratio\n            # if trim_right_ratio = 1.0, trim everything from right\n            padding_right = math.ceil(padding_total * self.trim_right_ratio)\n            padding_left = padding_total - padding_right\n            y = unpad1d(y, (padding_left, padding_right))\n        else:\n            # Asymmetric padding required for odd strides\n            padding_right = padding_total // 2\n            padding_left = padding_total - padding_right\n            y = unpad1d(y, (padding_left, padding_right))\n        return y\n\n\nclass SLSTM(nn.Module):\n    \"\"\"\n    LSTM without worrying about the hidden state, nor the layout of the data.\n    Expects input as convolutional layout.\n    \"\"\"\n    def __init__(self, dimension: int, num_layers: int = 2, skip: bool = True):\n        super().__init__()\n        self.skip = skip\n        self.lstm = nn.LSTM(dimension, dimension, num_layers)\n\n    def forward(self, x):\n        x = x.permute(2, 0, 1)\n        y, _ = self.lstm(x)\n        if self.skip:\n            y = y + x\n        y = y.permute(1, 2, 0)\n        return y\n\n\nclass SEANetResnetBlock(nn.Module):\n    \"\"\"Residual block from SEANet model.\n    Args:\n        dim (int): Dimension of the input/output\n        kernel_sizes (list): List of kernel sizes for the convolutions.\n        dilations (list): List of dilations for the convolutions.\n        activation (str): Activation function.\n        activation_params (dict): Parameters to provide to the activation function\n        norm (str): Normalization method.\n        norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution.\n        causal (bool): Whether to use fully causal convolution.\n        pad_mode (str): Padding mode for the convolutions.\n        compress (int): Reduced dimensionality in residual branches (from Demucs v3)\n        true_skip (bool): Whether to use true skip connection or a simple convolution as the skip connection.\n    \"\"\"\n    def __init__(self, dim: int, kernel_sizes: tp.List[int] = [3, 1], dilations: tp.List[int] = [1, 1],\n                 activation: str = 'ELU', activation_params: dict = {'alpha': 1.0},\n                 norm: str = 'weight_norm', norm_params: tp.Dict[str, tp.Any] = {}, causal: bool = False,\n                 pad_mode: str = 'reflect', compress: int = 2, true_skip: bool = True):\n        super().__init__()\n        assert len(kernel_sizes) == len(dilations), 'Number of kernel sizes should match number of dilations'\n        act = getattr(nn, activation)\n        hidden = dim // compress\n        block = []\n        for i, (kernel_size, dilation) in enumerate(zip(kernel_sizes, dilations)):\n            in_chs = dim if i == 0 else hidden\n            out_chs = dim if i == len(kernel_sizes) - 1 else hidden\n            block += [\n                act(**activation_params),\n                SConv1d(in_chs, out_chs, kernel_size=kernel_size, dilation=dilation,\n                        norm=norm, norm_kwargs=norm_params,\n                        causal=causal, pad_mode=pad_mode),\n            ]\n        self.block = nn.Sequential(*block)\n        self.shortcut: nn.Module\n        if true_skip:\n            self.shortcut = nn.Identity()\n        else:\n            self.shortcut = SConv1d(dim, dim, kernel_size=1, norm=norm, norm_kwargs=norm_params,\n                                    causal=causal, pad_mode=pad_mode)\n\n    def forward(self, x):\n        return self.shortcut(x) + self.block(x)\n\n\nclass SEANetEncoder(nn.Module):\n    \"\"\"SEANet encoder.\n    Args:\n        channels (int): Audio channels.\n        dimension (int): Intermediate representation dimension.\n        n_filters (int): Base width for the model.\n        n_residual_layers (int): nb of residual layers.\n        ratios (Sequence[int]): kernel size and stride ratios. The encoder uses downsampling ratios instead of\n            upsampling ratios, hence it will use the ratios in the reverse order to the ones specified here\n            that must match the decoder order\n        activation (str): Activation function.\n        activation_params (dict): Parameters to provide to the activation function\n        norm (str): Normalization method.\n        norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution.\n        kernel_size (int): Kernel size for the initial convolution.\n        last_kernel_size (int): Kernel size for the initial convolution.\n        residual_kernel_size (int): Kernel size for the residual layers.\n        dilation_base (int): How much to increase the dilation with each layer.\n        causal (bool): Whether to use fully causal convolution.\n        pad_mode (str): Padding mode for the convolutions.\n        true_skip (bool): Whether to use true skip connection or a simple\n            (streamable) convolution as the skip connection in the residual network blocks.\n        compress (int): Reduced dimensionality in residual branches (from Demucs v3).\n        lstm (int): Number of LSTM layers at the end of the encoder.\n    \"\"\"\n    def __init__(self, channels: int = 1, dimension: int = 128, n_filters: int = 32, n_residual_layers: int = 1,\n                 ratios: tp.List[int] = [8, 5, 4, 2], activation: str = 'ELU', activation_params: dict = {'alpha': 1.0},\n                 norm: str = 'weight_norm', norm_params: tp.Dict[str, tp.Any] = {}, kernel_size: int = 7,\n                 last_kernel_size: int = 7, residual_kernel_size: int = 3, dilation_base: int = 2, causal: bool = False,\n                 pad_mode: str = 'reflect', true_skip: bool = False, compress: int = 2, lstm: int = 2):\n        super().__init__()\n        self.channels = channels\n        self.dimension = dimension\n        self.n_filters = n_filters\n        self.ratios = list(reversed(ratios))\n        del ratios\n        self.n_residual_layers = n_residual_layers\n        self.hop_length = np.prod(self.ratios)\n\n        act = getattr(nn, activation)\n        mult = 1\n        model: tp.List[nn.Module] = [\n            SConv1d(channels, mult * n_filters, kernel_size, norm=norm, norm_kwargs=norm_params,\n                    causal=causal, pad_mode=pad_mode)\n        ]\n        # Downsample to raw audio scale\n        for i, ratio in enumerate(self.ratios):\n            # Add residual layers\n            for j in range(n_residual_layers):\n                model += [\n                    SEANetResnetBlock(mult * n_filters, kernel_sizes=[residual_kernel_size, 1],\n                                      dilations=[dilation_base ** j, 1],\n                                      norm=norm, norm_params=norm_params,\n                                      activation=activation, activation_params=activation_params,\n                                      causal=causal, pad_mode=pad_mode, compress=compress, true_skip=true_skip)]\n\n            # Add downsampling layers\n            model += [\n                act(**activation_params),\n                SConv1d(mult * n_filters, mult * n_filters * 2,\n                        kernel_size=ratio * 2, stride=ratio,\n                        norm=norm, norm_kwargs=norm_params,\n                        causal=causal, pad_mode=pad_mode),\n            ]\n            mult *= 2\n\n        if lstm:\n            model += [SLSTM(mult * n_filters, num_layers=lstm)]\n\n        model += [\n            act(**activation_params),\n            SConv1d(mult * n_filters, dimension, last_kernel_size, norm=norm, norm_kwargs=norm_params,\n                    causal=causal, pad_mode=pad_mode)\n        ]\n\n        self.model = nn.Sequential(*model)\n\n    def forward(self, x):\n        return self.model(x)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:07:24.517798Z","iopub.execute_input":"2024-06-07T04:07:24.518201Z","iopub.status.idle":"2024-06-07T04:07:24.706211Z","shell.execute_reply.started":"2024-06-07T04:07:24.518169Z","shell.execute_reply":"2024-06-07T04:07:24.705081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import itertools\nimport collections.abc\n\ndef _ntuple(n):\n    def parse(x):\n        if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):\n            return tuple(x)\n        return tuple(itertools.repeat(x, n))\n    return parse\n\nto_1tuple = _ntuple(1)\nto_2tuple = _ntuple(2)\nto_3tuple = _ntuple(3)\nto_4tuple = _ntuple(4)\nto_ntuple = _ntuple\n\nclass PatchEmbed(nn.Module):\n    \"\"\" 2D Image to Patch Embedding\n    \"\"\"\n    def __init__(\n            self,\n            img_size=224,\n            patch_size=16,\n            in_chans=3,\n            embed_dim=768,\n            norm_layer=None,\n            flatten=True,\n            bias=True,\n    ):\n        super().__init__()\n        img_size = to_2tuple(img_size)\n        patch_size = to_2tuple(patch_size)\n        self.img_size = img_size\n        self.patch_size = patch_size\n        self.grid_size = (img_size[0] // patch_size[0], img_size[1] // patch_size[1])\n        self.num_patches = self.grid_size[0] * self.grid_size[1]\n        self.flatten = flatten\n\n        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size, bias=bias)\n        self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()\n\n    def forward(self, x):\n        B, C, H, W = x.shape\n        assert H == self.img_size[0], f\"Input image height ({H}) doesn't match model ({self.img_size[0]}).\"\n        assert W == self.img_size[1], f\"Input image width ({W}) doesn't match model ({self.img_size[1]}).\"\n        x = self.proj(x)\n        if self.flatten:\n            x = x.flatten(2).transpose(1, 2)  # BCHW -> BNC\n        x = self.norm(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:07:24.728326Z","iopub.execute_input":"2024-06-07T04:07:24.728775Z","iopub.status.idle":"2024-06-07T04:07:24.741284Z","shell.execute_reply.started":"2024-06-07T04:07:24.728743Z","shell.execute_reply":"2024-06-07T04:07:24.739874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class VQSpectrogramEncoder(nn.Module):\n\n    def __init__(\n        self,\n        spec_n_mels: int,\n        vq_input_dim: int,\n        codebook_size: int,\n        codebook_dim: tp.Optional[int] = None,\n        commitment_weight = 0.25,\n        spec_size = (n_mels, spec_time_bins),\n        patch_size: int = 32,\n        encoder_type: tp.Literal['linear', 'seanet', 'patch_embed'] = 'patch_embed',\n    ):\n        super().__init__()\n        self._vq = VectorQuantization(\n            input_dim=vq_input_dim,\n            codebook_size=codebook_size,\n            codebook_dim=codebook_dim,\n            commitment_weight=commitment_weight,\n        )\n        self._encoder_type = encoder_type\n        if encoder_type == 'linear':\n            self._spec_enc = nn.Linear(spec_n_mels, vq_input_dim)\n        elif encoder_type == 'seanet':\n            self._spec_enc = SEANetEncoder(\n                channels=spec_n_mels,\n                dimension=vq_input_dim,\n                n_filters=128,\n                kernel_size=4,\n                ratios=[4, 2],\n                true_skip=True,  # use shortcut connection\n            )\n        elif encoder_type == 'patch_embed':\n            self._spec_enc = PatchEmbed(\n                img_size=spec_size,\n                patch_size=patch_size,\n                in_chans=1,\n                embed_dim=vq_input_dim\n            )\n\n    def forward(self, spec):\n        bs, n_mels, t = spec.shape\n        if self._encoder_type == 'linear':\n            x = self._spec_enc(spec.permute(0, 2, 1))  # (bs, n_mels, t) => (bs, t, n_mels) => (bs, t, vq_input_dim)\n        elif self._encoder_type == 'seanet':\n            x = self._spec_enc(spec)  # (bs, n_mels, t) =>  (bs, vq_input_dim, t)\n            x = rearrange(x, 'b d t -> b t d')  # (bs, vq_input_dim, t) => (bs, t, vq_input_dim)\n        elif self._encoder_type == 'patch_embed':\n            spec = rearrange(spec, 'b nm t -> b 1 nm t')\n            x = self._spec_enc(spec)  # (bs, 1, n_mels, t) =>  (bs, n_patches, vq_input_dim)\n        quantize, embed_ind, q_emb_loss = self._vq(x)\n        q_emb = quantize  # (bs, t, codebook_dim)\n        return q_emb, embed_ind, q_emb_loss","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:07:25.281181Z","iopub.execute_input":"2024-06-07T04:07:25.281586Z","iopub.status.idle":"2024-06-07T04:07:25.293150Z","shell.execute_reply.started":"2024-06-07T04:07:25.281555Z","shell.execute_reply":"2024-06-07T04:07:25.291642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### NanoGPT","metadata":{}},{"cell_type":"code","source":"class LayerNorm(nn.Module):\n    \"\"\" LayerNorm but with an optional bias. PyTorch doesn't support simply bias=False \"\"\"\n\n    def __init__(self, ndim, bias):\n        super().__init__()\n        self.weight = nn.Parameter(torch.ones(ndim))\n        self.bias = nn.Parameter(torch.zeros(ndim)) if bias else None\n\n    def forward(self, input):\n        return pt_F.layer_norm(input, self.weight.shape, self.weight, self.bias, 1e-5)\n\n\nclass MLP(nn.Module):\n\n    def __init__(self, emb_dim: int, dropout: float, bias: bool):\n        super().__init__()\n        self.c_fc    = nn.Linear(emb_dim, 4 * emb_dim, bias=bias)\n        self.gelu    = nn.GELU()\n        self.c_proj  = nn.Linear(4 * emb_dim, emb_dim, bias=bias)\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, x):\n        x = self.c_fc(x)\n        x = self.gelu(x)\n        x = self.c_proj(x)\n        x = self.dropout(x)\n        return x\n\n\nclass MultiheadSelfAttention(nn.Module):\n\n    def __init__(self, emb_dim: int, n_head: int, dropout: float, bias: bool, block_size: int):\n        super().__init__()\n        assert emb_dim % n_head == 0\n        # key, query, value projections for all heads, but in a batch\n        self.c_attn = nn.Linear(emb_dim, 3 * emb_dim, bias=bias)\n        # output projection\n        self.c_proj = nn.Linear(emb_dim, emb_dim, bias=bias)\n        # regularization\n        self.attn_dropout = nn.Dropout(dropout)\n        self.resid_dropout = nn.Dropout(dropout)\n        self.n_head = n_head\n        self.n_embd = emb_dim\n        self.dropout = dropout\n\n    def forward(self, x, attention_mask=None):\n        assert x.ndim == 3\n        B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd)\n        if attention_mask is not None:\n            assert attention_mask.ndim == 2\n            Bm, Tm = attention_mask.size()\n            assert B == Bm\n            assert T == Tm\n\n        # calculate query, key, values for all heads in batch and move head forward to be the batch dim\n        q, k, v  = self.c_attn(x).split(self.n_embd, dim=2)\n        k = k.view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)\n        q = q.view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)\n        v = v.view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)\n\n        # extend mask for multihead attention\n        if attention_mask is not None:\n            seq_lens = attention_mask.sum(dim=-1).long().tolist()\n            mask_entended = torch.zeros(Bm, self.n_head, Tm, Tm).long().to(x.device)  # (B, nh, T, T)\n            for batch_i, seq_len in enumerate(seq_lens):\n                mask_entended[batch_i, :, :seq_len, :seq_len] = 1\n        # self-attention; Self-attend: (B, nh, T, hs) x (B, nh, hs, T) -> (B, nh, T, T)\n        # manual implementation of attention\n        att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))\n        if attention_mask is not None:\n            assert att.shape == mask_entended.shape\n            att = att.masked_fill(mask_entended == 0, float('-inf'))\n        att = pt_F.softmax(att, dim=-1)\n        att = self.attn_dropout(att)\n        y = att @ v # (B, nh, T, T) x (B, nh, T, hs) -> (B, nh, T, hs)\n        y = y.transpose(1, 2).contiguous().view(B, T, C) # re-assemble all head outputs side by side\n\n        # output projection\n        y = self.resid_dropout(self.c_proj(y))\n        return y\n\n\nclass PTSelfAttention(nn.Module):\n\n    def __init__(self, emb_dim: int):\n        n_head = 1\n        super().__init__()\n        self._attn = nn.MultiheadAttention(\n            embed_dim=emb_dim,\n            num_heads=n_head,\n            batch_first=True\n        )\n\n    def forward(self, x, attention_mask):\n        assert x.ndim == 3\n        B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd)\n\n        assert attention_mask.ndim == 2\n        Bm, Tm = attention_mask.size()\n        assert B == Bm\n        assert T == Tm\n\n        seq_lens = attention_mask.sum(dim=-1).long().tolist()\n        mask_entended = torch.zeros(Bm, Tm, Tm).long().to(x.device)  # (B, T, T)\n        for batch_i, seq_len in enumerate(seq_lens):\n            mask_entended[batch_i, :seq_len, :seq_len] = 1\n\n        attn_output, attn_output_weights = self._attn(query=x, key=x, value=x, attn_mask=mask_entended == 0)\n        # attn_output, attn_output_weights = self._attn(query=x, key=x, value=x)\n        return attn_output\n\nclass Block(nn.Module):\n\n    def __init__(self, emb_dim: int, n_head: int, dropout: float, bias: bool, block_size: int):\n        super().__init__()\n        self.ln_1 = LayerNorm(emb_dim, bias=bias)\n        self.attn = MultiheadSelfAttention(\n            emb_dim, n_head, dropout, bias, block_size\n        )\n        # self.attn = PTSelfAttention(emb_dim)\n        assert emb_dim % n_head == 0\n        self.ln_2 = LayerNorm(emb_dim, bias=bias)\n        self.mlp = MLP(emb_dim=emb_dim, dropout=dropout, bias=bias)\n\n    def forward(self, x, mask=None):\n        x = x + self.attn(self.ln_1(x), mask)\n        x = x + self.mlp(self.ln_2(x))\n        return x\n\n\nclass GPTClassifier(nn.Module):\n\n    def __init__(\n        self,\n        emb_dim: int,\n        n_head: int,\n        dropout: float,\n        bias: bool,\n        block_size: int,\n        n_layer: int,\n        vocab_size: int,\n        n_classes: int,\n        loss_fn: nn.Module = nn.CrossEntropyLoss(),\n    ):\n        super().__init__()\n\n        self.block_size = block_size\n        self.transformer = nn.ModuleDict(dict(\n            wte = nn.Embedding(vocab_size, emb_dim),\n            wpe = nn.Embedding(block_size, emb_dim),\n            drop = nn.Dropout(dropout),\n            h = nn.ModuleList([\n                Block(\n                    emb_dim=emb_dim, n_head=n_head, dropout=dropout, bias=bias, block_size=block_size\n                )\n                for _ in range(n_layer)\n            ]),\n            ln_f = LayerNorm(emb_dim, bias=bias),\n        ))\n        self.cls_head = nn.Linear(emb_dim, n_classes, bias=False)\n\n        # init all weights\n        self.apply(self._init_weights)\n        # apply special scaled init to the residual projections, per GPT-2 paper\n        for pn, p in self.named_parameters():\n            if pn.endswith('c_proj.weight'):\n                torch.nn.init.normal_(p, mean=0.0, std=0.02/math.sqrt(2 * n_layer))\n\n        # report number of parameters\n        print(\"number of parameters: %.2fM\" % (self.get_num_params()/1e6,))\n\n        self._loss_fn = loss_fn\n\n    def get_num_params(self, non_embedding=True):\n        \"\"\"\n        Return the number of parameters in the model.\n        For non-embedding count (default), the position embeddings get subtracted.\n        The token embeddings would too, except due to the parameter sharing these\n        params are actually used as weights in the final layer, so we include them.\n        \"\"\"\n        n_params = sum(p.numel() for p in self.parameters())\n        if non_embedding:\n            n_params -= self.transformer.wpe.weight.numel()\n        return n_params\n\n    def _init_weights(self, module):\n        if isinstance(module, nn.Linear):\n            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)\n            if module.bias is not None:\n                torch.nn.init.zeros_(module.bias)\n        elif isinstance(module, nn.Embedding):\n            torch.nn.init.normal_(module.weight, mean=0.0, std=0.02)\n\n    def forward(self, input_ids, attention_mask=None, labels=None):\n        device = input_ids.device\n        # b, t, c = x.size()  # (batch_size, seq_len, feat_dim)\n        b, t = input_ids.size()  # (batch_size, longest_seq_len)\n        assert t <= self.block_size, f\"Cannot forward sequence of length {t}, block size is only {self.config.block_size}\"\n        pos = torch.arange(0, t, dtype=torch.long, device=device) # shape (t)\n\n         # forward the GPT model itself\n        tok_emb = self.transformer.wte(input_ids) # token embeddings of shape (b, t, n_embd)\n        pos_emb = self.transformer.wpe(pos) # position embeddings of shape (t, n_embd)\n        x = self.transformer.drop(tok_emb + pos_emb)\n        for block in self.transformer.h:\n            x_prev = x\n            x = block(x, attention_mask)\n        x = self.transformer.ln_f(x)\n\n        if attention_mask is not None:\n            sentence_lens = attention_mask.sum(axis=-1)\n            x_cls = x[\n                torch.arange(x.shape[0], device=device),\n                sentence_lens - 1,\n            ]  # (batch_size, seq_len, emb_dim) => (batch_size, emb_dim)\n        else:\n            x_cls = x[:, -1, :]  # (batch_size, seq_len, emb_dim) => (batch_size, emb_dim)\n        # x_cls = x.mean(dim=1)\n        if labels is not None:\n            # if we are given some desired targets also calculate the loss\n            pred_logits = self.cls_head(x_cls)  # (batch_size, emb_dim) => (batch_size, n_classes)\n            loss = self._loss_fn(pred_logits, labels.view(-1))\n            return pred_logits, loss\n        else:\n            # inference-time mini-optimization: only forward the lm_head on the very last position\n            pred_logits = self.cls_head(x_cls)  # (batch_size, emb_dim) => (batch_size, n_classes)\n            return pred_logits\n\n    def crop_block_size(self, block_size):\n        # model surgery to decrease the block size if necessary\n        # e.g. we may load the GPT2 pretrained model checkpoint (block size 1024)\n        # but want to use a smaller block size for some smaller, simpler model\n        assert block_size <= self.config.block_size\n        self.config.block_size = block_size\n        self.transformer.wpe.weight = nn.Parameter(self.transformer.wpe.weight[:block_size])\n        for block in self.transformer.h:\n            if hasattr(block.attn, 'bias'):\n                block.attn.bias = block.attn.bias[:,:,:block_size,:block_size]\n\n    def configure_optimizers(self, weight_decay, learning_rate, betas, device_type):\n        # start with all of the candidate parameters\n        param_dict = {pn: p for pn, p in self.named_parameters()}\n        # filter out those that do not require grad\n        param_dict = {pn: p for pn, p in param_dict.items() if p.requires_grad}\n        # create optim groups. Any parameters that is 2D will be weight decayed, otherwise no.\n        # i.e. all weight tensors in matmuls + embeddings decay, all biases and layernorms don't.\n        decay_params = [p for n, p in param_dict.items() if p.dim() >= 2]\n        nodecay_params = [p for n, p in param_dict.items() if p.dim() < 2]\n        optim_groups = [\n            {'params': decay_params, 'weight_decay': weight_decay},\n            {'params': nodecay_params, 'weight_decay': 0.0}\n        ]\n        num_decay_params = sum(p.numel() for p in decay_params)\n        num_nodecay_params = sum(p.numel() for p in nodecay_params)\n        print(f\"num decayed parameter tensors: {len(decay_params)}, with {num_decay_params:,} parameters\")\n        print(f\"num non-decayed parameter tensors: {len(nodecay_params)}, with {num_nodecay_params:,} parameters\")\n        # Create AdamW optimizer and use the fused version if it is available\n        fused_available = 'fused' in inspect.signature(torch.optim.AdamW).parameters\n        use_fused = fused_available and device_type == 'cuda'\n        extra_args = dict(fused=True) if use_fused else dict()\n        optimizer = torch.optim.AdamW(optim_groups, lr=learning_rate, betas=betas, **extra_args)\n        print(f\"using fused AdamW: {use_fused}\")\n\n        return optimizer","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:07:26.420198Z","iopub.execute_input":"2024-06-07T04:07:26.420606Z","iopub.status.idle":"2024-06-07T04:07:26.470189Z","shell.execute_reply.started":"2024-06-07T04:07:26.420575Z","shell.execute_reply":"2024-06-07T04:07:26.468748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### VQSpecGPTClassifier","metadata":{}},{"cell_type":"code","source":"def append_eos_token(x, eos_token, device=device):\n    return torch.cat([x, torch.ones(*x.shape[:-1], 1).long().to(device) * eos_token], dim=-1)\n\n\nclass VQSpecGPTClassifier(nn.Module):\n\n    def __init__(\n        self,\n        spec_vq_encoder,\n        gpt_cls,\n        eos_token,\n    ):\n        super().__init__()\n        self._spec_vq_encoder = spec_vq_encoder\n        self._gpt_cls = gpt_cls\n        self._eos_token = eos_token\n\n    def forward(self, spec, labels=None):\n        q_emb, embed_ind, q_emb_loss = self._spec_vq_encoder(spec)\n        embed_ind = append_eos_token(embed_ind, self._eos_token)\n        gpt_res = gpt_cls(input_ids=embed_ind, labels=labels, attention_mask=None)\n        if labels is not None:\n            logits, cls_loss = gpt_res\n        else:\n            logits = gpt_res\n            cls_loss = None\n        return logits, cls_loss, q_emb_loss, embed_ind, q_emb","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:07:30.922550Z","iopub.execute_input":"2024-06-07T04:07:30.922939Z","iopub.status.idle":"2024-06-07T04:07:30.931592Z","shell.execute_reply.started":"2024-06-07T04:07:30.922904Z","shell.execute_reply":"2024-06-07T04:07:30.930474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# VQ Spectorgram Encoder Configs\nVQ_INPUT_DIM = 256*2\nVQ_CODEBOOK_SIZE = 1024\nVQ_CODEBOOK_DIM = 256*2\n\n# NanoGPT Configs\nEOS_TOKEN = VQ_CODEBOOK_SIZE\nPAD_TOKEN = VQ_CODEBOOK_SIZE + 1\nMAX_SEQ_LEN: int = 100\nBLOCK_SIZE: int = MAX_SEQ_LEN\nN_LAYER: int = 6\nN_HEAD: int = 8\nEMB_DIM: int = 768\nDROPOUT: float = 0.05\nBIAS: bool = False\nVOCAB_SiZE: int = VQ_CODEBOOK_SIZE + 2  # + EOS_TOKEN + PAD_TOKEN\nN_CLASSES: int = len(SPECIES)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:09:43.743629Z","iopub.execute_input":"2024-06-07T04:09:43.744042Z","iopub.status.idle":"2024-06-07T04:09:43.751140Z","shell.execute_reply.started":"2024-06-07T04:09:43.744008Z","shell.execute_reply":"2024-06-07T04:09:43.749854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spec_vq_encoder = VQSpectrogramEncoder(\n    spec_n_mels=n_mels,\n    vq_input_dim=VQ_INPUT_DIM,\n    codebook_size=VQ_CODEBOOK_SIZE,\n    codebook_dim=VQ_CODEBOOK_DIM\n).to(device).eval()\n\ngpt_cls = GPTClassifier(\n    n_layer=N_LAYER,\n    n_head=N_HEAD,\n    emb_dim=EMB_DIM,\n    dropout=DROPOUT,\n    bias=BIAS,\n    vocab_size=VOCAB_SiZE,\n    block_size=BLOCK_SIZE,\n    n_classes=N_CLASSES,\n#     loss_fn=FocalLoss()\n).to(device).eval()\n\nmodel = VQSpecGPTClassifier(\n    spec_vq_encoder=spec_vq_encoder,\n    gpt_cls=gpt_cls,\n    eos_token=EOS_TOKEN,\n).to(device).eval()\n\nmodel.load_state_dict(model_state_dict)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:09:44.739510Z","iopub.execute_input":"2024-06-07T04:09:44.739977Z","iopub.status.idle":"2024-06-07T04:09:45.678145Z","shell.execute_reply.started":"2024-06-07T04:09:44.739941Z","shell.execute_reply":"2024-06-07T04:09:45.677109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## prediction","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef inference_single_file(filepath) -> pd.DataFrame:\n    # audio file processing\n    waveform, sr = torchaudio.load(filepath, normalize=True)\n    duration_sec = len(waveform.squeeze()) / sr\n    if duration_sec < 250:\n        end_secs = list(range(5, int(duration_sec) + 1, 5))\n    else:\n        end_secs = list(range(5, 241, 5))\n    bs = len(end_secs)\n\n    # cut waveform by 5sec\n    waveforms = torch.cat(\n        [waveform[:, int((end_sec-5)*sr):int(end_sec*sr)] for end_sec in end_secs],\n        dim=0\n    )  # (1, duration_sec x sr) => (bs, 5 x sr)\n    \n    # convert to spectrogram\n    specs = sepc_trans(waveforms.to(device))  # (bs, 5 x sr) => (bs, n_mels, spec_time_bins+1)\n    specs = specs[:, :, :spec_time_bins]  # (bs, n_mels, spec_time_bins+1) => (bs, n_mels, spec_time_bins)\n    specs = torchaudio.functional.amplitude_to_DB(\n        specs, multiplier=10, amin=1e-7, db_multiplier=0.0\n    )\n    specs_flat = rearrange(specs, 'b nm tb -> b (nm tb)')\n    specs = specs - torch.min(specs_flat, dim=1, keepdim=True)[0].unsqueeze(-1)\n    specs /= (torch.max(specs_flat, dim=1, keepdim=True)[0].unsqueeze(-1) - torch.min(specs_flat, dim=1, keepdim=True)[0].unsqueeze(-1))\n    specs -= 0.5\n    assert specs.shape == (len(end_secs), n_mels, spec_time_bins)\n    \n    # prediction by model\n    pred_logits, _, _, _, _ = model(specs)  # (bs, len(SPECIES))\n    pred_probs = pt_F.softmax(pred_logits, dim=-1).cpu().numpy()\n    df_pred_probs = pd.DataFrame(data=pred_probs, columns=SPECIES)\n\n    # prepare row_id DataFrame\n    filename_prefix_id = filepath.split('/')[-1].replace('.ogg', '')  # xxxxxx for unlabeled_soundscapes. soundscape_xxxxxx for test_saundscapes\n    df_row_ids = pd.Series([f'{filename_prefix_id}_{end_time}' for end_time in end_secs]).to_frame('row_id')\n\n    assert df_pred_probs.shape == (len(df_row_ids), len(SPECIES))\n    return pd.concat([df_row_ids, df_pred_probs], axis=1)\n\n\ndef rand_mock_inference(filepath) -> pd.DataFrame:\n    # audio file processing\n    end_secs = list(range(5, 241, 5))\n\n    # prepare row_id DataFrame\n    filename_prefix_id = filepath.split('/')[-1].replace('.ogg', '')  # xxxxxx for unlabeled_soundscapes. soundscape_xxxxxx for test_saundscapes\n    df_row_ids = pd.Series([f'{filename_prefix_id}_{end_time}' for end_time in end_secs]).to_frame('row_id')\n    \n    # rand pred\n    df_pred_probs = pd.DataFrame(data=pt_F.softmax(torch.rand(len(df_row_ids), len(SPECIES)), dim=-1), columns=SPECIES)\n\n    assert df_pred_probs.shape == (len(df_row_ids), len(SPECIES))\n    return pd.concat([df_row_ids, df_pred_probs], axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:09:46.092554Z","iopub.execute_input":"2024-06-07T04:09:46.093481Z","iopub.status.idle":"2024-06-07T04:09:46.109266Z","shell.execute_reply.started":"2024-06-07T04:09:46.093441Z","shell.execute_reply":"2024-06-07T04:09:46.108074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\ninference_single_file('/kaggle/input/birdclef-2024/unlabeled_soundscapes/1000308629.ogg').shape","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:09:46.642134Z","iopub.execute_input":"2024-06-07T04:09:46.642530Z","iopub.status.idle":"2024-06-07T04:09:49.813311Z","shell.execute_reply.started":"2024-06-07T04:09:46.642499Z","shell.execute_reply":"2024-06-07T04:09:49.812066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# TIMEOUT_SECS = 110 * 60\nTIMEOUT_SECS = 1e+8\n\ndf_pred_list = []\nskipped_files = []\n\nfor filepath in tqdm(df_test_files['file_path']):\n    elapsed_secs = (datetime.now() - start_dt).total_seconds()\n    if elapsed_secs > TIMEOUT_SECS:\n        # timeout\n        df_pred_list.append(rand_mock_inference(filepath))\n        skipped_files.append(filepath)\n    else:\n        df_pred_list.append(inference_single_file(filepath))\n\ndf_pred = pd.concat(df_pred_list, axis=0).reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:08:09.330699Z","iopub.execute_input":"2024-06-07T04:08:09.331065Z","iopub.status.idle":"2024-06-07T04:09:35.447467Z","shell.execute_reply.started":"2024-06-07T04:08:09.331024Z","shell.execute_reply":"2024-06-07T04:09:35.446303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(skipped_files)","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:09:51.836868Z","iopub.execute_input":"2024-06-07T04:09:51.837254Z","iopub.status.idle":"2024-06-07T04:09:51.844986Z","shell.execute_reply.started":"2024-06-07T04:09:51.837221Z","shell.execute_reply":"2024-06-07T04:09:51.843446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_pred","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:09:52.036912Z","iopub.execute_input":"2024-06-07T04:09:52.037323Z","iopub.status.idle":"2024-06-07T04:09:52.069572Z","shell.execute_reply.started":"2024-06-07T04:09:52.037289Z","shell.execute_reply":"2024-06-07T04:09:52.068465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## submission","metadata":{}},{"cell_type":"code","source":"df_pred.to_csv('submission.csv', index=False)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"END","metadata":{}},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"code","source":"df_pred[SPECIES].to_numpy().argmax(-1)[:100]","metadata":{"execution":{"iopub.status.busy":"2024-06-07T04:09:55.550561Z","iopub.execute_input":"2024-06-07T04:09:55.551018Z","iopub.status.idle":"2024-06-07T04:09:55.566205Z","shell.execute_reply.started":"2024-06-07T04:09:55.550986Z","shell.execute_reply":"2024-06-07T04:09:55.564833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}