Skip to content

Commit 993caf6

Browse files
authored
Refactor LSTM implementation and clean up code
Removed author information and unused print statements. Updated parameter names in docstrings for consistency.
1 parent 2ee5df1 commit 993caf6

1 file changed

Lines changed: 25 additions & 35 deletions

File tree

‎neural_network/lstm.py‎

Lines changed: 25 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,8 @@
1-
import numpy as np
2-
from numpy.random import Generator
3-
41
"""
5-
Author : Shashank Tyagi
6-
Email : tyagishashank118@gmail.com
7-
Description : This is a simple implementation of Long Short-Term Memory (LSTM)
8-
networks in Python.
2+
A simple implementation of Long Short-Term Memory (LSTM) networks in Python.
93
"""
4+
import numpy as np
5+
from numpy.random import Generator
106

117

128
class LongShortTermMemory:
@@ -46,10 +42,6 @@ def __init__(
4642
self.data_length: int = len(self.input_data)
4743
self.vocabulary_size: int = len(self.unique_chars)
4844

49-
# print(
50-
# f"Data length: {self.data_length}, Vocabulary size: {self.vocabulary_size}"
51-
# )
52-
5345
self.char_to_index: dict[str, int] = {
5446
c: i for i, c in enumerate(self.unique_chars)
5547
}
@@ -192,7 +184,7 @@ def sigmoid(self, input_array: np.ndarray, derivative: bool = False) -> np.ndarr
192184
"""
193185
Sigmoid activation function.
194186
195-
:param x: The input array.
187+
:param input_array: The input array.
196188
:param derivative: Whether to compute the derivative.
197189
:return: The sigmoid activation or its derivative.
198190
@@ -202,7 +194,7 @@ def sigmoid(self, input_array: np.ndarray, derivative: bool = False) -> np.ndarr
202194
True
203195
>>> np.round(output, 3)
204196
array([[0.731, 0.881, 0.953]])
205-
>>> derivative_output = lstm.sigmoid(output, derivative=True)
197+
>>> derivative_output = lstm.sigmoid(input_array=output, derivative=True)
206198
>>> np.round(derivative_output, 3)
207199
array([[0.197, 0.105, 0.045]])
208200
"""
@@ -214,17 +206,17 @@ def tanh(self, input_array: np.ndarray, derivative: bool = False) -> np.ndarray:
214206
"""
215207
Tanh activation function.
216208
217-
:param x: The input array.
209+
:param input_array: The input array.
218210
:param derivative: Whether to compute the derivative.
219211
:return: The tanh activation or its derivative.
220212
221213
>>> lstm = LongShortTermMemory("abcde" * 50, hidden_layer_size=10)
222-
>>> output = lstm.tanh(np.array([[1, 2, 3]]))
214+
>>> output = lstm.tanh(np.array(input_array=[[1, 2, 3]]))
223215
>>> isinstance(output, np.ndarray)
224216
True
225217
>>> np.round(output, 3)
226218
array([[0.762, 0.964, 0.995]])
227-
>>> derivative_output = lstm.tanh(output, derivative=True)
219+
>>> derivative_output = lstm.tanh(input_array=output, derivative=True)
228220
>>> np.round(derivative_output, 3)
229221
array([[0.42 , 0.071, 0.01 ]])
230222
"""
@@ -236,11 +228,11 @@ def softmax(self, input_array: np.ndarray) -> np.ndarray:
236228
"""
237229
Softmax activation function.
238230
239-
:param x: The input array.
231+
:param input_array: The input array.
240232
:return: The softmax activation.
241233
242234
>>> lstm = LongShortTermMemory("abcde" * 50, hidden_layer_size=10)
243-
>>> output = lstm.softmax(np.array([1, 2, 3]))
235+
>>> output = lstm.softmax(input_array=np.array([1, 2, 3]))
244236
>>> isinstance(output, np.ndarray)
245237
True
246238
>>> np.round(output, 3)
@@ -496,14 +488,13 @@ def test(self) -> str:
496488
if prediction == self.target_sequence[t]:
497489
accuracy += 1
498490

499-
# print(f"Ground Truth:\n{self.target_sequence}\n")
500-
# print(f"Predictions:\n{output}\n")
501-
# print(f"Accuracy: {round(accuracy * 100 / len(self.input_sequence), 2)}%")
502-
503491
return output
504492

505493

506-
if __name__ == "__main__":
494+
def test_with_sample_data() -> None:
495+
"""
496+
>>> test_with_sample_data()
497+
"""
507498
sample_data = """Long Short-Term Memory (LSTM) networks are a type
508499
of recurrent neural network (RNN) capable of learning "
509500
"order dependence in sequence prediction problems.
@@ -512,19 +503,18 @@ def test(self) -> str:
512503
LSTMs were introduced by Hochreiter and Schmidhuber in 1997, and were
513504
refined and "
514505
"popularized by many people in following work."""
515-
import doctest
516506

517-
doctest.testmod()
507+
stm_model = LongShortTermMemory(
508+
input_data=sample_data,
509+
hidden_layer_size=25,
510+
training_epochs=100,
511+
learning_rate=0.05,
512+
)
513+
lstm_model.train()
514+
lstm_model.test()
518515

519-
# lstm_model = LongShortTermMemory(
520-
# input_data=sample_data,
521-
# hidden_layer_size=25,
522-
# training_epochs=100,
523-
# learning_rate=0.05,
524-
# )
525516

526-
# #### Training #####
527-
# lstm_model.train()
517+
if __name__ == "__main__":
518+
import doctest
528519

529-
# #### Testing #####
530-
# lstm_model.test()
520+
doctest.testmod()

0 commit comments

Comments
 (0)