import os
import argparse
import logging
import socket

import torch
import torch.distributed as dist

from datasets import load_from_disk

from transformers import (
    AutoModelForCausalLM,
    AutoTokenizer,
    TrainingArguments,
    Trainer
)

from peft import LoraConfig, get_peft_model


# -------------------------
# Distributed helpers
# -------------------------

def get_dist_info():
    return {
        "rank": int(os.environ.get("RANK", 0)),
        "local_rank": int(os.environ.get("LOCAL_RANK", 0)),
        "world_size": int(os.environ.get("WORLD_SIZE", 1)),
    }


def is_rank0():
    return get_dist_info()["rank"] == 0


# -------------------------
# Logging
# -------------------------

def setup_logging(output_dir: str):

    os.makedirs(output_dir, exist_ok=True)

    dist_info = get_dist_info()
    rank = dist_info["rank"]

    log_file = os.path.join(
        output_dir,
        f"train-rank{rank}.log"
    )

    logging.basicConfig(
        level=logging.INFO if is_rank0() else logging.WARNING,
        format="%(asctime)s | %(levelname)s | %(message)s",
        handlers=[
            logging.FileHandler(log_file),
            logging.StreamHandler(),
        ],
    )

    logging.info("Logging initialized")
    logging.info(f"Hostname: {socket.gethostname()}")
    logging.info(f"Rank info: {dist_info}")

# -------------------------
# Argument parsing
# -------------------------

def env_or_default(env_key, default):
    return os.environ.get(env_key, default)


def parse_args():

    parser = argparse.ArgumentParser()

    parser.add_argument(
        "--model-path",
        type=str,
        default=env_or_default(
            "MODEL_PATH",
            "/mnt/models/kanana-nano-2.1b-base"
        ),
    )

    parser.add_argument(
        "--dataset",
        type=str,
        default=env_or_default(
            "DATASET_PATH",
            "/mnt/datasets/ultrachat_tokenized"
        ),
    )

    parser.add_argument(
        "--output-dir",
        type=str,
        default=env_or_default(
            "OUTPUT_DIR",
            "/mnt/output"
        ),
    )

    parser.add_argument(
        "--max-steps",
        type=int,
        default=int(env_or_default("MAX_STEPS", 500)),
    )

    parser.add_argument(
        "--logging-steps",
        type=int,
        default=int(env_or_default("LOGGING_STEPS", 10)),
    )

    parser.add_argument(
        "--save-steps",
        type=int,
        default=int(env_or_default("SAVE_STEPS", 500)),
    )

    parser.add_argument(
        "--bf16",
        action="store_true",
        default=env_or_default("BF16", "true").lower() == "true",
    )

    parser.add_argument(
        "--deepspeed",
        type=str,
        default=env_or_default("DEEPSPEED_CONFIG", None),
    )

    parser.add_argument(
        "--logging-dir",
        type=str,
        default=env_or_default("TENSORBOARD_LOGDIR", None),
    )

    parser.add_argument(
        "--lora-r",
        type=int,
        default=int(env_or_default("LORA_R", 4)),
    )

    parser.add_argument(
        "--lora-alpha",
        type=int,
        default=int(env_or_default("LORA_ALPHA", 8)),
    )

    parser.add_argument(
        "--lora-dropout",
        type=float,
        default=float(env_or_default("LORA_DROPOUT", 0.05)),
    )

    return parser.parse_args()


# -------------------------
# Main
# -------------------------

def main():

    args = parse_args()
    setup_logging(args.output_dir)
    dist_info = get_dist_info()
    torch.cuda.set_device(dist_info["local_rank"])

    if dist.is_initialized():
        dist.barrier()

    # -------------------------
    # Load model
    # -------------------------

    logging.info("Loading base model")

    base_model = AutoModelForCausalLM.from_pretrained(
        args.model_path,
        dtype=torch.bfloat16 if args.bf16 else torch.float16,
        trust_remote_code=True,
        local_files_only=True,
    )

    if dist.is_initialized():
        dist.barrier()

    # -------------------------
    # Apply LoRA
    # -------------------------

    logging.info("Applying LoRA")

    lora_cfg = LoraConfig(
        r=args.lora_r,
        lora_alpha=args.lora_alpha,
        lora_dropout=args.lora_dropout,
        task_type="CAUSAL_LM",
        target_modules=[
            "q_proj",
            "k_proj",
            "v_proj",
            "o_proj",
            "gate_proj",
            "up_proj",
            "down_proj",
        ],
    )

    model = get_peft_model(base_model, lora_cfg)

    if dist.is_initialized():
        dist.barrier()

    # -------------------------
    # Dataset
    # -------------------------
    
    train_dataset = load_from_disk(args.dataset)

    if dist.is_initialized():
        dist.barrier()

    # -------------------------
    # TrainingArguments
    # -------------------------
    training_args = TrainingArguments(
        output_dir=args.output_dir,
        logging_dir=args.logging_dir,
        bf16=args.bf16,
        logging_steps=args.logging_steps,
        save_steps=args.save_steps,
        max_steps=args.max_steps,
        deepspeed=args.deepspeed,
        report_to="tensorboard",
        disable_tqdm=not is_rank0(),
        dataloader_num_workers=0,
        remove_unused_columns=False,
        per_device_train_batch_size=1
    )

    # -------------------------
    # Trainer
    # -------------------------

    trainer = Trainer(
        model=model,
        args=training_args,
        train_dataset=train_dataset
    )

    # -------------------------
    # Train
    # -------------------------

    logging.info("Starting training")
    trainer.train()

    if is_rank0():
        logging.info("Training completed")

    if dist.is_available() and dist.is_initialized():
        dist.destroy_process_group()


# -------------------------

if __name__ == "__main__":
    main()