{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":91844,"databundleVersionId":11361821,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":361259,"sourceType":"modelInstanceVersion","isSourceIdPinned":false,"modelInstanceId":300361,"modelId":320917}],"dockerImageVersionId":31012,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"dad0bbf2","cell_type":"code","source":"import os\nimport math\nimport logging\nimport numpy as np\nimport pandas as pd\nimport librosa\nimport jax\nimport jax.numpy as jnp\nfrom flax import nnx\nimport orbax.checkpoint as ocp\nfrom sklearn.preprocessing import LabelEncoder\nimport tensorflow as tf\nfrom tqdm.auto import tqdm\nfrom pathlib import Path","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:32:36.933436Z","iopub.execute_input":"2025-04-28T00:32:36.933753Z","iopub.status.idle":"2025-04-28T00:32:58.501921Z","shell.execute_reply.started":"2025-04-28T00:32:36.933729Z","shell.execute_reply":"2025-04-28T00:32:58.500905Z"}},"outputs":[],"execution_count":null},{"id":"696e26dc","cell_type":"code","source":"import kagglehub","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:32:58.5034Z","iopub.execute_input":"2025-04-28T00:32:58.504287Z","iopub.status.idle":"2025-04-28T00:32:58.509315Z","shell.execute_reply.started":"2025-04-28T00:32:58.504259Z","shell.execute_reply":"2025-04-28T00:32:58.508279Z"}},"outputs":[],"execution_count":null},{"id":"ae54b71e","cell_type":"code","source":"class CFG:\n    \"\"\"\n    Configuration for BirdCLEF-2025 inference pipeline.\n    \"\"\"\n    test_soundscapes = '/kaggle/input/birdclef-2025/test_soundscapes'\n    submission_csv    = '/kaggle/input/birdclef-2025/sample_submission.csv'\n    taxonomy_csv      = '/kaggle/input/birdclef-2025/taxonomy.csv'\n    model_path        = '/kaggle/input/birdclef-cnn-baseline/flax/default/1'\n\n    FS          = 32000       # Sampling rate\n    WINDOW_SIZE = 5           # Segment duration (seconds)\n    N_MELS      = 128         # Mel bands\n    HOP_LENGTH  = 512         # STFT hop length\n    N_FRAMES    = math.ceil((WINDOW_SIZE * FS) / HOP_LENGTH)  # Time frames per segment\n\n    BATCH_SIZE  = 32          # Inference batch size\n\n# print(\"Downloading model via kagglehub...\")\n# model_path = kagglehub.model_download(\"nikhilpaleti/birdclef-cnn-baseline\")\n# print(f\"Model downloaded to: {model_path}\")\n# CFG.model_path = model_path\n\n# Check if all data is available\nprint(f\"Test soundscapes directory exists: {os.path.exists(CFG.test_soundscapes)}\")\nprint(f\"Taxonomy CSV exists: {os.path.exists(CFG.taxonomy_csv)}\")\nprint(f\"Submission CSV exists: {os.path.exists(CFG.submission_csv)}\")\nprint(f\"Model path exists: {os.path.exists(CFG.model_path)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:35:16.380124Z","iopub.execute_input":"2025-04-28T00:35:16.38045Z","iopub.status.idle":"2025-04-28T00:35:16.38917Z","shell.execute_reply.started":"2025-04-28T00:35:16.380427Z","shell.execute_reply":"2025-04-28T00:35:16.388102Z"}},"outputs":[],"execution_count":null},{"id":"1443f074","cell_type":"code","source":"def setup_logging():\n    logging.basicConfig(\n        format='%(asctime)s %(levelname)s: %(message)s',\n        level=logging.INFO\n    )\n    \ndef build_label_encoder(taxonomy_csv: str):\n    logging.info(f'Loading taxonomy from {taxonomy_csv}')\n    df = pd.read_csv(taxonomy_csv)\n    labels = sorted(df['primary_label'].astype(str).unique())\n    le = LabelEncoder().fit(labels)\n    num_classes = len(le.classes_)\n    logging.info(f'Found {num_classes} classes')\n    return le, num_classes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:35:16.589602Z","iopub.execute_input":"2025-04-28T00:35:16.589994Z","iopub.status.idle":"2025-04-28T00:35:16.596966Z","shell.execute_reply.started":"2025-04-28T00:35:16.589969Z","shell.execute_reply":"2025-04-28T00:35:16.595883Z"}},"outputs":[],"execution_count":null},{"id":"871ca75e","cell_type":"code","source":"def process_audio_segment(y: np.ndarray, cfg: CFG) -> np.ndarray:\n    \"\"\"\n    Compute log-mel spectrogram for a fixed-length audio segment.\n    \"\"\"\n    S = librosa.feature.melspectrogram(\n        y=y, sr=cfg.FS, n_mels=cfg.N_MELS,\n        hop_length=cfg.HOP_LENGTH, fmax=cfg.FS // 2\n    )\n    logS = librosa.power_to_db(S, ref=np.max)\n    if logS.shape[1] < cfg.N_FRAMES:\n        pad = cfg.N_FRAMES - logS.shape[1]\n        logS = np.pad(\n            logS,\n            ((0,0),(0,pad)),\n            mode='constant',\n            constant_values=logS.min()\n        )\n    else:\n        logS = logS[:, :cfg.N_FRAMES]\n    return logS.astype(np.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:35:16.804915Z","iopub.execute_input":"2025-04-28T00:35:16.805235Z","iopub.status.idle":"2025-04-28T00:35:16.81233Z","shell.execute_reply.started":"2025-04-28T00:35:16.805206Z","shell.execute_reply":"2025-04-28T00:35:16.811139Z"}},"outputs":[],"execution_count":null},{"id":"5f498137","cell_type":"code","source":"class AudioCNN(nnx.Module):\n    \"\"\"Basic CNN for log-mel spectrograms.\"\"\"\n    def __init__(self, num_classes: int, rngs: nnx.Rngs):\n        self.conv1 = nnx.Conv(1, 32, kernel_size=(3,3), rngs=rngs)\n        self.conv2 = nnx.Conv(32, 64, kernel_size=(3,3), rngs=rngs)\n        self.pool  = lambda x: nnx.avg_pool(x, window_shape=(2,2), strides=(2,2))\n        self.dense = nnx.Linear(159744, 128, rngs=rngs)\n        self.out   = nnx.Linear(128, num_classes, rngs=rngs)\n\n    def __call__(self, x: jnp.ndarray) -> jnp.ndarray:\n        x = self.pool(nnx.relu(self.conv1(x)))\n        x = self.pool(nnx.relu(self.conv2(x)))\n        x = x.reshape(x.shape[0], -1)\n        x = nnx.sigmoid(self.dense(x))\n        return self.out(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:35:17.005244Z","iopub.execute_input":"2025-04-28T00:35:17.005596Z","iopub.status.idle":"2025-04-28T00:35:17.01412Z","shell.execute_reply.started":"2025-04-28T00:35:17.005573Z","shell.execute_reply":"2025-04-28T00:35:17.012743Z"}},"outputs":[],"execution_count":null},{"id":"d27fbfd5","cell_type":"code","source":"def load_model(cfg: CFG, num_classes: int) -> nnx.Module:\n    logging.info('Restoring model checkpoint')\n    # Create abstract model to get graphdef and state spec\n    abstract = nnx.eval_shape(lambda: AudioCNN(num_classes, rngs=nnx.Rngs(0)))\n    graphdef, abstract_state = nnx.split(abstract)\n    ckpt = ocp.StandardCheckpointer()\n    restored_state = ckpt.restore(\n        os.path.join(cfg.model_path, 'model_state'),\n        abstract_state\n    )\n    model = nnx.merge(graphdef, restored_state)\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:35:17.216274Z","iopub.execute_input":"2025-04-28T00:35:17.216603Z","iopub.status.idle":"2025-04-28T00:35:17.223103Z","shell.execute_reply.started":"2025-04-28T00:35:17.216582Z","shell.execute_reply":"2025-04-28T00:35:17.222042Z"}},"outputs":[],"execution_count":null},{"id":"4a8f35ab","cell_type":"code","source":"def to_tf_inference_dataset(audio_paths, cfg: CFG):\n    \"\"\"\n    Build a tf.data.Dataset yielding (spectrogram, row_id) batches.\n    \"\"\"\n    def gen():\n        for path in audio_paths:\n            y, _ = librosa.load(str(path), sr=cfg.FS)\n            seg_len = cfg.FS * cfg.WINDOW_SIZE\n            n_segs = math.ceil(len(y) / seg_len)\n            soundscape = path.stem\n            for i in range(n_segs):\n                start = i * seg_len\n                end = start + seg_len\n                seg = y[start:end]\n                if len(seg) < seg_len:\n                    seg = np.pad(seg, (0, seg_len - len(seg)), mode='constant')\n                logS = process_audio_segment(seg, cfg)\n                row_id = f\"{soundscape}_{(i+1)*cfg.WINDOW_SIZE}\"\n                yield logS[..., None], row_id\n\n    output_signature = (\n        tf.TensorSpec((cfg.N_MELS, cfg.N_FRAMES, 1), tf.float32),\n        tf.TensorSpec((), tf.string)\n    )\n    ds = tf.data.Dataset.from_generator(\n        gen,\n        output_signature=output_signature\n    )\n    return ds.batch(cfg.BATCH_SIZE).prefetch(tf.data.AUTOTUNE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:35:17.38861Z","iopub.execute_input":"2025-04-28T00:35:17.389005Z","iopub.status.idle":"2025-04-28T00:35:17.39748Z","shell.execute_reply.started":"2025-04-28T00:35:17.388979Z","shell.execute_reply":"2025-04-28T00:35:17.396427Z"}},"outputs":[],"execution_count":null},{"id":"9d8b05f6","cell_type":"code","source":"def inference(cfg: CFG) -> pd.DataFrame:\n    setup_logging()\n    le, num_classes = build_label_encoder(cfg.taxonomy_csv)\n    classes = le.classes_.tolist()\n\n    model = load_model(cfg, num_classes)\n\n    # Prepare test files and dataset\n    audio_paths = sorted(Path(cfg.test_soundscapes).glob('*.ogg'))\n    logging.info(f'Found {len(audio_paths)} test soundscape files')\n    ds_inf = to_tf_inference_dataset(audio_paths, cfg)\n\n    all_row_ids = []\n    all_preds   = []\n\n    # Iterate batches\n    for specs_batch, ids_batch in tqdm(ds_inf.as_numpy_iterator(), desc='Inference'):  # specs_batch: (B,H,W,1), ids_batch: (B,)\n        # Run model\n        logits = model(jnp.array(specs_batch))\n        probs = jax.nn.sigmoid(logits)\n        probs_np = np.array(probs)\n        # Collect\n        for rid, p in zip(ids_batch, probs_np):\n            # rid from tf may be bytes\n            if isinstance(rid, bytes):\n                rid = rid.decode('utf-8')\n            all_row_ids.append(rid)\n            all_preds.append(p)\n\n    preds_arr = np.stack(all_preds, axis=0)\n    # Build DataFrame\n    df = pd.DataFrame(preds_arr, columns=classes)\n    df.insert(0, 'row_id', all_row_ids)\n\n    # Align with submission template\n    template = pd.read_csv(cfg.submission_csv)\n    df = (\n        df.set_index('row_id')\n          .reindex(template['row_id'])\n          .fillna(0.0)\n          .reset_index()\n    )\n    # Ensure all species columns\n    for col in template.columns[1:]:\n        if col not in df.columns:\n            df[col] = 0.0\n    df = df[template.columns]\n\n    # Save\n    df.to_csv('submission.csv', index=False)\n    logging.info('Saved submission.csv')\n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:35:17.542574Z","iopub.execute_input":"2025-04-28T00:35:17.542901Z","iopub.status.idle":"2025-04-28T00:35:17.553012Z","shell.execute_reply.started":"2025-04-28T00:35:17.542881Z","shell.execute_reply":"2025-04-28T00:35:17.551961Z"}},"outputs":[],"execution_count":null},{"id":"1fcb4676","cell_type":"code","source":"if __name__ == '__main__':\n    cfg = CFG()\n    submission_df = inference(cfg)\n    print(submission_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T00:35:22.00382Z","iopub.execute_input":"2025-04-28T00:35:22.00415Z","iopub.status.idle":"2025-04-28T00:35:23.818961Z","shell.execute_reply.started":"2025-04-28T00:35:22.00413Z","shell.execute_reply":"2025-04-28T00:35:23.817578Z"}},"outputs":[],"execution_count":null},{"id":"cfe61157-3bb6-41c5-9d0e-326032f6e439","cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}