YAML Metadata Warning:empty or missing yaml metadata in repo card

Check out the documentation for more information.


license: mit ---import numpy as np import tensorflow as tf from tensorflow.keras.models import Model from tensorflow.keras.layers import Input, LSTM, Embedding, Dense

Parameters

embedding_dim = 256 lstm_units = 512 max_seq_length = 20 vocab_size_source = 5000 # Adjust based on your dataset vocab_size_target = 5000

Encoder

encoder_inputs = Input(shape=(None,), name="encoder_inputs") encoder_embedding = Embedding(vocab_size_source, embedding_dim, mask_zero=True)(encoder_inputs) encoder_lstm = LSTM(lstm_units, return_state=True, name="encoder_lstm") encoder_outputs, state_h, state_c = encoder_lstm(encoder_embedding) encoder_states = [state_h, state_c]

Decoder

decoder_inputs = Input(shape=(None,), name="decoder_inputs") decoder_embedding = Embedding(vocab_size_target, embedding_dim, mask_zero=True)(decoder_inputs) decoder_lstm = LSTM(lstm_units, return_sequences=True, return_state=True, name="decoder_lstm") decoder_outputs, _, _ = decoder_lstm(decoder_embedding, initial_state=encoder_states) decoder_dense = Dense(vocab_size_target, activation='softmax', name="decoder_dense") decoder_outputs = decoder_dense(decoder_outputs)

Model

model = Model([encoder_inputs, decoder_inputs], decoder_outputs) model.compile(optimizer='adam', loss='sparse_categorical_crossentropy', metrics=['accuracy'])

Summary

model.summary()

Example: Dummy Data

batch_size = 64 num_samples = 10000

encoder_input_data = np.random.randint(1, vocab_size_source, (num_samples, max_seq_length)) decoder_input_data = np.random.randint(1, vocab_size_target, (num_samples, max_seq_length)) decoder_output_data = np.random.randint(1, vocab_size_target, (num_samples, max_seq_length))

Add one extra dimension for sparse categorical crossentropy

decoder_output_data = np.expand_dims(decoder_output_data, -1)

Training

model.fit( [encoder_input_data, decoder_input_data], decoder_output_data, batch_size=batch_size, epochs=10, validation_split=0.2 )

Save the model

model.save("lstm_translator.h5")

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support