#!/usr/bin/env python3
"""Full audio-tower FT: one S8 LR ascent, then two Gemini cosine-decay epochs."""
from __future__ import annotations

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

import numpy as np
import torch
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler, Subset

HERE = Path(__file__).resolve().parent
CODE = HERE.parent
sys.path.insert(0, str(CODE))
sys.path.insert(0, str(CODE / 'embedding_probe_study'))
from study_paths import read
from train_full_curriculum import TIMBRE_INDEX, IDENTITY_INDEX, fill_precomputed
from train_layered_curriculum import loss_terms
from timbre_targets import TimbreIndex
from clap_data import CLAPCollator, make_datasets
from model import CLAPMultiTask

ROOT = Path('/e/scratch/reformo/schuhmann1_moss/whisper_score_regression/clap_v2_full_ft_20261007')
NUMERIC = ('patches', 'patch_coord', 'patch_valid', 'mel20', 'scores', 'score_mask',
           'frame', 'frame_mask', 'timbre', 'identity', 'timbre_valid',
           'identity_valid', 'cps', 'cps_raw', 'cps_valid', 'event_starts',
           'event_ends', 'event_classes', 'event_valid', 'event_span_valid')


def batch_to_device(batch, device):
    return {key: batch[key].to(device, non_blocking=True) for key in NUMERIC}


def forward(core, batch):
    patches = {key: batch[key] for key in ('patches', 'patch_coord', 'patch_valid')}
    return core(patches, batch['mel20'], batch['event_starts'], batch['event_ends'])


def loader(dataset, collator, batch_size, workers, rank, world, seed, train, epoch=0):
    if train:
        sampler = DistributedSampler(dataset, num_replicas=world, rank=rank,
                                     shuffle=True, seed=seed)
        sampler.set_epoch(epoch)
    else:
        dataset = Subset(dataset, range(rank, len(dataset), world))
        sampler = None
    options = dict(batch_size=batch_size, sampler=sampler, shuffle=False,
                   num_workers=workers, pin_memory=True, collate_fn=collator)
    if workers:
        options.update(persistent_workers=True, multiprocessing_context='spawn')
    return DataLoader(dataset, **options)


@torch.no_grad()
def evaluate(core, data, collator, batch_size, workers, rank, world, device, weights):
    core.eval()
    batches = loader(data, collator, batch_size, workers, rank, world, 0, False)
    totals = torch.zeros(6, dtype=torch.float64, device=device)
    for raw in batches:
        batch = batch_to_device(raw, device)
        with torch.autocast('cuda', dtype=torch.bfloat16):
            output = forward(core, batch)
            loss = sum(loss_terms(output, batch, weights).values())
        n = len(batch['scores'])
        totals[0] += loss.float() * n
        totals[1] += n
        valid = batch['frame_mask'].bool()
        pred = output['frame'].float().sigmoid() >= .5
        truth = batch['frame'].bool()
        totals[2] += (pred & truth & valid).sum()
        totals[3] += (pred & ~truth & valid).sum()
        totals[4] += (~pred & truth & valid).sum()
        events = batch['event_valid'].bool()
        totals[5] += (output['event_class'].argmax(-1)[events] == batch['event_classes'][events]).sum()
    if world > 1:
        torch.distributed.all_reduce(totals)
    return {'loss': float(totals[0] / totals[1].clamp_min(1)),
            'clips': int(totals[1]),
            'frame_f1_0p5': float(2 * totals[2] / (2 * totals[2] + totals[3] + totals[4]).clamp_min(1)),
            'correct_events': int(totals[5])}


def save(out, core, optimizer, stage, epoch, next_batch, update, best, config, rank, world):
    rng = {'torch': torch.get_rng_state(), 'cuda': torch.cuda.get_rng_state(),
           'numpy': np.random.get_state(), 'python': random.getstate()}
    gathered = [None] * world
    if world > 1:
        torch.distributed.all_gather_object(gathered, rng)
    else:
        gathered[0] = rng
    if rank == 0:
        state = {'model': core.state_dict(), 'optimizer': optimizer.state_dict(),
                 'scheduler': {'type': 'S8 linear rise; Gemini 2-epoch cosine',
                               'completed_updates': update},
                 'stage': stage, 'epoch': epoch, 'next_batch': next_batch,
                 'update': update, 'best_gemini_validation_loss': best,
                 'world_size': world, 'rank_rng': gathered, 'config': config}
        temp = out / 'latest.pt.incomplete'
        torch.save(state, temp)
        os.replace(temp, out / 'latest.pt')
    if world > 1:
        torch.distributed.barrier()


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--model', choices=('clapv2_xs', 'clapv2_m'), required=True)
    parser.add_argument('--resume', type=Path)
    parser.add_argument('--smoke', action='store_true')
    args = parser.parse_args()
    study = read(Path('/e/scratch/reformo/schuhmann1_moss/whisper_score_regression/'
                      'embedding_probe_study_20261006/study.json'))
    spec = next(m for m in study['models'] if m['id'] == args.model)
    rank = int(os.environ.get('RANK', 0))
    world = int(os.environ.get('WORLD_SIZE', 1))
    local = int(os.environ.get('LOCAL_RANK', 0))
    torch.cuda.set_device(local)
    if world > 1:
        torch.distributed.init_process_group('nccl')
    device = torch.device('cuda', local)
    seed = 20261007
    torch.manual_seed(seed + rank)
    np.random.seed(seed + rank)
    random.seed(seed + rank)
    data, norm, taxonomy = make_datasets()
    collator = CLAPCollator(spec['model_name'], norm, taxonomy)
    core = CLAPMultiTask(spec, len(taxonomy['names'])).to(device)
    core.requires_grad_(True)
    if not all(p.requires_grad for p in core.audio.parameters()):
        raise RuntimeError('The complete CLAP audio tower must be trainable')
    peak_audio = 2e-5 if args.model == 'clapv2_xs' else 1e-5
    peak_head = 2e-4
    optimizer = torch.optim.AdamW([
        {'params': list(core.audio.parameters()), 'lr': peak_audio, 'peak_lr': peak_audio},
        {'params': [p for name, p in core.named_parameters() if not name.startswith('audio.')],
         'lr': peak_head, 'peak_lr': peak_head}], weight_decay=.01)
    for group in optimizer.param_groups:
        group['lr'] = group['peak_lr'] * .05
    model = DDP(core, device_ids=[local], find_unused_parameters=False) if world > 1 else core
    batch_size, accumulation = (8, 2) if args.model == 'clapv2_xs' else (4, 4)
    workers = 4
    out = ROOT / args.model
    out.mkdir(parents=True, exist_ok=True)
    s8_batches = math.ceil(math.ceil(len(data['S8_train']) / world) / batch_size)
    gem_batches = math.ceil(math.ceil(len(data['Gemini_train']) / world) / batch_size)
    s8_updates = math.ceil(s8_batches / accumulation)
    gem_updates = math.ceil(gem_batches / accumulation) * 2
    config = {'source': spec, 'source_checkpoint_bytes': Path(spec['checkpoint']).stat().st_size,
              'trainable_audio_parameters': sum(p.numel() for p in core.audio.parameters()),
              'trainable_total_parameters': sum(p.numel() for p in core.parameters()),
              's8_clips': len(data['S8_train']), 's8_validation_clips': len(data['S8_validation']),
              'gemini_train_clips': len(data['Gemini_train']),
              'gemini_validation_clips': len(data['Gemini_validation']),
              'gemini_test_clips': len(data['Gemini_test']),
              'batch_per_gpu': batch_size, 'accumulation': accumulation,
              'effective_batch': batch_size * world * accumulation,
              's8_updates': s8_updates, 'gemini_updates': gem_updates,
              'audio_peak_lr': peak_audio, 'head_peak_lr': peak_head,
              'warmup_initial_lr_ratio': .05, 'minimum_lr_ratio': .1,
              'temporal_patch_seconds': .16, 'frame_seconds': .02,
              'seed': seed, 'world_size': world,
              'schedule': 'one S8 epoch linear 0.05-to-1.0; two Gemini epochs cosine 1.0-to-0.1',
              'objective': 'same 192 masked Huber, CPS Huber, Orange cosine+Huber, '
                           '20ms BCE+Dice, onset BCE, duration Huber and event CE as Whisper'}
    if rank == 0:
        (out / 'config.json').write_text(json.dumps(config, indent=2) + '\n')
    weights = torch.as_tensor(taxonomy['class_weights'], device=device, dtype=torch.float32)
    indices = {'timbre': TimbreIndex(TIMBRE_INDEX, 128),
               'identity': TimbreIndex(IDENTITY_INDEX, 250)}
    stages = [('S8', 0), ('Gemini', 0), ('Gemini', 1)]
    begin_stage, begin_batch, update, best = 0, 0, 0, float('inf')
    if args.resume:
        checkpoint = torch.load(args.resume, map_location='cpu', weights_only=False)
        if checkpoint['world_size'] != world or checkpoint['config']['source']['id'] != args.model:
            raise RuntimeError('Resume checkpoint model/world mismatch')
        core.load_state_dict(checkpoint['model'], strict=True)
        optimizer.load_state_dict(checkpoint['optimizer'])
        begin_stage = checkpoint['stage']
        begin_batch = checkpoint['next_batch']
        update = checkpoint['update']
        best = checkpoint['best_gemini_validation_loss']
        if update <= s8_updates:
            factor = .05 + .95 * min(1., update / s8_updates)
        else:
            progress = min(1., (update - s8_updates) / max(1, gem_updates))
            factor = .1 + .9 * .5 * (1 + math.cos(math.pi * progress))
        for group in optimizer.param_groups:
            group['lr'] = group['peak_lr'] * factor
        rng = checkpoint['rank_rng'][rank]
        torch.set_rng_state(rng['torch']); torch.cuda.set_rng_state(rng['cuda'])
        np.random.set_state(rng['numpy']); random.setstate(rng['python'])
        del checkpoint
    if args.smoke:
        stages = stages[:1]
    for stage_number in range(begin_stage, len(stages)):
        name, epoch = stages[stage_number]
        dataset = data[name + '_train']
        train_loader = loader(dataset, collator, batch_size, workers, rank, world, seed,
                              True, stage_number)
        core.train()
        optimizer.zero_grad(set_to_none=True)
        for batch_index, raw in enumerate(train_loader):
            if stage_number == begin_stage and batch_index < begin_batch:
                continue
            if name == 'S8':
                fill_precomputed(raw, indices)
            batch = batch_to_device(raw, device)
            remain = len(train_loader) - (batch_index // accumulation) * accumulation
            divisor = min(accumulation, remain)
            with torch.autocast('cuda', dtype=torch.bfloat16):
                output = forward(model, batch)
                loss = sum(loss_terms(output, batch, weights).values()) / divisor
            if not torch.isfinite(loss):
                raise RuntimeError(f'Nonfinite loss at {name} {epoch} batch {batch_index}')
            loss.backward()
            if (batch_index + 1) % accumulation == 0 or batch_index + 1 == len(train_loader):
                torch.nn.utils.clip_grad_norm_(core.parameters(), 1.)
                optimizer.step()
                optimizer.zero_grad(set_to_none=True)
                update += 1
                if update <= s8_updates:
                    factor = .05 + .95 * min(1., update / s8_updates)
                else:
                    progress = min(1., (update - s8_updates) / max(1, gem_updates))
                    factor = .1 + .9 * .5 * (1 + math.cos(math.pi * progress))
                for group in optimizer.param_groups:
                    group['lr'] = group['peak_lr'] * factor
                if rank == 0 and update % 100 == 0:
                    print('TRAIN_PROGRESS', args.model, name, epoch + 1,
                          update, 'loss', round(float(loss.detach()) * divisor, 4), flush=True)
                if not args.smoke and update % 500 == 0:
                    save(out, core, optimizer, stage_number, epoch, batch_index + 1,
                         update, best, config, rank, world)
            if args.smoke and batch_index >= 1:
                break
        if args.smoke:
            if rank == 0:
                (out / 'SMOKE_OK.json').write_text(json.dumps({'model': args.model,
                    'audio_parameters': config['trainable_audio_parameters'],
                    'total_parameters': config['trainable_total_parameters'],
                    'loss': float(loss.detach()) * divisor, 'batch_per_gpu': batch_size}, indent=2) + '\n')
            break
        validation = data['S8_validation' if name == 'S8' else 'Gemini_validation']
        metrics = evaluate(core, validation, collator, batch_size, workers, rank, world,
                           device, weights)
        improved = name == 'Gemini' and metrics['loss'] < best
        if improved:
            best = metrics['loss']
        save(out, core, optimizer, stage_number + 1, 0, 0, update, best,
             config, rank, world)
        if rank == 0:
            with (out / 'metrics.jsonl').open('a') as f:
                f.write(json.dumps({'stage': name, 'epoch': epoch + 1,
                                    'update': update, 'validation': metrics,
                                    'best_gemini': improved}) + '\n')
            if improved:
                import shutil
                shutil.copy2(out / 'latest.pt', out / 'best.pt')
            print('EPOCH_DONE', args.model, name, epoch + 1, metrics, flush=True)
        if world > 1:
            torch.distributed.barrier()
        begin_batch = 0
    if not args.smoke:
        selected = torch.load(out / 'best.pt', map_location='cpu', weights_only=False)
        core.load_state_dict(selected['model'], strict=True)
        del selected
        test = evaluate(core, data['Gemini_test'], collator, batch_size, workers, rank,
                        world, device, weights)
        if rank == 0:
            (out / 'GEMINI_TEST.json').write_text(json.dumps(test, indent=2) + '\n')
            (out / 'COMPLETE.json').write_text(json.dumps({'stages': stages,
                'updates': update, 'best_gemini_validation_loss': best,
                'test': test}, indent=2) + '\n')
    if world > 1:
        torch.distributed.destroy_process_group()


if __name__ == '__main__':
    main()
