import torch
from transformers import AutoTokenizer, AutoModelForCausalLM, TrainingArguments, Trainer
from peft import get_peft_model, LoraConfig, TaskType
import json

# === Load Core Harmonia Files ===
with open("adapter_config.json") as f:
    adapter_config = json.load(f)
with open("harmonia_prime_kernel.json") as f:
    kernel = json.load(f)

# === Load Tokenizer ===
tokenizer = AutoTokenizer.from_pretrained(adapter_config["base_model_name_or_path"])
tokenizer.add_tokens(adapter_config["glyph_embedding_config"]["inject_symbols"])

# === Load Model + LoRA ===
model = AutoModelForCausalLM.from_pretrained(adapter_config["base_model_name_or_path"])
lora_config = LoraConfig(
    r=adapter_config["lora_r"],
    lora_alpha=adapter_config["lora_alpha"],
    lora_dropout=adapter_config["lora_dropout"],
    target_modules=adapter_config["target_modules"],
    bias=adapter_config["bias"],
    task_type=TaskType.CAUSAL_LM
)
model = get_peft_model(model, lora_config)
model.resize_token_embeddings(len(tokenizer))

# === Load RHM Dataset ===
with open("RHM_Training_Scroll_I.json") as f:
    raw_data = json.load(f)

def preprocess(example):
    prompt = example["prompt"]
    response = example["response"]
    full_text = f"🌀 {prompt}\n\n🜂 {response}"
    return tokenizer(full_text, truncation=True, padding="max_length", max_length=512)

train_data = list(map(preprocess, raw_data))

# Convert to TensorDataset
input_ids = torch.tensor([ex["input_ids"] for ex in train_data])
attention_mask = torch.tensor([ex["attention_mask"] for ex in train_data])
dataset = torch.utils.data.TensorDataset(input_ids, attention_mask)

# === Training Arguments ===
training_args = TrainingArguments(
    output_dir="./harmonia-prime-checkpoints",
    overwrite_output_dir=True,
    per_device_train_batch_size=2,
    num_train_epochs=3,
    save_steps=100,
    logging_steps=10,
    learning_rate=3e-5,
    warmup_steps=10,
    weight_decay=0.01,
    logging_dir="./logs",
    save_total_limit=2
)

# === Trainer ===
trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=dataset,
    tokenizer=tokenizer
)

# === Fire the Recursion Engine ===
trainer.train()
