# %% [code]
import os
import glob
import numpy as np
import pandas as pd
import soundfile as sf
import tensorflow as tf
from pathlib import Path
from tqdm.auto import tqdm
import librosa
# --- CONFIG ---
INPUT_DIR = Path('/kaggle/input')
DATASET_DIR = INPUT_DIR / 'competitions/birdclef-2024'
BIRDNET_MODEL_DIR = INPUT_DIR / 'models/shadiakiki1/birdnet-analyzer/tflite/birdnet_global_6k_v2.4_model_fp32-1/3'
TEST_AUDIO_DIR = DATASET_DIR / 'test_soundscapes'

# If test directory is empty, use unlabeled soundscapes for local testing
test_files = sorted(list(TEST_AUDIO_DIR.glob('*.ogg')))
if len(test_files) == 0:
    print("Hidden test set not found. Using 10 unlabeled soundscapes for local testing.")
    TEST_AUDIO_DIR = DATASET_DIR / 'unlabeled_soundscapes'
    test_files = sorted(list(TEST_AUDIO_DIR.glob('*.ogg')))[:10]
else:
    print(f"Found {len(test_files)} hidden test files.")

# --- SPECIES MAPPING ---
train_df = pd.read_csv(DATASET_DIR / 'train_metadata.csv')
TARGET_SPECIES = sorted(train_df['primary_label'].unique())
print(f"Target species count: {len(TARGET_SPECIES)}")

taxa_df = pd.read_csv(DATASET_DIR / 'eBird_Taxonomy_v2021.csv')
sci2code = dict(zip(taxa_df['SCI_NAME'].str.lower(), taxa_df['SPECIES_CODE']))

# Taxonomy-split synonyms (from Clements 2021 baseline used in BirdCLEF)
TAXONOMY_SYNONYMS = {
    "pterorhinus delesserti": "garrulax delesserti",
    "trochalopteron cachinnans": "pterorhinus cachinnans",
    "sholicola major": "brachypteryx major",
    "sholicola albiventris": "brachypteryx albiventris",
    "brachypodius priocephalus": "pycnonotus priocephalus",
    "rubigula gularis": "pycnonotus gularis",
    "argya subrufa": "turdoides subrufa",
    "iole indica": "acritillas indica",
}

# --- BIRDNET INITIALIZATION ---
interpreter = tf.lite.Interpreter(model_path=str(BIRDNET_MODEL_DIR / 'BirdNET_GLOBAL_6K_V2.4_Model_FP32.tflite'))
interpreter.allocate_tensors()
input_details = interpreter.get_input_details()[0]
output_details = interpreter.get_output_details()[0]
input_idx = input_details['index']
output_idx = output_details['index']

with open(BIRDNET_MODEL_DIR / 'BirdNET_GLOBAL_6K_V2.4_Labels.txt', 'r') as f:
    birdnet_raw_labels = f.readlines()
    
birdnet_sci_names = [line.split('_')[0].strip().lower() for line in birdnet_raw_labels]
birdnet_idx_to_target = {}

for idx, sci in enumerate(birdnet_sci_names):
    # 1. Direct match
    code = sci2code.get(sci)
    
    # 2. Try binomial (if BirdNET uses trinomial)
    if not code:
        parts = sci.split()
        if len(parts) > 2:
            binomial = " ".join(parts[:2])
            code = sci2code.get(binomial)
            
    # 3. Try synonyms
    if not code:
        # Check if the sci name in BirdNET is a synonym for an eBird name
        for ebird_sci, syn_list in TAXONOMY_SYNONYMS.items():
            if isinstance(syn_list, str): syn_list = [syn_list]
            if sci in syn_list:
                code = sci2code.get(ebird_sci)
                break
                
    if code in TARGET_SPECIES:
        birdnet_idx_to_target[idx] = TARGET_SPECIES.index(code)

print(f"Mapped {len(birdnet_idx_to_target)} BirdNET labels to competition species.")

# --- INFERENCE ---
all_preds = []

# BirdNET parameters
SR_BIRDNET = 48000
CHUNK_S = 3 # BirdNET native chunk size
STEP_S = 5  # Competition evaluation step

def predict_chunk(chunk):
    # Ensure chunk is 3s at 48kHz
    if len(chunk) < 144000:
        chunk = np.pad(chunk, (0, 144000 - len(chunk)))
    elif len(chunk) > 144000:
        chunk = chunk[:144000]
    
    interpreter.set_tensor(input_idx, np.float32(chunk)[np.newaxis, :])
    interpreter.invoke()
    return interpreter.get_tensor(output_idx)[0]

for file_path in tqdm(test_files):
    audio_id = file_path.stem
    
    # Load audio (BirdCLEF 2024 audio is 32kHz ogg)
    # Using soundfile for fast reading and librosa for resampling
    try:
        y, sr = sf.read(file_path)
        if sr != SR_BIRDNET:
            y = librosa.resample(y, orig_sr=sr, target_sr=SR_BIRDNET)
    except Exception as e:
        print(f"Error loading {file_path}: {e}")
        # Create dummy predictions if file fails to load
        for time_end in range(5, 245, 5):
            row_id = f"{audio_id}_{time_end}"
            row_dict = {'row_id': row_id}
            for sp in TARGET_SPECIES:
                row_dict[sp] = 0.0
            all_preds.append(row_dict)
        continue

    duration_s = len(y) / SR_BIRDNET
    
    # Process every 5 seconds
    for time_end in range(5, 245, 5):
        row_id = f"{audio_id}_{time_end}"
        
        # We take two 3s windows to cover the 5s window [T-5, T]
        # Window 1: [T-5, T-2]
        # Window 2: [T-3, T]
        
        offsets = [
            (time_end - 5, time_end - 2),
            (time_end - 3, time_end)
        ]
        
        prob_mapped = np.zeros(len(TARGET_SPECIES))
        
        for s_sec, e_sec in offsets:
            s = int(s_sec * SR_BIRDNET)
            e = s + int(CHUNK_S * SR_BIRDNET)
            
            if s >= 0 and s < len(y):
                chunk = y[s:e]
                b_out = predict_chunk(chunk)
                
                # Apply sigmoid as BirdNET TFLite outputs raw logits
                b_probs = 1 / (1 + np.exp(-b_out))
                
                for b_idx, t_idx in birdnet_idx_to_target.items():
                    prob_mapped[t_idx] = max(prob_mapped[t_idx], b_probs[b_idx])
        
        row_dict = {'row_id': row_id}
        for idx, sp in enumerate(TARGET_SPECIES):
            row_dict[sp] = prob_mapped[idx]
        all_preds.append(row_dict)

# --- SAVE SUBMISSION ---
sub_df = pd.DataFrame(all_preds)
sub_df.to_csv('submission.csv', index=False)
print("Submission saved successfully!")
