#!/usr/bin/env python3
"""
Train Gemma-4-E2B-it on Ayurveda dataset (local cache + local JSONL).
Optimized for RTX A4000 16GB VRAM + 20GB system RAM.
"""

import torch
import os

# ============================================================================
# 0. VERIFY PATHS (edit these if your files are elsewhere)
# ============================================================================
DATASET_PATH = r"C:\Users\Antplay\Downloads\gemma4_ayurveda_unsloth_500mb_FINAL.jsonl"
MODEL_NAME   = "unsloth/gemma-4-E2B-it"   # Will load from local HF cache
OUTPUT_DIR   = r"C:\Users\Antplay\Downloads\gemma4_ayurveda_lora_output"

print("Dataset:", DATASET_PATH)
print("Model:  ", MODEL_NAME, "(loading from local HF cache)")
print("Output: ", OUTPUT_DIR)

assert os.path.exists(DATASET_PATH), f"Dataset not found: {DATASET_PATH}"

# ============================================================================
# 1. LOAD DATASET (memory-efficient streaming)
# ============================================================================
from datasets import load_dataset

print("\n[1/5] Loading dataset...")

# For JSONL we can stream to keep RAM usage low, but standardize_sharegpt
# prefers a regular Dataset object. 404 MB JSONL -> ~1 GB RAM peak is fine.
dataset = load_dataset(
    "json",
    data_files=DATASET_PATH,
    split="train",
    streaming=False,   # Set True if you still hit RAM issues (slower)
)

print(f"Dataset loaded: {len(dataset):,} conversations")

# ============================================================================
# 2. STANDARDIZE TO SHAREGPT (converts messages -> text field)
# ============================================================================
from unsloth import standardize_sharegpt

print("\n[2/5] Standardizing ShareGPT format...")
dataset = standardize_sharegpt(dataset)
print("Standardize complete. Sample text length:", len(dataset[0]["text"]))

# ============================================================================
# 3. LOAD MODEL (from local HuggingFace cache)
# ============================================================================
from unsloth import FastLanguageModel

print("\n[3/5] Loading model from local cache...")

model, tokenizer = FastLanguageModel.from_pretrained(
    model_name           = MODEL_NAME,
    max_seq_length       = 3072,      # 3072 context as you wanted
    dtype                = None,      # Auto-detect bf16 on Ampere (A4000)
    load_in_4bit         = True,      # QLoRA 4-bit
    local_files_only     = True,      # Use locally downloaded model only
)

print(f"Model loaded. VRAM used: ~{torch.cuda.memory_allocated()/1024**3:.1f} GB")

# ============================================================================
# 4. ATTACH QLoRA ADAPTER (big rank for small model)
# ============================================================================
print("\n[4/5] Attaching LoRA adapter (r=64, alpha=128)...")

model = FastLanguageModel.get_peft_model(
    model,
    r                  = 64,          # High rank compensates for small 5B base
    lora_alpha         = 128,         # 2x rank
    lora_dropout       = 0,           # 0 dropout: 276k samples = no overfit risk
    bias               = "none",
    use_rslora         = False,
    use_gradient_checkpointing = "unsloth",  # Saves ~30% VRAM
    random_state       = 3407,
    target_modules     = [
        "q_proj", "k_proj", "v_proj", "o_proj",
        "gate_proj", "up_proj", "down_proj",
    ],
)

torch.cuda.empty_cache()
print(f"Adapter attached. VRAM after LoRA: ~{torch.cuda.memory_allocated()/1024**3:.1f} GB")

# ============================================================================
# 5. TRAINING ARGUMENTS (A4000 tuned)
# ============================================================================
from transformers import TrainingArguments

# Effective batch = 1 x 8 = 8
# 276k conversations / 8 = ~34,500 steps total for 1 epoch
# A4000 should do ~2-3 it/s -> ~4-6 hours

training_args = TrainingArguments(
    output_dir                  = OUTPUT_DIR,
    overwrite_output_dir        = True,

    num_train_epochs            = 1,
    per_device_train_batch_size = 1,          # MUST be 1 with 3072 ctx on A4000
    gradient_accumulation_steps = 8,          # Effective batch = 8

    learning_rate               = 3e-4,       # Slightly aggressive for 5B model
    lr_scheduler_type           = "cosine",
    warmup_steps                = 500,        # ~1.5% of total steps

    optim                       = "adamw_8bit",
    weight_decay                = 0.01,
    max_grad_norm               = 0.3,

    bf16                        = True,       # A4000 = Ampere = bf16 native
    fp16                        = False,

    logging_steps               = 50,
    save_strategy               = "steps",
    save_steps                  = 2000,       # Checkpoint every ~2k steps
    save_total_limit            = 2,          # Keep only last 2 checkpoints

    group_by_length             = True,       # Efficiency boost for variable lengths
    remove_unused_columns       = False,

    report_to                   = "none",     # Set "wandb" if you track
    seed                        = 3407,
)

# ============================================================================
# 6. TRAINER
# ============================================================================
from trl import SFTTrainer

print("\n[5/5] Initializing SFTTrainer...")

trainer = SFTTrainer(
    model             = model,
    tokenizer         = tokenizer,
    train_dataset     = dataset,
    dataset_text_field= "text",
    max_seq_length    = 3072,
    args              = training_args,
    packing           = False,          # NEVER pack instruction data
)

# ============================================================================
# 7. TRAIN
# ============================================================================
print("\n" + "="*60)
print("STARTING TRAINING")
print("="*60)
print(f"  Conversations : {len(dataset):,}")
print(f"  Epochs        : 1")
print(f"  Batch size    : 1 (effective: 8)")
print(f"  Context       : 3072")
print(f"  Learning rate : 3e-4 cosine")
print(f"  LoRA rank     : 64")
print(f"  Expected time : ~4-6 hours on RTX A4000")
print("="*60)

trainer.train()

# ============================================================================
# 8. SAVE
# ============================================================================
print("\n" + "="*60)
print("TRAINING COMPLETE - SAVING")
print("="*60)

# Save LoRA adapter only (~100-150 MB)
model.save_pretrained(os.path.join(OUTPUT_DIR, "adapter"))
tokenizer.save_pretrained(os.path.join(OUTPUT_DIR, "adapter"))
print(f"Adapter saved to: {os.path.join(OUTPUT_DIR, 'adapter')}")

# Optional: merge LoRA into base model for standalone inference
# Requires ~12 GB VRAM temporarily; skip if OOM
print("\nOptional: merging adapter into base model (16-bit)...")
try:
    model.save_pretrained_merged(
        os.path.join(OUTPUT_DIR, "merged_16bit"),
        tokenizer,
        save_method = "merged_16bit",
    )
    print(f"Merged model saved to: {os.path.join(OUTPUT_DIR, 'merged_16bit')}")
except RuntimeError as e:
    print(f"Merge skipped (OOM or error): {e}")
    print("You can still inference using the adapter + base model.")

print("\nAll done! Output folder:", OUTPUT_DIR)
