Download export_train_data.py from WWsCa/vetnet-train-package: direct link, hf CLI and curl.
- Browser
- Download file 2.91 kB
-
https://huggingface.co/WWsCa/vetnet-train-package/resolve/main/export_train_data.py
- Command line
-
hf download hf://WWsCa/vetnet-train-package/export_train_data.py
-
curl -L -o export_train_data.py https://huggingface.co/WWsCa/vetnet-train-package/resolve/main/export_train_data.py
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() | |