import os, math, json, argparse, time
import torch
from torch import nn
from torch.cuda.amp import autocast, GradScaler
from torch.optim import AdamW
from transformers import BertConfig, BertForMaskedLM, get_linear_schedule_with_warmup
from tokenizers import Tokenizer

# 引入你之前的DataLoader实现
from torch_dataset_loader import JsonlShardDataset, StreamConfig, collate_dynamic_padding

def init_distributed():
    if "RANK" in os.environ and "WORLD_SIZE" in os.environ:
        torch.distributed.init_process_group(backend="nccl")
        local_rank = int(os.environ.get("LOCAL_RANK", "0"))
        torch.cuda.set_device(local_rank)
        return True, local_rank
    return False, 0

def build_model(tokenizer_dir: str, max_position_embeddings: int = 512):
    tok_path_json = os.path.join(tokenizer_dir, "tokenizer.json")
    tok = Tokenizer.from_file(tok_path_json)
    vocab_size = tok.get_vocab_size()
    cfg = BertConfig(
        vocab_size=vocab_size,
        hidden_size=512,
        num_hidden_layers=6,
        num_attention_heads=8,
        intermediate_size=2048,
        max_position_embeddings=max_position_embeddings,
        hidden_act="gelu",
        layer_norm_eps=1e-12,
        type_vocab_size=2,
        pad_token_id=0,
    )
    model = BertForMaskedLM(cfg)
    return model

def create_loader(shards_index: str, batch_size: int, num_workers: int, shuffle_buffer: int):
    ds = JsonlShardDataset(StreamConfig(
        shards_index=shards_index,
        shuffle_shards=True,
        shuffle_records_buffer=shuffle_buffer,
        seed=42,
        repeat=1,
    ))
    loader = torch.utils.data.DataLoader(
        ds,
        batch_size=batch_size,
        num_workers=num_workers,
        collate_fn=collate_dynamic_padding,
        pin_memory=True,
        prefetch_factor=2 if num_workers > 0 else None,
        persistent_workers=True if num_workers > 0 else False,
    )
    return loader

def save_ckpt(model, optimizer, scheduler, scaler, step, out_dir, is_main):
    if not is_main:
        return
    os.makedirs(out_dir, exist_ok=True)
    ckpt_dir = os.path.join(out_dir, f"step-{step}")
    os.makedirs(ckpt_dir, exist_ok=True)
    model.save_pretrained(ckpt_dir, safe_serialization=True)
    torch.save({
        "step": step,
        "optimizer": optimizer.state_dict(),
        "scheduler": scheduler.state_dict(),
        "scaler": scaler.state_dict(),
    }, os.path.join(ckpt_dir, "trainer_state.pt"))

def parse_args():
    p = argparse.ArgumentParser()
    p.add_argument("--shards_index", required=True)
    p.add_argument("--tokenizer_dir", required=True)
    p.add_argument("--output_dir", required=True)
    p.add_argument("--max_steps", type=int, default=100000)
    p.add_argument("--batch_size", type=int, default=32)
    p.add_argument("--num_workers", type=int, default=4)
    p.add_argument("--lr", type=float, default=2e-4)
    p.add_argument("--weight_decay", type=float, default=0.01)
    p.add_argument("--warmup_ratio", type=float, default=0.06)
    p.add_argument("--grad_accum", type=int, default=2)
    p.add_argument("--log_every", type=int, default=100)
    p.add_argument("--save_every", type=int, default=5000)
    p.add_argument("--max_position_embeddings", type=int, default=512)
    return p.parse_args()

def main():
    args = parse_args()
    is_dist, local_rank = init_distributed()
    device = torch.device(f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu")
    is_main = (not is_dist) or (int(os.environ.get("RANK", "0")) == 0)

    model = build_model(args.tokenizer_dir, args.max_position_embeddings).to(device)
    if is_dist:
        model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], output_device=local_rank)

    loader = create_loader(args.shards_index, args.batch_size, args.num_workers, shuffle_buffer=10000)

    no_decay = ["bias", "LayerNorm.weight"]
    params = [
        {"params": [p for n,p in model.named_parameters() if not any(nd in n for nd in no_decay)], "weight_decay": args.weight_decay},
        {"params": [p for n,p in model.named_parameters() if any(nd in n for nd in no_decay)], "weight_decay": 0.0},
    ]
    optimizer = AdamW(params, lr=args.lr, betas=(0.9, 0.999), eps=1e-8)
    total_updates = args.max_steps
    warmup_steps = max(1, int(total_updates * args.warmup_ratio))
    scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=warmup_steps, num_training_steps=total_updates)

    scaler = GradScaler(enabled=True)
    model.train()

    step = 0
    micro_step = 0
    running_loss = 0.0
    start = time.time()
    while step < args.max_steps:
        for batch in loader:
            micro_step += 1
            for k in list(batch.keys()):
                batch[k] = batch[k].to(device, non_blocking=True)

            with autocast(enabled=True):
                outputs = model(
                    input_ids=batch["input_ids"],
                    attention_mask=batch["attention_mask"],
                    token_type_ids=batch["token_type_ids"],
                    labels=batch["mlm_labels"],
                )
                loss = outputs.loss / args.grad_accum

            scaler.scale(loss).backward()

            if micro_step % args.grad_accum == 0:
                scaler.unscale_(optimizer)
                torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
                scaler.step(optimizer)
                scaler.update()
                optimizer.zero_grad(set_to_none=True)
                scheduler.step()
                step += 1

                running_loss += outputs.loss.item()
                if is_main and step % args.log_every == 0:
                    elapsed = time.time() - start
                    print(f"step {step} | loss {running_loss/args.log_every:.4f} | lr {scheduler.get_last_lr()[0]:.2e} | elapsed {elapsed:.1f}s")
                    running_loss = 0.0
                    start = time.time()

                if step % args.save_every == 0:
                    save_ckpt(model.module if isinstance(model, nn.parallel.DistributedDataParallel) else model,
                              optimizer, scheduler, scaler, step, args.output_dir, is_main)

                if step >= args.max_steps:
                    break

    # final save
    save_ckpt(model.module if isinstance(model, nn.parallel.DistributedDataParallel) else model,
              optimizer, scheduler, scaler, step, args.output_dir, is_main)
    if is_dist:
        torch.distributed.destroy_process_group()

if __name__ == "__main__":
    main()
