X-Orange's picture
Upload folder using huggingface_hub
6550ac5 verified
Raw
History Blame Contribute Delete
3.06 kB
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))