| import glob, json | |
| from Classifier_Model import ClassifierModel | |
| import os | |
| import argparse | |
| parser = argparse.ArgumentParser(description="刘狗模型") | |
| parser.add_argument("--jsonl_path", type=str, default="sample", help="The path of the JSONL file to be processed") | |
| parser.add_argument("--jsonl_key", type=str, default="content", help="Keys of valid fields in JSON data") | |
| parser.add_argument("--save_path", type=str, default="save", help="The path where the file is stored") | |
| parser.add_argument("--save_size", type=int, default=160_000, help="Number of lines in the stored file") | |
| parser.add_argument("--batch_size", type=int, default=400, help="The number of data items the model processes at once") | |
| parser.add_argument("--max_length", type=int, default=1024, help="Maximum length of the data") | |
| parser.add_argument("--target_score", type=int, default=3, help="The minimum passing score for the text") | |
| parser.add_argument("--model_path", type=str, default="new_model.pth", help="The path of the model") | |
| parser.add_argument("--tokenizer_path", type=str, default="local_tokenizer", help="The path of the tokenizer") | |
| parser.add_argument("--device", type=str, default="cuda", help="Device for loading the model") | |
| parser.add_argument("--dtype", type=str, default="fp16", help="Dtype for loading the model") | |
| args = parser.parse_args() | |
| jsonl_path = args.jsonl_path | |
| jsonl_key = args.jsonl_key | |
| save_path = args.save_path | |
| save_size = args.save_size | |
| batch_size = args.batch_size | |
| max_length = args.max_length | |
| target_score = args.target_score | |
| model_path = args.model_path | |
| tokenizer_path = args.tokenizer_path | |
| device = args.device | |
| dtype = args.dtype | |
| os.makedirs(save_path, exist_ok=True) | |
| all_text = [] | |
| batch_text = [] | |
| f_idx = 0 | |
| cm = ClassifierModel( | |
| model_path=model_path, | |
| tokenizer_path=tokenizer_path, | |
| device=device, | |
| dtype=dtype | |
| ) | |
| parquet_files = glob.glob(os.path.join(jsonl_path, "*.jsonl"))[: ] | |
| for file in parquet_files: | |
| with open(file, "r", encoding="utf-8") as rf: | |
| print(f"file {os.path.basename(file)}") | |
| for line in rf.readlines(): | |
| text = json.loads(line)[jsonl_key] | |
| batch_text.append(text) | |
| if len(batch_text)%(batch_size*100)==0: | |
| scores = cm.compute(batch_text, batch_size=batch_size, max_length=1024) | |
| for i, j in zip(batch_text, scores): | |
| if j>target_score: | |
| all_text.append(json.dumps({"content": i, "score": j}, ensure_ascii=False)) | |
| batch_text = [] | |
| if len(all_text)>=save_size: | |
| with open(os.path.join(save_path, f"file_{f_idx}.jsonl"), "w", encoding="utf-8") as wf: | |
| wf.write("\n".join(all_text[: save_size])) | |
| all_text = all_text[save_size: ] | |
| f_idx += 1 | |
| if len(all_text) > 0: | |
| with open(os.path.join(save_path, f"file_{f_idx}.jsonl"), "w", encoding="utf-8") as wf: | |
| wf.write("\n".join(all_text)) | |