Extends layer_embedding() for language models with a reverse projection
from the embedding dimension to the vocabulary dimension. Call the layer
with reverse = TRUE to perform this projection. By default, the reverse
projection uses the transpose of the embedding matrix, tying the two sets of
weights. With tie_weights = FALSE, it uses a separate trainable variable.
This layer has no bias terms.
Usage
layer_reversible_embedding(
object,
input_dim,
output_dim,
tie_weights = TRUE,
embeddings_initializer = "uniform",
embeddings_regularizer = NULL,
embeddings_constraint = NULL,
mask_zero = FALSE,
reverse_dtype = NULL,
logit_soft_cap = NULL,
...
)Arguments
- object
Object to compose the layer with. A tensor, array, or sequential model.
- input_dim
Integer. Size of the vocabulary, i.e. maximum integer index + 1.
- output_dim
Integer. Dimension of the dense embedding.
- tie_weights
Boolean. Whether the embedding and reverse-projection matrices share the same weights. Defaults to
TRUE.- embeddings_initializer
Initializer for the
embeddingsmatrix (seekeras3::initializer_*).- embeddings_regularizer
Regularizer function applied to the
embeddingsmatrix (seekeras3::regularizer_*).- embeddings_constraint
Constraint function applied to the
embeddingsmatrix (seekeras3::constraint_*).- mask_zero
Boolean, whether or not the input value 0 is a special "padding" value that should be masked out. This is useful when using recurrent layers which may take variable length input. If this is
TRUE, then all subsequent layers in the model need to support masking or an exception will be raised. Ifmask_zerois set toTRUE, as a consequence, index 0 cannot be used in the vocabulary (input_dimshould equal size of vocabulary + 1).- reverse_dtype
Dtype used for the reverse-projection computation. Defaults to the layer's compute dtype.
- logit_soft_cap
Optional positive number. When set, reverse-projection logits are scaled by
tanh(logits / logit_soft_cap) * logit_soft_cap. This narrows the range of output logits and can improve training.- ...
For forward/backward compatibility.
Value
The return value depends on the value provided for the first argument.
If object is:
a
keras_model_sequential(), then the layer is added to the sequential model (which is modified in place). To enable piping, the sequential model is also returned, invisibly.a
keras_input(), then the output tensor from callinglayer(input)is returned.NULLor missing, then aLayerinstance is returned.
Examples
embedding <- layer_reversible_embedding(input_dim = 8, output_dim = 4)
token_ids <- op_array(matrix(c(0L, 1L, 2L), nrow = 1), dtype = "int32")
hidden_states <- embedding(token_ids)
logits <- embedding(hidden_states, reverse = TRUE)
shape(logits)See also
Other core layers: layer_dense() layer_einsum_dense() layer_embedding() layer_identity() layer_lambda() layer_masking()
Other layers: Layer() layer_activation() layer_activation_elu() layer_activation_leaky_relu() layer_activation_parametric_relu() layer_activation_relu() layer_activation_softmax() layer_activity_regularization() layer_adaptive_average_pooling_1d() layer_adaptive_average_pooling_2d() layer_adaptive_average_pooling_3d() layer_adaptive_max_pooling_1d() layer_adaptive_max_pooling_2d() layer_adaptive_max_pooling_3d() layer_add() layer_additive_attention() layer_alpha_dropout() layer_attention() layer_aug_mix() layer_auto_contrast() layer_average() layer_average_pooling_1d() layer_average_pooling_2d() layer_average_pooling_3d() layer_batch_normalization() layer_bidirectional() layer_category_encoding() layer_center_crop() layer_concatenate() layer_contrast_limited_adaptive_histogram_equalization() layer_conv_1d() layer_conv_1d_transpose() layer_conv_2d() layer_conv_2d_transpose() layer_conv_3d() layer_conv_3d_transpose() layer_conv_lstm_1d() layer_conv_lstm_2d() layer_conv_lstm_3d() layer_cropping_1d() layer_cropping_2d() layer_cropping_3d() layer_cut_mix() layer_dense() layer_depthwise_conv_1d() layer_depthwise_conv_2d() layer_discretization() layer_dot() layer_dropout() layer_einsum_dense() layer_embedding() layer_equalization() layer_feature_space() layer_flatten() layer_flax_module_wrapper() layer_gaussian_dropout() layer_gaussian_noise() layer_global_average_pooling_1d() layer_global_average_pooling_2d() layer_global_average_pooling_3d() layer_global_max_pooling_1d() layer_global_max_pooling_2d() layer_global_max_pooling_3d() layer_group_normalization() layer_group_query_attention() layer_gru() layer_hashed_crossing() layer_hashing() layer_identity() layer_integer_lookup() layer_jax_model_wrapper() layer_lambda() layer_layer_normalization() layer_lstm() layer_masking() layer_max_num_bounding_boxes() layer_max_pooling_1d() layer_max_pooling_2d() layer_max_pooling_3d() layer_maximum() layer_mel_spectrogram() layer_minimum() layer_mix_up() layer_multi_head_attention() layer_multiply() layer_normalization() layer_permute() layer_rand_augment() layer_random_brightness() layer_random_color_degeneration() layer_random_color_jitter() layer_random_contrast() layer_random_crop() layer_random_elastic_transform() layer_random_erasing() layer_random_flip() layer_random_gaussian_blur() layer_random_grayscale() layer_random_hue() layer_random_invert() layer_random_perspective() layer_random_posterization() layer_random_rotation() layer_random_saturation() layer_random_sharpness() layer_random_shear() layer_random_translation() layer_random_zoom() layer_repeat_vector() layer_rescaling() layer_reshape() layer_resizing() layer_rms_normalization() layer_rnn() layer_separable_conv_1d() layer_separable_conv_2d() layer_simple_rnn() layer_solarization() layer_spatial_dropout_1d() layer_spatial_dropout_2d() layer_spatial_dropout_3d() layer_spectral_normalization() layer_stft_spectrogram() layer_string_lookup() layer_subtract() layer_text_vectorization() layer_tfsm() layer_time_distributed() layer_torch_module_wrapper() layer_unit_normalization() layer_upsampling_1d() layer_upsampling_2d() layer_upsampling_3d() layer_zero_padding_1d() layer_zero_padding_2d() layer_zero_padding_3d() rnn_cell_gru() rnn_cell_lstm() rnn_cell_simple() rnn_cells_stack()