File size: 2,906 Bytes
ef4736a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
导出领域自适应预训练(DAPT)数据 — 原始文本格式
按文档分组拼接,模型通过 next-token-prediction 内化兽医知识
"""
import json
import sqlite3
import random
import re
from pathlib import Path

random.seed(42)

DB_PATH = Path(r"C:\Users\HP\Documents\666666666666666666\VetCopilot-LCPS\vetcopilot-backend\vetcopilot.db")
OUTPUT = Path("train_data.jsonl")
VAL_OUTPUT = Path("train_data_val.jsonl")
TRAIN_SPLIT = 0.9
MIN_CHUNK_LEN = 100

DOC_PREFIX = "\n\n【兽医文献】\n"
DOC_SUFFIX = "\n\n---\n"


def clean_text(text: str) -> str:
    """清洗文本:去多余空行、统一空白"""
    text = re.sub(r'\n{3,}', '\n\n', text)
    text = re.sub(r'[ \t]{3,}', '  ', text)
    return text.strip()


def main():
    conn = sqlite3.connect(str(DB_PATH))
    conn.row_factory = sqlite3.Row

    rows = conn.execute("""
        SELECT d.filename, d.title, c.content
        FROM knowledge_documents d
        JOIN knowledge_chunks c ON c.document_id = d.id
        WHERE d.status = 'ready'
        ORDER BY d.filename, c.chunk_index
    """).fetchall()
    conn.close()

    print(f"查询到 {len(rows)} 条记录")

    # 按文档分组拼接为长文本
    doc_texts = []
    current_fn = None
    current_parts = []

    for r in rows:
        fn = r["filename"]
        content = r["content"].strip()
        if len(content) < MIN_CHUNK_LEN:
            continue

        if fn != current_fn:
            # 保存上一个文档
            if current_parts:
                title = r["title"] or ""
                full = DOC_PREFIX + f"标题: {title}\n\n" + "\n\n".join(current_parts) + DOC_SUFFIX
                doc_texts.append(clean_text(full))
            current_fn = fn
            current_parts = [content]
        else:
            current_parts.append(content)

    # 保存最后一个文档
    if current_parts:
        full = DOC_PREFIX + "\n\n".join(current_parts) + DOC_SUFFIX
        doc_texts.append(clean_text(full))

    print(f"拼成 {len(doc_texts)} 个完整文档")

    # 打乱后按 9:1 分割
    random.shuffle(doc_texts)
    split_idx = int(len(doc_texts) * TRAIN_SPLIT)
    train_docs = doc_texts[:split_idx]
    val_docs = doc_texts[split_idx:]

    # 写入 JSONL(每条是一个文档的完整文本)
    def write_jsonl(docs, path):
        with open(path, "w", encoding="utf-8") as f:
            for text in docs:
                f.write(json.dumps({"text": text}, ensure_ascii=False) + "\n")

    write_jsonl(train_docs, OUTPUT)
    write_jsonl(val_docs, VAL_OUTPUT)

    total_chars = sum(len(t) for t in doc_texts)
    print(f"\n训练文档: {len(train_docs)}, 验证文档: {len(val_docs)}")
    print(f"总字符数: {total_chars:,}")
    print(f"训练集: {OUTPUT.stat().st_size/1024/1024:.1f} MB")
    print(f"验证集: {VAL_OUTPUT.stat().st_size/1024/1024:.1f} MB")


if __name__ == "__main__":
    main()