Skip to contents

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 embeddings matrix (see keras3::initializer_*).

embeddings_regularizer

Regularizer function applied to the embeddings matrix (see keras3::regularizer_*).

embeddings_constraint

Constraint function applied to the embeddings matrix (see keras3::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. If mask_zero is set to TRUE, as a consequence, index 0 cannot be used in the vocabulary (input_dim should 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 calling layer(input) is returned.

  • NULL or missing, then a Layer instance is returned.

Details

Call arguments

  • inputs: Tensor inputs to the layer.

  • reverse: Boolean. Whether to project from output_dim to input_dim instead of performing a normal embedding lookup. Defaults to FALSE.

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)

## shape(1, 3, 8)

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()