| |
| import gradio as gr |
| import sys |
| from io import StringIO |
| from datasets import load_dataset |
| |
| |
| |
| |
| from transformers import AutoTokenizer |
| from transformers import BertTokenizer |
| from transformers import BertForTokenClassification |
|
|
| tokenizer = AutoTokenizer.from_pretrained("bert-base-chinese", use_fast=True) |
|
|
| import torch |
| import numpy as np |
| from sklearn.metrics import classification_report |
|
|
|
|
| |
| class PrintRedirect: |
| def __init__(self): |
| self.buffer = StringIO() |
| self.stdout = sys.stdout |
| self.stderr = sys.stderr |
|
|
| def write(self, msg): |
| self.buffer.write(msg) |
| self.stdout.write(msg) |
|
|
| def flush(self): |
| self.buffer.flush() |
| self.stdout.flush() |
| self.stderr.flush() |
|
|
| def get_log(self): |
| return self.buffer.getvalue() |
|
|
| |
| def isatty(self): |
| return self.stdout.isatty() |
|
|
| |
| def fileno(self): |
| return self.stdout.fileno() |
|
|
| def close(self): |
| self.buffer.close() |
| self.stdout.close() |
|
|
| |
| redirect = PrintRedirect() |
| sys.stdout = redirect |
| |
| |
|
|
|
|
| |
| DATA_FILE = "fengshi_jiapu_train.jsonl" |
| MODEL_SAVE_PATH = "./fengshi_ner_final" |
| BASE_MODEL = "bert-base-chinese" |
| MAX_LENGTH = 64 |
| BATCH_SIZE = 4 |
| EPOCHS = 5 |
|
|
|
|
| |
| label_list = ["O", "B-PERSON", "I-PERSON", "B-TIME", "I-TIME", "B-LOCATION"] |
| label2id = {label: i for i, label in enumerate(label_list)} |
| id2label = {i: label for i, label in enumerate(label_list)} |
|
|
|
|
| |
| def preprocess_function(examples): |
| texts = examples["text"] |
| labels = examples["labels"] |
|
|
| |
| tokenized_inputs = BertTokenizer.from_pretrained(BASE_MODEL).__call__( |
| texts, |
| truncation=True, |
| padding="max_length", |
| max_length=MAX_LENGTH, |
| is_split_into_words=False |
| ) |
|
|
| |
| final_labels = [] |
| for i, label in enumerate(labels): |
| word_ids = tokenized_inputs.word_ids(batch_index=i) |
| label_ids = [] |
| previous_word_idx = None |
| for word_idx in word_ids: |
| |
| if word_idx is None or word_idx != previous_word_idx: |
| label_ids.append(label2id[label[word_idx]] if word_idx < len(label) else -100) |
| else: |
| label_ids.append(-100) |
| previous_word_idx = word_idx |
| final_labels.append(label_ids) |
|
|
| tokenized_inputs["labels"] = final_labels |
| return tokenized_inputs |
|
|
|
|
| |
| def run_training(): |
| |
| print("=== 开始加载数据 ===") |
| dataset = load_dataset("json", data_files=DATA_FILE) |
| print(f"数据加载完成,训练集样本数:{len(dataset['train'])}") |
|
|
| |
| print("=== 加载模型和Tokenizer ===") |
| tokenizer = BertTokenizer.from_pretrained(BASE_MODEL) |
| model = BertForTokenClassification.from_pretrained( |
| BASE_MODEL, |
| num_labels=len(label_list), |
| id2label=id2label, |
| label2id=label2id |
| ) |
|
|
| |
| print("=== 预处理数据 ===") |
| tokenized_dataset = dataset.map(preprocess_function, batched=True) |
|
|
| |
| print("=== 配置训练参数 ===") |
| training_args = TrainingArguments( |
| output_dir="./results", |
| per_device_train_batch_size=BATCH_SIZE, |
| num_train_epochs=EPOCHS, |
| logging_steps=10, |
| save_steps=100, |
| evaluation_strategy="no", |
| report_to="none", |
| learning_rate=2e-5, |
| ) |
|
|
| |
| data_collator = DataCollatorForTokenClassification(tokenizer=tokenizer) |
|
|
| |
| trainer = Trainer( |
| model=model, |
| args=training_args, |
| train_dataset=tokenized_dataset["train"], |
| data_collator=data_collator, |
| ) |
|
|
| |
| print("=== 开始训练 ===") |
| trainer.train() |
|
|
| |
| print("=== 保存模型 ===") |
| model.save_pretrained(MODEL_SAVE_PATH) |
| tokenizer.save_pretrained(MODEL_SAVE_PATH) |
| print(f"模型已保存到:{MODEL_SAVE_PATH}") |
|
|
| |
| return redirect.get_log() |
|
|
|
|
| |
| with gr.Blocks(title="NER训练脚本") as demo: |
| gr.Markdown("# 家谱NER训练任务") |
| log_output = gr.Textbox(label="训练日志(实时输出)", lines=20) |
| start_btn = gr.Button("启动训练") |
| start_btn.click(fn=run_training, outputs=log_output) |
|
|
|
|
| |
| if __name__ == "__main__": |
| demo.launch( |
| server_port=7860, |
| server_name="0.0.0.0", |
| show_error=True, |
| quiet=False |
| ) |