File size: 6,320 Bytes
932bc69 | 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 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 | """CPU-only tests: renewal uses salloc; cancellation/failure never renews."""
import shlex
import subprocess
import tempfile
import unittest
from pathlib import Path
from unittest.mock import MagicMock, patch
from ngram_interactive_supervisor import Supervisor
class SupervisorTests(unittest.TestCase):
def controller(self, root):
item = Supervisor.__new__(Supervisor)
item.root = root
item.source = root / "source"
item.job = "1"
item.ssh = ["ssh"]
item.login = "login"
item.status = MagicMock()
item.child = MagicMock()
item.child.poll.return_value = 1
item.child.returncode = 1
item.maintenance = None
item.start_segment = MagicMock(return_value=root / "segments/1")
return item
def test_timeout_renews(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(return_value="TIMEOUT")
controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal"))
with self.assertRaisesRegex(RuntimeError, "reached renewal"):
controller.run()
controller.allocate.assert_called_once()
def test_cancellation_does_not_renew(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(return_value="CANCELLED")
controller.allocate = MagicMock()
with self.assertRaisesRegex(RuntimeError, "not renewing"):
controller.run()
controller.allocate.assert_not_called()
def test_long_completing_then_timeout_renews(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(
side_effect=["COMPLETING"] * 20 + ["UNKNOWN", "TIMEOUT"]
)
controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal"))
with patch("ngram_interactive_supervisor.time.sleep") as sleep:
with self.assertRaisesRegex(RuntimeError, "reached renewal"):
controller.run()
self.assertEqual(sleep.call_count, 21)
controller.allocate.assert_called_once()
def test_accounting_timeout_overrides_completing(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.remote = MagicMock(side_effect=["COMPLETING", "TIMEOUT|"])
self.assertEqual(controller.job_state(), "TIMEOUT")
def test_expired_job_missing_from_squeue_uses_accounting(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.remote = MagicMock(side_effect=[
subprocess.CalledProcessError(1, "squeue"), "TIMEOUT|"
])
self.assertEqual(controller.job_state(), "TIMEOUT")
def test_transient_query_error_retries(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(side_effect=[
subprocess.TimeoutExpired("ssh", 30), "TIMEOUT"
])
with patch("ngram_interactive_supervisor.time.sleep"):
self.assertEqual(controller.wait_for_allocation_exit(), "TIMEOUT")
def test_dead_trainer_in_running_allocation_does_not_renew(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
controller.job_state = MagicMock(return_value="RUNNING")
controller.allocate = MagicMock()
with patch("ngram_interactive_supervisor.time.monotonic", side_effect=[0, 61]):
self.assertEqual(controller.wait_for_allocation_exit(), "RUNNING")
controller.allocate.assert_not_called()
def test_restart_after_timeout_does_not_relaunch_old_segment(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
(controller.root / "segments" / controller.job).mkdir(parents=True)
controller.job_state = MagicMock(return_value="TIMEOUT")
controller.allocate = MagicMock(side_effect=RuntimeError("test reached renewal"))
with self.assertRaisesRegex(RuntimeError, "reached renewal"):
controller.run()
controller.start_segment.assert_not_called()
controller.allocate.assert_called_once()
def test_restart_after_cancellation_does_not_renew(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
(controller.root / "segments" / controller.job).mkdir(parents=True)
controller.job_state = MagicMock(return_value="CANCELLED")
controller.allocate = MagicMock()
with self.assertRaisesRegex(RuntimeError, "not renewing"):
controller.run()
controller.start_segment.assert_not_called()
controller.allocate.assert_not_called()
def test_allocation_command_and_fairshare(self):
with tempfile.TemporaryDirectory() as directory:
controller = self.controller(Path(directory))
(controller.root / "allocations").mkdir()
controller.remote = MagicMock(return_value="low|user|||||0.1|\nbest|user|||||0.9|")
def launch(command, **kwargs):
allocation = shlex.split(command[-1])
self.assertEqual(allocation[0], "salloc")
for flag in ("--account=best", "--nodes=1", "--gpus-per-node=4", "--time=04:00:00", "--qos=interactive"):
self.assertIn(flag, allocation)
(Path(allocation[-1]) / "job_id").write_text("90001\n")
return MagicMock()
with patch("ngram_interactive_supervisor.subprocess.Popen", side_effect=launch):
controller.allocate()
self.assertEqual(controller.job, "90001")
if __name__ == "__main__":
unittest.main()
|