public class TrainableWordEmbedding extends Embedding<java.lang.String> implements WordEmbedding
TrainableWordEmbedding is an implementation of WordEmbedding and Embedding based on a SimpleVocabulary. This WordEmbedding is ideal when there
are no pre-trained embeddings available.| Modifier and Type | Class and Description |
|---|---|
static class |
TrainableWordEmbedding.Builder
A builder for a
TrainableWordEmbedding. |
Embedding.BaseBuilder<T,B extends Embedding.BaseBuilder<T,B>>, Embedding.DefaultEmbedding, Embedding.DefaultItemdataType, embedder, embedding, embeddingSize, fallthroughEmbedding, numItems, sparseGrad, unembedderinputNames, inputShapes| Constructor and Description |
|---|
TrainableWordEmbedding(NDArray embedding,
java.util.List<java.lang.String> items)
Constructs a pretrained embedding.
|
TrainableWordEmbedding(NDArray embedding,
java.util.List<java.lang.String> items,
boolean sparseGrad)
Constructs a pretrained embedding.
|
TrainableWordEmbedding(SimpleVocabulary simpleVocabulary,
int embeddingSize)
Constructs a new instance of
TrainableWordEmbedding from a SimpleVocabulary
and a given embedding size. |
TrainableWordEmbedding(TrainableWordEmbedding.Builder builder)
Constructs a new instance of
TrainableWordEmbedding from the TrainableWordEmbedding.Builder. |
| Modifier and Type | Method and Description |
|---|---|
static TrainableWordEmbedding.Builder |
builder()
Creates a builder to build an
Embedding. |
java.lang.String |
decode(byte[] byteArray)
Decodes the given byte array into an object of input parameter type.
|
NDArray |
embedWord(NDManager manager,
int index)
Embeds the word after preprocessed using
WordEmbedding.preprocessWordToEmbed(String). |
byte[] |
encode(java.lang.String input)
Encodes an object of input type into a byte array.
|
int |
preprocessWordToEmbed(java.lang.String word)
Pre-processes the word to embed into an array to pass into the model.
|
java.lang.String |
unembedWord(NDArray word)
Returns the closest matching word for the given index.
|
boolean |
vocabularyContains(java.lang.String word)
Returns whether an embedding exists for a word.
|
embed, embed, forward, getDirectParameters, getOutputShapes, getParameterShape, hasItem, loadParameters, saveParameters, unembedgetChildren, initialize, toStringbeforeInitialize, cast, clear, describeInput, getParameters, isInitialized, readInputShapes, saveInputShapes, setInitializer, setInitializerclone, equals, finalize, getClass, hashCode, notify, notifyAll, wait, wait, waitembedWordforward, validateLayoutpublic TrainableWordEmbedding(TrainableWordEmbedding.Builder builder)
TrainableWordEmbedding from the TrainableWordEmbedding.Builder.builder - the TrainableWordEmbedding.Builderpublic TrainableWordEmbedding(SimpleVocabulary simpleVocabulary, int embeddingSize)
TrainableWordEmbedding from a SimpleVocabulary
and a given embedding size.simpleVocabulary - a SimpleVocabulary to get tokens fromembeddingSize - the required embedding sizepublic TrainableWordEmbedding(NDArray embedding, java.util.List<java.lang.String> items)
embedding - the embedding arrayitems - the items in the embedding (in matching order to the embedding array)public TrainableWordEmbedding(NDArray embedding, java.util.List<java.lang.String> items, boolean sparseGrad)
embedding - the embedding arrayitems - the items in the embedding (in matching order to the embedding array)sparseGrad - whether to compute row sparse gradient in the backward calculationpublic boolean vocabularyContains(java.lang.String word)
vocabularyContains in interface WordEmbeddingword - the word to checkpublic int preprocessWordToEmbed(java.lang.String word)
Make sure to call WordEmbedding.embedWord(NDManager, int) after this.
preprocessWordToEmbed in interface WordEmbeddingword - the word to embedpublic NDArray embedWord(NDManager manager, int index)
WordEmbedding.preprocessWordToEmbed(String).embedWord in interface WordEmbeddingmanager - the manager for the embedding arrayindex - the index of the word to embedpublic java.lang.String unembedWord(NDArray word)
unembedWord in interface WordEmbeddingword - the word embedding to find the matching string word for.public byte[] encode(java.lang.String input)
Embedding objects.encode in interface AbstractIndexedEmbedding<java.lang.String>input - the input object to be encodedpublic java.lang.String decode(byte[] byteArray)
decode in interface AbstractIndexedEmbedding<java.lang.String>byteArray - the byte array to be decodedpublic static TrainableWordEmbedding.Builder builder()
Embedding.