(self)
| 25 | class TestCSVSaver(unittest.TestCase): |
| 26 | |
| 27 | def test_saved_content(self): |
| 28 | with tempfile.TemporaryDirectory() as tempdir: |
| 29 | saver = CSVSaver(output_dir=tempdir, filename="predictions.csv", delimiter="\t") |
| 30 | meta_data = {"filename_or_obj": ["testfile" + str(i) for i in range(8)]} |
| 31 | saver.save_batch(torch.zeros(8), meta_data) |
| 32 | saver.finalize() |
| 33 | filepath = os.path.join(tempdir, "predictions.csv") |
| 34 | self.assertTrue(os.path.exists(filepath)) |
| 35 | with open(filepath) as f: |
| 36 | reader = csv.reader(f, delimiter="\t") |
| 37 | i = 0 |
| 38 | for row in reader: |
| 39 | self.assertEqual(row[0], "testfile" + str(i)) |
| 40 | self.assertEqual(np.array(row[1:]).astype(np.float32), 0.0) |
| 41 | i += 1 |
| 42 | self.assertEqual(i, 8) |
| 43 | |
| 44 | |
| 45 | if __name__ == "__main__": |
nothing calls this directly
no test coverage detected