#!/usr/bin/env python3
"""
Whisper Large-V3 Fine-Tuner for NVIDIA GPUs
Optimized for VRAM efficiency with LoRA adapter training
"""

import torch
import json
import os
from pathlib import Path
from torch.utils.data import DataLoader, Dataset
from transformers import WhisperProcessor, WhisperForConditionalGeneration, get_linear_schedule_with_warmup
from peft import get_peft_model, LoraConfig, TaskType
from dataclasses import dataclass
import soundfile as sf
import librosa
import argparse
import logging

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

@dataclass
class TrainingConfig:
    model_id: str = "./whisper-large-v3"
    output_dir: str = "./whisper-finetuned"
    learning_rate: float = 1e-4
    num_epochs: int = 3
    batch_size: int = 8  # Adjust based on your GPU VRAM
    gradient_accumulation_steps: int = 4  # Effective batch = 8 * 4 = 32
    warmup_steps: int = 500
    save_steps: int = 1000
    eval_steps: int = 1000
    max_audio_length: int = 30  # seconds
    lora_rank: int = 32
    lora_alpha: int = 64
    
class AudioDataset(Dataset):
    """Load audio files from directory or CSV data manifest"""
    
    def __init__(self, data_dir, processor, max_duration=30):
        self.processor = processor
        self.max_duration = max_duration
        self.audio_files = []
        self.transcripts = []
        self.sample_rate = 16000
        
        # Simple: load all .wav and .mp3 files from data_dir
        data_path = Path(data_dir)
        if data_path.is_dir():
            audio_extensions = {'.wav', '.mp3', '.ogg', '.flac', '.m4a'}
            for ext in audio_extensions:
                self.audio_files.extend(data_path.glob(f"*{ext}"))
        
        logger.info(f"Found {len(self.audio_files)} audio files")
        
        # If CSV with transcriptions exists, load them
        csv_path = data_path / "transcriptions.csv"
        if csv_path.exists():
            with open(csv_path) as f:
                for line in f:
                    parts = line.strip().split(",", 1)
                    if len(parts) == 2:
                        filename, text = parts
                        self.transcripts.append(text)
    
    def __len__(self):
        return len(self.audio_files)
    
    def __getitem__(self, idx):
        audio_path = self.audio_files[idx]
        
        # Load audio at 16kHz
        speech, sr = librosa.load(str(audio_path), sr=self.sample_rate)
        
        # Truncate to max_duration
        max_samples = int(self.sample_rate * self.max_duration)
        if len(speech) > max_samples:
            speech = speech[:max_samples]
        
        # Process for model
        input_features = self.processor(
            speech,
            sampling_rate=self.sample_rate,
            return_tensors="pt"
        ).input_features[0]
        
        # If no transcription, use empty string
        # (In real training, you'd want your Whisper to generate transcripts first, or have manual ones)
        transcript = self.transcripts[idx] if idx < len(self.transcripts) else ""
        
        # Encode text
        labels = self.processor.tokenizer(text=transcript).input_ids
        
        return {
            "input_features": input_features,
            "labels": torch.tensor(labels, dtype=torch.long)
        }

def setup_lora_model(model, config):
    """Configure LoRA adapters for efficient fine-tuning"""
    lora_config = LoraConfig(
        r=config.lora_rank,
        lora_alpha=config.lora_alpha,
        target_modules=["q_proj", "v_proj"],  # Whisper attention layers
        lora_dropout=0.1,
        bias="none",
        task_type=TaskType.SEQ_2_SEQ_LM,
    )
    
    model = get_peft_model(model, lora_config)
    model.print_trainable_parameters()
    return model

def train(config: TrainingConfig, data_dir: str):
    """Train whisper model with LoRA"""
    
    # Device
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    logger.info(f"Using device: {device}")
    
    # Load processor and model
    logger.info(f"Loading model from {config.model_id}")
    processor = WhisperProcessor.from_pretrained(config.model_id)
    model = WhisperForConditionalGeneration.from_pretrained(config.model_id).to(device)
    
    # Enable gradient checkpointing to save memory
    model.gradient_checkpointing_enable()
    
    # Setup LoRA adapters
    model = setup_lora_model(model, config)
    model.to(device)
    
    # Dataset and DataLoader
    logger.info(f"Loading dataset from {data_dir}")
    dataset = AudioDataset(data_dir, processor, config.max_audio_length)
    dataloader = DataLoader(dataset, batch_size=config.batch_size, shuffle=True)
    
    # Optimizer and scheduler
    optimizer = torch.optim.AdamW(model.parameters(), lr=config.learning_rate)
    num_training_steps = len(dataloader) * config.num_epochs
    scheduler = get_linear_schedule_with_warmup(
        optimizer,
        num_warmup_steps=config.warmup_steps,
        num_training_steps=num_training_steps
    )
    
    # Training loop
    model.train()
    global_step = 0
    
    for epoch in range(config.num_epochs):
        logger.info(f"Epoch {epoch+1}/{config.num_epochs}")
        
        for step, batch in enumerate(dataloader):
            # Move to device
            input_features = batch["input_features"].to(device)
            labels = batch["labels"].to(device)
            
            # Forward pass
            outputs = model(
                input_features=input_features,
                labels=labels
            )
            loss = outputs.loss
            
            # Backward pass
            loss.backward()
            
            # Gradient accumulation
            if (step + 1) % config.gradient_accumulation_steps == 0:
                optimizer.step()
                scheduler.step()
                optimizer.zero_grad()
                global_step += 1
                
                if global_step % 10 == 0:
                    logger.info(f"Step {global_step}, Loss: {loss.item():.4f}")
                
                if global_step % config.save_steps == 0:
                    logger.info(f"Saving model at step {global_step}")
                    model.save_pretrained(f"{config.output_dir}/checkpoint-{global_step}")
                    processor.save_pretrained(f"{config.output_dir}/checkpoint-{global_step}")
    
    # Save final model
    logger.info(f"Saving final model to {config.output_dir}")
    model.save_pretrained(config.output_dir)
    processor.save_pretrained(config.output_dir)
    logger.info("Training complete!")

if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--data_dir", type=str, default="./training_data",
                       help="Directory with audio files for training")
    parser.add_argument("--model_id", type=str, default="./whisper-large-v3",
                       help="Path to pretrained Whisper model")
    parser.add_argument("--epochs", type=int, default=3)
    parser.add_argument("--batch_size", type=int, default=8)
    parser.add_argument("--learning_rate", type=float, default=1e-4)
    
    args = parser.parse_args()
    
    config = TrainingConfig(
        model_id=args.model_id,
        num_epochs=args.epochs,
        batch_size=args.batch_size,
        learning_rate=args.learning_rate
    )
    
    train(config, args.data_dir)
