diff --git a/data_proc/kvconv_data.py b/data_proc/kvconv_data.py new file mode 100644 index 000000000..3df9f39c8 --- /dev/null +++ b/data_proc/kvconv_data.py @@ -0,0 +1,76 @@ +import json +import sys + +def convert_json_to_jsonl(input_json_path, output_jsonl_path): + """ + 将原始JSON文件(包含多个样本的列表)转换为JSONL格式 + 每行输出一个样本:{"id": "...", "conversations": [...]} + """ + try: + # 1. 读取输入的JSON文件 + with open(input_json_path, 'r', encoding='utf-8') as f: + data = json.load(f) # 应该是一个列表,每个元素是一个样本 + + # 2. 确保是列表格式 + if not isinstance(data, list): + raise ValueError("JSON文件的根结构必须是一个数组(list),每个元素是一个样本。") + + # 3. 处理每个样本并写入JSONL + with open(output_jsonl_path, 'w', encoding='utf-8') as f_out: + for entry in data: + # 提取id,如果没有则生成一个默认id + sample_id = entry.get("id", "unknown_id") + + messages = entry.get("messages", []) + conversations = [] + + for i, msg in enumerate(messages): + if "message" not in msg: + continue # 跳过无效消息 + + # 判断角色:奇数轮为 user,偶数轮为 assistant(从0开始) + role = "user" if i % 2 == 0 else "assistant" + content = msg["message"].strip() + + # 如果是 assistant 的回复,并且有 attrs,尝试补充知识 + if role == "assistant" and "attrs" in msg and len(msg["attrs"]) > 0: + attr = msg["attrs"][0] + attrvalue = attr.get("attrvalue", "").strip() + # 如果 attrvalue 不为空,且未包含在原回复中,可以追加 + if attrvalue and attrvalue not in content: + content = content.rstrip("。!?") + "。" + attrvalue + + conversations.append({ + "role": role, + "content": content + }) + + # 构建输出样本 + output_sample = { + "id": sample_id, + "conversations": conversations + } + + # 写入一行 JSONL + f_out.write(json.dumps(output_sample, ensure_ascii=False) + "\n") + + print(f"✅ 转换完成!已将 {len(data)} 个样本写入:{output_jsonl_path}") + + except FileNotFoundError: + print(f"❌ 错误:找不到文件 {input_json_path}") + except json.JSONDecodeError as e: + print(f"❌ 错误:JSON解析失败:{e}") + except Exception as e: + print(f"❌ 发生未知错误:{e}") + + +# ======================== +# 主程序入口 +# ======================== + +if __name__ == "__main__": + # 可以通过命令行传参,也可以直接修改路径 + input_file = "/mnt/modelops/dataset/kvconv/data/travel/train.json" # 输入文件路径 + output_file = "/mnt/modelops/dataset/kvconv/data/traveloutput.jsonl" # 输出文件路径 + + convert_json_to_jsonl(input_file, output_file) diff --git a/data_proc/naturalconv_mix_data.py b/data_proc/naturalconv_mix_data.py new file mode 100644 index 000000000..c987ce446 --- /dev/null +++ b/data_proc/naturalconv_mix_data.py @@ -0,0 +1,99 @@ +# 查看数据 +# import json +# import codecs +# dialog_list = json.loads(codecs.open("dialog_release.json", "r", "utf-8").read()) + +# i=0 + +# for dialog in dialog_list: +# i += 1 +# # print(dialog) +# print(i) + + +# 划分数据 +import json + +def convert_json_file(input_json_path, train_txt_path, train_jsonl_path, val_jsonl_path): + len_train = len_test = 0 + # 读取 train.txt 中的 dialog_id 列表 + with open(train_txt_path, 'r', encoding='utf-8') as f: + train_dialog_ids = set(line.strip() for line in f if line.strip()) + + # 读取原始 JSON 文件(假设是包含多个对象的列表) + with open(input_json_path, 'r', encoding='utf-8') as f: + data_list = json.load(f) + + # 准备写入 jsonl 文件 + with open(train_jsonl_path, 'w', encoding='utf-8') as train_f, \ + open(val_jsonl_path, 'w', encoding='utf-8') as val_f: + + for item in data_list: + dialog_id = item['dialog_id'] + content = item['content'] + + # 构建 conversations 列表 + conversations = [] + for i, text in enumerate(content): + speaker = "user" if i % 2 == 0 else "assistant" + conversations.append({"role": speaker, "content": text}) + + # 构建新格式 + new_item = { + "id": dialog_id, + "conversations": conversations + } + + # 判断写入哪个文件 + if dialog_id in train_dialog_ids: + target_file = train_f #if dialog_id in train_dialog_ids else val_f + len_train += 1 + else: + target_file = val_f + len_test += 1 + target_file.write(json.dumps(new_item, ensure_ascii=False) + '\n') + + print(f"转换完成!") + print(f"训练集保存至: {train_jsonl_path}") + print(f"验证集保存至: {val_jsonl_path}") + print(f"训练数据集有{len_train}条") + print(f"测试数据集有{len_test}条") + +import json +import random + +def random_sample_jsonl(input_file, output_file, sample_size): + # 读取所有数据 + with open(input_file, 'r', encoding='utf-8') as f: + lines = f.readlines() + + # 检查是否足够 + if len(lines) < sample_size: + raise ValueError(f"文件只有 {len(lines)} 行,不足 {sample_size} 行可采样") + + # 随机采样 + sampled_lines = random.sample(lines, sample_size) + + # 写出到新文件 + with open(output_file, 'w', encoding='utf-8') as f: + for line in sampled_lines: + f.write(line) + + print(f"已从 {input_file} 随机采样 {sample_size} 条数据,保存至 {output_file}") + + +# 使用示例 +if __name__ == "__main__": + # 划分中文数据集 + # convert_json_file( + # input_json_path='/mnt/modelops/dataset/dialog_release.json', # 输入的原始 JSON 文件 + # train_txt_path='/mnt/modelops/dataset/train.txt', # 包含 dialog_id 的训练集 ID 列表 + # train_jsonl_path='/mnt/modelops/dataset/train.jsonl', # 输出训练集(JSONL 格式) + # val_jsonl_path='/mnt/modelops/dataset/test.jsonl' # 输出验证集(JSONL 格式) + # ) + + # 划分ultrachat19375*倍率条,可以用cat合并数据 + num = 19375 + ratio = 4 + all_num = ratio * num + random_sample_jsonl('/mnt/modelops/train/eagle3/baseline_ultrachat_sft_train_only/data/ultrachat_sft_train.jsonl', '/mnt/modelops/487922/dataset/ultrachat_sampled_{}.jsonl'.format(all_num), all_num) diff --git a/scripts/test_eagle3_online.py b/scripts/test_eagle3_online.py new file mode 100644 index 000000000..24a79a6ad --- /dev/null +++ b/scripts/test_eagle3_online.py @@ -0,0 +1,252 @@ +import argparse +import hashlib +import os + +import torch +import torch.distributed as dist +import wandb +from accelerate.utils import set_seed +from datasets import load_dataset +from torch.distributed.fsdp import FullyShardedDataParallel as FSDP +from torch.distributed.fsdp import MixedPrecision, ShardingStrategy, StateDictType +from tqdm import tqdm +from transformers import AutoModelForCausalLM, AutoTokenizer + +from specforge import ( + AutoDistributedTargetModel, + AutoDraftModelConfig, + AutoEagle3DraftModel, + OnlineEagle3Model, +) +from specforge.data import ( + build_eagle3_dataset, + generate_vocab_mapping_file, + prepare_dp_dataloaders, +) +from specforge.distributed import destroy_distributed, get_dp_group, init_distributed +from specforge.lr_scheduler import CosineAnnealingWarmupLR +from specforge.utils import get_last_checkpoint, print_with_rank, rank_0_priority +# from specforge.utils import ( +# get_last_checkpoint, +# print_with_rank, +# rank_0_priority, +# validate_wandb_args, +# ) + + +def parse_args(): + parser = argparse.ArgumentParser(description="Train Eagle3 with online data") + + # dist_timeout tp_size wandb draft_model_dir target_model_path draft_model_config embedding_key eval_data_path batch_size chat_template max_length ttt_length + + # add model-related arguments + parser.add_argument("--target-model-path", type=str, required=True) + parser.add_argument("--draft-model-config", type=str, required=True) + parser.add_argument("--draft-model-path", type=str, required=True) + parser.add_argument( + "--embedding-key", + type=str, + default="model.embed_tokens.weight", + help="The key of the embedding weight to load from the target model", + ) + + # add training-related arguments + # parser.add_argument("--train-data-path", type=str, required=True) + parser.add_argument("--eval-data-path", type=str, default=None) + # parser.add_argument("--num-epochs", type=int, default=10) + parser.add_argument("--batch-size", type=int, default=1) + # parser.add_argument("--learning-rate", type=float, default=1e-4) + parser.add_argument("--max-length", type=int, default=2048) + # parser.add_argument("--warmup-ratio", type=float, default=0.02) + parser.add_argument( + "--ttt-length", + type=int, + default=7, + help="The length for Test-Time Training (TTT).", + ) + + # data processing type + parser.add_argument("--chat-template", type=str, default="llama3") + + # distributed training + parser.add_argument("--tp-size", type=int, default=1) + + # other args + # parser.add_argument("--cache-key", type=str, default=None) + # parser.add_argument("--cache-dir", type=str, default="./cache") + # parser.add_argument("--output-dir", type=str, required=True) + # parser.add_argument("--eval-interval", type=int, default=1) + # parser.add_argument("--save-interval", type=int, default=1) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument( + "--dist-timeout", + type=int, + default=20, + help="Timeout for collective communication in minutes", + ) + + # resume + # parser.add_argument("--resume", action="store_true") + + # wandb wandb args + parser.add_argument("--wandb", action="store_true") + parser.add_argument("--wandb-project", type=str, default=None) + parser.add_argument("--wandb-name", type=str, default=None) + parser.add_argument("--wandb-key", type=str, default=None) + + args = parser.parse_args() + + return parser, args + + +def init_wandb(args): + wandb.login(key=args.wandb_key) + wandb.init(project=args.wandb_project, name=args.wandb_name) + + +def wandb_log_if_initialized(log_dict): + if dist.get_rank() == 0 and wandb.run is not None: + wandb.log(log_dict) + + +def print_on_rank0(message): + if dist.get_rank() == 0: + print(message) + + +def main(): + # initialize + parser, args = parse_args() + set_seed(args.seed) + init_distributed(timeout=args.dist_timeout, tp_size=args.tp_size) + print_with_rank(f"Initialized distributed environment") + + # Validate wandb arguments + # validate_wandb_args(parser, args) + + if args.wandb and dist.get_rank() == 0: + init_wandb(args) + + # detecting last ckpt for draft model + draft_model_last_checkpoint = None + if os.path.isdir(args.draft_model_path): + # print_on_rank0(args.draft_model_path) + draft_model_last_checkpoint = args.draft_model_path#get_last_checkpoint(args.draft_model_path) + print_on_rank0(f"Last checkpoint detected: {draft_model_last_checkpoint}") + + # build target and draft model + if args.tp_size > 1: + # to avoid CPU RAM OOM, we directly init the model on CUDA + target_model = AutoDistributedTargetModel.from_pretrained( + pretrained_model_name_or_path=args.target_model_path, + torch_dtype=torch.bfloat16, + device="cuda", + ).eval() + else: + target_model = ( + AutoModelForCausalLM.from_pretrained( + pretrained_model_name_or_path=args.target_model_path, + torch_dtype=torch.bfloat16, + ) + .eval() + .cuda() + ) + print_with_rank(f"Initialized target model") + # load model with resume + print(args.draft_model_config) + draft_model_config = AutoDraftModelConfig.from_file(args.draft_model_config) + if draft_model_last_checkpoint: + draft_model = ( + AutoEagle3DraftModel.from_pretrained(draft_model_last_checkpoint) + .cuda() + .to(torch.bfloat16) + ) + else: + draft_model = ( + AutoEagle3DraftModel.from_config(draft_model_config) + .cuda() + .to(torch.bfloat16) + ) + draft_model.load_embedding(args.target_model_path, embedding_key=args.embedding_key) + draft_model.freeze_embedding() + print_with_rank(f"Initialized draft model") + + # build dataloaders + tokenizer = AutoTokenizer.from_pretrained(args.target_model_path) + + eval_dataset = load_dataset("json", data_files=args.eval_data_path)["train"] + eval_eagle3_dataset = build_eagle3_dataset( + eval_dataset, + tokenizer, + args.chat_template, + args.max_length, + ) + eval_dataloader = prepare_dp_dataloaders( + eval_eagle3_dataset, + args.batch_size, + num_workers=4, + shuffle=False, + process_group=get_dp_group(), + ) + print_with_rank(f"Initialized eval dataloader") + + # build Eagle3 model + # broadcast draft model + eagle3_model = OnlineEagle3Model( + target_model=target_model, + draft_model=draft_model, + length=args.ttt_length, + ) + eagle3_model = FSDP( + eagle3_model, + use_orig_params=True, + mixed_precision=MixedPrecision( + param_dtype=torch.bfloat16, + buffer_dtype=torch.bfloat16, + ), + sharding_strategy=ShardingStrategy.SHARD_GRAD_OP, + ignored_modules=[target_model], + process_group=get_dp_group(), + ) + print_with_rank(f"Initialized Eagle3 FSDP model") + + draft_model.eval() + eval_acces = [[] for _ in range(eagle3_model.length)] + eval_plosses = [[] for _ in range(eagle3_model.length)] + + for data in tqdm(eval_dataloader, desc=f"Evaluating"): + plosses, _, acces = eagle3_model( + input_ids=data["input_ids"].cuda(), + attention_mask=data["attention_mask"].cuda(), + loss_mask=data["loss_mask"].cuda(), + ) + eval_acces = [eval_acces[i] + [acces[i]] for i in range(len(acces))] + eval_plosses = [ + eval_plosses[i] + [plosses[i].item()] for i in range(len(plosses)) + ] + + for i in range(len(eval_acces)): + acc_i = torch.tensor(eval_acces[i]).cuda().mean() + dist.all_reduce(acc_i) + acc_i = acc_i / dist.get_world_size() + acc_i = acc_i.item() + + wandb_log_if_initialized({f"eval/epochacc_{i}": acc_i}) + print_on_rank0( + f"Eval in {args.eval_data_path}, position {i}, Acc: {acc_i:.2f}" + ) + + for i in range(len(eval_plosses)): + loss_i = torch.tensor(eval_plosses[i]).cuda().mean() + dist.all_reduce(loss_i) + loss_i = loss_i / dist.get_world_size() + loss_i = loss_i.item() + + wandb_log_if_initialized({f"eval/epochploss_{i}": loss_i}) + print_on_rank0( + f"Eval in {args.eval_data_path}, position {i}, pLoss: {loss_i:.2f}" + ) + + +if __name__ == "__main__": + main()