""" 数据源契约定义。 这个包集中放 data 层对外承诺的数据源形状。 Pipeline 只依赖这里的协议,具体数据集负责显式实现对应协议。 """ from dataclasses import dataclass from typing import Callable, Protocol import keras import tensorflow as tf @dataclass class TokenizerBundle: """文本任务的推理配套资源 它描述“模型之外,文本推理还需要什么”,同时也承载页面展示词表信息所需的数据。 """ tokenizer: Callable decode: Callable end_of_text: int vocab_size: int vocab_path: str = "" class TextGenerationDataSource(Protocol): """文本生成数据源需要提供文档、token 数据和分词资源。""" data_dir: str sequence_length: int batch_size: int validation_batches: int def doc_ds(self) -> tf.data.Dataset: """返回原始文档数据集""" ... def tokens_ds(self) -> tf.data.Dataset: """返回 tokenized 数据集""" ... def tokenizer_bundle(self) -> TokenizerBundle: """返回分词器信息""" ... def stat(self, seq_length: int | None = None) -> None: """打印数据集统计信息""" from deep_learning.data.common import collect_stats info = self.tokenizer_bundle() stats = collect_stats( name=self.__class__.__name__, loader=self.doc_ds, tokenizer=info.tokenizer ) stats.print_report(seq_length=seq_length) class SupervisedDataSource(Protocol): """有监督任务数据源需要提供训练数据和样例测试能力。""" def training_ds(self): ... def test_examples(self, model: keras.Model) -> None: ... __all__ = ["SupervisedDataSource", "TextGenerationDataSource", "TokenizerBundle"]