-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain_caesar_rnn.py
More file actions
103 lines (86 loc) · 3.69 KB
/
Copy pathtrain_caesar_rnn.py
File metadata and controls
103 lines (86 loc) · 3.69 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
import argparse
from pathlib import Path
from caesar_rnn.config import DataConfig, ExperimentConfig, ModelConfig, TrainingConfig
from caesar_rnn.data import build_dataloaders
from caesar_rnn.inference import evaluate_examples
from caesar_rnn.model import CaesarDecoderRNN
from caesar_rnn.trainer import save_checkpoint, train_model
from caesar_rnn.utils import RUSSIAN_ALPHABET, get_device, set_seed
from caesar_rnn.vocab import Vocab
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description="Train an RNN to decode Caesar cipher text.")
parser.add_argument("--shift", type=int, default=5, help="Caesar shift used to encrypt the phrases.")
parser.add_argument("--train-size", type=int, default=4000, help="Number of training samples.")
parser.add_argument("--test-size", type=int, default=1000, help="Number of test samples.")
parser.add_argument("--epochs", type=int, default=5, help="Training epochs.")
parser.add_argument("--batch-size", type=int, default=128, help="Mini-batch size.")
parser.add_argument("--embed-dim", type=int, default=48, help="Embedding dimension.")
parser.add_argument("--hidden-dim", type=int, default=96, help="Hidden state dimension.")
parser.add_argument("--learning-rate", type=float, default=0.003, help="Adam learning rate.")
parser.add_argument("--seed", type=int, default=42, help="Random seed.")
parser.add_argument(
"--artifacts-dir",
type=Path,
default=Path("artifacts"),
help="Directory where metrics and model weights are stored.",
)
return parser
def config_from_args(args: argparse.Namespace) -> ExperimentConfig:
return ExperimentConfig(
data=DataConfig(
shift=args.shift,
train_size=args.train_size,
test_size=args.test_size,
),
model=ModelConfig(
embed_dim=args.embed_dim,
hidden_dim=args.hidden_dim,
),
training=TrainingConfig(
epochs=args.epochs,
batch_size=args.batch_size,
learning_rate=args.learning_rate,
seed=args.seed,
artifacts_dir=args.artifacts_dir,
),
)
def main() -> None:
args = build_parser().parse_args()
config = config_from_args(args)
set_seed(config.training.seed)
device = get_device()
vocab = Vocab.build()
train_loader, test_loader = build_dataloaders(config.data, config.training, vocab)
model = CaesarDecoderRNN(vocab_size=vocab.size, pad_idx=vocab.pad_idx, config=config.model).to(device)
print(f"Alphabet: {RUSSIAN_ALPHABET}")
print(f"Shift: {config.data.shift}")
print(f"Device: {device}")
print()
history = train_model(
model=model,
train_loader=train_loader,
test_loader=test_loader,
training_config=config.training,
device=device,
pad_idx=vocab.pad_idx,
)
model_path = config.training.artifacts_dir / "caesar_decoder_rnn.pt"
metrics_path = config.training.artifacts_dir / "metrics.json"
save_checkpoint(model, vocab, config, model_path, metrics_path, history)
demo_samples = [
"машинное обучение",
"рекуррентная сеть",
"шифр цезаря работает",
]
examples = evaluate_examples(model, vocab, config.data.shift, device, demo_samples)
print()
print("Examples:")
for example in examples:
print(f" original : {example['original']}")
print(f" encrypted: {example['encrypted']}")
print(f" predicted: {example['predicted']}")
print()
print(f"Saved model to: {model_path}")
print(f"Saved metrics to: {metrics_path}")
if __name__ == "__main__":
main()