| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312 |
- """
- 文本语料加载器接口
- 功能:
- - 从预存文件加载文本语料
- - 支持按领域(domain)组织语料
- - 将语料保存到文件
- 数据格式约定:
- - 语料文件:.jsonl 格式,每行 {"text": "...", "metadata": {...}}
- - 领域目录:database/corpus/{domain}/texts.jsonl
- 目录结构:
- database/
- ├── corpus/ # 文本语料
- │ ├── news_en/
- │ │ └── texts.jsonl
- │ ├── news_zh/
- │ ├── academic/
- │ ├── code/
- │ ├── dialogue/
- │ └── literature/
- └── splits/ # 数据集划分(可选)
- └── train_val_test.json
- """
- import json
- import os
- from pathlib import Path
- from typing import Dict, List, Optional, Tuple, Iterator
- from dataclasses import dataclass, field
- @dataclass
- class CorpusMeta:
- """语料库元数据"""
- domain: str
- language: str
- num_texts: int
- avg_length: float
- source: str = ""
- extra: Dict = field(default_factory=dict)
- class CorpusLoader:
- """
- 从预存文件加载文本语料
- """
- def __init__(self, corpus_dir: str = "database/corpus"):
- self.base_dir = Path(corpus_dir)
- self.base_dir.mkdir(parents=True, exist_ok=True)
- def get_domain_path(self, domain: str) -> Path:
- """获取领域目录路径"""
- return self.base_dir / domain
- def get_texts_path(self, domain: str) -> Path:
- """获取语料文件路径"""
- return self.get_domain_path(domain) / "texts.jsonl"
- def load_texts(
- self,
- domain: str,
- limit: Optional[int] = None
- ) -> List[str]:
- """
- 加载某个领域的文本列表
- Args:
- domain: 领域标识(如 "news_en", "academic")
- limit: 最多加载多少条(用于快速测试)
- Returns:
- texts: 文本列表
- """
- path = self.get_texts_path(domain)
- if not path.exists():
- raise FileNotFoundError(f"Corpus not found: {path}")
- texts = []
- with open(path, "r", encoding="utf-8") as f:
- for i, line in enumerate(f):
- if limit and i >= limit:
- break
- record = json.loads(line.strip())
- texts.append(record.get("text", ""))
- return texts
- def load_texts_with_metadata(
- self,
- domain: str,
- limit: Optional[int] = None
- ) -> List[Tuple[str, Dict]]:
- """
- 加载文本及其元数据
- Returns:
- [(text, metadata), ...]
- """
- path = self.get_texts_path(domain)
- if not path.exists():
- raise FileNotFoundError(f"Corpus not found: {path}")
- result = []
- with open(path, "r", encoding="utf-8") as f:
- for i, line in enumerate(f):
- if limit and i >= limit:
- break
- record = json.loads(line.strip())
- result.append((
- record.get("text", ""),
- record.get("metadata", {})
- ))
- return result
- def load_all_domains(
- self,
- domains: Optional[List[str]] = None,
- limit_per_domain: Optional[int] = None
- ) -> Dict[str, List[str]]:
- """
- 加载多个领域的文本
- Args:
- domains: 领域列表,None 表示加载所有
- limit_per_domain: 每个领域最多加载多少条
- Returns:
- {domain: [texts], ...}
- """
- if domains is None:
- domains = self.list_available_domains()
- result = {}
- for domain in domains:
- result[domain] = self.load_texts(domain, limit_per_domain)
- return result
- def list_available_domains(self) -> List[str]:
- """列出所有可用的领域"""
- domains = []
- for path in self.base_dir.iterdir():
- if path.is_dir() and self.get_texts_path(path.name).exists():
- domains.append(path.name)
- return sorted(domains)
- def save_texts(
- self,
- texts: List[str],
- domain: str,
- metadata_list: Optional[List[Dict]] = None,
- source: str = ""
- ) -> None:
- """
- 保存文本到语料文件
- Args:
- texts: 文本列表
- domain: 领域标识
- metadata_list: 可选的元数据列表
- source: 数据来源说明
- """
- path = self.get_texts_path(domain)
- path.parent.mkdir(parents=True, exist_ok=True)
- with open(path, "w", encoding="utf-8") as f:
- for i, text in enumerate(texts):
- record = {
- "text": text,
- "metadata": {
- "index": i,
- "source": source,
- **(metadata_list[i] if metadata_list and i < len(metadata_list) else {})
- }
- }
- f.write(json.dumps(record, ensure_ascii=False) + "\n")
- def get_corpus_stats(self, domain: str) -> Dict:
- """获取语料库统计信息"""
- texts = self.load_texts(domain)
- lengths = [len(t) for t in texts]
- char_lengths = [len(t) for t in texts]
- return {
- "domain": domain,
- "num_texts": len(texts),
- "total_chars": sum(char_lengths),
- "avg_chars": sum(char_lengths) / len(texts) if texts else 0,
- "min_chars": min(char_lengths) if texts else 0,
- "max_chars": max(char_lengths) if texts else 0,
- }
- def compute_corpus_statistics(self) -> Dict[str, Dict]:
- """计算所有领域的统计信息"""
- result = {}
- for domain in self.list_available_domains():
- result[domain] = self.get_corpus_stats(domain)
- return result
- class CorpusBuilder:
- """
- 语料构建器 - 用于从原始数据构建标准格式语料
- (当需要预处理原始数据时使用)
- """
- def __init__(self, output_dir: str = "database/corpus"):
- self.output_dir = Path(output_dir)
- def from_text_files(
- self,
- input_dir: str,
- domain: str,
- pattern: str = "*.txt",
- encoding: str = "utf-8"
- ) -> int:
- """
- 从目录中的文本文件构建语料
- Args:
- input_dir: 输入目录
- domain: 领域标识
- pattern: 文件匹配模式
- encoding: 文件编码
- Returns:
- 加载的文本数量
- """
- input_path = Path(input_dir)
- texts = []
- for file_path in input_path.glob(pattern):
- with open(file_path, "r", encoding=encoding) as f:
- content = f.read().strip()
- if content:
- texts.append(content)
- loader = CorpusLoader(self.output_dir)
- loader.save_texts(texts, domain, source=f"from_text_files:{input_dir}")
- return len(texts)
- def from_jsonl(
- self,
- input_path: str,
- domain: str,
- text_field: str = "text"
- ) -> int:
- """
- 从 JSONL 文件复制/转换语料
- Args:
- input_path: 输入 JSONL 路径
- domain: 领域标识
- text_field: 文本字段名
- Returns:
- 加载的文本数量
- """
- texts = []
- with open(input_path, "r", encoding="utf-8") as f:
- for line in f:
- record = json.loads(line.strip())
- texts.append(record.get(text_field, ""))
- loader = CorpusLoader(self.output_dir)
- loader.save_texts(texts, domain, source=f"from_jsonl:{input_path}")
- return len(texts)
- # ============= 便捷函数 =============
- def load_texts(
- domain: str,
- corpus_dir: str = "database/corpus",
- limit: Optional[int] = None
- ) -> List[str]:
- """便捷函数:加载单个领域的文本"""
- loader = CorpusLoader(corpus_dir)
- return loader.load_texts(domain, limit)
- def load_all_texts(
- domains: Optional[List[str]] = None,
- corpus_dir: str = "database/corpus",
- limit_per_domain: Optional[int] = None
- ) -> Dict[str, List[str]]:
- """便捷函数:加载多个领域的文本"""
- loader = CorpusLoader(corpus_dir)
- return loader.load_all_domains(domains, limit_per_domain)
- def save_texts(
- texts: List[str],
- domain: str,
- corpus_dir: str = "database/corpus",
- source: str = ""
- ) -> None:
- """便捷函数:保存文本语料"""
- loader = CorpusLoader(corpus_dir)
- loader.save_texts(texts, domain, source=source)
- def list_domains(corpus_dir: str = "database/corpus") -> List[str]:
- """便捷函数:列出所有可用领域"""
- loader = CorpusLoader(corpus_dir)
- return loader.list_available_domains()
|