#!/usr/bin/env python3
"""
Qwen 3.6-27B Fine-Tuner for NVIDIA GPUs
Optimized for VRAM efficiency with LoRA adapter training
Works with any text dataset (CSV, JSON, or plain text files)
"""

import torch
import json
import os
from pathlib import Path
from torch.utils.data import DataLoader, Dataset
from transformers import AutoTokenizer, AutoModelForCausalLM, get_linear_schedule_with_warmup
from peft import get_peft_model, LoraConfig, TaskType
from dataclasses import dataclass
import argparse
import logging

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

@dataclass
class TrainingConfig:
    model_id: str = "./qwen3.6-27b"
    output_dir: str = "./qwen-finetuned"
    learning_rate: float = 1e-4
    num_epochs: int = 3
    batch_size: int = 4  # Adjust based on GPU VRAM
    gradient_accumulation_steps: int = 4  # Effective batch = 4 * 4 = 16
    warmup_steps: int = 500
    save_steps: int = 1000
    eval_steps: int = 1000
    max_seq_length: int = 2048
    lora_rank: int = 64
    lora_alpha: int = 128

class TextDataset(Dataset):
    """Load text data from CSV, JSON, or plain text files"""
    
    def __init__(self, data_dir, tokenizer, max_length=2048):
        self.tokenizer = tokenizer
        self.max_length = max_length
        self.texts = []
        
        data_path = Path(data_dir)
        
        # Load from CSV (columns: text, or instruction, input, output)
        csv_files = list(data_path.glob("*.csv"))
        for csv_file in csv_files:
            try:
                import pandas as pd
                df = pd.read_csv(csv_file)
                if 'text' in df.columns:
                    self.texts.extend(df['text'].tolist())
                elif 'instruction' in df.columns and 'input' in df.columns:
                    for _, row in df.iterrows():
                        text = f"{row['instruction']}\n{row['input']}\n{row.get('output', '')}"
                        self.texts.append(text)
            except Exception as e:
                logger.warning(f"Could not load {csv_file}: {e}")
        
        # Load from JSON (list of dicts with 'text' key)
        json_files = list(data_path.glob("*.json"))
        for json_file in json_files:
            try:
                with open(json_file) as f:
                    data = json.load(f)
                    if isinstance(data, list):
                        for item in data:
                            if isinstance(item, dict) and 'text' in item:
                                self.texts.append(item['text'])
            except Exception as e:
                logger.warning(f"Could not load {json_file}: {e}")
        
        # Load from plain text files
        txt_files = list(data_path.glob("*.txt"))
        for txt_file in txt_files:
            try:
                with open(txt_file) as f:
                    self.texts.append(f.read())
            except Exception as e:
                logger.warning(f"Could not load {txt_file}: {e}")
        
        logger.info(f"Loaded {len(self.texts)} text samples")
    
    def __len__(self):
        return len(self.texts)
    
    def __getitem__(self, idx):
        text = self.texts[idx]
        
        # Tokenize
        encoding = self.tokenizer(
            text,
            max_length=self.max_length,
            truncation=True,
            padding='max_length',
            return_tensors='pt'
        )
        
        input_ids = encoding['input_ids'][0]
        labels = input_ids.clone()
        
        return {
            'input_ids': input_ids,
            'labels': labels,
            'attention_mask': encoding['attention_mask'][0],
        }

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"],  # Qwen attention layers
        lora_dropout=0.1,
        bias="none",
        task_type=TaskType.CAUSAL_LM,
    )
    
    model = get_peft_model(model, lora_config)
    model.print_trainable_parameters()
    return model

def train(config: TrainingConfig, data_dir: str):
    """Train Qwen model with LoRA"""
    
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    logger.info(f"Using device: {device}")
    
    # Load tokenizer and model
    logger.info(f"Loading model from {config.model_id}")
    tokenizer = AutoTokenizer.from_pretrained(config.model_id)
    model = AutoModelForCausalLM.from_pretrained(
        config.model_id,
        torch_dtype=torch.float16,
        device_map="auto"
    )
    
    # Enable gradient checkpointing to save memory
    model.gradient_checkpointing_enable()
    
    # Setup LoRA adapters
    model = setup_lora_model(model, config)
    
    # Dataset and DataLoader
    logger.info(f"Loading dataset from {data_dir}")
    dataset = TextDataset(data_dir, tokenizer, config.max_seq_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_ids = batch["input_ids"].to(device)
            labels = batch["labels"].to(device)
            attention_mask = batch["attention_mask"].to(device)
            
            # Forward pass
            outputs = model(
                input_ids=input_ids,
                labels=labels,
                attention_mask=attention_mask
            )
            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}")
                    tokenizer.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)
    tokenizer.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 text files for training")
    parser.add_argument("--model_id", type=str, default="./qwen3.6-27b",
                       help="Path to pretrained Qwen model")
    parser.add_argument("--epochs", type=int, default=3)
    parser.add_argument("--batch_size", type=int, default=4)
    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)
