|
2 | 2 | from typing import Dict, Tuple, TYPE_CHECKING |
3 | 3 |
|
4 | 4 | import torch |
| 5 | +from allennlp.common import Params |
5 | 6 |
|
6 | 7 | from allennlp.common.checks import ConfigurationError |
7 | 8 | from allennlp.data import TokenIndexer, Token |
| 9 | +from allennlp.modules import TextFieldEmbedder |
8 | 10 | from allennlp.modules.scalar_mix import ScalarMix |
| 11 | +from allennlp.modules.text_field_embedders import BasicTextFieldEmbedder |
| 12 | +from allennlp.modules.token_embedders import EmptyEmbedder |
9 | 13 | from allennlp.modules.token_embedders.token_embedder import TokenEmbedder |
10 | 14 | from allennlp.nn.util import ( |
11 | 15 | remove_sentence_boundaries, |
@@ -76,15 +80,29 @@ def __init__( |
76 | 80 |
|
77 | 81 | # Extract the name of the tokens that the LM was trained on. |
78 | 82 | text_field_embedder = dict_config["model"]["text_field_embedder"] |
79 | | - token_names = list(text_field_embedder["token_embedders"].keys()) |
80 | | - if len(token_names) != 1: |
81 | | - # We don't currently support embedding with language models trained with multiple |
82 | | - # embedded indices. |
83 | | - # |
84 | | - # Note: We only care about embedded indices. This does not include "tokens" which |
85 | | - # is just used to compute the loss in LanguageModel. |
86 | | - raise ConfigurationError(f"LM from {archive_file} trained with multiple embedders!") |
87 | | - self._token_name = token_names[0] |
| 83 | + text_field_embedder = TextFieldEmbedder.from_params(Params(text_field_embedder)) |
| 84 | + if not isinstance(text_field_embedder, BasicTextFieldEmbedder): |
| 85 | + raise ConfigurationError( |
| 86 | + f"Language model from {archive_file} uses a non-standard TextFieldEmbedder!" |
| 87 | + ) |
| 88 | + non_empty_embedders = [ |
| 89 | + name |
| 90 | + for name, token_embedder in text_field_embedder._token_embedders.items() |
| 91 | + if not isinstance(token_embedder, EmptyEmbedder) |
| 92 | + ] |
| 93 | + |
| 94 | + if len(non_empty_embedders) == 0: |
| 95 | + # Only empty embedders were contained in the language model |
| 96 | + # We need at least one non-empty embedder in the language model |
| 97 | + raise ConfigurationError( |
| 98 | + f"Language model from {archive_file} trained with only empty embedders!" |
| 99 | + ) |
| 100 | + elif len(non_empty_embedders) > 1: |
| 101 | + raise ConfigurationError( |
| 102 | + f"Language model from {archive_file} trained with multiple non-empty embedders!" |
| 103 | + ) |
| 104 | + |
| 105 | + self._token_name = non_empty_embedders[0] |
88 | 106 |
|
89 | 107 | # TODO(brendanr): Find a way to remove this hack. The issue fundamentally is that the |
90 | 108 | # BasicTextFieldEmbedder concatenates multiple embedded representations. When a |
|
0 commit comments