#!/usr/bin/env python3
"""
构造预训练模型评估数据
将 validation.txt 转换为适合预训练评估的 JSONL 格式
"""

import os
import json
import random
from typing import List, Dict, Tuple
from tokenizers import Tokenizer
import argparse

def load_tokenizer(tokenizer_path: str) -> Tokenizer:
    """加载分词器"""
    tokenizer = Tokenizer.from_file(os.path.join(tokenizer_path, "tokenizer.json"))
    return tokenizer

def get_special_ids(tokenizer: Tokenizer) -> Dict[str, int]:
    """获取特殊token的ID"""
    vocab = tokenizer.get_vocab()
    return {
        "cls": vocab.get("[CLS]", 101),
        "sep": vocab.get("[SEP]", 102),
        "pad": vocab.get("[PAD]", 0),
        "mask": vocab.get("[MASK]", 103),
        "unk": vocab.get("[UNK]", 100)
    }

def apply_mlm(input_ids: List[int], special: Dict[str, int], vocab_size: int, 
              mlm_prob: float, rng: random.Random) -> Tuple[List[int], List[int]]:
    """应用掩码语言建模"""
    mlm_input_ids = input_ids.copy()
    mlm_labels = [-100] * len(input_ids)
    
    # 找到可以掩码的位置（排除特殊token）
    candidate_positions = []
    for i, token_id in enumerate(input_ids):
        if token_id not in [special["cls"], special["sep"], special["pad"]]:
            candidate_positions.append(i)
    
    if not candidate_positions:
        return mlm_input_ids, mlm_labels
    
    # 随机选择要掩码的token
    num_to_mask = max(1, int(round(len(candidate_positions) * mlm_prob)))
    masked_positions = rng.sample(candidate_positions, min(num_to_mask, len(candidate_positions)))
    
    for pos in masked_positions:
        original_token = input_ids[pos]
        mlm_labels[pos] = original_token
        
        # 80% 替换为 [MASK]
        # 10% 替换为随机token
        # 10% 保持原样
        rand = rng.random()
        if rand < 0.8:
            mlm_input_ids[pos] = special["mask"]
        elif rand < 0.9:
            mlm_input_ids[pos] = rng.randint(0, vocab_size - 1)
        # else: 保持原样
    
    return mlm_input_ids, mlm_labels

def create_validation_samples(validation_file: str, tokenizer: Tokenizer, 
                            special: Dict[str, int], max_seq_len: int = 128,
                            mlm_prob: float = 0.15, num_samples: int = 1000) -> List[Dict]:
    """创建验证样本"""
    
    # 读取验证文本
    with open(validation_file, 'r', encoding='utf-8') as f:
        texts = [line.strip() for line in f if line.strip()]
    
    print(f"📖 读取验证文本: {len(texts)} 条")
    
    # 随机采样
    rng = random.Random(42)
    sampled_texts = rng.sample(texts, min(num_samples, len(texts)))
    
    samples = []
    vocab_size = tokenizer.get_vocab_size()
    
    for i, text in enumerate(sampled_texts):
        # 分词
        tokens = tokenizer.encode(text)
        input_ids = tokens.ids[:max_seq_len]  # 截断到最大长度
        
        # 应用MLM
        mlm_input_ids, mlm_labels = apply_mlm(
            input_ids, special, vocab_size, mlm_prob, rng
        )
        
        # 填充到固定长度
        while len(mlm_input_ids) < max_seq_len:
            mlm_input_ids.append(special["pad"])
            mlm_labels.append(-100)
        
        # 创建注意力掩码
        attention_mask = [1 if token != special["pad"] else 0 for token in mlm_input_ids]
        
        # 创建token类型ID（简化处理，都设为0）
        token_type_ids = [0] * max_seq_len
        
        # 创建NSP标签（随机）
        next_sentence_label = rng.randint(0, 1)
        
        sample = {
            "input_ids": mlm_input_ids,
            "attention_mask": attention_mask,
            "token_type_ids": token_type_ids,
            "labels": mlm_labels,
            "next_sentence_label": next_sentence_label,
            "original_text": text
        }
        
        samples.append(sample)
        
        if (i + 1) % 100 == 0:
            print(f"  已处理 {i + 1}/{len(sampled_texts)} 个样本")
    
    return samples

def main():
    parser = argparse.ArgumentParser(description="构造预训练模型评估数据")
    parser.add_argument("--validation_file", required=True, help="验证文本文件路径")
    parser.add_argument("--tokenizer_path", required=True, help="分词器路径")
    parser.add_argument("--output_file", required=True, help="输出JSONL文件路径")
    parser.add_argument("--max_seq_len", type=int, default=128, help="最大序列长度")
    parser.add_argument("--mlm_prob", type=float, default=0.15, help="MLM掩码概率")
    parser.add_argument("--num_samples", type=int, default=1000, help="样本数量")
    
    args = parser.parse_args()
    
    # 检查输入文件
    if not os.path.exists(args.validation_file):
        print(f"❌ 验证文件不存在: {args.validation_file}")
        return
    
    if not os.path.exists(args.tokenizer_path):
        print(f"❌ 分词器路径不存在: {args.tokenizer_path}")
        return
    
    # 加载分词器
    print("🔄 加载分词器...")
    tokenizer = load_tokenizer(args.tokenizer_path)
    special = get_special_ids(tokenizer)
    
    # 创建验证样本
    print("�� 创建验证样本...")
    samples = create_validation_samples(
        args.validation_file, tokenizer, special,
        args.max_seq_len, args.mlm_prob, args.num_samples
    )
    
    # 保存为JSONL
    print(f"�� 保存到: {args.output_file}")
    with open(args.output_file, 'w', encoding='utf-8') as f:
        for sample in samples:
            f.write(json.dumps(sample, ensure_ascii=False) + '\n')
    
    print(f"✅ 完成! 生成了 {len(samples)} 个验证样本")
    print(f"📊 样本统计:")
    print(f"   - 最大序列长度: {args.max_seq_len}")
    print(f"   - MLM掩码概率: {args.mlm_prob}")
    print(f"   - 总样本数: {len(samples)}")

if __name__ == "__main__":
    main()
