fzd10 / train.py
fdbw's picture
Update train.py
4b56279 verified
Raw
History Blame Contribute Delete
5.67 kB
# 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 # 保留日志输出
)