Spaces:
Sleeping
Sleeping
| """ | |
| 数据源契约定义。 | |
| 这个包集中放 data 层对外承诺的数据源形状。 | |
| Pipeline 只依赖这里的协议,具体数据集负责显式实现对应协议。 | |
| """ | |
| from dataclasses import dataclass | |
| from typing import Callable, Protocol | |
| import keras | |
| import tensorflow as tf | |
| 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"] | |