File size: 3,041 Bytes
a5fd608
 
07cb7d3
a5fd608
 
 
 
 
 
 
 
07cb7d3
 
 
 
a5fd608
 
 
07cb7d3
a5fd608
 
 
 
 
 
 
 
 
 
 
 
 
 
07cb7d3
a5fd608
 
 
 
 
07cb7d3
 
 
 
a5fd608
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
07cb7d3
a5fd608
 
 
 
 
 
 
 
 
 
 
07cb7d3
 
a5fd608
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
"""Wiki 数据集主模块

实现 WikiDataset 类,继承自 TextGenerationDataSource。
"""

import pathlib
from dataclasses import dataclass, field
from typing import Optional

import tensorflow as tf

from deep_learning.data.spec import TextGenerationDataSource, TokenizerBundle
from .loader import doc_load
from .transformer import transform
from .tokenizer import sentence_piece, character_vectorization


@dataclass
class WikiDataset(TextGenerationDataSource):
    """Wiki 数据集

    将文档加载、分词、统计等功能绑定在一起的数据集类。

    Usage:
        dataset = WikiDataset(
            data_dir="~/data/wiki/mini_c4",
            tokenizer_type="sentence_piece"  # 或 "character"
        )

        # 获取文档数据集
        doc_ds = dataset.doc_ds()

        # 获取 token 数据集
        tokens_ds = dataset.tokens_ds()

        # 打印统计信息
        dataset.stat(seq_length=256)
    """

    data_dir: str
    sequence_length: int = 256
    batch_size: int = 32
    validation_batches: int = 0
    glob_pattern: str = "*"
    tokenizer_type: str = "sentence_piece"

    _data_path: pathlib.Path = field(init=False, repr=False)
    _tokenizer_bundle: Optional[TokenizerBundle] = field(
        init=False, repr=False, default=None
    )

    def __post_init__(self):
        self._data_path = pathlib.Path(self.data_dir).expanduser()

    def _load_tokenizer(self):
        """懒加载分词器"""
        if self._tokenizer_bundle is None:
            if self.tokenizer_type == "sentence_piece":
                tokenizer, end_of_text, decode = sentence_piece()
            elif self.tokenizer_type == "character":
                tokenizer, end_of_text, decode = character_vectorization()
            else:
                raise ValueError(f"Unknown tokenizer type: {self.tokenizer_type}")

            vocab_size = tokenizer.vocabulary_size()
            self._tokenizer_bundle = TokenizerBundle(
                tokenizer=tokenizer,
                decode=decode,
                end_of_text=end_of_text,
                vocab_size=vocab_size
            )

    def doc_ds(self) -> tf.data.Dataset:
        """返回原始文档数据集

        Returns:
            TensorFlow Dataset,每个元素是一个文档字符串
        """
        return doc_load(self._data_path, glob_pattern=self.glob_pattern)

    def tokens_ds(self) -> tf.data.Dataset:
        """返回 tokenized 数据集

        Returns:
            TensorFlow Dataset,每个元素是 (input_ids, target_ids) 对
        """
        self._load_tokenizer()
        ds = self.doc_ds()
        return transform(
            ds=ds,
            tokenizer=self._tokenizer_bundle.tokenizer,
            end_of_text=self._tokenizer_bundle.end_of_text,
            sequence_length=self.sequence_length,
            batch_size=self.batch_size
        )

    def tokenizer_bundle(self) -> TokenizerBundle:
        """返回分词器信息"""
        self._load_tokenizer()
        return self._tokenizer_bundle