Models¶
Neural network architectures for sequence-to-sequence translation.
Overview¶
TorchLingo provides two model architectures:
| Model | Architecture | Use Case |
|---|---|---|
SimpleTransformer |
Transformer with sinusoidal positional encoding | Modern, best quality |
SimpleSeq2SeqLSTM |
LSTM encoder-decoder | Classic, simpler |
Quick Comparison¶
flowchart LR
subgraph LSTM
A1[Sequential] --> B1[Hidden State]
B1 --> C1[One at a time]
end
subgraph Transformer
A2[Parallel] --> B2[Attention]
B2 --> C2[All at once]
end
| Feature | Transformer | LSTM |
|---|---|---|
| Training speed | Fast (parallel) | Slow (sequential) |
| Long sequences | Handles well | Struggles |
| Memory | O(n²) | O(n) |
| Parameters | More | Fewer |
| Quality | Better | Good |
Submodules¶
-
Transformer
Modern encoder-decoder with multi-head attention and sinusoidal positional encoding.
-
LSTM
Classic sequence-to-sequence with LSTM cells.
-
Positional Encoding
Sinusoidal positional encoding implementation.
Quick Start¶
Transformer¶
from torchlingo.models import SimpleTransformer
model = SimpleTransformer(
src_vocab_size=10000,
tgt_vocab_size=10000,
d_model=512,
n_heads=8,
num_encoder_layers=6,
num_decoder_layers=6,
)
# Forward pass
logits = model(src_batch, tgt_batch)
LSTM¶
from torchlingo.models import SimpleSeq2SeqLSTM
model = SimpleSeq2SeqLSTM(
src_vocab_size=10000,
tgt_vocab_size=10000,
emb_dim=256,
hidden_dim=512,
num_layers=2,
)
# Forward pass
logits = model(src_batch, tgt_batch)
Common Interface¶
Both models share a similar interface:
# Training forward pass
logits = model(src, tgt) # [batch, tgt_len, vocab_size]
# Encode only (for inference)
memory = model.encode(src) # [batch, src_len, d_model]
# Decode with memory (for inference)
logits = model.decode(tgt, memory)
Model Sizing Guide¶
Tiny (Testing/Demo)¶
config = Config(
d_model=64,
n_heads=2,
num_encoder_layers=1,
num_decoder_layers=1,
)
# ~500K parameters
Small (Learning)¶
config = Config(
d_model=256,
n_heads=8,
num_encoder_layers=4,
num_decoder_layers=4,
)
# ~15M parameters
Medium (Production)¶
config = Config(
d_model=512,
n_heads=8,
num_encoder_layers=6,
num_decoder_layers=6,
)
# ~65M parameters
Saving and Loading¶
import torch
# Save model
torch.save(model.state_dict(), "model.pt")
# Load model
model = SimpleTransformer(src_vocab_size, tgt_vocab_size)
model.load_state_dict(torch.load("model.pt"))
# For complete checkpoints (recommended)
torch.save({
'model_state_dict': model.state_dict(),
'optimizer_state_dict': optimizer.state_dict(),
'epoch': epoch,
'loss': loss,
}, "checkpoint.pt")