# Source: content/notes/deep-learning/rnns-and-sequence-models.md
# Independent CPU example; use the curriculum environment.
# See /notes/ml/#example-environment or /notes/deep-learning/#example-environment.

import torch
from torch import nn
from torch.nn.utils.rnn import pack_padded_sequence

torch.manual_seed(7)
torch.set_num_threads(1)
B, T, D = 256, 9, 3
lengths = torch.randint(3, T + 1, (B,))
valid = torch.arange(T)[None, :] < lengths[:, None]
x = torch.randn(B, T, D) * valid[..., None]
y = (x[:, :, 0].sum(1) > 0).long()

class SequenceClassifier(nn.Module):
    def __init__(self):
        super().__init__()
        self.rnn = nn.LSTM(D, 12, batch_first=True, bidirectional=True)
        self.head = nn.Linear(24, 2)

    def forward(self, values, sizes):
        packed = pack_padded_sequence(values, sizes.cpu(), batch_first=True,
                                      enforce_sorted=False)
        _, (h, _) = self.rnn(packed)
        final = h.reshape(1, 2, len(sizes), 12)[-1]
        return self.head(torch.cat((final[0], final[1]), dim=-1))

model = SequenceClassifier()
opt = torch.optim.Adam(model.parameters(), lr=0.02)
criterion = nn.CrossEntropyLoss()
model.train()
initial = criterion(model(x[:192], lengths[:192]), y[:192]).item()
for _ in range(90):
    opt.zero_grad(set_to_none=True)
    loss = criterion(model(x[:192], lengths[:192]), y[:192])
    loss.backward()
    nn.utils.clip_grad_norm_(model.parameters(), 1.0)
    opt.step()
model.eval()
with torch.no_grad():
    final_loss = criterion(model(x[:192], lengths[:192]), y[:192]).item()
    scores = model(x[192:], lengths[192:])
    accuracy = (scores.argmax(1) == y[192:]).float().mean().item()
    poisoned = x.clone()
    poisoned[~valid] = 1000.0
    torch.testing.assert_close(model(poisoned, lengths), model(x, lengths))
assert final_loss < initial * 0.5
assert accuracy > 0.70
print({"initial_loss": initial, "final_loss": final_loss, "test_accuracy": accuracy})
