#!/usr/bin/env python3
"""Identical Flash holdout and public-benchmark protocols to the Whisper study."""
from __future__ import annotations

import argparse
import json
import math
import sys
from pathlib import Path

import numpy as np
import torch
from torch.utils.data import DataLoader, Dataset

HERE = Path(__file__).resolve().parent
CODE = HERE.parent
sys.path.insert(0, str(CODE))
sys.path.insert(0, str(CODE / 'embedding_probe_study'))
from clap_data import CLAPCollator, GEMINI_ROOT
from model import CLAPMultiTask
from gemini_finetune.data import GeminiDataset
from evaluate_public_benchmarks import decode, parquet_samples, emolia_samples
from study_data import SourceData
from matched_acted import run as acted
from predict_benchmarks import score_metrics
from study_paths import read, write
from train_probes import evaluate as full_evaluate

ROOT = Path('/e/scratch/reformo/schuhmann1_moss/whisper_score_regression/clap_v2_full_ft_20261007')
RELEASE = Path('/e/scratch/reformo/schuhmann1_moss/whisper_score_regression/hf_layered_release/model')


class PublicAudio(Dataset):
    def __init__(self, samples):
        self.samples = samples

    def __len__(self):
        return len(self.samples)

    def __getitem__(self, index):
        sample = self.samples[index]
        wav, duration, truncated = decode(sample)
        frames = min(1500, math.ceil(len(wav) / 320))
        return {'key': sample.key, 'uid': sample.key, 'domain': 'public',
                'audio': wav, 'scores': np.zeros(192, np.float32),
                'score_mask': np.zeros(192, bool),
                'frame': np.zeros(frames, np.float32),
                'frame_mask': np.zeros(frames, np.float32),
                'timbre': np.zeros(128, np.float32),
                'identity': np.zeros(250, np.float32),
                'timbre_valid': False, 'identity_valid': False,
                'cps': None, 'events_frames': [], 'event_class_ids': [],
                'duration_s': duration, 'truncated_30s': truncated}


class P3Data(Dataset):
    def __init__(self, split):
        self.source = SourceData('p3_' + split)

    def __len__(self):
        return len(self.source)

    def __getitem__(self, index):
        key, wave, _meta, targets = self.source[index]
        return {**targets, 'key': key, 'uid': key, 'domain': 'public',
                'enabled': True, 'audio': wave, 'cps': None,
                'timbre_valid': bool(targets['timbre_valid']),
                'identity_valid': bool(targets['identity_valid']),
                'events_frames': list(zip(targets['event_starts'], targets['event_ends'])),
                'event_class_ids': targets['event_classes']}


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--model', choices=('clapv2_xs', 'clapv2_m'), required=True)
    args = parser.parse_args()
    device = torch.device('cuda:0')
    torch.cuda.set_device(device)
    torch.set_num_threads(2)
    spec = next(s for s in read(Path('/e/scratch/reformo/schuhmann1_moss/'
                  'whisper_score_regression/embedding_probe_study_20261006/study.json'))['models']
                if s['id'] == args.model)
    out = ROOT / args.model
    if not (out / 'COMPLETE.json').exists():
        raise RuntimeError('Training is not complete')
    taxonomy = read(RELEASE / 'classes.json')
    norm = read(GEMINI_ROOT / 'checkpoint_normalization.json')
    core = CLAPMultiTask(spec, len(taxonomy['names']), pretrained=False).to(device).eval()
    state = torch.load(out / 'best.pt', map_location='cpu', weights_only=False)
    core.load_state_dict(state['model'], strict=True)
    collator = CLAPCollator(spec['model_name'], norm, taxonomy)
    batch_size = 8 if args.model == 'clapv2_xs' else 4

    def make_loader(data):
        return DataLoader(data, batch_size=batch_size, shuffle=False, num_workers=4,
                          multiprocessing_context='spawn', persistent_workers=True,
                          pin_memory=True, collate_fn=collator)

    def predict(model, batch):
        patches = {key: batch[key] for key in ('patches', 'patch_coord', 'patch_valid')}
        with torch.autocast('cuda', torch.bfloat16):
            return model(patches, batch['mel20'], batch['event_starts'],
                         batch['event_ends'], predict_events=True)

    for domain, dataset in [('gemini_test', GeminiDataset(GEMINI_ROOT, 'test')),
                            ('p3_test', P3Data('test'))]:
        dest = out / (domain + '_metrics.json')
        if dest.exists():
            continue
        metrics = full_evaluate(core, dataset, device,
                                torch.ones(len(taxonomy['names']), device=device),
                                True, make_loader, predict)
        metrics.update(model=args.model, phase='S8 + Gemini full audio-tower FT',
                       domain=domain, native_patch_seconds=.16,
                       frame_grid_seconds=.02,
                       boundary_head='Native patch features + original intra-patch 20 ms mel')
        write(dest, metrics)
        print('HOLDOUT_DONE', args.model, domain, metrics['clips'], flush=True)

    predictions = {}
    with torch.inference_mode():
        for kind in ('emonet', 'emolia-emo', 'emolia-dim', 'crema', 'ravdess'):
            scores_path = out / (kind + '_predictions.npz')
            rows_path = out / (kind + '_predictions.jsonl')
            if scores_path.exists() and rows_path.exists():
                scores = np.load(scores_path)['scores']
                rows = [json.loads(line) for line in rows_path.open()]
            else:
                samples = (parquet_samples(kind) if kind in ('emonet', 'crema', 'ravdess')
                           else emolia_samples(kind))
                scores_parts = []
                audio_meta = []
                for batch in make_loader(PublicAudio(samples)):
                    audio_meta.extend(zip(batch['duration_s'], batch['truncated_30s']))
                    patch = {key: batch[key].to(device) for key in
                             ('patches', 'patch_coord', 'patch_valid')}
                    mel = batch['mel20'].to(device)
                    with torch.autocast('cuda', torch.bfloat16):
                        output = core(patch, mel)
                    scores_parts.append(output['scores'].float().cpu().numpy())
                scores = np.concatenate(scores_parts).astype(np.float32)
                scores = scores * np.asarray(norm['score_std'], np.float32) + np.asarray(norm['score_mean'], np.float32)
                rows = [{**sample.metadata, 'key': sample.key,
                         'duration_s': duration, 'truncated_30s': truncated}
                        for sample, (duration, truncated) in zip(samples, audio_meta)]
                np.savez_compressed(scores_path, scores=scores)
                rows_path.write_text(''.join(json.dumps(row, ensure_ascii=False) + '\n' for row in rows))
            predictions[kind] = (scores, rows)
            print('PUBLIC_PREDICTIONS', args.model, kind, len(rows), flush=True)
    score_metrics(predictions, args.model, out)
    for kind in ('crema', 'ravdess'):
        scores, rows = predictions[kind]
        acted(kind, scores, rows, out)
    write(out / 'EVAL_COMPLETE.json', {'model': args.model,
          'public': list(predictions), 'gemini_test_clips': 3544,
          'same_score_normalization_as_whisper': True,
          'same_acted_nested_actor_cv_as_whisper': True})
    print('EVAL_COMPLETE', args.model, flush=True)


if __name__ == '__main__':
    main()
