import unittest
import torch
from torch.nn import functional as F
from example import model,batch,train,experiment,TRAIN,HOLDOUT
class TrainingTests(unittest.TestCase):
    def test_labels_are_shifted_once_inside_gpt2(self):
        net=model().eval(); data=batch(['red means rouge','blue'])
        output=net(**data)
        manual=F.cross_entropy(output.logits[:,:-1].reshape(-1,9),data['labels'][:,1:].reshape(-1),ignore_index=-100)
        self.assertAlmostEqual(float(output.loss.detach()),float(manual.detach()),places=6)
    def test_only_padding_is_ignored(self):
        data=batch(['red means rouge','blue'])
        self.assertEqual(data['labels'][1].tolist(),[3,1,-100,-100])
        self.assertEqual(data['input_ids'][1].tolist(),[3,1,0,0])
    def test_parameters_change(self):
        net=model(); before=net.transformer.wte.weight.detach().clone()
        train(net,batch(TRAIN),steps=1)
        self.assertFalse(torch.equal(before,net.transformer.wte.weight))
    def test_fit_and_save_reload(self):
        result=experiment()
        self.assertLess(result['train_after'],result['train_before']/2)
        self.assertTrue(result['reloaded_logits_equal'])
        self.assertEqual(set(TRAIN)&set(HOLDOUT),set())
    def test_bad_batch_and_steps(self):
        for texts in ([],[''],['not-in-vocabulary']):
            with self.assertRaises(ValueError): batch(texts)
        with self.assertRaises(ValueError): train(model(),batch(TRAIN),steps=0)
if __name__=='__main__': unittest.main()
