corpus.py 8.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312
  1. """
  2. 文本语料加载器接口
  3. 功能:
  4. - 从预存文件加载文本语料
  5. - 支持按领域(domain)组织语料
  6. - 将语料保存到文件
  7. 数据格式约定:
  8. - 语料文件:.jsonl 格式,每行 {"text": "...", "metadata": {...}}
  9. - 领域目录:database/corpus/{domain}/texts.jsonl
  10. 目录结构:
  11. database/
  12. ├── corpus/ # 文本语料
  13. │ ├── news_en/
  14. │ │ └── texts.jsonl
  15. │ ├── news_zh/
  16. │ ├── academic/
  17. │ ├── code/
  18. │ ├── dialogue/
  19. │ └── literature/
  20. └── splits/ # 数据集划分(可选)
  21. └── train_val_test.json
  22. """
  23. import json
  24. import os
  25. from pathlib import Path
  26. from typing import Dict, List, Optional, Tuple, Iterator
  27. from dataclasses import dataclass, field
  28. @dataclass
  29. class CorpusMeta:
  30. """语料库元数据"""
  31. domain: str
  32. language: str
  33. num_texts: int
  34. avg_length: float
  35. source: str = ""
  36. extra: Dict = field(default_factory=dict)
  37. class CorpusLoader:
  38. """
  39. 从预存文件加载文本语料
  40. """
  41. def __init__(self, corpus_dir: str = "database/corpus"):
  42. self.base_dir = Path(corpus_dir)
  43. self.base_dir.mkdir(parents=True, exist_ok=True)
  44. def get_domain_path(self, domain: str) -> Path:
  45. """获取领域目录路径"""
  46. return self.base_dir / domain
  47. def get_texts_path(self, domain: str) -> Path:
  48. """获取语料文件路径"""
  49. return self.get_domain_path(domain) / "texts.jsonl"
  50. def load_texts(
  51. self,
  52. domain: str,
  53. limit: Optional[int] = None
  54. ) -> List[str]:
  55. """
  56. 加载某个领域的文本列表
  57. Args:
  58. domain: 领域标识(如 "news_en", "academic")
  59. limit: 最多加载多少条(用于快速测试)
  60. Returns:
  61. texts: 文本列表
  62. """
  63. path = self.get_texts_path(domain)
  64. if not path.exists():
  65. raise FileNotFoundError(f"Corpus not found: {path}")
  66. texts = []
  67. with open(path, "r", encoding="utf-8") as f:
  68. for i, line in enumerate(f):
  69. if limit and i >= limit:
  70. break
  71. record = json.loads(line.strip())
  72. texts.append(record.get("text", ""))
  73. return texts
  74. def load_texts_with_metadata(
  75. self,
  76. domain: str,
  77. limit: Optional[int] = None
  78. ) -> List[Tuple[str, Dict]]:
  79. """
  80. 加载文本及其元数据
  81. Returns:
  82. [(text, metadata), ...]
  83. """
  84. path = self.get_texts_path(domain)
  85. if not path.exists():
  86. raise FileNotFoundError(f"Corpus not found: {path}")
  87. result = []
  88. with open(path, "r", encoding="utf-8") as f:
  89. for i, line in enumerate(f):
  90. if limit and i >= limit:
  91. break
  92. record = json.loads(line.strip())
  93. result.append((
  94. record.get("text", ""),
  95. record.get("metadata", {})
  96. ))
  97. return result
  98. def load_all_domains(
  99. self,
  100. domains: Optional[List[str]] = None,
  101. limit_per_domain: Optional[int] = None
  102. ) -> Dict[str, List[str]]:
  103. """
  104. 加载多个领域的文本
  105. Args:
  106. domains: 领域列表,None 表示加载所有
  107. limit_per_domain: 每个领域最多加载多少条
  108. Returns:
  109. {domain: [texts], ...}
  110. """
  111. if domains is None:
  112. domains = self.list_available_domains()
  113. result = {}
  114. for domain in domains:
  115. result[domain] = self.load_texts(domain, limit_per_domain)
  116. return result
  117. def list_available_domains(self) -> List[str]:
  118. """列出所有可用的领域"""
  119. domains = []
  120. for path in self.base_dir.iterdir():
  121. if path.is_dir() and self.get_texts_path(path.name).exists():
  122. domains.append(path.name)
  123. return sorted(domains)
  124. def save_texts(
  125. self,
  126. texts: List[str],
  127. domain: str,
  128. metadata_list: Optional[List[Dict]] = None,
  129. source: str = ""
  130. ) -> None:
  131. """
  132. 保存文本到语料文件
  133. Args:
  134. texts: 文本列表
  135. domain: 领域标识
  136. metadata_list: 可选的元数据列表
  137. source: 数据来源说明
  138. """
  139. path = self.get_texts_path(domain)
  140. path.parent.mkdir(parents=True, exist_ok=True)
  141. with open(path, "w", encoding="utf-8") as f:
  142. for i, text in enumerate(texts):
  143. record = {
  144. "text": text,
  145. "metadata": {
  146. "index": i,
  147. "source": source,
  148. **(metadata_list[i] if metadata_list and i < len(metadata_list) else {})
  149. }
  150. }
  151. f.write(json.dumps(record, ensure_ascii=False) + "\n")
  152. def get_corpus_stats(self, domain: str) -> Dict:
  153. """获取语料库统计信息"""
  154. texts = self.load_texts(domain)
  155. lengths = [len(t) for t in texts]
  156. char_lengths = [len(t) for t in texts]
  157. return {
  158. "domain": domain,
  159. "num_texts": len(texts),
  160. "total_chars": sum(char_lengths),
  161. "avg_chars": sum(char_lengths) / len(texts) if texts else 0,
  162. "min_chars": min(char_lengths) if texts else 0,
  163. "max_chars": max(char_lengths) if texts else 0,
  164. }
  165. def compute_corpus_statistics(self) -> Dict[str, Dict]:
  166. """计算所有领域的统计信息"""
  167. result = {}
  168. for domain in self.list_available_domains():
  169. result[domain] = self.get_corpus_stats(domain)
  170. return result
  171. class CorpusBuilder:
  172. """
  173. 语料构建器 - 用于从原始数据构建标准格式语料
  174. (当需要预处理原始数据时使用)
  175. """
  176. def __init__(self, output_dir: str = "database/corpus"):
  177. self.output_dir = Path(output_dir)
  178. def from_text_files(
  179. self,
  180. input_dir: str,
  181. domain: str,
  182. pattern: str = "*.txt",
  183. encoding: str = "utf-8"
  184. ) -> int:
  185. """
  186. 从目录中的文本文件构建语料
  187. Args:
  188. input_dir: 输入目录
  189. domain: 领域标识
  190. pattern: 文件匹配模式
  191. encoding: 文件编码
  192. Returns:
  193. 加载的文本数量
  194. """
  195. input_path = Path(input_dir)
  196. texts = []
  197. for file_path in input_path.glob(pattern):
  198. with open(file_path, "r", encoding=encoding) as f:
  199. content = f.read().strip()
  200. if content:
  201. texts.append(content)
  202. loader = CorpusLoader(self.output_dir)
  203. loader.save_texts(texts, domain, source=f"from_text_files:{input_dir}")
  204. return len(texts)
  205. def from_jsonl(
  206. self,
  207. input_path: str,
  208. domain: str,
  209. text_field: str = "text"
  210. ) -> int:
  211. """
  212. 从 JSONL 文件复制/转换语料
  213. Args:
  214. input_path: 输入 JSONL 路径
  215. domain: 领域标识
  216. text_field: 文本字段名
  217. Returns:
  218. 加载的文本数量
  219. """
  220. texts = []
  221. with open(input_path, "r", encoding="utf-8") as f:
  222. for line in f:
  223. record = json.loads(line.strip())
  224. texts.append(record.get(text_field, ""))
  225. loader = CorpusLoader(self.output_dir)
  226. loader.save_texts(texts, domain, source=f"from_jsonl:{input_path}")
  227. return len(texts)
  228. # ============= 便捷函数 =============
  229. def load_texts(
  230. domain: str,
  231. corpus_dir: str = "database/corpus",
  232. limit: Optional[int] = None
  233. ) -> List[str]:
  234. """便捷函数:加载单个领域的文本"""
  235. loader = CorpusLoader(corpus_dir)
  236. return loader.load_texts(domain, limit)
  237. def load_all_texts(
  238. domains: Optional[List[str]] = None,
  239. corpus_dir: str = "database/corpus",
  240. limit_per_domain: Optional[int] = None
  241. ) -> Dict[str, List[str]]:
  242. """便捷函数:加载多个领域的文本"""
  243. loader = CorpusLoader(corpus_dir)
  244. return loader.load_all_domains(domains, limit_per_domain)
  245. def save_texts(
  246. texts: List[str],
  247. domain: str,
  248. corpus_dir: str = "database/corpus",
  249. source: str = ""
  250. ) -> None:
  251. """便捷函数:保存文本语料"""
  252. loader = CorpusLoader(corpus_dir)
  253. loader.save_texts(texts, domain, source=source)
  254. def list_domains(corpus_dir: str = "database/corpus") -> List[str]:
  255. """便捷函数:列出所有可用领域"""
  256. loader = CorpusLoader(corpus_dir)
  257. return loader.list_available_domains()