""" 文本语料加载器接口 功能: - 从预存文件加载文本语料 - 支持按领域(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()