Repository navigation
Expand file tree
/
Copy pathtest_models.py
More file actions
143 lines (111 loc) · 4.26 KB
/
Copy pathtest_models.py
File metadata and controls
143 lines (111 loc) · 4.26 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
"""
Test script to verify model implementations work correctly
"""
import sys
import os
# Add current directory to path
sys.path.append(os.path.dirname(os.path.abspath(__file__)))
def test_imports():
"""Test if all required modules can be imported"""
try:
import torch
print("✓ PyTorch imported successfully")
import numpy as np
print("✓ NumPy imported successfully")
# Test our model modules
from taskB import EncoderRNN, DecoderRNN, Seq2seqRNN, HierEncoderRNN, Decoder2RNN
print("✓ Model classes imported successfully")
return True
except ImportError as e:
print(f"✗ Import error: {e}")
return False
def test_basic_models():
"""Test basic model instantiation"""
try:
from taskB import EncoderRNN, DecoderRNN, Seq2seqRNN
vocab_size = 1000
embed_dim = 300
hidden_dim = 300
# Test EncoderRNN
encoder = EncoderRNN(vocab_size, embed_dim, hidden_dim)
print("✓ EncoderRNN created successfully")
# Test DecoderRNN
decoder = DecoderRNN(vocab_size, embed_dim, hidden_dim)
print("✓ DecoderRNN created successfully")
# Test Seq2seqRNN
model = Seq2seqRNN(vocab_size, embed_dim, hidden_dim)
print("✓ Seq2seqRNN created successfully")
return True
except Exception as e:
print(f"✗ Model creation error: {e}")
return False
def test_advanced_models():
"""Test advanced model instantiation"""
try:
from taskB import HierEncoderRNN, Decoder2RNN, Seq2seqRNN
vocab_size = 1000
embed_dim = 300
hidden_dim = 300
# Test HierEncoderRNN
hier_encoder = HierEncoderRNN(vocab_size, embed_dim, hidden_dim)
print("✓ HierEncoderRNN created successfully")
# Test Decoder2RNN
decoder2 = Decoder2RNN(vocab_size, embed_dim, hidden_dim)
print("✓ Decoder2RNN created successfully")
# Test Seq2seqRNN with advanced components
model_hier = Seq2seqRNN(vocab_size, embed_dim, hidden_dim, encoder_type='hierarchical')
print("✓ Seq2seqRNN with hierarchical encoder created successfully")
model_dual = Seq2seqRNN(vocab_size, embed_dim, hidden_dim, decoder_type='dual')
print("✓ Seq2seqRNN with dual decoder created successfully")
return True
except Exception as e:
print(f"✗ Advanced model creation error: {e}")
return False
def test_forward_pass():
"""Test forward pass with dummy data"""
try:
import torch
from taskB import Seq2seqRNN
vocab_size = 1000
embed_dim = 300
hidden_dim = 300
batch_size = 2
seq_len = 10
# Create model
model = Seq2seqRNN(vocab_size, embed_dim, hidden_dim)
# Create dummy data
src = torch.randint(0, vocab_size, (batch_size, seq_len))
tgt = torch.randint(0, vocab_size, (batch_size, seq_len))
src_lengths = [seq_len] * batch_size
# Forward pass
output = model(src, src_lengths, tgt)
print(f"✓ Forward pass successful, output shape: {output.shape}")
return True
except Exception as e:
print(f"✗ Forward pass error: {e}")
return False
def main():
"""Run all tests"""
print("Testing NLP Assignment 2 Models")
print("=" * 40)
tests = [
("Import Test", test_imports),
("Basic Models Test", test_basic_models),
("Advanced Models Test", test_advanced_models),
("Forward Pass Test", test_forward_pass)
]
results = []
for test_name, test_func in tests:
print(f"\n{test_name}:")
result = test_func()
results.append((test_name, result))
print("\n" + "=" * 40)
print("Test Summary:")
for test_name, result in results:
status = "PASS" if result else "FAIL"
print(f" {test_name}: {status}")
all_passed = all(result for _, result in results)
print(f"\nOverall: {'ALL TESTS PASSED' if all_passed else 'SOME TESTS FAILED'}")
return all_passed
if __name__ == "__main__":
main()