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