SubwordField
| 172 | |
| 173 | |
| 174 | class SubwordField(Field): |
| 175 | """SubwordField""" |
| 176 | def __init__(self, *args, **kwargs): |
| 177 | self.fix_len = kwargs.pop('fix_len') if 'fix_len' in kwargs else 0 |
| 178 | super(SubwordField, self).__init__(*args, **kwargs) |
| 179 | |
| 180 | def build(self, corpus, min_freq=1): |
| 181 | """Create vocab based on corpus""" |
| 182 | if hasattr(self, 'vocab'): |
| 183 | return |
| 184 | sequences = getattr(corpus, self.name) |
| 185 | counter = Counter(piece for seq in sequences for token in seq for piece in self.preprocess(token)) |
| 186 | self.vocab = Vocab(counter, min_freq, self.specials, self.unk_index) |
| 187 | |
| 188 | def transform(self, sequences): |
| 189 | """Sequences transform function, such as converting word to id, adding bos tags to sequences, etc.""" |
| 190 | sequences = [[self.preprocess(token) for token in seq] for seq in sequences] |
| 191 | if self.fix_len <= 0: |
| 192 | self.fix_len = max(len(token) for seq in sequences for token in seq) |
| 193 | if self.use_vocab: |
| 194 | sequences = [[[self.vocab[i] for i in token] for token in seq] for seq in sequences] |
| 195 | if self.bos: |
| 196 | sequences = [[[self.bos_index]] + seq for seq in sequences] |
| 197 | if self.eos: |
| 198 | sequences = [seq + [[self.eos_index]] for seq in sequences] |
| 199 | sequences = [ |
| 200 | nn.pad_sequence([np.array(ids[:self.fix_len], dtype=np.int64) for ids in seq], self.pad_index, self.fix_len) |
| 201 | for seq in sequences |
| 202 | ] |
| 203 | |
| 204 | return sequences |
| 205 | |
| 206 | |
| 207 | class ErnieField(Field): |