vetnet-train-package / export_train_data.py
WWsCa's picture
Upload export_train_data.py with huggingface_hub
ef4736a verified
Raw History Blame Contribute Delete
2.91 kB
"""
导出领域自适应预训练(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()