#!/usr/bin/env python3
"""
Train Qwen2.5-27B on software development data via QLoRA (CPU-only).

Features:
- Checkpoint resumption (survives crashes/restarts)
- Progress logging to file
- CPU-only (no GPU needed)
- Optimized for Ryzen 9 9900X + 60GB RAM

Usage:
  python3 train_dev_model.py
  python3 train_dev_model.py --resume  # Resume from last checkpoint
"""

import os
import sys
import json
import argparse
import time
import logging
from pathlib import Path
from datetime import datetime

# Force CPU-only PyTorch
os.environ["CUDA_VISIBLE_DEVICES"] = ""
os.environ["CUDA_DEVICE_ORDER"] = ""

import torch
from torch.utils.data import Dataset

# Load model from local path
MODEL_NAME = "/home/vincent/projects/bitcoin-ai/model"
DATA_FILE = Path(__file__).parent.parent / "data" / "processed" / "dev_finetuning_massive.json"
OUTPUT_DIR = Path(__file__).parent.parent / "output" / "dev_ai_model"
LOG_FILE = Path(__file__).parent.parent / "output" / "training.log"

# Training config
MAX_STEPS = 3000  # More steps for thorough training
BATCH_SIZE = 1
GRADIENT_ACCUM = 4
LEARNING_RATE = 2e-4
WARMUP_STEPS = 300
SAVE_INTERVAL = 200  # Save checkpoint every N steps

# Setup logging
logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(message)s",
    handlers=[
        logging.FileHandler(LOG_FILE),
        logging.StreamHandler(sys.stdout)
    ]
)
logger = logging.getLogger(__name__)


class DevDataset(Dataset):
    """Load instruction-tuning pairs for code training."""
    
    def __init__(self, data_path, tokenizer, max_length=2048):
        with open(data_path, "r") as f:
            self.data = json.load(f)
        self.tokenizer = tokenizer
        self.max_length = max_length
        logger.info(f"Loaded {len(self.data)} training pairs")
    
    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, idx):
        item = self.data[idx]
        messages = item["messages"]
        
        # Build chat format
        prompt = self.tokenizer.apply_chat_template(
            messages, tokenize=False, add_generation_prompt=False
        )
        
        tokenized = self.tokenizer(
            prompt,
            truncation=True,
            max_length=self.max_length,
            padding="max_length",
            return_tensors="pt"
        )
        
        input_ids = tokenized["input_ids"].squeeze()
        attention_mask = tokenized["attention_mask"].squeeze()
        
        # Labels = input_ids (for causal LM, mask padding with -100)
        labels = input_ids.clone()
        labels[attention_mask == 0] = -100
        
        return {
            "input_ids": input_ids,
            "attention_mask": attention_mask,
            "labels": labels,
        }


def find_latest_checkpoint(output_dir):
    """Find the latest checkpoint directory for resumption."""
    output_dir = Path(output_dir)
    checkpoints = [d for d in output_dir.iterdir() if d.is_dir() and d.name.startswith("checkpoint-")]
    if not checkpoints:
        return None
    
    # Extract step number and sort
    def step_num(d):
        return int(d.name.split("-")[1])
    
    latest = sorted(checkpoints, key=step_num)[-1]
    logger.info(f"Found checkpoint: {latest.name}")
    return latest


def main(resume=False):
    logger.info("=" * 60)
    logger.info("Dev-AI Training - Qwen2.5-27B QLoRA (CPU)")
    logger.info("=" * 60)
    
    # Check data exists
    if not DATA_FILE.exists():
        logger.error(f"Dataset not found: {DATA_FILE}")
        logger.error("Run scripts/collect_code_dataset.py first")
        sys.exit(1)
    
    # Load model
    logger.info(f"Loading model from {MODEL_NAME}...")
    from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
    
    bnb_config = BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_quant_type="nf4",
        bnb_4bit_compute_dtype=torch.float32,
        bnb_4bit_use_double_quant=False,
    )
    
    tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
    tokenizer.pad_token = tokenizer.eos_token
    tokenizer.padding_side = "right"
    
    model = AutoModelForCausalLM.from_pretrained(
        MODEL_NAME,
        quantization_config=bnb_config,
        device_map="cpu",
        torch_dtype=torch.float32,
        low_cpu_mem_usage=True,
    )
    logger.info(f"Model loaded: {sum(p.numel() for p in model.parameters()) / 1e9:.1f}B parameters")
    
    # Setup LoRA
    from peft import LoraConfig, get_peft_model, TaskType
    
    lora_config = LoraConfig(
        r=16,
        lora_alpha=32,
        target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
        lora_dropout=0.05,
        bias="none",
        task_type=TaskType.CAUSAL_LM,
    )
    
    model = get_peft_model(model, lora_config)
    model.print_trainable_parameters()
    
    # Load dataset
    dataset = DevDataset(DATA_FILE, tokenizer)
    
    # Training setup
    from transformers import (
        TrainingArguments,
        Trainer,
        DataCollatorForSeq2Seq,
    )
    
    # Resume checkpoint
    resume_from = None
    if resume:
        resume_from = find_latest_checkpoint(OUTPUT_DIR)
        if resume_from:
            logger.info(f"Resuming from {resume_from}")
    
    training_args = TrainingArguments(
        output_dir=str(OUTPUT_DIR),
        num_train_epochs=1,
        max_steps=MAX_STEPS,
        per_device_train_batch_size=BATCH_SIZE,
        gradient_accumulation_steps=GRADIENT_ACCUM,
        learning_rate=LEARNING_RATE,
        fp16=False,
        bf16=False,
        logging_steps=10,
        save_steps=SAVE_INTERVAL,
        save_total_limit=10,  # Keep last 10 checkpoints
        optim="adamw_torch",
        warmup_steps=WARMUP_STEPS,
        report_to="none",
        dataloader_num_workers=0,
        ddp_find_unused_parameters=False,
        resume_from_checkpoint=resume_from,
    )
    
    data_collator = DataCollatorForSeq2Seq(tokenizer=tokenizer, model=model)
    
    # Custom callback to log progress
    class ProgressCallback:
        def __init__(self):
            self.start_time = time.time()
            self.last_step_time = None
            self.last_step_num = None
        
        def on_step_end(self, args, state, control, **kwargs):
            elapsed = time.time() - self.start_time
            steps_done = state.global_step
            
            if self.last_step_time:
                step_time = time.time() - self.last_step_time
            else:
                step_time = 0
            
            self.last_step_time = time.time()
            
            if steps_done % 50 == 0:
                eta_seconds = (MAX_STEPS - steps_done) * (elapsed / max(steps_done, 1))
                eta_hours = eta_seconds / 3600
                logger.info(
                    f"Step {steps_done}/{MAX_STEPS} | "
                    f"Loss: {state.log_history[-1].get('loss', 0):.4f} | "
                    f"Step time: {step_time:.1f}s | "
                    f"ETA: {eta_hours:.1f}h"
                )
    
    trainer = Trainer(
        model=model,
        args=training_args,
        train_dataset=dataset,
        data_collator=data_collator,
    )
    
    trainer.add_callback(ProgressCallback())
    
    logger.info(f"Starting training: {MAX_STEPS} steps, batch={BATCH_SIZE}, accum={GRADIENT_ACCUM}")
    logger.info(f"Checkpoint interval: every {SAVE_INTERVAL} steps")
    logger.info(f"Log file: {LOG_FILE}")
    
    trainer.train(resume_from_checkpoint=resume_from)
    
    # Save final adapter
    final_path = OUTPUT_DIR / "final_adapter"
    model.save_pretrained(str(final_path))
    tokenizer.save_pretrained(str(final_path))
    
    logger.info("=" * 60)
    logger.info(f"TRAINING COMPLETE")
    logger.info(f"Final adapter saved to: {final_path}")
    logger.info(f"Total time: {(time.time() - ProgressCallback().start_time) / 3600:.1f} hours")
    logger.info("=" * 60)


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument("--resume", action="store_true", help="Resume from latest checkpoint")
    args = parser.parse_args()
    main(resume=args.resume)
