File size: 5,667 Bytes
c2d946e 9dc72fc 4ee3bd0 d9af4cc 4b56279 4ee3bd0 d0f19ae 4ee3bd0 9dc72fc c2d946e 4446cb9 c2d946e 4446cb9 c2d946e 4446cb9 c2d946e 4446cb9 c2d946e 4446cb9 c2d946e 9dc72fc c2d946e 9dc72fc c2d946e 9dc72fc c2d946e 9dc72fc c2d946e 9dc72fc c2d946e 9dc72fc c2d946e 9dc72fc c2d946e 9dc72fc c2d946e 9dc72fc 4446cb9 c2d946e 4446cb9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 | # 1. 引入必要库(包含Gradio)
import gradio as gr
import sys
from io import StringIO
from datasets import load_dataset
#from transformers import (
# BertTokenizer, BertForTokenClassification,
# Trainer, TrainingArguments, DataCollatorForTokenClassification
#)
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
# 2. 修复日志重定向类(补充isatty等缺失方法)
class PrintRedirect:
def __init__(self):
self.buffer = StringIO()
self.stdout = sys.stdout # 保留原输出到Logs面板
self.stderr = sys.stderr
def write(self, msg):
self.buffer.write(msg)
self.stdout.write(msg) # 同时打印到Logs
def flush(self):
self.buffer.flush()
self.stdout.flush()
self.stderr.flush()
def get_log(self):
return self.buffer.getvalue()
# 关键修复:补充uvicorn需要的isatty方法
def isatty(self):
return self.stdout.isatty() # 复用原stdout的isatty判断
# 可选:补充其他可能缺失的方法,避免后续报错
def fileno(self):
return self.stdout.fileno()
def close(self):
self.buffer.close()
self.stdout.close()
# 初始化日志重定向(仅重定向stdout,不重定向stderr,避免干扰uvicorn)
redirect = PrintRedirect()
sys.stdout = redirect
# 注释掉stderr重定向,避免日志模块冲突
# sys.stderr = redirect
# 3. 配置参数(根据你的需求调整)
DATA_FILE = "fengshi_jiapu_train.jsonl" # 数据文件(已上传到Space)
MODEL_SAVE_PATH = "./fengshi_ner_final"
BASE_MODEL = "bert-base-chinese"
MAX_LENGTH = 64
BATCH_SIZE = 4
EPOCHS = 5
# 4. 标签映射(NER任务必备)
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)}
# 5. 数据预处理函数
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 # 你的数据是完整文本,不是分词后的列表
)
# 处理标签(适配BERT的子词切分)
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:
# 子词或padding的标签设为-100(Trainer会自动忽略)
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
# 6. 训练主函数(Gradio触发)
def run_training():
# 加载数据
print("=== 开始加载数据 ===")
dataset = load_dataset("json", data_files=DATA_FILE)
print(f"数据加载完成,训练集样本数:{len(dataset['train'])}")
# 加载tokenizer和模型
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, # 每10步打印日志
save_steps=100,
evaluation_strategy="no", # 若有验证集可改为"epoch"
report_to="none", # 禁用wandb,避免依赖
learning_rate=2e-5,
)
# 数据整理器
data_collator = DataCollatorForTokenClassification(tokenizer=tokenizer)
# 初始化Trainer
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()
# 7. Gradio界面(满足Spaces的SDK要求,点击按钮启动训练)
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)
# 8. 启动Gradio(添加show_error=True,便于排查剩余问题)
if __name__ == "__main__":
demo.launch(
server_port=7860,
server_name="0.0.0.0",
show_error=True, # 显示启动错误
quiet=False # 保留日志输出
) |