import os
import torch
import numpy as np
from torch.utils.data import IterableDataset
from transformers import (
    LlamaConfig, 
    LlamaForCausalLM, 
    Trainer, 
    TrainingArguments
)

# --- HYPERPARAMETERS ---
CONTEXT_LENGTH = 512
MICRO_BATCH_SIZE = 4
GRADIENT_ACCUMULATION_STEPS = 64
LEARNING_RATE = 3e-4
TOTAL_STEPS = 152_587  

# --- 1. MODEL ARCHITECTURE (350M Llama) ---
config = LlamaConfig(
    vocab_size=50304,
    hidden_size=1024,
    intermediate_size=2816,
    num_hidden_layers=24,
    num_attention_heads=16,
    num_key_value_heads=16,
    max_position_embeddings=CONTEXT_LENGTH,
    tie_word_embeddings=True,
    rms_norm_eps=1e-5,
    rope_theta=10000.0,
    attn_implementation="flash_attention_2"
)

print(f"Setting up 350M Llama model from scratch...")
model = LlamaForCausalLM(config)
model_params = sum(p.numel() for p in model.parameters())
print(f"Model parameters: {model_params / 1e6:.2f} M")
print(f"Effective batch size: {MICRO_BATCH_SIZE * GRADIENT_ACCUMULATION_STEPS} sequences")
print(f"Tokens per step: {MICRO_BATCH_SIZE * GRADIENT_ACCUMULATION_STEPS * CONTEXT_LENGTH:,}")
print(f"Total training tokens: {TOTAL_STEPS * MICRO_BATCH_SIZE * GRADIENT_ACCUMULATION_STEPS * CONTEXT_LENGTH / 1e9:.2f} B")


# --- 2. DATA LOADER (Linear 1-Epoch-Streaming) ---
class BinaryDataset(IterableDataset):
    def __init__(self, bin_path, seq_len):
        self.bin_path = bin_path
        self.seq_len = seq_len
        self.file_size = os.path.getsize(bin_path)
        self.total_tokens = self.file_size // 4 
        self.total_sequences = self.total_tokens // seq_len

    def __iter__(self):
        worker_info = torch.utils.data.get_worker_info()
        
        if worker_info is None:
            start_seq = 0
            end_seq = self.total_sequences
        else:
            per_worker = self.total_sequences // worker_info.num_workers
            worker_id = worker_info.id
            start_seq = worker_id * per_worker
            end_seq = start_seq + per_worker if worker_id < worker_info.num_workers - 1 else self.total_sequences

        mmap = np.memmap(self.bin_path, dtype=np.uint32, mode='r')

        for seq_idx in range(start_seq, end_seq):
            start_token = seq_idx * self.seq_len
            end_token = start_token + self.seq_len + 1
            
            if end_token > self.total_tokens:
                break
                
            tokens_slice = mmap[start_token:end_token].astype(np.int64)
            tokens = torch.tensor(tokens_slice, dtype=torch.long)
            
            yield {
                "input_ids": tokens[:-1],
                "labels": tokens[1:]
            }

train_dataset = BinaryDataset("train.bin", CONTEXT_LENGTH)
val_dataset = BinaryDataset("val.bin", CONTEXT_LENGTH)

# --- 3. TRAINING ARGUMENTS ---
training_args = TrainingArguments(
    output_dir="./apex-2-350m-run",
    per_device_train_batch_size=MICRO_BATCH_SIZE,
    per_device_eval_batch_size=MICRO_BATCH_SIZE,
    gradient_accumulation_steps=GRADIENT_ACCUMULATION_STEPS,
    max_steps=TOTAL_STEPS,
    logging_steps=50,
    save_steps=2000,
    save_total_limit=3,
    eval_steps=2000,
    evaluation_strategy="steps",
    optim="adamw_torch_fused", 
    learning_rate=LEARNING_RATE,
    weight_decay=0.1,
    lr_scheduler_type="cosine",
    warmup_steps=2000,
    bf16=True,
    fp16=False,
    dataloader_num_workers=4,
    report_to="none",
    max_grad_norm=1.0,
)

# --- 4. START TRAINING ---
trainer = Trainer(
    model=model,
    args=training_args,
    train_dataset=train_dataset,
    eval_dataset=val_dataset,
)

print("Starting training of Apex 2 350M...")
trainer.train()