#!/usr/bin/env python3
"""
CPU-optimized QLoRA fine-tuning for Bitcoin AI assistant.

Uses 4-bit quantization and LoRA adapters for efficient fine-tuning
on the Ryzen 9 9900X CPU.
"""

import os
import json
import torch
from pathlib import Path

# Force CPU training before importing transformers
os.environ["CUDA_VISIBLE_DEVICES"] = ""
from transformers import (
    AutoModelForCausalLM,
    AutoTokenizer,
    BitsAndBytesConfig,
    TrainingArguments,
    Trainer,
    DataCollatorForSeq2Seq,
)
from peft import (
    LoraConfig,
    get_peft_model,
    prepare_model_for_kbit_training,
)
from datasets import Dataset

# Configuration
MODEL_NAME = "../model"  # Local path to downloaded model
DATA_FILE = Path(__file__).parent.parent / "data" / "bitcoin_finetuning.json"
OUTPUT_DIR = Path(__file__).parent.parent / "output" / "bitcoin_ai_model"
MAX_STEPS = 2000  # Adjust based on dataset size

def load_model_cpu():
    """Load model optimized for CPU with 4-bit quantization."""
    bnb_config = BitsAndBytesConfig(
        load_in_4bit=True,
        bnb_4bit_quant_type="nf4",
        bnb_4bit_compute_dtype=torch.float32,
        bnb_4bit_use_double_quant=True,
    )
    
    print("Loading base model (this may take a while on CPU)...")
    model = AutoModelForCausalLM.from_pretrained(
        MODEL_NAME,
        quantization_config=bnb_config,
        device_map="cpu",
        torch_dtype=torch.float32,
    )
    
    model = prepare_model_for_kbit_training(model)
    return model

def add_lora(model):
    """Add LoRA adapters for efficient fine-tuning."""
    lora_config = LoraConfig(
        r=16,
        lora_alpha=32,
        target_modules=["q_proj", "v_proj", "k_proj", "o_proj"],
        lora_dropout=0.05,
        bias="none",
        task_type="CAUSAL_LM",
    )
    
    return get_peft_model(model, lora_config)

def prepare_dataset():
    """Load and format training data."""
    if not DATA_FILE.exists():
        raise FileNotFoundError(f"Training data not found: {DATA_FILE}")
    
    with open(DATA_FILE, 'r') as f:
        data = json.load(f)
    
    # Convert to HuggingFace Dataset format
    dataset = Dataset.from_list(data)
    return dataset

def tokenize_function(examples, tokenizer, max_length=512):
    """Tokenize the instruction-response pairs."""
    inputs = [
        f"### Instruction:\n{inst}\n\n### Response:\n{resp}\n"
        for inst, resp in zip(examples["instruction"], examples["output"])
    ]
    
    # Tokenize inputs
    model_inputs = tokenizer(
        inputs,
        max_length=max_length,
        truncation=True,
        padding="max_length",
        return_tensors="pt"
    )
    
    # For causal LM, labels are the same as input_ids
    model_inputs["labels"] = model_inputs["input_ids"].clone()
    
    return model_inputs

def train():
    """Run the fine-tuning process."""
    print("Starting Bitcoin AI fine-tuning...")
    print(f"Base model: {MODEL_NAME}")
    print(f"Training steps: {MAX_STEPS}")
    print(f"Output directory: {OUTPUT_DIR}")
    
    # Create output directory
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    
    # Load components
    model = load_model_cpu()
    model = add_lora(model)
    tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
    dataset = prepare_dataset()
    
    # Tokenize dataset
    print("Tokenizing dataset...")
    tokenized_dataset = dataset.map(
        lambda x: tokenize_function(x, tokenizer),
        batched=True,
        batch_size=4,
        remove_columns=dataset.column_names
    )
    
    # Training arguments optimized for CPU
    training_args = TrainingArguments(
        output_dir=str(OUTPUT_DIR),
        num_train_epochs=1,
        max_steps=MAX_STEPS,
        per_device_train_batch_size=1,
        gradient_accumulation_steps=4,
        learning_rate=2e-4,
        fp16=False,
        bf16=False,
        logging_steps=10,
        save_steps=100,
        save_total_limit=3,
        optim="adamw_torch",
        warmup_steps=200,
        report_to="none",
        # Force CPU training
        dataloader_num_workers=0,
        ddp_find_unused_parameters=False,
    )
    
    # Data collator
    data_collator = DataCollatorForSeq2Seq(tokenizer=tokenizer, model=model)
    
    # Initialize trainer
    trainer = Trainer(
        model=model,
        args=training_args,
        train_dataset=tokenized_dataset,
        data_collator=data_collator,
    )
    
    print("Configuration ready. Starting training...")
    print("This will run on CPU and take ~7-14 days")
    
    # Start training
    trainer.train()
    
    # Save the fine-tuned model
    model.save_pretrained(str(OUTPUT_DIR / "final_model"))
    tokenizer.save_pretrained(str(OUTPUT_DIR / "final_model"))
    
    print("Training complete! Model saved.")

if __name__ == "__main__":
    train()