# 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 # 保留日志输出 )